基于HF TF模型构建BERT知识蒸馏Distiller类报错排查
Hugging Face TF版BERT知识蒸馏TypeError修复方案
根因定位
- 你直接复用的Keras官方Distiller类默认适配输入输出均为纯张量的原生Keras模型,和Hugging Face Transformers的TF模型输出格式完全不兼容
TFAutoModelForSequenceClassification前向传播返回的是TFSequenceClassifierOutput结构化数据类对象,不是损失函数可直接接收的logits张量,把整个对象传入SparseCategoricalCrossentropy作为y_pred参数时,会触发内部类型校验失败- 报错提示"Expected any non-tensor type, but got a tensor instead"属于框架抛出的误导性提示,本质不是张量类型错误,是传入了不符合接口要求的自定义类对象
前置常识
知识蒸馏的核心计算分为两部分:一是学生模型输出和真实标签的硬标签损失,二是学生输出和教师输出(经温度平滑的软分布)的蒸馏损失,两部分加权求和得到总损失。所有损失、指标计算环节,都必须传入形状为
(batch_size, num_labels)的纯logits张量,不能传入Hugging Face自定义的输出类对象。
- Hugging Face所有TF分类模型的前向返回对象,都可以通过
.logits属性拿到和原生Keras分类模型输出格式完全一致的预测分数张量 - 你通过
to_tf_dataset生成的字典格式输入完全适配Hugging Face模型,不需要做任何输入格式修改,仅需处理模型输出即可 - 训练完成的教师模型做前向推理时,要放在梯度追踪上下文外,避免不必要的显存占用和计算开销
分步修复操作
1. 重写适配Hugging Face TF模型的Distiller类
核心修改点是所有教师、学生模型的前向调用后,追加.logits取值操作,只把纯张量传入损失和指标计算逻辑,可直接复用以下代码:
import tensorflow as tf from tensorflow import keras class Distiller(keras.Model): def __init__(self, student, teacher, temperature=2.0, alpha=0.1): super().__init__() self.student = student self.teacher = teacher self.temperature = temperature self.alpha = alpha # 硬标签损失权重,1-alpha为蒸馏损失权重 # 硬标签损失和训练教师模型时保持一致 self.hard_loss_fn = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True) # 蒸馏损失用KL散度计算软分布差异 self.distill_loss_fn = tf.keras.losses.KLDivergence() def compile(self, optimizer, metrics, **kwargs): super().compile(optimizer=optimizer, metrics=metrics, **kwargs) # 自定义训练过程指标 self.student_loss_tracker = keras.metrics.Mean(name="student_loss") self.distill_loss_tracker = keras.metrics.Mean(name="distill_loss") self.total_loss_tracker = keras.metrics.Mean(name="total_loss") @property def metrics(self): return [ self.total_loss_tracker, self.student_loss_tracker, self.distill_loss_tracker, ] + self.compiled_metrics.metrics def train_step(self, data): x, y = data # 教师前向:关闭梯度追踪,仅提取logits teacher_logits = self.teacher(x, training=False).logits with tf.GradientTape() as tape: # 学生前向:开启梯度追踪,提取logits student_logits = self.student(x, training=True).logits # 计算硬标签损失 student_loss = self.hard_loss_fn(y, student_logits) # 计算蒸馏损失:温度缩放后计算KL散度,乘以温度平方做梯度校正 distill_loss = self.distill_loss_fn( tf.nn.softmax(teacher_logits / self.temperature, axis=-1), tf.nn.softmax(student_logits / self.temperature, axis=-1) ) * (self.temperature ** 2) # 加权得到总损失 total_loss = self.alpha * student_loss + (1 - self.alpha) * distill_loss # 反向传播仅更新学生模型参数 trainable_vars = self.student.trainable_variables gradients = tape.gradient(total_loss, trainable_vars) self.optimizer.apply_gradients(zip(gradients, trainable_vars)) # 更新指标状态 self.compiled_metrics.update_state(y, student_logits) self.student_loss_tracker.update_state(student_loss) self.distill_loss_tracker.update_state(distill_loss) self.total_loss_tracker.update_state(total_loss) return {m.name: m.result() for m in self.metrics} def test_step(self, data): # 验证/测试步逻辑,不做反向传播 x, y = data student_logits = self.student(x, training=False).logits student_loss = self.hard_loss_fn(y, student_logits) self.compiled_metrics.update_state(y, student_logits) self.student_loss_tracker.update_state(student_loss) self.total_loss_tracker.update_state(student_loss) return {m.name: m.result() for m in self.metrics}
2. 初始化并编译Distiller
教师模型训练完成后不需要额外设置trainable=False,train_step逻辑已经保证教师参数不会被更新,不要给教师模型单独编译避免浪费显存:
# teacher_model为你训练好的95%准确率的4分类教师模型 # student_model为你初始化的同checkpoint或更小尺寸的学生模型 distiller = Distiller(student=student_model, teacher=teacher_model, temperature=2.0, alpha=0.3) distiller.compile( optimizer=tf.keras.optimizers.Adam(learning_rate=5e-5), # 可搭配你之前用的学习率调度器 metrics=[tf.keras.metrics.SparseCategoricalAccuracy(name="accuracy")] )
3. 启动蒸馏训练
直接传入你之前通过to_tf_dataset生成的训练、验证集即可,不需要修改数据集格式:
distill_history = distiller.fit( tf_train_set, validation_data=tf_val_set, epochs=5 # 根据验证集效果调整轮次 )
常见踩坑提示
- 所有涉及模型预测分数的环节(损失计算、指标计算、推理预测),都要从返回对象中取
.logits属性,不要直接传入整个TFSequenceClassifierOutput对象 - 蒸馏损失计算时必须乘以
temperature^2做梯度缩放,否则温度越高梯度值越小,会导致模型收敛极慢 - 如果学生模型是参数量更小的变体(如DistilBERT、TinyBERT),可将alpha(硬标签权重)调低到0.1~0.3,让模型更多学习教师的软标签分布,蒸馏效果更好
- 不要手动冻结教师模型层、不要给教师模型单独编译,否则会额外占用显存,甚至导致训练逻辑异常
内容的提问来源于stack exchange,提问作者ChessMateK
相关产品推荐
相关产品推荐

