3.7 回调函数:训练过程的监控哨


文档摘要

3.7 回调函数:训练过程的监控哨 本节摘要:回调是 fit 留给你的官方干预口:在 epoch 或 batch 的边界节拍上被唤醒,能读日志、改参数、停训练、存检查点。本节先给出完整的唤醒节拍表,再实战三件内置武器——ModelCheckpoint 存档、EarlyStopping 早停、ReduceLROnPlateau 降学习率,最后给出自定义回调模板。3.6 节的早停与调度器都是回调的用户,这一节揭开它们共同的底座。 读完你应当能做到 阅读完本节,你应当能够: 画出回调各方法与训练节拍的对应关系,知道每类干预该挂在哪个钩子; 用 ModelCheckpoint 在验证指标改善时自动存最优权重; 组合 EarlyStopping 与 checkpoint,实现"停了还是最好的";

3.7 回调函数:训练过程的监控哨

本节摘要:回调是 fit 留给你的官方干预口:在 epoch 或 batch 的边界节拍上被唤醒,能读日志、改参数、停训练、存检查点。本节先给出完整的唤醒节拍表,再实战三件内置武器——ModelCheckpoint 存档、EarlyStopping 早停、ReduceLROnPlateau 降学习率,最后给出自定义回调模板。3.6 节的早停与调度器都是回调的用户,这一节揭开它们共同的底座。

读完你应当能做到

阅读完本节,你应当能够:

  1. 画出回调各方法与训练节拍的对应关系,知道每类干预该挂在哪个钩子;
  2. 用 ModelCheckpoint 在验证指标改善时自动存最优权重;
  3. 组合 EarlyStopping 与 checkpoint,实现"停了还是最好的";
  4. 写自定义回调在训练中读取梯度或动态调整逻辑。

唤醒节拍表

回调对象传入 fit 的 callbacks 参数列表后,Keras 在训练的关键节拍上依次调用它的钩子方法。节拍从粗到细:训练级(on_train_begin 或 end)、epoch 级(on_epoch_begin 或 end)、批次级(on_train_batch_begin 或 end)、验证级(on_test_begin 或 end 等)。logs 字典携带当前指标快照——训练级与验证级节拍里正是靠它读取 val_loss。模型、优化器都能经 self.model 访问,这意味着回调有权改学习率、有权停训练、有权写任何变量。

图 15 回调唤醒节拍与 fit 主循环的关系

图 15 回调唤醒节拍与 fit 主循环的关系

三件内置武器

武器一,ModelCheckpoint:按条件把权重(或整模型)写到磁盘。最常用的姿势是"验证指标每次改善才存"——训练跑完,磁盘上永远是最优的那一步:

import tensorflow as tf import numpy as np from sklearn.datasets import fetch_california_housing housing = fetch_california_housing() X, y = housing.data.astype("float32"), housing.target.astype("float32") m = tf.keras.Sequential([tf.keras.layers.Input(shape=(8,)), tf.keras.layers.Dense(64, activation="relu"), tf.keras.layers.Dense(1)]) m.compile(optimizer="adam", loss="mse") ckpt = tf.keras.callbacks.ModelCheckpoint( filepath="best.weights.h5", save_weights_only=True, # 只存权重:小、快、可移植 monitor="val_loss", save_best_only=True, # 指标改善才覆盖 ) hist = m.fit(X[:8000], y[:8000], validation_data=(X[8000:9000], y[8000:9000]), epochs=40, batch_size=64, verbose=0, callbacks=[ckpt]) print("best weights on disk") # 输出:best weights on disk # 训练结束后加载: m.load_weights("best.weights.h5") final = m.evaluate(X[9000:10000], y[9000:10000], verbose=0) print(f"test loss with best weights: {final:.3f}") # 输出示例:test loss with best weights: 0.573

武器二与武器三在 3.6 与 3.4 节已实战过:EarlyStopping 监控 val_loss、patience 到期停止并回滚最佳权重;ReduceLROnPlateau 监控改善停滞、自动降学习率。三者常一起出场,各管一件事:checkpoint 管存、early stop 管停、plateau 管调。

combo = [ tf.keras.callbacks.ModelCheckpoint("best.weights.h5", save_weights_only=True, monitor="val_loss", save_best_only=True), tf.keras.callbacks.EarlyStopping(monitor="val_loss", patience=8, restore_best_weights=True), tf.keras.callbacks.ReduceLROnPlateau(monitor="val_loss", factor=0.5, patience=3), ] m.compile(optimizer=tf.keras.optimizers.Adam(0.005), loss="mse") h = m.fit(X[:8000], y[:8000], validation_data=(X[8000:9000], y[8000:9000]), epochs=200, batch_size=64, verbose=0, callbacks=combo) print(f"ran {len(h.history['loss'])} epochs, best val {min(h.history['val_loss']):.3f}") # 参考输出:ran 37 epochs, best val 0.512 # 三件套联动:降档救过几次,八轮耐心耗尽后停机并回滚

自定义回调:把 fit 的内部摊给你看

自定义回调继承 Callback 类,按需覆盖钩子。一个实用例子:每 N 轮打印学习率与训练验证差距,早停预警一目了然:

class ProgressWatch(tf.keras.callbacks.Callback): def __init__(self, every=5): super().__init__() self.every = every def on_epoch_end(self, epoch, logs=None): logs = logs or {} if (epoch + 1) % self.every == 0: gap = logs.get("val_loss", 0) - logs.get("loss", 0) lr = float(self.model.optimizer.learning_rate) print(f"epoch {epoch + 1}: gap {gap:.3f}, lr {lr:.5f}") watch = ProgressWatch(every=10) m.fit(X[:8000], y[:8000], validation_data=(X[8000:9000], y[8000:9000]), epochs=30, batch_size=64, verbose=0, callbacks=[watch]) # 输出示例: # epoch 10: gap 0.062, lr 0.00500 # epoch 20: gap 0.091, lr 0.00250 —— plateau 已降过一档 # epoch 30: gap 0.044, lr 0.00250 # 差距走扩就是过拟合预警,lr 变化印证调度器在起作用

批次级钩子做梯度观测也顺手——在 on_train_batch_end 里读优化器的 iterations 变量做采样,配合 5.1 节的手写循环可以交叉验证训练行为。唯一的纪律是批次级钩子必须轻:每个批次都被叫醒,里面放重活会拖垮整条训练。

⚠️ 常见坑:EarlyStopping 监控的指标在 logs 里不存在(比如监控 val_acc 但没在 metrics 里装配它),回调静默不工作或直接报错。挂回调前确认监控名与 compile 的 metrics、验证配置对得上。

💡 关键直觉:回调是 fit 主动让出的操作台。凡是"训练途中要做的事"——存档、刹车、调参、观测——都先想"能不能写成回调",而不是把 fit 拆成手写循环。

本节要点回顾

  • 节拍表:训练级、epoch 级、批次级、验证级,logs 携带指标快照。
  • checkpoint:monitor 加 save_best_only,磁盘上永远是历史最优。
  • 三件套组合:存、停、调各司其职,参数互不冲突。
  • 自定义回调:继承 Callback 覆盖钩子,可读日志、可改优化器、可停训练。
  • 批次级钩子轻量化:高频唤醒的钩子里不放重活。

下节解决训练结束后的交接:保存与加载。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U