周二下午,焊接机器人工作站。

"2 号机器人第 4 轴又跳了热保护,"维保老周蹲在电柜边,手背试了试伺服电机外壳,"上周跳了两次,每次停 15 分钟降温,产线节拍直接断。拆开看,轴承没坏、油脂也没干,就是'莫名其妙热'。"

我打开上位机导出的关节温度日志。

"这里面记了啥?"我问。

"每 10 秒一条,"老周说,"六个关节,每个关节的电机绕组温度、减速机壳体温度、电流、转速、负载率、环境温度、风扇状态、运行模式。全记着。但系统只做'温度>85℃就报警停机',跟电闸一样,烧到头才跳。"

"最头疼的是故障分不清,"老周补一句,"过热分好几种啊:轴承预紧过紧是慢热、绕组匝间弱是快热、润滑失效是温升带波动、散热风道堵是平台期上不去但一直高。现在报警就写个'过热',维修组来了还得拆半天猜。"

"我就想干一件事,"老周说,"把过去两个月六关节温度数据灌进去,模型不光说'要过热了',还得说'是哪种过热',比如'轴承预紧类概率 81%',我直接拿扳手去调预紧,不用拆电机。"

"关节过热不是突变,是'体温曲线有指纹',"我接话,"像人发烧,病毒感冒是陡升,慢性炎症是缓升带波动,中暑是平台高温。用 pandas 做窗口特征构造(温升速率、温差比、波动熵、负载热滞后),numpy 算滑动统计,scikit-learn 做多分类(RandomForest + 校准),scipy 做各类别置信区间+温升曲线拟合显著性,matplotlib 画六关节温度热力图+故障概率雷达+特征重要性+ROC多分类+温升指纹曲线+故障决策网络,networkx 建'温度特征→故障类型'链路。"

"对,"老周点头,"要能说清'为什么是轴承预紧不是绕组问题'。我拿去跟维修班长说'第4轴电机温升 0.8℃/min,减速机壳温比电机低 2℃,波动小,典型预紧过紧',班长就认。"

"用 pandas 做窗口特征,sklearn RF 多分类+概率校准,scipy 做曲线拟合 p 值,matplotlib 出 6 图,networkx 画链路,存 results/,"我开工程,"数据自包含,合成六关节 2 周 10 秒级数据,4 类过热故障,下载就能跑。"

敲了行原型:

# 目标: 从关节温度时序中识别过热故障类型(多分类+概率)

# 方法: 滑动窗口特征 + RF多分类 + 概率校准 + 温升指纹

# 输出: 故障类型概率 + 指纹曲线 + 维修建议

 

"完整版 OOP 封好,"我说,"数据加载器、窗口特征工程器、过热分类器、指纹分析器、可视化器、决策链路,输出故障类型概率+6图+报告。"

老周凑近看:"那以后看报告:J4 电机当前 78℃,10 分钟内升到 86℃概率 92%,故障类型:轴承预紧过紧 81% / 绕组弱 9% / 润滑失效 7% / 风道堵塞 3%。建议:停机调轴承预紧,无需拆绕组。提前 10 分钟预警,不用等跳保护。"

"对,"我接话,"机器人维护不是'烧了再修',是'看体温曲线认出病因'。数字孪生里挂关节热指纹节点,这套就是维保的'关节体温计+化验单'——不用等 85℃才跳电闸。"

 

一、实际应用场景(真实痛点)

 

场景设定:六轴工业机器人长期运行后,关节伺服电机及减速机出现过热。传统阈值报警(>85℃停机)属于"事后保护",且只报"过热"不区分类型。实际过热至少分四类:

 

1. 轴承预紧过紧:温升平缓、波动小、电机与壳体温差小

2. 绕组匝间弱:温升快、电流略高、波动小

3. 润滑失效:温升带周期性波动、壳体温度偏高

4. 散热风道堵塞:平台期高温、风扇状态异常、温升停滞但绝对值高

 

现场原话(叙事化):

 

"不是我们没装温度保护,"老周说,"85℃一跳,跟保险丝一样。但每次跳完,维修组拆开像开盲盒:有时候是轴承,有时候是风道,有时候啥都没坏就是热。上周 2 号机跳了两次,第一次拆了半天发现预紧紧了,第二次又怀疑绕组,结果也没事。停机时间比故障本身还贵。"

"最亏的是数据,"老周说,"10 秒一条,六关节全记,两个月攒了快百万条。就躺数据库里,系统只会比大小。我就想让它'认病',不是'报温度'。"

核心矛盾:"单阈值停机 + 故障不分类 + 数据沉睡" 与 "时序指纹提取 + 多分类概率校准 + 故障类型定位 + 提前预警维修" 之间的断层。

 

二、痛点分析(映射到滨州职业学院《先进制造技术》课程模型)

 

《先进制造技术》模块 本篇痛点对应

工业机器人技术基础:伺服驱动与关节传动维护 关节电机/减速机过热故障识别

智能制造与数字孪生:设备健康监控与预测性维护 关节热指纹建模 + PHM

先进制造技术基础:精度与可靠性理论 热变形对重复定位精度的影响

柔性制造系统FMS与先进生产管理:设备运维管理 按故障类型精准派工

 

一句话总结:我们需要一个"机器人关节温度时序→滑动窗口特征+RandomForest多分类+概率校准+温升指纹分析程序",用 

"pandas" 做窗口特征工程,

"numpy" 算滑动统计,

"scikit-learn" RF多分类+CalibratedClassifierCV,

"scipy" 做温升曲线拟合显著性+置信区间,

"matplotlib" 画六关节热力图+概率雷达+特征重要性+多分类ROC+指纹曲线+决策网络,

"networkx" 建推理链路,实现从"85℃才停机"到"数据驱动的故障类型识别+提前预警+精准派工"。

 

三、核心逻辑讲解(大白话)

 

3.1 问题本质:把关节想成"会发烧的人"

 

把机器人关节想成一个人:

 

* 电机绕组温度 = 腋下体温

* 减速机壳体温度 = 皮肤表面温度

* 温升速率 = 烧得有多急

* 波动情况 = 是不是一阵阵发冷发热

* 电机/壳体温差 = 体内外温差

* 你的目标 = 不是"超过 37.5℃就叫发烧",而是"看发烧曲线,判断是病毒/细菌/慢性炎症/中暑"

 

故障类型 体温指纹

轴承预紧过紧 缓升、稳、内外温差小

绕组匝间弱 陡升、电流高、很稳

润滑失效 升+周期波动、壳温高

风道堵塞 平台高温、风扇异常、不再涨

 

3.2 业务逻辑 → 代码映射

 

读取机器人关节温度时序(10秒级)

   │

   ▼ JointDataLoader (pandas)

读取表:

  时间戳, 机器人编号, 关节号(J1~J6),

  电机绕组温度, 减速机壳温, 电流A, 转速rpm,

  负载率%, 环境温度, 风扇状态(0/1), 运行模式

   │

   ▼ JointWindowFeature (pandas + numpy)

滑动窗口特征(窗口=30点=5分钟):

  温升速率 = polyfit(t, 电机温度, 1)斜率

  电机壳温差 = 电机温 - 壳温

  温度波动熵 = -∑p·log(p) 分箱后

  负载热滞后 = corr(负载, 温度, lag=3)

  电流温比 = 电流/电机温

  平台度 = 末段方差(判断是否停滞高温)

   │

   ▼ OverheatClassifier (sklearn)

多分类:

  RandomForestClassifier → 4类过热故障

  CalibratedClassifierCV → 概率校准

  交叉验证 → 每类F1 + 宏平均

  特征重要性 → 哪类特征最 discriminative

   │

   ▼ ThermalFingerprint (scipy)

温升指纹分析:

  对每关节做线性/指数拟合, 输出p值

  计算各类别预测概率的95% CI

  识别"指纹形态"匹配度

   │

   ▼ JointVisualizer (matplotlib + networkx)

可视化:

  1. 六关节温度热力图(时间×关节)

  2. 故障类型概率雷达图

  3. 特征重要性条形图

  4. 多分类 ROC(一对其余)

  5. 温升指纹曲线(4类典型曲线叠加)

  6. 故障识别决策网络

   │

   ▼ SyntheticJointData (numpy)

合成数据:

  6关节 × 2周 × 10秒级

  4类故障注入, 可复现

 

3.3 为什么不能"85℃停机"

 

视角 问题

单阈值 85℃ 烧到头才跳,已损伤绝缘

只报"过热" 维修开盲盒,拆错方向

温升指纹 看曲线形态提前 10 分钟认病

多分类概率 "预紧过紧 81%"直接派工

精准维修 调预紧/查绕组/换油脂/清风道

 

3.4 分析前后对比

 

维度 传统方式 本程序

触发时机 85℃停机 温升异常即预警(提前10min)

输出内容 "过热" 4类概率分布

维修动作 全拆排查 按主类精准处理

置信度 无 校准概率+CI

 

四、OOP 代码实现

 

4.1 项目结构

 

robot_joint_thermal/

├── robot_joint_thermal/

│ ├── __init__.py

│ ├── joint_data_loader.py # 数据加载

│ ├── joint_window_feature.py # 滑动窗口特征

│ ├── overheat_classifier.py # 多分类+校准

│ ├── thermal_fingerprint.py # 温升指纹(scipy)

│ ├── joint_visualizer.py # 可视化

│ └── synthetic_joint_data.py # 合成数据

├── tests/

│ ├── __init__.py

│ └── test_joint_thermal.py

├── results/

│ ├── joint_heatmap.png

│ ├── fault_radar.png

│ ├── feature_importance.png

│ ├── roc_multiclass.png

│ ├── thermal_fingerprint.png

│ ├── decision_network.png

│ ├── fault_detail.csv

│ └── thermal_report.txt

└── run_joint_thermal.py

 

4.2 核心源码

 

<details>

 

<summary></summary>

 

"""机器人关节温度数据加载器"""

 

import pandas as pd

from pathlib import Path

from typing import Optional

 

 

class JointDataLoader:

    """读取关节温度时序(10秒级)"""

 

    def __init__(self, filepath: str = "joint_temp_data.csv",

                 encoding: str = "utf-8"):

        self.filepath = Path(filepath)

        self.encoding = encoding

 

    def load(self) -> pd.DataFrame:

        if not self.filepath.exists():

            raise FileNotFoundError(self.filepath)

        df = pd.read_csv(self.filepath, encoding=self.encoding)

        req = ["timestamp", "robot_id", "joint_id",

               "motor_temp_c", "housing_temp_c"]

        miss = [c for c in req if c not in df.columns]

        if miss:

            raise ValueError(f"缺列: {miss}")

        num_cols = ["motor_temp_c", "housing_temp_c", "current_a",

                   "speed_rpm", "load_pct", "ambient_temp_c"]

        for c in num_cols:

            if c in df.columns:

                df[c] = pd.to_numeric(df[c], errors="coerce")

        df["timestamp"] = pd.to_datetime(df["timestamp"], errors="coerce")

        df = df.dropna(subset=["motor_temp_c"]).reset_index(subset=None, drop=True)

        return df

 

    def summary(self, df: pd.DataFrame) -> str:

        s = f"记录数: {len(df)}\n"

        s += f"机器人: {df['robot_id'].nunique()} 台\n"

        s += f"关节: {sorted(df['joint_id'].unique())}\n"

        s += (f"电机温度范围: {df['motor_temp_c'].min():.1f} ~ "

             f"{df['motor_temp_c'].max():.1f} ℃")

        return s

 

</details>

 

<details>

 

<summary></summary>

 

"""滑动窗口特征工程 (pandas + numpy)"""

 

import numpy as np

import pandas as pd

from typing import List

 

 

class JointWindowFeature:

    """按关节做滑动窗口特征提取"""

 

    def __init__(self, window: int = 30, lag: int = 3):

        self.window = window # 30点=5分钟(10秒级)

        self.lag = lag # 负载热滞后步数

 

    def _entropy(self, arr: np.ndarray, bins: int = 8) -> float:

        hist, _ = np.histogram(arr, bins=bins)

        p = hist / (hist.sum() + 1e-9)

        p = p[p > 0]

        return float(-np.sum(p * np.log(p)))

 

    def build(self, df: pd.DataFrame) -> pd.DataFrame:

        df = df.sort_values(["robot_id", "joint_id", "timestamp"]).reset_index(drop=True)

        rows = []

        for (rid, jid), g in df.groupby(["robot_id", "joint_id"]):

            g = g.reset_index(drop=True)

            mt = g["motor_temp_c"].values.astype(float)

            ht = g["housing_temp_c"].values.astype(float)

            cur = g["current_a"].values if "current_a" in g else np.ones(len(g))

            load = g["load_pct"].values if "load_pct" in g else np.ones(len(g))

            t = np.arange(len(g)).astype(float)

            for i in range(self.window, len(g) + 1):

                w_mt = mt[i-self.window:i]

                w_ht = ht[i-self.window:i]

                w_cur = cur[i-self.window:i]

                w_load = load[i-self.window:i]

                w_t = t[i-self.window:i]

                # 温升速率 ℃/min (10秒/点 → ×6)

                slope, _ = np.polyfit(w_t, w_mt, 1)

                rise_rate = slope * 6.0

                # 电机壳温差

                diff = w_mt[-1] - w_ht[-1]

                # 波动熵

                ent = self._entropy(w_mt)

                # 负载热滞后相关

                if i > self.lag:

                    lag_corr = np.corrcoef(load[i-self.window-self.lag:i-self.lag],

                                          w_mt)[0, 1]

                    if np.isnan(lag_corr):

                        lag_corr = 0.0

                else:

                    lag_corr = 0.0

                # 电流温比

                cur_temp_ratio = w_cur[-1] / (w_mt[-1] + 1e-9)

                # 平台度: 末1/3方差

                tail = w_mt[-10:]

                plateau = float(np.var(tail))

                # 末值

                rows.append({

                    "robot_id": rid,

                    "joint_id": jid,

                    "ts_end": g["timestamp"].iloc[i-1],

                    "motor_temp_end": float(w_mt[-1]),

                    "housing_temp_end": float(w_ht[-1]),

                    "rise_rate_c_per_min": float(rise_rate),

                    "motor_housing_diff": float(diff),

                    "temp_entropy": ent,

                    "load_thermal_lag_corr": float(lag_corr),

                    "current_temp_ratio": float(cur_temp_ratio),

                    "plateau_var": plateau,

                    "ambient_temp_c": float(g["ambient_temp_c"].iloc[i-1]) if "ambient_temp_c" in g else 25.0,

                    "fan_status": int(g["fan_status"].iloc[i-1]) if "fan_status" in g else 1,

                })

        feat = pd.DataFrame(rows)

        return feat

 

    def attach_label(self, feat: pd.DataFrame,

                    label_csv: pd.DataFrame = None,

                    rule_based: bool = True) -> pd.DataFrame:

        """

        若提供标注则用标注; 否则用物理规则生成弱标签(合成数据场景)

        4类: 0正常,1预紧过紧,2绕组弱,3润滑失效,4风道堵塞

        """

        if label_csv is not None:

            return feat.merge(label_csv, on=["robot_id", "joint_id"], how="left")

        if not rule_based:

            return feat

        out = feat.copy()

 

        def classify(r):

            if r["motor_temp_end"] < 70 and r["rise_rate_c_per_min"] < 0.5:

                return 0

            # 风道堵塞: 平台高温 + 风扇异常

            if r["plateau_var"] < 0.3 and r["motor_temp_end"] > 78 and r["fan_status"] == 0:

                return 4

            # 绕组弱: 陡升 + 电流温比高

            if r["rise_rate_c_per_min"] > 1.5 and r["current_temp_ratio"] > 0.25:

                return 2

            # 润滑失效: 波动熵高 + 壳温高

            if r["temp_entropy"] > 2.0 and r["motor_housing_diff"] < 3:

                return 3

            # 预紧过紧: 缓升 + 温差小 + 稳

            if 0.3 <= r["rise_rate_c_per_min"] <= 1.2 and r["motor_housing_diff"] < 4:

                return 1

            return 1 if r["motor_temp_end"] > 72 else 0

 

        out["fault_type"] = out.apply(classify, axis=1)

        return out

 

    def get_feature_columns(self) -> List[str]:

        return [

            "motor_temp_end", "housing_temp_end", "rise_rate_c_per_min",

            "motor_housing_diff", "temp_entropy", "load_thermal_lag_corr",

            "current_temp_ratio", "plateau_var", "ambient_temp_c",

        ]

 

</details>

 

<details>

 

<summary></summary>

 

"""过热故障多分类 (scikit-learn)"""

 

import numpy as np

from typing import Dict, List

from sklearn.ensemble import RandomForestClassifier

from sklearn.calibration import CalibratedClassifierCV

from sklearn.model_selection import cross_val_predict, StratifiedKFold

from sklearn.metrics import classification_report, roc_curve, auc

from sklearn.preprocessing import label_binarize

 

 

class OverheatClassifier:

    """4类过热故障 + 概率校准"""

 

    LABELS = {0: "normal", 1: "bearing_preload",

              2: "winding_weak", 3: "lubrication_fail",

              4: "duct_block"}

 

    def __init__(self, random_state: int = 42):

        self.random_state = random_state

        self.rf_ = None

        self.calib_ = None

 

    def fit(self, X: np.ndarray, y: np.ndarray) -> RandomForestClassifier:

        self.rf_ = RandomForestClassifier(

            n_estimators=300,

            max_depth=12,

            min_samples_leaf=3,

            class_weight="balanced",

            random_state=self.random_state,

            n_jobs=-1,

        )

        self.rf_.fit(X, y)

        return self.rf_

 

    def calibrate(self, X: np.ndarray, y: np.ndarray,

                 method: str = "isotonic") -> CalibratedClassifierCV:

        self.calib_ = CalibratedClassifierCV(self.rf_, method=method, cv=3)

        self.calib_.fit(X, y)

        return self.calib_

 

    def predict(self, X: np.ndarray) -> np.ndarray:

        m = self.calib_ or self.rf_

        return m.predict(X)

 

    def predict_proba(self, X: np.ndarray) -> np.ndarray:

        m = self.calib_ or self.rf_

        return m.predict_proba(X)

 

    def feature_importance(self, feature_names: List[str]) -> Dict:

        imp = dict(zip(feature_names, self.rf_.feature_importances_))

        return dict(sorted(imp.items(), key=lambda x: x[1], reverse=True))

 

    def evaluate(self, X: np.ndarray, y: np.ndarray) -> str:

        y_pred = self.rf_.predict(X)

        return classification_report(y, y_pred,

                                    target_names=[self.LABELS[i] for i in sorted(set(y))],

                                    zero_division=0)

 

    def roc_multiclass(self, X: np.ndarray, y: np.ndarray,

                      n_classes: int = 5) -> Dict:

        y_bin = label_binarize(y, classes=list(range(n_classes)))

        # 用未校准RF概率做ROC(演示)

        proba = self.rf_.predict_proba(X)

        res = {}

        for i in range(n_classes):

            fpr, tpr, _ = roc_curve(y_bin[:, i], proba[:, i])

            res[i] = {

                "label": self.LABELS[i],

                "fpr": fpr.tolist(), "tpr": tpr.tolist(),

                "auc": float(auc(fpr, tpr)),

            }

        return res

 

    def cross_val_macro_f1(self, X: np.ndarray, y: np.ndarray,

                          cv: int = 5) -> Dict:

        skf = StratifiedKFold(n_splits=cv, shuffle=True,

                             random_state=self.random_state)

        from sklearn.metrics import f1_score

        preds = cross_val_predict(self.rf_, X, y, cv=skf, method="predict")

        # 注意: 此处用全量模型近似, 演示用

        macro = f1_score(y, self.rf_.predict(X), average="macro")

        weighted = f1_score(y, self.rf_.predict(X), average="weighted")

        return {"macro_f1": float(macro), "weighted_f1": float(weighted)}

 

</details>

 

<details>

 

<summary></summary>

 

"""温升指纹分析 (scipy)"""

 

import numpy as np

from typing import Dict

from scipy import stats

from scipy.optimize import curve_fit

 

 

class ThermalFingerprint:

    """对温升曲线做指纹拟合"""

 

    def __init__(self):

        pass

 

    def fit_linear(self, t: np.ndarray, temp: np.ndarray) -> Dict:

        res = stats.linregress(t, temp)

        return {

            "slope": float(res.slope),

            "intercept": float(res.intercept),

            "r2": float(res.rvalue ** 2),

            "p_value": float(res.pvalue),

        }

 

    def fit_exponential(self, t: np.ndarray, temp: np.ndarray) -> Dict:

        def exp_func(x, a, b, c):

            return a * (1 - np.exp(-b * x)) + c

        try:

            popt, _ = curve_fit(exp_func, t, temp,

                              maxfev=5000,

                              p0=(10, 0.01, temp.min()))

            a, b, c = popt

            pred = exp_func(t, a, b, c)

            ss_res = np.sum((temp - pred) ** 2)

            ss_tot = np.sum((temp - temp.mean()) ** 2)

            r2 = 1 - ss_res / (ss_tot + 1e-9)

            # 显著性用线性残差做近似检验

            _, p = stats.ttest_1samp(temp - pred, 0)

            return {"asym": float(a), "rate": float(b),

                   "base": float(c), "r2": float(r2), "p_value": float(p)}

        except Exception:

            return {"asym": 0, "rate": 0, "base": float(temp.mean()),

                   "r2": 0, "p_value": 1.0}

 

    def classify_fingerprint(self, info: Dict,

                           rise_rate: float,

                           plateau_var: float) -> str:

        """根据拟合形态匹配故障指纹"""

        if info["r2"] > 0.95 and rise_rate > 1.5:

            return "winding_weak"

        elif info["r2"] > 0.9 and plateau_var < 0.3:

            return "duct_block"

        elif info["r2"] > 0.85 and rise_rate < 1.2:

            return "bearing_preload"

        else:

            return "lubrication_fail"

 

    def prob_confidence_interval(self, proba: np.ndarray,

                               confidence: float = 0.95,

                               n_boot: int = 500) -> Dict:

        """Bootstrap 计算主类概率置信区间"""

        rng = np.random.RandomState(0)

        cls = int(np.argmax(proba))

        p = proba[cls]

        samples = []

        for _ in range(n_boot):

            idx = rng.randint(0, len(proba) if len(proba.shape) == 1 else proba.shape[0],

                            size=min(20, proba.shape[0] if len(proba.shape) == 1 else proba.shape[0]))

            if len(proba.shape) == 1:

                samples.append(proba[cls])

            else:

                samples.append(np.mean(proba[:, cls]))

        alpha = 1 - confidence

        return {

            "class": cls,

            "point_prob": float(p),

            "ci_lower": float(np.percentile(samples, alpha/2*100)),

            "ci_upper": float(np.percentile(samples, (1-alpha/2)*100)),

        }

 

</details>

 

<details>

 

<summary></summary>

 

"""可视化 (matplotlib + networkx)"""

 

import numpy as np

import pandas as pd

import matplotlib.pyplot as plt

from pathlib import Path

import networkx as nx

from typing import Dict, List

 

plt.rcParams["font.sans-serif"] = ["SimHei", "DejaVu Sans"]

plt.rcParams["axes.unicode_minus"] = False

 

LABELS_CN = {

    0: "正常", 1: "轴承预紧过紧", 2: "绕组匝间弱",

    3: "润滑失效", 4: "风道堵塞",

}

 

 

class JointVisualizer:

    def __init__(self, results_dir: str = "results"):

        self.results_dir = Path(results_dir)

        self.results_dir.mkdir(exist_ok=True)

 

    def joint_heatmap(self, df: pd.DataFrame):

        """六关节温度热力图(时间×关节)"""

        pivot = df.pivot_table(index="joint_id", columns="ts_end",

                             values="motor_temp_end", aggfunc="mean")

        # 取前60列避免太密

        pivot = pivot.iloc[:, :60]

        fig, ax = plt.subplots(figsize=(13, 5))

        im = ax.imshow(pivot.values, cmap="inferno", aspect="auto")

        ax.set_yticks(range(len(pivot.index)))

        ax.set_yticklabels(pivot.index, fontsize=10)

        ax.set_xlabel("时间窗序号", fontsize=12)

        ax.set_ylabel("关节", fontsize=12)

        ax.set_title("六关节电机温度热力图", fontsize=13, fontweight="bold")

        plt.colorbar(im, ax=ax, label="电机温度 ℃")

        plt.tight_layout()

        plt.savefig(self.results_dir/"joint_heatmap.png",

                   dpi=150, bbox_inches="tight")

        plt.close()

 

    def fault_radar(self, proba: np.ndarray,

                   labels: List[int]):

        """故障类型概率雷达图"""

        fig = plt.figure(figsize=(8, 8))

        ax = plt.subplot(111, polar=True)

        n = len(labels)

        angles = np.linspace(0, 2*np.pi, n, endpoint=False).tolist()

        angles += angles[:1]

        vals = [proba[i] for i in labels]

        vals += vals[:1]

        names = [LABELS_CN[i] for i in labels] + [LABELS_CN[labels[0]]]

        ax.plot(angles, vals, "b-", linewidth=2.5)

        ax.fill(angles, vals, alpha=0.25, color="#3498DB")

        ax.set_xticks(angles[:-1])

        ax.set_xticklabels([LABELS_CN[i] for i in labels], fontsize=10)

        ax.set_ylim(0, 1)

        ax.set_title("过热故障类型概率雷达",

                    fontsize=13, fontweight="bold", pad=20)

        plt.tight_layout()

        plt.savefig(self.results_dir/"fault_radar.png",

                   dpi=150, bbox_inches="tight")

        plt.close()

 

    def feature_importance(self, importance: Dict):

        fig, ax = plt.subplots(figsize=(10, 6))

        names = list(importance.keys())[:8]

        vals = [importance[n] for n in names]

        ax.barh(range(len(names)), vals[::-1], color=plt.cm.viridis(np.array(vals[::-1])/max(vals)),

               edgecolor="black", height=0.6)

        ax.set_yticks(range(len(names)))

        ax.set_yticklabels(names[::-1], fontsize=10)

        ax.set_xlabel("特征重要性")

        ax.set_title("过热故障识别驱动特征", fontsize=13, fontweight="bold")

        ax.grid(axis="x", alpha=0.3)

        plt.tight_layout()

        plt.savefig(self.results_dir/"feature_importance.png",

                   dpi=150, bbox_inches="tight")

        plt.close()

 

    def roc_multiclass_plot(self, roc_data: Dict):

        fig, ax = plt.subplots(figsize=(9, 9))

        colors = ["#3498DB", "#E74C3C", "#27AE60", "#F39C12", "#9B59B6"]

        for i, (k, v) in enumerate(roc

利用AI解决实际问题,如果你觉得这个工具好用,欢迎关注长安牧笛!

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐