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

如何获取TensorFlow/Keras触发早停Callback时对应的训练Epoch编号

获取TensorFlow回调触发训练终止时对应Epoch编号的方法

方案1:直接使用fit返回的History对象读取(最简单)

你调用model.fit()得到的hist对象自带已完成Epoch的记录,不需要修改原有回调逻辑,只需在训练结束后读取即可:

# 训练结束后添加如下代码
# hist.epoch 是所有已完成Epoch的编号列表,默认从0开始计数,加1可匹配日志输出的编号
stopped_epoch = hist.epoch[-1] + 1
print(f"训练触发终止时的Epoch编号:{stopped_epoch}")

你给出的示例运行后,上述代码会输出训练触发终止时的Epoch编号:66,和运行日志完全一致。

方案2:在自定义回调中新增属性记录(灵活性更高)

如果需要在回调触发终止时即时获取Epoch编号,或者需要在回调内部做后续逻辑处理,可以修改自定义回调类,新增属性存储当前Epoch值:

class stopAtLossValue(tf.keras.callbacks.Callback):
    def __init__(self):
        super().__init__()
        self.current_epoch = 0

    def on_epoch_begin(self, epoch, logs=None):
        # 每个Epoch开始时更新当前Epoch编号
        self.current_epoch = epoch

    def on_batch_end(self, batch, logs={}):
        eps = 0.01 
        if logs.get('loss') <= eps:
            self.model.stop_training = True
            # 触发停止时可直接读取,和日志匹配需要加1
            print(f"训练在第{self.current_epoch + 1}轮触发终止")

完整可运行修改后示例代码

import numpy as np
import tensorflow as tf
from tensorflow import keras 


class stopAtLossValue(tf.keras.callbacks.Callback):
    def __init__(self):
        super().__init__()
        self.current_epoch = 0

    def on_epoch_begin(self, epoch, logs=None):
        self.current_epoch = epoch

    def on_batch_end(self, batch, logs={}):
        eps = 0.01 
        if logs.get('loss') <= eps:
            self.model.stop_training = True
            print(f"\n训练在第{self.current_epoch + 1}轮触发终止,当前loss:{logs.get('loss'):.4f}")
                    
training_input=  np.random.random([30,10])
training_output = np.random.random([30,1])

model = tf.keras.Sequential([  
    tf.keras.layers.Flatten(input_shape=(10,)),
    tf.keras.layers.Dense(15,activation=tf.keras.activations.linear),
    tf.keras.layers.Dense(15, activation='relu'),
    tf.keras.layers.Dense(1) 
])   
                               
model.compile(loss="mse",optimizer = tf.keras.optimizers.Adam(learning_rate=0.01))

hist = model.fit(training_input, training_output, epochs=100, batch_size=100,  verbose=1, callbacks=[stopAtLossValue()])

# 方案1的读取方式
stopped_epoch = hist.epoch[-1] + 1
print(f"通过History对象读取到的终止Epoch编号:{stopped_epoch}")

内容的提问来源于stack exchange,提问作者sergey_208

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 09:06:03