You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

鸢尾花分类器train.py代码改造:保存混淆矩阵与训练模型

鸢尾花分类器:添加混淆矩阵保存与模型持久化功能

前置依赖安装

先确保安装所需依赖:

pip install matplotlib scikit-learn joblib

修改后的train.py完整代码示例

import os
import joblib
import matplotlib.pyplot as plt
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay

# 创建outputs目录(不存在则自动创建)
output_dir = "outputs"
os.makedirs(output_dir, exist_ok=True)

# 加载鸢尾花数据集并拆分训练/测试集
iris = load_iris()
X_train, X_test, y_train, y_test = train_test_split(iris.data, iris.target, test_size=0.2, random_state=42)

# 训练分类模型(可替换为你正在使用的模型)
model = RandomForestClassifier(n_estimators=100, random_state=42)
model.fit(X_train, y_train)

# 1. 生成并保存混淆矩阵为PNG图片
y_pred = model.predict(X_test)
cm = confusion_matrix(y_test, y_pred)
disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=iris.target_names)
disp.plot(cmap=plt.cm.Blues)
plt.title("Iris Classification Confusion Matrix")
plt.savefig(os.path.join(output_dir, "confusion_matrix.png"))
plt.close()  # 关闭绘图释放内存资源

# 2. 使用joblib保存训练完成的模型
model_path = os.path.join(output_dir, "iris_classifier_model.joblib")
joblib.dump(model, model_path)

print(f"混淆矩阵已保存至: {os.path.join(output_dir, 'confusion_matrix.png')}")
print(f"模型已保存至: {model_path}")

关键代码说明

  • 目录处理:os.makedirs(output_dir, exist_ok=True)确保outputs目录存在,避免保存文件时出现路径不存在的报错。
  • 混淆矩阵保存:借助sklearn.metrics的ConfusionMatrixDisplay快速可视化混淆矩阵,通过plt.savefig()将图片写入指定目录,最后用plt.close()释放内存。
  • 模型持久化:joblib.dump(model, model_path)将训练好的模型序列化存储,后续可通过joblib.load(model_path)直接加载模型进行预测。

内容的提问来源于stack exchange,提问作者Esther

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.01 20:34:50