🤗 Transformers 图像分类知识蒸馏实战:用 Trainer 将 ViT 教师模型蒸馏到 MobileNetV2
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
本指南演示如何基于 🤗 Transformers 的
TrainerAPI,将一个在 Beans 数据集上微调好的 ViT 图像分类模型(教师)蒸馏到一个随机初始化的 MobileNetV2(学生)模型。你将学会:如何准备与预处理数据集、如何通过重写Trainer.compute_loss()实现基于 KL 散度的蒸馏损失、如何配置TrainingArguments并完成训练、评估与推送到 Hugging Face Hub,以及如何用基线对比验证蒸馏的真实收益。
知识蒸馏(Knowledge Distillation)核心概念
知识蒸馏是一种模型压缩与知识迁移技术:从更大、更复杂的模型(教师,teacher)向更小、更简单的模型(学生,student)传递知识。该思想最早由 Hinton 等人在论文Distilling the Knowledge in a Neural Network中提出。其核心直觉是:教师模型输出的**软目标(soft target,即经过温度缩放的概率分布)**比硬标签携带更丰富的类间相似性信息——例如一张豆子图片在"健康"类上概率高、在"锈病"类上概率中等,这种分布形态正是学生模型需要模仿的知识。
本指南执行的是任务特定的知识蒸馏(task-specific distillation):教师与学生在同一任务(图像分类)上对齐。流程为:
- 取一个在某任务上完成训练的预训练教师模型(本例为
merve/beans-vit-224,基于google/vit-base-patch16-224-in21k在 Beans 数据集上微调而来); - 随机初始化一个学生模型(本例为 MobileNetV2),同样面向图像分类;
- 训练学生模型,使其输出分布与教师输出分布的差异最小化,从而模仿教师的行为。
蒸馏核心损失由两项加权合成:学生与教师分布之间的KL 散度蒸馏损失与学生的真实标签交叉熵损失,前者传递教师知识,后者保证学生不偏离真实任务。
说明:原指南发布于 2023 年,文档中的部分 API(如
image_processor、eval_strategy、report_to参数名)在仓库当前版本中已有演进。本文以当前仓库源码为准,对示例做了适配(如使用processing_class替代旧式image_processor传参、统一使用eval_strategy与report_to="tensorboard"),并标注了文档原文与当前实现的差异,便于读者对照。
环境准备与依赖安装
蒸馏与评估过程所需库如下:
pip install transformers datasets accelerate tensorboard evaluate --upgradetransformers:模型、Trainer、TrainingArguments与图像处理器等核心组件;datasets:加载与处理 Beans 数据集;accelerate:Trainer的底层训练加速后端(设备放置、混合精度等);tensorboard:训练日志可视化(对应TrainingArguments中的report_to="tensorboard");evaluate:加载accuracy等评估指标。
仓库中examples/pytorch/image-classification目录(run_image_classification.py)提供了可直接运行的图像分类训练脚本,可结合本指南对照学习。本文示例为单机场景;如需分布式训练背景知识,可参考仓库中的 分布式训练示例。
加载数据集与图像预处理
使用datasets加载 Beans 数据集:
from datasets import load_dataset dataset = load_dataset("beans")Beans 是一个植物病害图像分类数据集,包含train、validation、test三个划分,标签为三类叶片病害(如角斑病、锈病、健康)。
图像预处理直接复用教师模型的处理器即可。本例中教师(ViT)与学生(MobileNetV2)的处理器在相同输入分辨率下返回相同输出,因此任选其一均可,这里使用教师的处理器:
from transformers import AutoImageProcessor teacher_processor = AutoImageProcessor.from_pretrained("merve/beans-vit-224") def process(examples): processed_inputs = teacher_processor(examples["image"]) return processed_inputs processed_datasets = dataset.map(process, batched=True)dataset.map(process, batched=True)会对数据集的每个划分(train/validation/test)批量应用预处理:AutoImageProcessor内部完成图像缩放、归一化(ViT 默认使用 ImageNet 均值/方差)并输出模型所需的pixel_values张量。
在当前仓库中,
AutoImageProcessor已被逐步统一到AutoProcessor(见 processing_utils.py 中的处理器自动加载逻辑),但图像分类场景下两者均可使用。
设计蒸馏训练器:重写 Trainer 的 compute_loss
我们的目标:让随机初始化的 MobileNet 模仿微调后的 ViT。实现方式是继承Trainer并重写compute_loss(),在每一步训练中:
- 分别获取教师与学生的 logits 输出;
- 用
temperature(温度)缩放 logits 得到软目标(soft target),温度控制各软目标的重要程度; - 用
lambda(蒸馏损失权重)衡量蒸馏损失在总损失中的占比; - 用KL 散度损失计算学生与教师分布的差异。
关于 KL 散度:给定两个分布 P 与 Q,KL 散度描述"用 Q 表示 P 额外需要多少信息"。若两者完全相同,则 KL 散度为 0。在蒸馏语境下,我们最小化"学生分布表示教师分布"所需的额外信息量,从而让学生分布逼近教师分布。
示例实现如下:
from transformers import TrainingArguments, Trainer from accelerate import Accelerator import torch import torch.nn as nn import torch.nn.functional as F class ImageDistilTrainer(Trainer): def __init__(self, teacher_model=None, student_model=None, temperature=None, lambda_param=None, *args, **kwargs): super().__init__(model=student_model, *args, **kwargs) self.teacher = teacher_model self.student = student_model self.loss_function = nn.KLDivLoss(reduction="batchmean") device = Accelerator().device self.teacher.to(device) self.teacher.eval() self.temperature = temperature self.lambda_param = lambda_param def compute_loss(self, student, inputs, return_outputs=False): student_output = self.student(**inputs) with torch.no_grad(): teacher_output = self.teacher(**inputs) # 计算教师与学生的软目标 soft_teacher = F.softmax(teacher_output.logits / self.temperature, dim=-1) soft_student = F.log_softmax(student_output.logits / self.temperature, dim=-1) # 蒸馏损失(乘以 temperature 的平方以补偿缩放) distillation_loss = self.loss_function(soft_student, soft_teacher) * (self.temperature ** 2) # 真实标签损失(由学生模型前向返回) student_target_loss = student_output.loss # 最终损失 = 加权组合 loss = (1. - self.lambda_param) * student_target_loss + self.lambda_param * distillation_loss return (loss, student_output) if return_outputs else loss实现要点:
nn.KLDivLoss(reduction="batchmean"):PyTorch 要求输入为 log 概率(学生对数软目标log_softmax),目标为概率(教师软目标softmax);batchmean按 batch 求平均,配合温度平方缩放后数值稳定;- 温度缩放与补偿:
logits / temperature使分布更平滑(温度越高越均匀,突出类间软信息);同时蒸馏损失乘以temperature ** 2,以抵消温度缩放带来的梯度幅度变化,这是 Hinton 论文中的标准做法; torch.no_grad():教师模型仅用于前向推理,不参与反向传播,因此置于no_grad上下文并设为eval()模式(关闭 dropout 等);- 总损失:
(1 - lambda) * 交叉熵 + lambda * KL,lambda=0.5表示两者等权; - 本示例中取
temperature=5、lambda=0.5,读者可自行调参对比效果。
底层原理:Trainer 如何调用 compute_loss
在 trainer.py 中,基类Trainer.compute_loss()默认将模型前向返回的loss(outputs["loss"])作为训练损失直接返回。而在训练主循环training_step(trainer.py)与评估循环evaluation_loop(trainer.py)中,均通过self.compute_loss(...)获取损失;当return_outputs=True时返回(loss, outputs)元组,其中outputs[1:]即 logits,供评估阶段计算指标。因此,子类只需重写compute_loss,即可在不改动训练/评估循环的前提下自定义损失函数——这正是ImageDistilTrainer能无缝接入Trainer的原因。
教师侧损失来源:ViTForImageClassification.forward()(modeling_vit.py)在传入labels时,通过self.loss_function(labels, logits, self.config)计算交叉熵并放进输出对象的loss字段;MobileNetV2ForImageClassification.forward()(modeling_mobilenet_v2.py)同样在labels存在时返回loss。这就是student_output.loss(真实标签损失)的直接来源。
登录 Hub 并配置训练参数
登录 Hugging Face Hub,以便通过Trainer将模型推送到 Hub:
from huggingface_hub import notebook_login notebook_login()配置TrainingArguments、教师模型与学生模型:
from transformers import AutoModelForImageClassification, MobileNetV2Config, MobileNetV2ForImageClassification training_args = TrainingArguments( output_dir="my-awesome-model", num_train_epochs=30, fp16=True, logging_strategy="epoch", eval_strategy="epoch", save_strategy="epoch", load_best_model_at_end=True, metric_for_best_model="accuracy", report_to="tensorboard", push_to_hub=True, hub_strategy="every_save", hub_model_id=repo_name, ) num_labels = len(processed_datasets["train"].features["labels"].names) # 初始化教师模型(在 beans 上微调过的 ViT) teacher_model = AutoModelForImageClassification.from_pretrained( "merve/beans-vit-224", num_labels=num_labels, ignore_mismatched_sizes=True ) # 从零训练学生模型:随机初始化 MobileNetV2 student_config = MobileNetV2Config() student_config.num_labels = num_labels student_model = MobileNetV2ForImageClassification(student_config)TrainingArguments关键参数说明:
| 参数 | 值 | 作用 |
|---|---|---|
output_dir | "my-awesome-model" | 检查点与模型保存目录 |
num_train_epochs | 30 | 训练轮数 |
fp16 | True | 启用半精度混合精度训练(需 GPU 支持) |
logging_strategy/eval_strategy/save_strategy | "epoch" | 每个 epoch 记录日志、评估、保存检查点 |
load_best_model_at_end | True | 训练结束时加载验证集最优检查点 |
metric_for_best_model | "accuracy" | 以 accuracy 作为选优指标 |
report_to | "tensorboard" | 日志上报到 TensorBoard |
push_to_hub/hub_strategy/hub_model_id | True/"every_save"/repo_name | 每次保存检查点即推送 Hub(需已登录且repo_name已定义) |
版本适配提示:原文档使用的
eval_strategy参数在当前仓库为推荐名称;report_to原文为"trackio",本文统一采用生态更通用的"tensorboard"。若你的环境支持 trackio 等第三方跟踪器,可按需替换。
MobileNetV2Config 关键字段(仓库源码)
从 configuration_mobilenet_v2.py 可见,MobileNetV2Config()默认即对应google/mobilenet_v2_1.0_224结构,核心字段包括:
num_channels(默认3):输入图像通道数;image_size(默认224):输入分辨率,与教师 ViT 的 224 一致,因此两模型处理器可互换;depth_multiplier(默认1.0):通道数缩放系数(宽度因子);expand_ratio(默认6.0):倒残差块首层输出通道 = 输入通道 × 扩展比;output_stride(默认32):输入/输出特征图空间分辨率比值,设为 8 或 16 时使用空洞卷积;hidden_act(默认"relu6"):激活函数;classifier_dropout_prob(默认0.8):分类头 dropout 概率;num_labels:分类类别数,需手动设置为数据集的标签数(本示例代码中通过student_config.num_labels = num_labels覆盖)。
MobileNetV2ForImageClassification在 modeling_mobilenet_v2.py 中实现:主干网络输出池化特征后,经 dropout 与线性分类头得到 logits;当config.num_labels == 1时计算回归损失(MSE),大于 1 时计算分类交叉熵损失。教师模型加载时使用ignore_mismatched_sizes=True,用于忽略分类头尺寸与 Hub 上原模型不一致的层,从而适配 Beans 的 3 类标签。
定义评估指标
compute_metrics函数用于在训练过程中计算模型的accuracy:
import evaluate import numpy as np accuracy = evaluate.load("accuracy") def compute_metrics(eval_pred): predictions, labels = eval_pred acc = accuracy.compute(references=labels, predictions=np.argmax(predictions, axis=1)) return {"accuracy": acc["accuracy"]}eval_pred来自评估循环:预测 logits 通过np.argmax(..., axis=1)转为类别索引,再与真实标签比对得到准确率。该指标同时被metric_for_best_model="accuracy"用于挑选最优检查点。
初始化蒸馏训练器并开始训练
使用上面定义的训练参数初始化Trainer,同时初始化数据整理器(data collator):
from transformers import DefaultDataCollator data_collator = DefaultDataCollator() trainer = ImageDistilTrainer( student_model=student_model, teacher_model=teacher_model, training_args=training_args, train_dataset=processed_datasets["train"], eval_dataset=processed_datasets["validation"], data_collator=data_collator, processing_class=teacher_processor, compute_metrics=compute_metrics, temperature=5, lambda_param=0.5 )版本适配提示:
Trainer构造时传递处理器的参数,在当前仓库中为processing_class(旧版本为image_processor或tokenizer)。ImageDistilTrainer.__init__通过**kwargs透传给基类Trainer,因此processing_class、compute_metrics等参数会正常生效。
开始训练:
trainer.train()在测试集上评估:
trainer.evaluate(processed_datasets["test"])实验结果与蒸馏效率验证
据原指南报告:蒸馏后的 MobileNet 在测试集上达到72% 准确率;作为蒸馏有效性的健康性检查(sanity check),使用相同超参数在 Beans 数据集上从零训练 MobileNet(无教师监督),测试集准确率仅为63%。两者对比说明:
- 蒸馏带来的约 9 个百分点的提升,直接归因于教师软目标传递的知识;
- 蒸馏是在相同训练预算(epoch、batch、学习率一致)下实现的,排除了超参数差异的干扰。
蒸馏后的训练日志与检查点可参考 Hub 仓库(教师-学生对merve/vit-mobilenet-beans-224,从零训练的 MobileNetV2 见merve/resnet-mobilenet-beans-5系列)。
请注意:上述数字为原指南作者在特定数据集划分与超参数下的实测结果,仅供参考;在不同环境复现时数值可能略有波动。建议读者自行尝试不同的预训练教师、学生架构与蒸馏参数(
temperature、lambda),并对比"从零训练基线"以验证蒸馏收益。
进阶方向:进一步探索
- 更换教师/学生架构:教师可替换为其他在 Hub 上微调过的 ViT 变体或更大的 CNN;学生可尝试 MobileNetV3、EfficientNet 等轻量架构,验证压缩比与精度权衡;
- 调整蒸馏超参数:增大
temperature使软目标更平滑、更强调类间关系;调节lambda平衡蒸馏损失与标签损失; - 结合硬标签与软目标:本示例已同时使用两者;还可引入特征层蒸馏(对齐中间特征图)、attention 蒸馏等变体;
- 生产部署:蒸馏后的轻量学生模型更易部署。仓库提供了模型导出支持(见 exporters)与 ONNX 等格式转换能力,可将训练好的 MobileNetV2 导出用于推理服务。
总结
本文以完整可运行的代码路径,演示了基于 🤗 TransformersTrainer的图像分类知识蒸馏:通过继承Trainer并重写compute_loss(),将 ViT 教师的软目标知识以 KL 散度形式注入 MobileNetV2 学生模型,并结合真实标签损失进行联合优化。文中同时给出了数据预处理、训练配置、评估指标、结果对比等完整闭环,并补充了Trainer.compute_loss调用链、ViTForImageClassification与MobileNetV2ForImageClassification损失计算实现、MobileNetV2Config关键字段等仓库源码级依据。掌握该方法后,你可以将其推广到任意图像分类任务与任意"教师-学生"模型组合,实现高效、低成本的模型压缩与知识迁移。
核心文件索引
- 蒸馏训练器基类:trainer.py(
compute_loss默认实现与重写约定) - 教师模型实现:modeling_vit.py(
ViTForImageClassification) - 学生模型实现:modeling_mobilenet_v2.py(
MobileNetV2ForImageClassification) - 学生模型配置:configuration_mobilenet_v2.py(
MobileNetV2Config) - 图像分类参考脚本:run_image_classification.py
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考