如何在Keras多轮训练中保留Epoch编号以实现TensorBoard连续绘图?
解决Keras中TensorBoard训练步数不延续的问题
我之前也遇到过这个烦人的问题——每次重启训练,TensorBoard里的曲线就从头开始画,根本没法连贯看趋势。其实在Keras里有几种靠谱的解决方式,不用切换到原生TF那么麻烦:
方法一:自定义全局步数变量+自定义回调
这种方式最灵活,能完全掌控步数的计数逻辑:
- 先定义一个全局的步数变量,用来跨训练轮次跟踪累计步数:
import tensorflow as tf global_step = tf.Variable(0, trainable=False, dtype=tf.int64)
- 自定义一个回调,在每个batch训练完成后更新这个全局步数:
class UpdateGlobalStep(tf.keras.callbacks.Callback): def on_train_batch_end(self, batch, logs=None): global global_step global_step.assign_add(1)
- 重写TensorBoard回调,让它使用我们的全局步数来记录日志,替代Keras默认的步数生成逻辑:
class CustomTensorBoard(tf.keras.callbacks.TensorBoard): def __init__(self, log_dir, **kwargs): super().__init__(log_dir, **kwargs) self.global_step = global_step def on_train_batch_end(self, batch, logs=None): logs = logs or {} # 手动写入标量日志,绑定全局步数 with tf.summary.create_file_writer(self.log_dir).as_default(): for name, value in logs.items(): if name in ['loss', 'accuracy']: # 根据你的模型指标自行调整 tf.summary.scalar(name, value, step=self.global_step) super().on_train_batch_end(batch, logs)
- 训练时要记得加载之前的模型权重和
global_step的数值,同时把两个回调都加入训练流程:
# 用Checkpoint保存和恢复模型与全局步数 checkpoint = tf.train.Checkpoint(model=your_model, global_step=global_step) checkpoint.restore(tf.train.latest_checkpoint('./checkpoints')) # 启动训练,initial_epoch设为上一次结束的epoch数 your_model.fit( train_data, epochs=100, initial_epoch=last_recorded_epoch, callbacks=[UpdateGlobalStep(), CustomTensorBoard('./logs')] )
方法二:利用initial_epoch+复用日志目录(简单但有局限)
如果你的训练只是中断后续训,没有多次从头启动的情况,可以试试这个轻量化方法:
- 每次续训时,指定
initial_epoch参数为上一次训练结束的epoch数 - 保持TensorBoard的
log_dir不变,不要每次训练都更换目录
这种方式下,Keras会自动延续epoch的计数,但注意步数(step)仍会从每个epoch的0开始,图表里的step是epoch内的相对步数,而非全局累计步数。如果需要全局累计步数,这种方法就不适用了。
方法三:直接使用TensorFlow原生Summary Writer(贴近原生TF逻辑)
如果你熟悉原生TF的写法,也可以完全绕过Keras的TensorBoard回调,手动控制日志写入:
writer = tf.summary.create_file_writer('./logs') class CustomSummaryCallback(tf.keras.callbacks.Callback): def __init__(self, writer): self.writer = writer self.global_step = global_step def on_train_batch_end(self, batch, logs=None): logs = logs or {} with self.writer.as_default(): tf.summary.scalar('loss', logs['loss'], step=self.global_step) tf.summary.scalar('accuracy', logs['accuracy'], step=self.global_step) self.global_step.assign_add(1) # 训练时加入自定义回调 your_model.fit( train_data, epochs=100, initial_epoch=last_recorded_epoch, callbacks=[CustomSummaryCallback(writer)] )
这种方式和原生TF的逻辑完全一致,适合需要高度定制日志内容的场景。
总结一下:如果需要全局累计步数,方法一和方法三是最可靠的;如果只是需要epoch计数延续,方法二足够用。另外一定要记得保存和恢复global_step变量,不然还是会从头开始计数。
内容的提问来源于stack exchange,提问作者Tlokuus
相关产品推荐
相关产品推荐

