知识蒸馏训练必须用tf.keras.model子类化而非sequential,因其需在自定义train_step中同步前向传播教师与学生模型、手动实现带温度缩放的kl散度损失、屏蔽教师梯度,并确保教师调用时training=false以稳定bn状态。

知识蒸馏训练必须用 tf.keras.Model 子类化而非 Sequential
因为蒸馏需要同时前向传播教师模型和学生模型,并自定义损失组合(如软目标交叉熵 + 真实标签交叉熵),Sequential 模型无法灵活接入多个输入/输出分支。子类化 tf.keras.Model 才能重写 train_step,控制教师 logits 的复用与梯度屏蔽。
常见错误是试图在 compile() 中塞进两个 loss——这会导致 tf.keras 自动对每个输出分配 loss,但蒸馏中教师 logits 是中间计算结果,不是模型输出层。必须绕过默认训练流程。
- 教师模型设为
trainable=False,且调用时加training=False确保 BN 层不更新 - 学生模型的
train_step中,需显式调用教师模型:teacher_logits = self.teacher(x, training=False) - 温度缩放(temperature scaling)必须手动应用:用
tf.nn.softmax(teacher_logits / T),不能依赖外部 softmax 层
train_step 里怎么算软目标 KL 散度损失
TensorFlow 没有开箱即用的“带温度的 KL 散度”损失函数,直接用 tf.keras.losses.KLDivergence() 会出错:它要求输入是概率分布(已 softmax),但若直接传入未缩放的 logits,数值不稳定且梯度不准。
正确做法是手动实现 soft target loss:先对教师 logits 做温度缩放 + softmax,再对学生 logits 做同样缩放 + log_softmax,最后点乘求和。等价于 KL(p_teacher || p_student) 的稳定形式。
def distillation_loss(y_true, y_student, y_teacher, temperature=3.0):
p_teacher = tf.nn.softmax(y_teacher / temperature)
log_p_student = tf.nn.log_softmax(y_student / temperature)
return -tf.reduce_sum(p_teacher * log_p_student) * (temperature ** 2) / tf.cast(tf.shape(y_true)[0], tf.float32)
- 注意最后除以 batch size,否则 loss 值随 batch 变化,影响超参稳定性
- 乘以
temperature ** 2是为了补偿缩放带来的梯度衰减(推导自 KL 展开式) - 真实标签 loss(如
sparse_categorical_crossentropy)应保持原始 logits 输入,不加温度
教师模型输出 logits 还是概率?为什么不能用 predict()
必须用 logits。教师模型若最后一层带 softmax,输出的是概率,再做一次 softmax(logits / T) 就相当于双重非线性变换,破坏了蒸馏所需的平滑软标签特性。而且 predict() 会触发完整 inference 流程(含数据预处理、batch padding 等),无法嵌入到 student 的 train_step 中参与梯度图构建。
- 确保教师模型最后一层是 Dense(无激活),或加载权重时去掉原 softmax 层
- 调用教师模型必须走
__call__(即括号调用),而非predict()或evaluate() - 如果教师是 Hugging Face 模型(如
TFAutoModelForSequenceClassification),需取logits属性,不是probabilities
蒸馏训练时 batch size 和 learning rate 怎么调
蒸馏本身不改变数据 pipeline,但软目标 loss 对 batch size 更敏感:小 batch 下 teacher logits 方差大,soft label 噪声高;大 batch 虽稳定,但内存压力翻倍(要同时存 teacher + student 的中间激活)。
learning rate 通常要比纯学生训练高 1.5–2 倍,因为软目标提供了更密集的梯度信号,收敛更快,但过高会导致 student 忽略真实标签。
- 推荐起始 batch size ≥ 64(GPU 显存允许下),低于 32 时建议加 label smoothing(
label_smoothing=0.1)缓解噪声 - student 的初始 lr 设为
1e-3(Adam),比常规训练高一档;teacher 不更新参数,lr 无关 - 可加 warmup:前 10% step 线性增 lr,避免初期 soft loss 主导导致 student 过早坍缩
最容易被忽略的是教师模型的 BN 层状态——即使 trainable=False,BN 的 running_mean / running_var 仍可能在 training=True 下被意外更新。务必确认所有 BN 层调用时都显式传入 training=False。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











