You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在Keras多轮训练中保留Epoch编号以实现TensorBoard连续绘图?

解决Keras中TensorBoard训练步数不延续的问题

我之前也遇到过这个烦人的问题——每次重启训练,TensorBoard里的曲线就从头开始画,根本没法连贯看趋势。其实在Keras里有几种靠谱的解决方式,不用切换到原生TF那么麻烦:

方法一:自定义全局步数变量+自定义回调

这种方式最灵活,能完全掌控步数的计数逻辑:

  1. 先定义一个全局的步数变量,用来跨训练轮次跟踪累计步数:
import tensorflow as tf
global_step = tf.Variable(0, trainable=False, dtype=tf.int64)
  1. 自定义一个回调,在每个batch训练完成后更新这个全局步数:
class UpdateGlobalStep(tf.keras.callbacks.Callback):
    def on_train_batch_end(self, batch, logs=None):
        global global_step
        global_step.assign_add(1)
  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)
  1. 训练时要记得加载之前的模型权重和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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.26 10:08:51