直接继承tf.keras.callbacks.Callback类并重写on_train_batch_end等钩子方法即可;需注意batch级回调需显式设steps_per_epoch或用tf.data.Dataset,logs为只读字典,多GPU下仅chief worker执行,save_weights报错常因eager模式与格式不兼容。

怎么继承 tf.keras.callbacks.Callback 写自定义回调
直接继承 tf.keras.callbacks.Callback 类,重写需要的生命周期方法即可。它本身是个空类,不强制实现任何方法,但只有在训练/验证流程中被调用到的方法才会生效——比如你只写了 on_train_batch_end,那其他阶段(如 epoch 开始)就不会触发任何逻辑。
常见误操作是把自定义逻辑写在 __init__ 里就以为能运行,其实初始化只是准备,真正介入训练流靠的是这些钩子方法:
-
on_train_begin:整个model.fit()开始前执行一次 -
on_batch_begin/on_batch_end:每个 batch 前后(含 train/val) -
on_epoch_end:每个 epoch 完成后,logs参数里有loss、accuracy等指标值 -
on_train_end:训练彻底结束时(哪怕被EarlyStopping中断也会触发)
为什么 on_epoch_end 拿不到验证指标?
因为默认情况下,model.fit() 不会自动跑验证,除非你传了 validation_data 或 validation_split。即使加了,验证指标名也取决于你编译模型时用的指标名——比如你用 metrics=['acc'],日志里就是 val_acc,不是 val_accuracy;用 metrics=['accuracy'] 才是后者。
更隐蔽的问题是:如果你在 on_epoch_end 里直接读 logs['val_loss'],但某次 epoch 没跑验证(比如设置了 validation_freq=2),这个 key 就不存在,会抛 KeyError。
安全写法是:
def on_epoch_end(self, epoch, logs=None):
if logs is None:
logs = {}
val_loss = logs.get('val_loss') # 用 .get() 避免 KeyError
if val_loss is not None:
print(f"Epoch {epoch}: val_loss = {val_loss:.4f}")
如何让自定义 Callback 支持早停 + 模型保存联动?
Callback 之间默认独立,但你可以主动读取或修改 self.model,甚至调用 self.model.stop_training = True 来手动中断训练——这和 EarlyStopping 的原理一样。
典型联动场景:当验证 loss 连续 3 轮没下降,就保存当前最优权重,并终止训练:
class EarlySaveCallback(tf.keras.callbacks.Callback):
def __init__(self, filepath, patience=3):
super().__init__()
self.filepath = filepath
self.patience = patience
self.best_loss = float('inf')
self.wait = 0
<pre class="brush:php;toolbar:false;">def on_epoch_end(self, epoch, logs=None):
val_loss = logs.get('val_loss')
if val_loss is None:
return
if val_loss = self.patience:
self.model.stop_training = True # 主动终止
注意:self.model.save_weights() 保存的是权重,不是完整模型;如果想存结构+权重,得用 tf.keras.models.save_model(),但必须确保模型已 build(即第一次前向后)。
为什么自定义 Callback 在多 GPU 或 TPU 上行为异常?
因为 tf.distribute.Strategy 下,on_batch_end 等方法会在每个设备上都执行一次,而不是只在 host device 上执行。如果你在回调里写了文件写入、打印或更新全局变量,就会重复执行多次,甚至引发竞争或 I/O 冲突。
解决办法只有两个:
- 用
tf.distribute.get_strategy().cluster_resolver判断是否为主节点(但 TPU 不一定暴露该信息) - 更稳妥的是:所有副作用操作(如
print、open().write、np.save)只在on_train_begin和on_train_end里做,或者加锁 + 主机判断 - TensorFlow 官方推荐方式是:把中间结果存在
self.model.history.history或自定义属性里,等on_train_end统一处理
最常被忽略的一点:Callback 的实例在分布式训练中不是跨设备共享的,每个 replica 拿到的是自己副本,所以不能依赖实例属性做跨设备聚合计算。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











