直接继承tf.keras.losses.loss子类不够,因其call()仅提供y_true和y_pred,无法访问模型中间输出、多输入张量或gradienttape上下文;需在tf.gradienttape内手动构建损失并求导以实现灵活控制。

为什么直接用 tf.keras.losses.Loss 子类还不够
因为很多真实场景的损失需要访问模型中间输出、多输入张量,甚至要和梯度计算耦合(比如对抗训练、梯度惩罚、自监督对比损失)。这时候单纯继承 tf.keras.losses.Loss 并重写 call() 会卡住——它只给 y_true 和 y_pred,拿不到 model 的内部变量或 GradientTape 上下文。
在 tf.GradientTape 里手写损失并求导才是可控路径
核心思路:把损失计算逻辑放进 with tf.GradientTape() as tape: 块内,显式调用模型、拼接运算、构造标量 loss,再用 tape.gradient(loss, model.trainable_variables) 拿梯度。这样你完全掌控每一步张量来源和依赖关系。
常见错误现象:ValueError: Cannot differentiate a constant 或梯度为 None —— 多半是损失没对上可训练变量,或用了 tf.constant / .numpy() 中断了计算图。
- 损失必须是标量(
shape=()),不能是[batch_size];用tf.reduce_mean()或tf.reduce_sum()收尾 - 所有参与 loss 计算的张量(包括模型输出、中间特征)必须来自同一
tf.GradientTape范围,且未被.numpy()、tf.print()(默认不参与图)、tf.stop_gradient()截断 - 如果要用到标签之外的额外数据(如样本权重、掩码、负样本索引),得通过函数参数传入,别试图从全局变量或 dataset.batch() 外部“偷”
model(x) 和 tape.watch() 到底要不要用
绝大多数情况下不用手动 tape.watch(model.trainable_variables) —— tf.GradientTape 默认追踪所有可训练变量,只要它们参与了 loss 计算。但要注意:
- 如果你在 tape 内部调用的是
model(x, training=True),那模型权重自动被追踪,无需watch() - 如果你手动拆解了模型层(比如
layer1(x); layer2(h1)),且某些层权重没被自动纳入(例如自定义tf.keras.layers.Layer里忘了设self.trainable_weights),才需显式tape.watch(layer.trainable_variables) - 绝对不要
tape.watch(x)(输入数据)除非你真要优化输入(如风格迁移、对抗样本生成);否则浪费内存还易出错
一个带中间特征和梯度惩罚的实操例子
假设你要实现一个带 L2 梯度惩罚的对比损失:loss = contrastive_loss + λ × ||∇ₓ f(x)||²。关键点在于梯度惩罚项必须在 tape 内对输入 x 求导,再平方求均值。
with tf.GradientTape() as tape:
# 正向:获取模型输出和中间特征
z = model.encoder(x, training=True) # shape [B, D]
logits = model.classifier(z) # shape [B, C]
<pre class="brush:python;toolbar:false;"># 自定义对比部分(例如 InfoNCE)
loss_contrast = -tf.reduce_mean(
tf.gather_nd(logits, tf.stack([tf.range(B), y_true], axis=1))
- tf.math.log(tf.reduce_sum(tf.exp(logits), axis=1))
)
# 梯度惩罚:对输入 x 求 z 的梯度(注意是 z,不是 logits)
with tf.GradientTape() as inner_tape:
inner_tape.watch(x)
z_inner = model.encoder(x, training=True)
grad_z_x = inner_tape.gradient(z_inner, x) # shape [B, H, W, C]
loss_grad_penalty = tf.reduce_mean(tf.square(grad_z_x))
total_loss = loss_contrast + 1e-4 * loss_grad_penalty对模型参数求总 loss 的梯度
gradients = tape.gradient(total_loss, model.trainable_variables) optimizer.apply_gradients(zip(gradients, model.trainable_variables))
这里容易被忽略的是:内层 tape 必须 watch(x),外层 tape 不用 watch 模型变量(自动追踪),而 grad_z_x 的 shape 必须和 x 一致才能做 tf.square;若 encoder 输出是向量,grad_z_x 就是雅可比矩阵,得用 tf.norm(..., axis=-1) 等方式降维,否则会报 shape mismatch。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











