Keras K折交叉验证如何统一保存并绘制全量训练历史
K折交叉验证全量训练历史统一存储与绘制实现
直接append Keras返回的History对象无法直接序列化持久化,批量处理指标也不方便。History对象的核心指标数据存在它的history属性中,是纯Python字典结构(key为指标名,value为对应指标在每个epoch的数值列表),提取这部分数据统一归集,就可以实现单文件存储、批量绘图。
改造训练循环,结构化归集所有折训练记录
调整原有列表存储逻辑,每折训练完成后提取可序列化的指标字典,附加折号标记统一存入列表:
import pickle import matplotlib.pyplot as plt import numpy as np from sklearn.model_selection import KFold import tensorflow as tf num_folds = 10 # 固定随机种子保证可复现 kfold = KFold(n_splits=num_folds, shuffle=True, random_state=42) # 单列表承载全部折的训练记录 all_folds_history = [] for fold_no, (train_idx, val_idx) in enumerate(kfold.split(X, label), 1): print("---"*20) print(f"开始训练第 {fold_no}/{num_folds} 折") history = siamese.fit( [tf.gather(X[:,0], train_idx),tf.gather(X[:,1], train_idx)], tf.gather(label, train_idx), validation_data=([tf.gather(X[:,0], val_idx),tf.gather(X[:,1], val_idx)], tf.gather(label, val_idx)), batch_size=batch_size, epochs=epochs, ) # 仅存储可序列化的纯指标数据 fold_record = { "fold_id": fold_no, "metrics": history.history } all_folds_history.append(fold_record)
单文件持久化保存所有训练记录
不需要拆分多个独立文件,直接将归集好的全量记录序列化存储为单个文件即可:
# 写入单个文件 with open("kfold_all_training_history.pkl", "wb") as f: pickle.dump(all_folds_history, f) # 后续使用时仅需加载这一个文件 with open("kfold_all_training_history.pkl", "rb") as f: loaded_history = pickle.load(f)
如果偏好文本格式存储,也可以将字典转存为json格式,逻辑完全一致。
批量绘制所有折的训练指标曲线
加载单文件中的全量记录后,即可一次性绘制所有折的指标变化,也可以计算跨折平均指标做整体效果评估:
plt.figure(figsize=(12, 6)) epochs_range = range(1, epochs + 1) # 逐折绘制单折指标 for record in all_folds_history: metrics = record["metrics"] # 单折训练损失,淡色显示 plt.plot(epochs_range, metrics["loss"], color="#1f77b4", linestyle="--", alpha=0.25) # 单折验证损失,淡色显示 plt.plot(epochs_range, metrics["val_loss"], color="#ff7f0e", linestyle="-", alpha=0.25) # 绘制跨折平均指标,加粗突出 avg_train_loss = np.mean([rec["metrics"]["loss"] for rec in all_folds_history], axis=0) avg_val_loss = np.mean([rec["metrics"]["val_loss"] for rec in all_folds_history], axis=0) plt.plot(epochs_range, avg_train_loss, color="#1f77b4", linestyle="--", linewidth=2, label="平均训练损失") plt.plot(epochs_range, avg_val_loss, color="#ff7f0e", linestyle="-", linewidth=2, label="平均验证损失") plt.xlabel("训练轮次(Epoch)") plt.ylabel("损失值(Loss)") plt.title("10折交叉验证训练/验证损失曲线") plt.legend() plt.grid(alpha=0.3) plt.show()
准确率、AUC等其他指标的绘制逻辑完全相同,替换对应指标的key即可。
内容的提问来源于stack exchange,提问作者Ali ZareShahi
相关产品推荐
相关产品推荐

