python的先进制造技术工业场景模拟第六十六篇:读取机器人关节温度数据集,训练模型,识别关节电机过热故障类型。
周二下午,焊接机器人工作站。
"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解决实际问题,如果你觉得这个工具好用,欢迎关注长安牧笛!
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐


所有评论(0)