如何在TensorFlow中延续训练并续写至同一历史文件
TensorFlow模型断点续训、历史记录续写及全局最优模型保存方案
一、训练历史文件可以续写,操作方法如下
训练时model.fit()返回的History对象会记录每轮的指标数据,我们可以把这些数据序列化保存,续训时加载旧历史并合并新数据,实现完整追踪训练演化过程。
1. 首次训练保存历史
训练完成后将历史数据存为JSON或CSV文件,示例用JSON:
import json import tensorflow as tf # 示例模型构建 model = tf.keras.Sequential([ tf.keras.layers.Dense(64, activation='relu', input_shape=(10,)), tf.keras.layers.Dense(1) ]) model.compile(optimizer='adam', loss='mse') # 首次训练 history = model.fit(X_train, y_train, epochs=10) # 保存历史到文件 with open('training_history.json', 'w') as f: json.dump(history.history, f) # 同时保存当前模型,方便后续续训 model.save('current_model.h5')
2. 续训时加载并续写历史
加载之前的模型和历史文件,续训后合并新旧数据再保存:
import json import tensorflow as tf # 加载之前的模型 model = tf.keras.models.load_model('current_model.h5') # 加载旧历史数据 with open('training_history.json', 'r') as f: old_history = json.load(f) # 用新数据续训,记录新的训练历史 new_history = model.fit(X_new_train, y_new_train, epochs=5) # 合并历史:把每个指标的新数据追加到旧列表末尾 for metric in old_history.keys(): old_history[metric].extend(new_history.history[metric]) # 覆盖保存合并后的历史文件 with open('training_history.json', 'w') as f: json.dump(old_history, f) # 同时更新当前模型的保存 model.save('current_model.h5')
如果习惯用CSV,逻辑类似:首次训练后用pandas把history.history转为DataFrame存成CSV,续训后读取CSV,将新的指标数据追加进去再保存。
二、保存全局最优模型的实现方法
默认的ModelCheckpoint只会保存单次训练中的最优模型,要实现全局最优(跨多次续训的最优),可以自定义回调函数,追踪历史上的最佳性能,只有当当前模型超过历史最优时才保存。
1. 自定义全局最优保存回调
import numpy as np import tensorflow as tf from tensorflow.keras.callbacks import Callback class GlobalBestCheckpoint(Callback): def __init__(self, save_path, monitor='val_loss', mode='min'): super().__init__() self.save_path = save_path self.monitor = monitor # 监控的指标,比如val_loss、val_accuracy self.mode = mode # min表示指标越小越好,max表示越大越好 # 初始化全局最优值:min模式设为无穷大,max模式设为负无穷 self.best_score = np.inf if mode == 'min' else -np.inf # 从文件加载之前的最优值,避免续训时重置 try: with open('best_score.txt', 'r') as f: self.best_score = float(f.read()) except FileNotFoundError: pass def on_epoch_end(self, epoch, logs=None): current_score = logs.get(self.monitor) if not current_score: return # 判断当前性能是否优于全局最优 if (self.mode == 'min' and current_score < self.best_score) or \ (self.mode == 'max' and current_score > self.best_score): self.best_score = current_score # 保存最优模型 self.model.save(self.save_path) # 同时保存当前最优值,方便后续续训加载 with open('best_score.txt', 'w') as f: f.write(str(self.best_score)) print(f"\n全局最优模型已更新,保存至 {self.save_path}")
2. 训练(含续训)时使用该回调
不管是首次训练还是续训,都加载这个回调,它会自动跨训练周期追踪全局最优:
# 首次训练示例 model = tf.keras.Sequential([ tf.keras.layers.Dense(64, activation='relu', input_shape=(10,)), tf.keras.layers.Dense(1) ]) model.compile(optimizer='adam', loss='mse', metrics=['mae']) # 初始化全局最优回调,监控验证损失,越小越好 global_checkpoint = GlobalBestCheckpoint('global_best_model.h5', monitor='val_loss', mode='min') # 首次训练,带验证集 history = model.fit( X_train, y_train, epochs=10, validation_data=(X_val, y_val), callbacks=[global_checkpoint] ) model.save('current_model.h5') # 续训示例 model = tf.keras.models.load_model('current_model.h5') # 同样加载全局最优回调,它会自动读取之前保存的best_score.txt global_checkpoint = GlobalBestCheckpoint('global_best_model.h5', monitor='val_loss', mode='min') # 用新数据续训 new_history = model.fit( X_new_train, y_new_train, epochs=5, validation_data=(X_new_val, y_new_val), callbacks=[global_checkpoint] ) model.save('current_model.h5')
这样即使某次续训的数据集质量差,只要模型性能没超过历史最优,就不会覆盖之前的全局最优模型,完全满足你的需求。
内容的提问来源于stack exchange,提问作者Jonathan Roy
相关产品推荐
相关产品推荐

