You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.28 10:03:21