工具框架(一):Hugging Face Transformers 蒸馏

作为”工程化流程”工具篇的第一篇,本文聚焦如何用 Hugging Face(HF)Transformers 生态落地知识蒸馏。前置:你应已理解蒸馏的基本损失(软标签 + 硬标签加权和,Hinton et al. 2015),并知道 DistilBERT 是 HF 出品的经典蒸馏产物(Sanh et al. 2019)。本文不重复理论,只讲”怎么在 HF 里做”。

一、HF 生态与蒸馏的关系

HF Transformers 本身不是”蒸馏框架”,而是提供预训练模型、Tokenizer 与 Trainer 训练循环的基础设施,蒸馏可以在其上”拼装”出来。最具代表性的成品就是 DistilBERT(Sanh et al. 2019),其训练目标是一个三重损失(triple loss)

  1. 软标签蒸馏损失:学生与教师在温度 $T$ 下输出分布的 KL 散度;
  2. 语言建模损失(MLM):学生自身的掩码语言建模任务损失(用于任务无关预训练蒸馏);
  3. 余弦嵌入损失:对齐学生与教师最后隐状态的余弦距离,是一种”基于特征”的蒸馏项。

其中软标签部分正是 Hinton et al. 2015 的蒸馏损失;余弦项则把”中间表征”也作为知识迁移,与后文 FitNets(Romero et al. 2015)的 hint 思想相通。

二、基于 Trainer 的蒸馏(子类化 compute_loss)

最常用做法:子类化 Trainer,重写 compute_loss,在同一批数据上同时跑学生和冻结的教师,计算加权和损失。伪代码骨架:

from transformers import Trainer, TrainingArguments
import torch.nn.functional as F

teacher.eval()  # 冻结,不反传

class DistilTrainer(Trainer):
    def compute_loss(self, model, inputs, return_outputs=False, **kw):
        labels = inputs["labels"]
        with torch.no_grad():
            t_logits = teacher(**inputs).logits
        out = model(**inputs)
        T, alpha = 2.0, 0.5
        soft = F.kl_div(
            F.log_softmax(out.logits / T, dim=-1),
            F.softmax(t_logits / T, dim=-1),
            reduction="batchmean",
        ) * (T * T)
        hard = F.cross_entropy(out.logits, labels)
        loss = alpha * soft + (1.0 - alpha) * hard
        return (loss, out) if return_outputs else loss

要点:teacher 需与 student 在同一设备;alpha 平衡软/硬损失;温度 $T$ 软化分布、放大暗知识。T*T 系数用于补偿因温度缩放带来的梯度量级变化(Hinton et al. 2015 原论文给出的标准做法)。

三、no_trainer 脚本式示例

HF 历史上在 examples/distillation 中提供过完整的 DistilBERT 训练脚本 train.py,用命令行参数直接控制三重损失的权重,例如:

python train.py --student_type distilbert \
  --teacher_type bert --teacher_name bert-base-uncased \
  --alpha_ce 5.0 --alpha_mlm 2.0 --alpha_cos 1.0 \
  --mlm --freeze_pos_embs --dump_path ./my_distill

其中 --alpha_ce(硬标签 CE 权重)、--alpha_mlm(MLM 权重)、--alpha_cos(余弦对齐权重)对应上一节的三重损失。该脚本还建议用教师的部分层初始化学生权重,以稳定收敛。若要完全掌控训练循环,也可不依赖 Trainer,自行写 for batch in dataloader 的反向传播——逻辑与上面 compute_loss 一致。

四、社区实践与生成式蒸馏

除上述手写方式,社区还有更上层的封装:

  • TRL 库:HF 的 TRL 提供 GKDTrainer(Generalized Knowledge Distillation Trainer),面向生成式/指令模型做 token 级蒸馏,对应 Agarwal et al. 2024 提出的 GKD 思路(on-policy 蒸馏,让学生在”自己生成的样本”上与教师对齐)——这属于进阶内容,本文仅点名工具;
  • Distil 系列模型*:除 DistilBERT 外,社区还有 DistilGPT2、DistilRoBERTa 等,均沿用”学生层数减半 + 蒸馏”的范式,可作为现成学生起点;
  • 特征级蒸馏:在 Transformer 上,还可对齐中间隐状态与注意力矩阵(TinyBERT 式做法),但需自定义对齐层与损失。

五、落地注意事项

  1. Tokenizer 一致性:学生与教师必须共享同一 Tokenizer(DistilBERT 与 BERT 共享 bert-base-uncased tokenizer),否则软标签空间错位;
  2. 学生初始化:用教师对应层初始化学生,可显著加快收敛;
  3. 温度与 alpha 调参:温度过低退化为硬标签、过高抹平暗知识,需按任务搜索;
  4. 评估闭环:蒸馏后用阶段五指标(精度保留率、体积、速度)验收,避免”训出来但不达标”。

延伸阅读 / 参考文献

  • Hinton, Vinyals, Dean (2015),《Distilling the Knowledge in a Neural Network》, arXiv:1503.02531 —— 软标签温度 $T$ 与蒸馏损失加权。
  • Romero et al. (2015),《FitNets: Hints for Thin Deep Nets》, arXiv:1412.6550, ICLR 2015 —— 特征/中间层蒸馏(余弦对齐项的同源思想)。
  • Sanh et al. (2019),《DistilBERT, a distilled version of BERT》, arXiv:1910.01108 —— HF 蒸馏旗舰案例与三重损失。
  • Agarwal et al. (2024),《On-Policy Distillation of Language Models…》(GKD), arXiv:2306.13649, ICLR 2024 —— TRL GKDTrainer 的理论源头。
本文累计阅读