5.1 自定义训练循环:亲手接管 GradientTape 本节摘要:GradientTape 是 TensorFlow 自动微分的操作台:前向计算被它录制,gradient 调用沿录像回放求导。手写训练循环不是炫技——GAN 的交替更新、对比学习的成对采样、多任务的自定义梯度,fit 都装不下,只有接管循环才能实现。本节拆解 tape 的录制范围与求导机制,给出一个与 fit 行为对齐的完整循环模板(指标累计、检查点、进度输出齐全),并以 4.4 节的 GAN 为例演示非标流程的写法。 读完你应当能做到 阅读完本节,你应当能够: 解释 tape 的录制边界:什么被记录、什么被漏掉、persistent 何时需要; 写出完整的自定义训练循环,逐件补齐 fit 自动提供的能力;
本节摘要:GradientTape 是 TensorFlow 自动微分的操作台:前向计算被它录制,gradient 调用沿录像回放求导。手写训练循环不是炫技——GAN 的交替更新、对比学习的成对采样、多任务的自定义梯度,fit 都装不下,只有接管循环才能实现。本节拆解 tape 的录制范围与求导机制,给出一个与 fit 行为对齐的完整循环模板(指标累计、检查点、进度输出齐全),并以 4.4 节的 GAN 为例演示非标流程的写法。
阅读完本节,你应当能够:
GradientTape 的规矩用排程语言说:进 tape 上下文的每次张量运算都被登记进"录像带",gradient 调用时框架沿录像反向回放,逐算子套用它的求导说明书(1.4 节提过每个算子自带说明书)。两个常被漏掉的机制细节:录制只覆盖上下文内的运算——tape 外算的前向(或用 tape 外变量算的部分)不被记录,gradient 返回 None;默认整盘录像只许回放一次——要对同一前向对多个对象求导(如对抗训练同时对真伪两个头求导),要开 persistent=True 并手动释放。
import tensorflow as tf w = tf.Variable(3.0) b = tf.Variable(1.0) # 细节一:录制范围 x = tf.constant(2.0) with tf.GradientTape() as tape: y = w * x # 在 tape 内:被录制 z = y + b # 在 tape 外:这一步不在录像带里 g = tape.gradient(z, [w, b]) print("grad w:", float(g[0]), "grad b:", g[1]) # 输出:grad w: 2.0 grad b: None # b 的梯度是 None——加 b 的那步在 tape 外,无法回放 # 细节二:persistent 与多次求导 with tf.GradientTape(persistent=True) as tape: y = w * x g_w = tape.gradient(y, w) g_x = tape.gradient(y, x) # 对常量求输入梯度,研究分析常用 del tape # 手动释放,否则资源滞留 print("grad w:", float(g_w), "grad x:", float(g_x)) # 输出:grad w: 2.0 grad x: 3.0 # 同一次录制回放两次:w 的梯度 2,x 的梯度 3(y 对 x 的导数即 w)
watch_accessed_variables 默认为 True:上下文内读到的可训练变量自动入带。显式控制用 tape.watch(t)——对非常量张量求导时必须手动 watch。
手写循环与 fit 的差距不在梯度计算(那五行),而在 fit 顺带做的工程活:指标累计、epoch 重置、验证只读、检查点、进度显示。模板把它们逐件摆出来:
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") ds = tf.data.Dataset.from_tensor_slices((X[:8000], y[:8000])).shuffle(8000).batch(64) val = tf.data.Dataset.from_tensor_slices((X[8000:9000], y[8000:9000])).batch(256) model = tf.keras.Sequential([tf.keras.layers.Input(shape=(8,)), tf.keras.layers.Dense(64, activation="relu"), tf.keras.layers.Dense(1)]) opt = tf.optimizers.Adam(1e-3) loss_fn = tf.keras.losses.MeanSquaredError() train_mae = tf.keras.metrics.MeanAbsoluteError() # 指标对象自己记账 val_mae = tf.keras.metrics.MeanAbsoluteError() ckpt = tf.train.Checkpoint(model=model, optimizer=opt) manager = tf.train.CheckpointManager(ckpt, "ckpt_dir", max_to_keep=2) @tf.function # 训练步整体跟踪成图 def train_step(xb, yb): with tf.GradientTape() as tape: pred = model(xb, training=True) loss = loss_fn(yb, pred) grads = tape.gradient(loss, model.trainable_variables) opt.apply_gradients(zip(grads, model.trainable_variables)) train_mae.update_state(yb, pred) # 每批记账 return loss @tf.function def val_step(xb, yb): pred = model(xb, training=False) # 只读前向 val_mae.update_state(yb, pred) for epoch in range(5): train_mae.reset_state(); val_mae.reset_state() for xb, yb in ds: loss = train_step(xb, yb) for xb, yb in val: val_step(xb, yb) if epoch % 2 == 0: manager.save() # 阶段性存档 print(f"epoch {epoch + 1}: train mae {train_mae.result():.3f}, " f"val mae {val_mae.result():.3f}") # 输出示例: # epoch 1: train mae 0.842, val mae 0.801 # epoch 3: train mae 0.561, val mae 0.548 # epoch 5: train mae 0.524, val mae 0.523 # 收敛轨迹与 fit 一致——机制相同,只是排程权在你手里
模板里有三处值得停顿的对应关系:指标对象用 update_state 与 reset_state(3.3 节的机制原样复用);训练步被 @tf.function 整体跟踪(1.5 节的图模式);检查点显式管理(3.7 节 ModelCheckpoint 的底层)。
接管前先证明你排的和 Keras 排的一样。对齐手段:同结构模型、同数据、同优化器与学习率、同批次数,各跑一个 epoch,比对损失轨迹:
# fit 版(同一模型重新编译以重置变量) model2 = tf.keras.Sequential([tf.keras.layers.Input(shape=(8,)), tf.keras.layers.Dense(64, activation="relu"), tf.keras.layers.Dense(1)]) model2.compile(optimizer=tf.optimizers.Adam(1e-3), loss=tf.keras.losses.MeanSquaredError()) h = model2.fit(X[:8000], y[:8000], epochs=1, batch_size=64, verbose=0) print(f"fit-style loss {h.history['loss'][0]:.4f}") # 手写版首 epoch 的训练损失 model3 = tf.keras.Sequential([tf.keras.layers.Input(shape=(8,)), tf.keras.layers.Dense(64, activation="relu"), tf.keras.layers.Dense(1)]) opt3 = tf.optimizers.Adam(1e-3) ds3 = tf.data.Dataset.from_tensor_slices((X[:8000], y[:8000])).batch(64) acc_loss, n = 0.0, 0 for xb, yb in ds3: with tf.GradientTape() as tape: l = loss_fn(yb, model3(xb, training=True)) grads = tape.gradient(l, model3.trainable_variables) opt3.apply_gradients(zip(grads, model3.trainable_variables)) acc_loss += float(l); n += 1 print(f"loop-style loss {acc_loss / n:.4f}") # 两次输出在初始化噪声内接近(如 0.93 对 0.91)—— # 排程换人,机制未变;差异只来自随机初始化
4.4 节的 GAN 已经演示了双 tape 双优化器的交替模板。再补一个常见形状——多任务共享骨干、各任务独立头,梯度合并后一次更新:
shared = tf.keras.layers.Dense(32, activation="relu") head_a = tf.keras.layers.Dense(1) # 任务 A:回归头 head_b = tf.keras.layers.Dense(3) # 任务 B:三分类头 xb = tf.random.normal([32, 8]) ya = tf.random.normal([32, 1]) yb = tf.one_hot(tf.random.uniform([32], 0, 3, dtype=tf.int32), 3) with tf.GradientTape() as tape: h = shared(xb) loss = tf.reduce_mean(tf.square(ya - head_a(h))) + \ tf.reduce_mean(tf.nn.softmax_cross_entropy_with_logits(yb, head_b(h))) grads = tape.gradient(loss, shared.trainable_variables + head_a.trainable_variables + head_b.trainable_variables) print("multi-task grads:", sum(1 for g in grads if g is not None)) # 输出:multi-task grads: 6 # 三层各两件套(权重与偏置)全拿到梯度,一次 apply 更新
非标流程的通用心法:先把"一个批次内要更新几次、每次谁冻结谁更新"写成伪码节拍,再逐拍翻译成 tape 与 apply_gradients 的组合。节拍清楚了,代码只是翻译。
⚠️ 常见坑:在循环外创建指标、在循环内忘了 reset_state——验证指标第二轮从上一轮的累计值接着涨,数值越滚越大还以为学崩了。reset_state 与 epoch 边界绑定,模板里那两行不是装饰。
💡 关键直觉:fit 是"标准节拍"的现成排程,手写循环是把节拍表拿到自己手里。接管前先明确你要改哪一拍——只改节拍,别把没坏的部分一起重写。
下节把单机循环扩容到多设备。