如何获取TensorFlow模型权重保存对应的训练轮次(epoch)?
如何确定TensorFlow模型权重对应的训练轮次?
问题场景
你通过ModelCheckpoint保存模型权重(仅保存权重),加载时仅输出CheckpointLoadStatus对象,无法直接获取该权重对应的训练轮次,需要解决这个问题以继续训练。
你的输出信息:
<tensorflow.python.checkpoint.checkpoint.CheckpointLoadStatus object at 0x7b3e7e6e5000>
相关代码:
# 保存权重 checkpoint_filepath = '/content/pamap2-fedavg-128' checkpoint_callback = tf.keras.callbacks.ModelCheckpoint( filepath=checkpoint_filepath, save_weights_only=True, monitor='acc', mode='auto', save_best_only=False, verbose=1, ) # 加载权重 checkpoint_filepath = '/content/pamap2-fedavg-128' print(global_model.load_weights(checkpoint_filepath))
原因分析
当前代码中filepath未包含训练轮次(epoch)的占位符,每次保存权重都会覆盖之前的文件;且仅保存权重时,TensorFlow默认不会在checkpoint文件中存储epoch元数据,因此无法直接从现有权重文件读取对应轮次。
解决方案
1. 从训练日志中查找
因为你设置了verbose=1,训练时ModelCheckpoint会在控制台输出类似内容:
Epoch 00010: saving model to /content/pamap2-fedavg-128
找到最后一次输出该信息的日志,对应的epoch就是当前权重的保存轮次。
2. 修改保存代码(避免后续再出现此问题)
修改ModelCheckpoint的filepath,加入{epoch}占位符,让保存的文件名直接包含epoch号:
checkpoint_filepath = '/content/pamap2-fedavg-128-epoch-{epoch:02d}' checkpoint_callback = tf.keras.callbacks.ModelCheckpoint( filepath=checkpoint_filepath, save_weights_only=True, monitor='acc', mode='auto', save_best_only=False, verbose=1, )
保存后的文件会命名为pamap2-fedavg-128-epoch-10这类格式,直接就能看到对应的训练轮次。
3. 额外记录epoch到单独文件
如果需要更灵活的记录方式,可以在训练循环中额外保存当前epoch到文本文件:
# 在训练循环的epoch结束后执行 with open('/content/latest_epoch.txt', 'w') as f: f.write(str(current_epoch))
后续加载权重时,读取该文件即可获取对应的epoch。
总结
如果是已保存的权重,优先去训练日志中查找最后一次保存的epoch;如果日志丢失,可能只能从已知的训练节点重新开始。后续务必修改保存代码,在文件名中加入epoch占位符,或者额外记录epoch信息,避免再次出现此类问题。
内容的提问来源于stack exchange,提问作者Ariaeimehr
相关产品推荐
相关产品推荐

