如何获取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
相关产品推荐
相关产品推荐

