我最早意识到model.fit()不是万能钥匙,是在做一个多任务模型。模型有 9 个输出头,每个任务的 loss 权重需要动态调整,其中一个任务的梯度还要做全局归一化——这些东西在fit()的固定流程里根本塞不进去。后来我把训练循环拆开重写,发现真正的深度学习工作流不是"调参-看指标-结束",而是从建模、训练、验证到部署,每一环都有可以精细控制的地方。TensorFlow Keras 的高级 API 魅力也正在于此:它不是让你丢掉model.fit(),而是让你在需要的时候,有底气和能力写出自己的训练逻辑。
这篇文章主要面向已经跑通过model.fit()、想进一步掌握自定义训练循环、函数式建模、分布式并行和混合精度等现代 Keras 工作流的深度学习开发者。我会按自己的实际踩坑顺序来写——先讲清楚fit()到底封装了什么、什么场景必须超越它;再讲三种建模 API 的选型逻辑;然后手把手拆一个自定义训练循环;接着讲提速和分布式的细节;最后给一份从fit()迁移到高级工作流的排错清单。
1. 被 model.fit() 隐藏的工作流:先搞清楚"超越"到底在超什么
1.1 model.fit() 背后蒙住你的那一层黑盒
很多人对 Keras 的认知停留在"搭积木 + fit",但fit()远不是一句话那么简单。我后来翻源码才意识到,它一揽子帮你干了这些事:数据批次自动迭代(numpy 数组会被 shuffle、按 batch 切分、按 epoch 重复);自动切换training=True/False的前向逻辑;自动调用反向传播去计算所有trainable_variables的梯度;自动执行优化器的apply_gradients;自动维护并重置各类指标;自动把回调事件派发到on_epoch_begin、on_batch_end这些钩子里;甚至验证集的评估也是在同一个流程里顺手就做的。
这些封装的视角是"让训练这件事的默认路径最好用",而不是"让训练逻辑任意可定制"。所以当你只需要一个常规监督学习模型,fit()确实是最省心的入口。真正的麻烦在于:一旦你的训练逻辑和这条固定流水线冲突,你会发现没有地方插入你想要的逻辑。你没法在一个 batch 内部对样本做 MixUp 后把混合前的梯度也保留下来,没法在判别器和生成器之间切换更新频率,也没法在某个 epoch 之后动态改写损失权重。
1.2 什么场景下 fit 真的不够用
我自己总结了几类必须跳出fit()的典型场景,不完整但很常见:
多损失动态组合。多任务学习里,任务 A 的 loss 可能是任务 B loss 的注意力权重;或者你今天要用加权 Focal Loss,明天要切成对比 Loss,这种动态计算在
fit()里只能靠自定义 loss 函数勉强做,但代价是代码变得很绕。对抗式训练。GAN 的结构是两个子网络交替更新,判别器每步更新一次,生成器可能两步才更新一次,而且两个网络共享数据流但不共享梯度。
fit()拿这种结构没办法。Batch 内部的自定义操作。SimCLR 这类对比学习需要在 batch 内计算一个大的相似度矩阵并对对角线以外的负样本做归一化;MixUp 需要把两批样本线性插值。这些操作本质上是"在 step 级别对数据做变换",
fit()默认流程无法插入。精细的梯度控制。比如只冻结前 5 层、只更新
BatchNorm的非 gamma/beta 参数、或者对梯度做全局范数裁剪(不是简单 clip_norm,而是跨层联合裁剪)。特殊的优化策略。Lookahead、LARS、SWA 这类优化器,需要定期拷贝参数快照或做参数平均,
fit()里的标准optimizer.apply_gradients流程没办法优雅地插入这些额外步骤。
1.3 分清"必须手写循环"和"换个 API 就行"
我见过不少同学一上来就把训练循环重写了,其实很多东西根本用不着手写。我的判断标准是分三档:
- 大约 80% 的场景,
model.fit()就够用。比如普通图像分类、结构化数据拟合、单任务回归。 - 大约 15% 的场景,用函数式 API 的多输入多输出 + 自定义损失 + 自定义回调可以解决,不用动训练循环。
- 只有剩下 5% 的场景,才需要真正手写
train_step。
这条决策线很重要,因为手写循环意味着你自己负责处理验证集评估、checkpoint 保存、早停、TensorBoard 日志等一系列工程问题。如果你只是为了"显得高级"而重写训练循环,通常只会给自己增加一堆 Debug 成本。
2. Keras 三种建模 API 的选型逻辑:建模方式决定后续灵活性的上限
2.1 Sequential:适合固定层堆叠的基线模型
Sequential是最简单的一种,适合层与层之间是线性堆叠的模型。它的优点一眼就能看明白:代码最短,可读性最好,别人接手你的代码几乎没有理解成本。
model = tf.keras.Sequential([ layers.Dense(64, activation="relu"), layers.Dense(64, activation="relu"), layers.Dense(10, activation="softmax"), ])但它也是最"死"的:输入输出只有一个,网络结构必须是一条链。你没法让某个分支跳过一层,也没法画出一个残差连接。所以我的建议是:只有当你确定这就是一个基线模型时,才用Sequential。一旦模型里出现分叉、合并、共享层,就不要再硬塞了。
2.2 函数式 API:拓扑结构灵活但仍是"静态图"风格
函数式 API 是绝大多数人应该默认选择的建模方式。它的核心思路是:把每一层的输出当作普通张量传来传去,最后用Model(inputs=..., outputs=...)把整个计算图绑定起来。
inputs = tf.keras.Input(shape=(784,)) x = layers.Dense(64, activation="relu")(inputs) x = layers.Dense(64, activation="relu")(x) outputs = layers.Dense(10, activation="softmax")(x) model = tf.keras.Model(inputs=inputs, outputs=outputs)同样是全连接网络,函数式 API 的写法比Sequential多了几行,但换来的是拓扑自由:你可以做残差连接、共享层、多输入多输出、特征拼接。
input_a = tf.keras.Input(shape=(32,)) input_b = tf.keras.Input(shape=(64,)) shared = layers.Dense(16, activation="relu") a_out = shared(input_a) b_out = shared(input_b) concat = layers.concatenate([a_out, b_out]) outputs = layers.Dense(1)(concat) model = tf.keras.Model(inputs=[input_a, input_b], outputs=outputs)函数式 API 之所以叫"函数式",是因为它构造的是一个显式的静态计算图,层与层之间的依赖关系在模型构建那一刻就确定了。这意味着 Keras 可以做很多自动优化:比如自动推导 shape、在模型保存时序列化完整结构、通过plot_model()画出清晰的拓扑图。什么时候轮到 Model 子类化?当你发现某个结构无法用静态画图的方式表达时。
2.3 Model 子类化:动态计算图的自由度
Model 子类化是三种方式里最灵活的,它把模型定义变成了一段普通的 Python 代码。你继承tf.keras.Model,在__init__里定义子模块,在call()里写完整的前向逻辑。
class MyModel(tf.keras.Model): def __init__(self): super().__init__() self.dense1 = layers.Dense(64, activation="relu") self.dense2 = layers.Dense(64, activation="relu") self.out = layers.Dense(10, activation="softmax") def call(self, inputs, training=False): x = self.dense1(inputs) x = self.dense2(x) return self.out(x)这种方式的优点很直接:前向逻辑里可以用if、for、循环来动态控制计算路径,甚至可以依赖外部输入来决定走哪条分支。它特别适合需要运行时决策的模型,比如根据输入的序列长度决定 RNN 展开多少步,或者根据当前训练阶段决定是否启用辅助分类头。
但自由是要付出代价的。函数式 API 替你维护了一张完整的计算图图结构,Keras 知道每个中间张量的来源;而子类化模型的内部结构对 Keras 是"黑盒",它只知道输入输出,没法自动做结构层面的可视化或序列化。你需要用model.get_layer()或手动记录子层来解决这些问题。
三种建模方式的选择经验:
| 建模方式 | 灵活性 | 可读性 | 静态结构可视化 | 适用场景 |
|---|---|---|---|---|
| Sequential | 低 | 高 | 支持 | 线性堆叠的基线模型 |
| 函数式 API | 中 | 高 | 支持 | 残差、多输入输出、共享层 |
| Model 子类化 | 高 | 中 | 不支持 | 动态计算路径、复杂控制流 |
选型的核心逻辑是:先满足灵活性下限,再去选一个尽可能简单的建模方式。能用Sequential就不升级,但一旦模型出现分叉或共享层,立刻切到函数式 API;只有函数式 API 表达不了时才用子类化。
3. 自定义训练循环的完整骨架:从梯度带到底层指标更新
3.1 tf.GradientTape 的作用域与梯度计算原理
自定义训练循环的核心是tf.GradientTape。你可以把它理解成一台"计算录像机":在with块里执行的所有张量操作,只要与可训练变量相关,梯度的传播路径都会被录下来。等 forward 结束,调用tape.gradient(loss, model.trainable_variables)就会根据录像带反向回放,算出每个变量对应的梯度。
optimizer = tf.keras.optimizers.Adam(learning_rate=1e-3) loss_fn = tf.keras.losses.SparseCategoricalCrossentropy() train_acc = tf.keras.metrics.SparseCategoricalAccuracy() def train_step(x, y): with tf.GradientTape() as tape: logits = model(x, training=True) loss = loss_fn(y, logits) grads = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables)) train_acc.update_state(y, logits) return loss有一个容易忽略的点:GradientTape默认只录制tf.Variable之间的操作路径,默认不会把tf.Tensor的中间计算全部保留下来。如果你需要计算某个中间张量对输入的梯度,得手动调用tape.watch(x)。另外,GradientTape 每次调用后资源就释放了,如果你需要同时算两路梯度(比如 GAN 里生成器和判别器各自的反向),最稳妥的做法是开两个独立的tape,分别做前向和反向。
3.2 一个可落地的 train_step 函数加上验证循环
有了train_step,还需要验证循环。验证循环不需要梯度,所以不需要包GradientTape,但要注意training=False和两个开关的组合,它会影响 BatchNorm 和 Dropout 的行为:
training=True时,BatchNorm 使用当前 batch 的均值和方差更新全局统计量,Dropout 启用。training=False时,BatchNorm 使用累积的全局统计量,Dropout 关闭。
完整的训练 + 验证逻辑长这样:
@tf.function def test_step(x, y): logits = model(x, training=False) loss = loss_fn(y, logits) val_acc.update_state(y, logits) return loss for epoch in range(epochs): for x_batch, y_batch in train_ds: train_step(x_batch, y_batch) for x_val, y_val in val_ds: test_step(x_val, y_val) print(f"epoch {epoch}: acc={train_acc.result().numpy():.4f}, val_acc={val_acc.result().numpy():.4f}") train_acc.reset_states() val_acc.reset_states()这段代码虽然短,但已经具备了一个训练循环的完整骨架:训练步、验证步、指标统计、epoch 日志。接下来要做的所有高级工作流都在这个骨架上扩展。
3.3 指标状态管理:为什么每个 epoch 要 reset_states
新手写自定义循环最容易踩的坑就是忘了reset_states()。tf.keras.metrics.Accuracy这类指标是有状态的,它的状态会累积所有历史 batch 的统计量。如果你不在每个 epoch 结束时重置,下一轮的准确率会把上一轮的数字也算进去,数字会一直平滑但永远不会"新鲜"。
这里我习惯用result()取当前值、reset_states()清空状态、再在下个 epoch 重新累积。写fit()时这些是自动完成的,手写循环里必须自己来。另一个问题是"训练指标热启动":如果你的启发式学习率调整依赖验证 loss,但训练 loss 还是上一轮的累积值,早停判断就会出错。所以我建议在每一轮验证之前就重置训练指标,保证日志里的训练指标只反映当前 epoch。
3.4 接上 checkpoint 与早停,自己把"回调"干活的能力补回来
既然绕过了fit()的自动回调分发,就不能两手空空。我最常用的方案是配合tf.train.Checkpoint做断点保存:
ckpt = tf.train.Checkpoint(step=tf.Variable(1), optimizer=optimizer, net=model) manager = tf.train.CheckpointManager(ckpt, "./tf_ckpts", max_to_keep=3) for epoch in range(epochs): for x_batch, y_batch in train_ds: train_step(x_batch, y_batch) ckpt.step.assign_add(1) if int(ckpt.step) % 1000 == 0: manager.save() # 每个 epoch 结束时也保存一次 manager.save()早停逻辑同样可以手动实现:记录当前最佳验证指标,连续若干 epoch 没有改善就 break,同时恢复上一次最佳 checkpoint。这套逻辑并不复杂,但你必须真的写出来——这就是"超越 fit"的代价和收益。
4. 让自定义训练更快的工程细节:tf.function、混合精度与数据管道
4.1 用 @tf.function 把训练步编译成静态图
上面示例里的train_step如果直接运行,是 Python 逐行解释执行的,速度可能比fit()慢不少。要让自定义循环真正跑出fit()的水平,需要把训练步用@tf.function装饰起来,让 TensorFlow 把它编译成静态计算图。
@tf.function def train_step(x, y): with tf.GradientTape() as tape: logits = model(x, training=True) loss = loss_fn(y, logits) grads = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables)) train_acc.update_state(y, logits) return loss初看只是加了个装饰器,但背后的变化非常大:TensorFlow 会把model.call()里的 Python 操作转换成图节点,通常能带来数倍的提速。这里有一个关键约束:被tf.function包裹的函数里,Python 原生对象(比如dict、list、int)都必须遵守 TensorFlow 的类型规范。最常见的坑是if tensor_value:这种判断没法直接用,你得改用tf.cond或tf.while_loop;另外print也需要用tf.print,否则只能在首次执行时打印一次。
我自己习惯的做法是:先不装饰,用纯 Python 跑通逻辑;确认没问题后再加@tf.function做加速。这样能在定位问题时少浪费很多时间。
4.2 混合精度策略:在自定义循环中别忘了 loss scaling
现代 GPU(如 V100、A100、H100)对 FP16 的吞吐量明显高于 FP32。Keras 在fit()里开启混合精度,只需要设置全局策略:
from tensorflow.keras import mixed_precision policy = mixed_precision.set_global_policy("mixed_float16")但在自定义训练循环里,你还需要处理一个细节:loss scaling。FP16 的数值范围有限,梯度在反向传播时可能会下溢到 0。TensorFlow 的优化器可以自动做 loss scaling,但需要你把 loss 放大后再求梯度,更新完成后再把梯度缩小回去。在自定义循环里,手动设置优化器:
optimizer = tf.keras.optimizers.Adam(learning_rate=1e-3) optimizer = mixed_precision.LossScaleOptimizer(optimizer) def train_step(x, y): with tf.GradientTape() as tape: logits = model(x, training=True) loss = loss_fn(y, logits) loss = optimizer.get_scaled_loss(loss) scaled_grads = tape.gradient(loss, model.trainable_variables) grads = optimizer.get_unscaled_gradients(scaled_grads) optimizer.apply_gradients(zip(grads, model.trainable_variables))get_scaled_loss会在内部放大梯度,get_unscaled_gradients再把它缩回来。这个逻辑在fit()里是自动的,自定义循环里必须显式写,不然混合精度模式下的模型会有很大的概率出现梯度消失,表现就是 loss 卡住不降。
4.3 tf.data 管道:手动 shuffle、batch、prefetch 的正确姿势
自定义训练配合tf.data.Dataset几乎是必须的,因为fit()内部默认的数据迭代方式只适合小数据集,真正的大规模训练必须把数据管道做好。
train_ds = tf.data.Dataset.from_tensor_slices((x_train, y_train)) train_ds = train_ds.shuffle(10000).batch(256).prefetch(tf.data.AUTOTUNE)这三个操作分别对应三个层面的优化:shuffle打破样本顺序避免模型学到伪模式;batch把样本组装成矩阵提高 GPU 利用率;prefetch让 CPU 在 GPU 计算的同时提前准备下一批数据,避免 GPU 空闲等待。
还有一个容易被忽视的点:prefetch最好放在所有数据变换之后,它的作用是把数据准备和模型执行解耦,放到越靠后效果越好。如果你用了map做图像解码和数据增强,那map之后就马上接prefetch是最优选择。
4.4 分布式策略:MirroredStrategy 下自定义循环只需改三行
当单卡显存不够,或者多卡可以显著缩短训练时间时,你会用tf.distribute.MirroredStrategy。它的用法非常机械化:在创建模型、优化器、Dataset 之前声明一个 strategy 上下文,然后将它们都放进with strategy.scope():中。
strategy = tf.distribute.MirroredStrategy() with strategy.scope(): model = create_model() optimizer = tf.keras.optimizers.Adam(learning_rate=1e-3) train_ds = ... train_ds = strategy.experimental_distribute_dataset(train_ds)自定义循环中需要把原本的 train_step 再用strategy.run包一层:
@tf.function def distributed_train_step(x, y): per_replica_losses = strategy.run(train_step, args=(x, y)) return strategy.reduce(tf.distribute.ReduceOp.SUM, per_replica_losses, axis=None)strategy.run会把同一个函数复制到每个 GPU 上,每个 GPU 各自计算一个 batch 的梯度,最后再统一合并。值得留意的是,BatchNorm 在分布式下会同步全局统计量,如果你模型里大量使用 BatchNorm,需要把reduction相关的细节确认好,否则不同卡的 BN 统计量不一致会导致验证指标抖动。
5. 从 fit 迁移到自定义循环后的排错清单与回调复用
5.1 常见运行时报错:梯度为 None、验证指标不更新、Dataset 重复迭代
从fit()切到自定义循环,第一批报错几乎可以开一个"经典大礼包"。
梯度为 None:最常见的原因是某个层的参数没有参与 loss 的计算。比如你定义了一个自编码器,但 forward 时只用了 encoder,decoder 参数没有参与到 loss 路径里,
tape.gradient对没参与计算的变量会返回None。还有一种是模型里的某些变量被tf.stop_gradient断开了。遇到None梯度时,建议逐个检查zip(grads, model.trainable_variables),过滤出grad is None的变量名。验证指标不更新:十有八九是你在
test_step里也用了@tf.function,但指标对象的update_state在静态图下只在第一次 trace 时执行了一次。这种问题很隐蔽,我的排查思路是暂时去掉test_step的装饰器,如果指标开始更新,再检查是不是变量捕获或控制流导致tf.function只 trace 了一次。Dataset 重复迭代报错:手写循环时经常想重复遍历一个 Dataset,但 Dataset 在 Python 里通常只能消费一次。解决办法是每次重开一个
iter,或者用dataset.repeat()配合steps_per_epoch控制长度。更稳妥的方案是每个 epoch 都重新构建 Dataset,这样 shuffle 和 prefetch 的效果也更好。
5.2 复用 Keras 回调:LambdaCallback 这个万能接口
手写训练循环后,Keras 原生的EarlyStopping、ReduceLROnPlateau、ModelCheckpoint依然可以用。一个比较省力的方案是把这些回调的效果自己封装进循环里,但我个人更推荐保留 Keras 回调机制:手写循环时,你可以用tf.keras.callbacks.LambdaCallback来对接已有回调。
from tensorflow.keras.callbacks import LambdaCallback, EarlyStopping def on_epoch_end(epoch, logs): logs["custom_lr"] = optimizer.learning_rate.numpy() logs["train_acc"] = train_acc.result().numpy() callback = LambdaCallback(on_epoch_end=on_epoch_end)然后你在每个 epoch 结束时手动构造一个logs字典传给各个回调。这样你既保留了EarlyStopping的优雅逻辑,又不耽误自定义循环的灵活性。不过我要提醒一句:这些回调在fit()里的触发时序是固定的,手写循环里你需要自己保证在每个 epoch 结束或每个 batch 结束时调用对应接口,顺序不对照样会出现"早停没生效"这类诡异问题。
5.3 渐进式迁移:先换数据管道,再换回调,最后换训练循环
如果你想从fit()平滑过渡到高级工作流,我不建议直接把所有逻辑删了重写。比较稳妥的路线是:
- 先用
tf.data.Dataset替换fit()里的 numpy 数据输入,跑通数据管道。 - 把监控需求用自定义回调接住,提前摸清各个回调事件触发的时机。
- 保持
fit()不变,但内部把model替换成一个子类化模型或函数式模型,确认逻辑正确。 - 最后才把
fit()替换成自定义训练循环,此时需要移植的只剩下训练步和验证步两部分。
这样每一步都有可验证的中间状态,出问题时很容易定位到底是数据的问题、模型结构的问题还是训练策略的问题。
写到这里,我自己也很感慨:model.fit()是 Keras 最友好的名片,但一张名片毕竟撑不起一整个深度学习项目。从自定义 train_step 到混合精度再到分布式训练,每一步都需要自己动手去补 Keras 的封装留下的"空位"。这些空位不是设计缺陷,而是刻意给你留的控制接口。当你把训练循环真正握在手里之后,再回头看fit(),你会更理解它的便利,也更清楚自己的模型到底需要哪种工作流。
最后分享一个自己的使用习惯:生产项目里,我会给训练代码建一个engine.py,把train_step、test_step、checkpoint 管理、日志汇总都封装成类,主脚本只负责配置参数和启动训练。这样既享受了自定义训练的自由度,又不会因为逻辑散落各处而难以维护。你要是有兴趣,下次可以聊聊怎么给这套自定义训练引擎加上完整的实验管理(指标追踪、配置记录、模型版本),它会把整个工作流的可靠性再拉高一个档次。