3.3 compile、fit、evaluate:训练三步的内部调度 本节摘要:compile 把损失、优化器、指标装配成一张训练步规程;fit 把训练步跟踪成图并逐批调度,途中自动完成数据迭代、梯度回放、变量写入、指标累计与回调唤醒;evaluate 则是关掉求导与更新的只读执行。本节沿一次完整的训练调用逐层拆解内部时序,并教你用自定义指标与 verbose 输出验证每一步确实发生了。这是本章枢纽:理解了 fit 的内部,后面所有配件的接入点都不言自明。 读完你应当能做到 阅读完本节,你应当能够: 逐条列出 fit 一次迭代内部发生的事件序列; 解释指标对象与损失的区别:瞬时值与跨批次累计值; 用自定义指标接入 fit,验证指标的更新与重置时机;
本节摘要:compile 把损失、优化器、指标装配成一张训练步规程;fit 把训练步跟踪成图并逐批调度,途中自动完成数据迭代、梯度回放、变量写入、指标累计与回调唤醒;evaluate 则是关掉求导与更新的只读执行。本节沿一次完整的训练调用逐层拆解内部时序,并教你用自定义指标与 verbose 输出验证每一步确实发生了。这是本章枢纽:理解了 fit 的内部,后面所有配件的接入点都不言自明。
阅读完本节,你应当能够:
把 fit 的每轮拆开,内部事件按时序是:从数据管道拉一批;调用训练步函数(已被 tf.function 跟踪成图);训练步内依次做前向算预测、按 compile 装配的损失函数算误差、GradientTape 沿图求梯度、优化器按策略更新变量;然后各指标对象用本批的预测与真值做累计更新;回调按节拍被唤醒(见 3.7 节);进度条按 verbose 设置输出。跑完一个 epoch,指标被重置,验证阶段以只读模式跑一遍并产出 val 指标。

进度条上每轮滚动的 loss 数字,常被误读为"当前批次的损失"。它是 Metric 对象的跨批次累计值——内部用变量维护滑动统计,每批调一次 update_state,epoch 结束 reset_state。这个机制解释了两个现象:进度条数字平滑下降而单批损失可能剧烈波动;epoch 摘要里的 loss 与手算单批损失对不上。验证机制最直接的办法是挂一个自定义指标:
import tensorflow as tf import numpy as np class BatchCounter(tf.keras.metrics.Metric): """统计被喂过的批次数,验证 fit 的调度节拍""" def __init__(self, name="batch_counter", **kwargs): super().__init__(name=name, **kwargs) self.count = self.add_weight(name="count", initializer="zeros") def update_state(self, y_true, y_pred, sample_weight=None): self.count.assign_add(1.0) # 每批被叫醒一次 def result(self): return self.count def reset_state(self): self.count.assign(0.0) # epoch 边界归零 X = np.random.rand(320, 8).astype("float32") y = np.random.rand(320).astype("float32") model = tf.keras.Sequential([tf.keras.layers.Input(shape=(8,)), tf.keras.layers.Dense(1)]) model.compile(optimizer="adam", loss="mse", metrics=[BatchCounter()]) model.fit(X, y, epochs=2, batch_size=64, verbose=2) # 输出: # Epoch 1/2 # 5/5 - 1s - loss: 0.3317 - batch_counter: 5.0000 # Epoch 2/2 # 5/5 - 0s - loss: 0.2958 - batch_counter: 5.0000 # 320 除以 64 等于 5 批:计数器每轮都是 5, # 且第二轮从 0 重新累计——reset_state 在 epoch 边界生效
计数器每轮精确等于批次数、第二轮不累加,这两个观测把"逐批调度、边界重置"两个机制坐实了。
compile 的三个参数各自对应训练步的一个槽位:optimizer 填更新槽,loss 填误差槽,metrics 填观测槽。装配发生在 compile 时刻,而不是 fit 时刻——所以改损失要重新 compile。loss 可以传字符串(映射到内置函数)或可调用对象;多输出模型可传字典按输出名分别指定:
inputs = tf.keras.Input(shape=(8,)) out = tf.keras.layers.Dense(1, name="price")(inputs) dual = tf.keras.Model(inputs, out) dual.compile( optimizer="rmsprop", loss={"price": "mse"}, # 按输出名指定损失 metrics={"price": [tf.keras.metrics.MAE]}, # 平均绝对误差一并观测 ) print("compiled with per-output slots") # 输出:compiled with per-output slots # 多输出结构的装配方式,第 4 章多任务场景会复用
validation_data 与 validation_split 都能产出 val 指标,但调度细节不同。validation_split 是从输入数组尾部切走一段(不洗牌!),所以喂进来前必须先自己打乱;validation_data 接受独立的数组或管道,验证管道按第 2 章标准配(不打乱、大批次、prefetch)。验证阶段以只读模式跑训练步:求导与变量写入被关闸,前向与指标照常——这正对应 evaluate 的语义。
# 正确的验证姿势:数据先洗,管道再喂 n = len(X) idx = np.random.permutation(n) Xs, ys = X[idx], y[idx] cut = int(n * 0.8) model.compile(optimizer="adam", loss="mse") h = model.fit(Xs[:cut], ys[:cut], validation_data=(Xs[cut:], ys[cut:]), epochs=3, batch_size=64, verbose=0) print([round(v, 3) for v in h.history["val_loss"]]) # 输出示例:[0.312, 0.301, 0.298] # 若不先打乱就 validation_split,切走的是未洗数据尾部, # 分布偏斜会让 val_loss 异常偏高
⚠️ 常见坑:validation_split 直接用在未打乱的时序数据上,切走的是最近时段——验证集分布与训练分布系统性错位,val 曲线看起来"糟"其实只是偏。时序数据改用时间前后切分加 validation_data。
💡 关键直觉:fit 是被编排好的"逐批三步舞",你能插队的唯一入口是回调与自定义指标——前者在舞步之间动手,后者在舞步之内记账。3.6 与 3.7 节的武器全部挂在这两个入口上。
下节细看训练步里的更新槽:优化器。