TensorFlow迁移学习中保存与恢复history变量的方法求助
解决方案
方法1:序列化History对象(修复dill/%store失效问题)
TensorFlow的History核心属性是history(训练指标字典)和epoch(训练轮次列表),直接序列化整个对象容易因版本兼容失效,拆分保存关键数据即可:
- 保存流程:
import pickle # 训练完成后,保存指标字典 with open('history_dict.pkl', 'wb') as f: pickle.dump(history.history, f) # 单独保存最后一轮epoch数,避免后续解析冗余 with open('last_epoch.txt', 'w') as f: f.write(str(history.epoch[-1]))
- 恢复流程:
import pickle from tensorflow.keras.callbacks import History # 初始化空History对象 restored_history = History() # 加载指标字典 with open('history_dict.pkl', 'rb') as f: restored_history.history = pickle.load(f) # 加载最后一轮epoch,生成完整epoch列表 with open('last_epoch.txt', 'r') as f: last_epoch = int(f.read()) restored_history.epoch = list(range(last_epoch + 1)) # 直接用于后续微调训练 fine_tune_epochs = 10 total_epochs = len(restored_history.epoch) + fine_tune_epochs history_tuned = model.fit( train_set, validation_data=dev_set, initial_epoch=restored_history.epoch[-1], epochs=total_epochs, verbose=2, callbacks=callbacks )
- 合并多轮训练的history(生成连贯曲线):
# 合并指标字典 for key in restored_history.history: restored_history.history[key].extend(history_tuned.history[key]) # 更新epoch列表 restored_history.epoch.extend(range(restored_history.epoch[-1]+1, total_epochs)) # 重新保存合并后的数据 with open('merged_history_dict.pkl', 'wb') as f: pickle.dump(restored_history.history, f) with open('last_epoch.txt', 'w') as f: f.write(str(total_epochs - 1))
方法2:基于CSVLogger的轻量方案
无需完全还原History对象,提取CSV日志中的关键信息即可继续训练,同时合并日志生成连贯曲线:
- 从CSV读取信息并继续训练:
import pandas as pd df = pd.read_csv('demo/logs/hist.log') last_epoch = df.shape[0] - 1 # CSV从epoch 0开始记录,行数-1即为最后一轮epoch数 fine_tune_epochs = 10 total_epochs = last_epoch + 1 + fine_tune_epochs history_tuned = model.fit( train_set, validation_data=dev_set, initial_epoch=last_epoch, epochs=total_epochs, verbose=2, callbacks=callbacks )
- 合并多轮训练日志(生成连贯曲线):
# 读取初始训练和微调的日志 df_initial = pd.read_csv('demo/logs/hist.log') df_fine_tune = pd.read_csv('demo/logs/fine_tune_hist.log') # 修正微调日志的epoch编号(叠加初始训练的总轮次) df_fine_tune['epoch'] = df_fine_tune['epoch'] + df_initial.shape[0] # 合并并保存 merged_df = pd.concat([df_initial, df_fine_tune], ignore_index=True) merged_df.to_csv('demo/logs/merged_hist.log', index=False)
后续绘图直接使用merged_df的指标列即可得到连贯曲线。
额外提示
- 使用pickle保存时,确保恢复环境的TensorFlow版本与保存时一致,避免类结构不兼容。
- 保存模型权重的同时,建议用文本文件记录
base_model.trainable状态及解冻的层范围,恢复时可快速还原模型配置。
内容的提问来源于stack exchange,提问作者Marios Constantinou
相关产品推荐
相关产品推荐

