鸢尾花分类器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
相关产品推荐
相关产品推荐

