[特殊字符] Transformers 图像分类知识蒸馏实战:用 Trainer 将 ViT 教师模型蒸馏到 MobileNetV2
2026/9/10 9:36:54 网站建设 项目流程

🤗 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):教师与学生在同一任务(图像分类)上对齐。流程为:

  1. 取一个在某任务上完成训练的预训练教师模型(本例为merve/beans-vit-224,基于google/vit-base-patch16-224-in21k在 Beans 数据集上微调而来);
  2. 随机初始化一个学生模型(本例为 MobileNetV2),同样面向图像分类;
  3. 训练学生模型,使其输出分布与教师输出分布的差异最小化,从而模仿教师的行为。

蒸馏核心损失由两项加权合成:学生与教师分布之间的KL 散度蒸馏损失与学生的真实标签交叉熵损失,前者传递教师知识,后者保证学生不偏离真实任务。

说明:原指南发布于 2023 年,文档中的部分 API(如image_processoreval_strategyreport_to参数名)在仓库当前版本中已有演进。本文以当前仓库源码为准,对示例做了适配(如使用processing_class替代旧式image_processor传参、统一使用eval_strategyreport_to="tensorboard"),并标注了文档原文与当前实现的差异,便于读者对照。

环境准备与依赖安装

蒸馏与评估过程所需库如下:

pip install transformers datasets accelerate tensorboard evaluate --upgrade
  • transformers:模型、TrainerTrainingArguments与图像处理器等核心组件;
  • datasets:加载与处理 Beans 数据集;
  • accelerateTrainer的底层训练加速后端(设备放置、混合精度等);
  • 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 是一个植物病害图像分类数据集,包含trainvalidationtest三个划分,标签为三类叶片病害(如角斑病、锈病、健康)。

图像预处理直接复用教师模型的处理器即可。本例中教师(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(),在每一步训练中:

  1. 分别获取教师与学生的 logits 输出;
  2. temperature(温度)缩放 logits 得到软目标(soft target),温度控制各软目标的重要程度;
  3. lambda(蒸馏损失权重)衡量蒸馏损失在总损失中的占比;
  4. 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 * KLlambda=0.5表示两者等权;
  • 本示例中取temperature=5lambda=0.5,读者可自行调参对比效果。

底层原理:Trainer 如何调用 compute_loss

在 trainer.py 中,基类Trainer.compute_loss()默认将模型前向返回的lossoutputs["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_epochs30训练轮数
fp16True启用半精度混合精度训练(需 GPU 支持)
logging_strategy/eval_strategy/save_strategy"epoch"每个 epoch 记录日志、评估、保存检查点
load_best_model_at_endTrue训练结束时加载验证集最优检查点
metric_for_best_model"accuracy"以 accuracy 作为选优指标
report_to"tensorboard"日志上报到 TensorBoard
push_to_hub/hub_strategy/hub_model_idTrue/"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_processortokenizer)。ImageDistilTrainer.__init__通过**kwargs透传给基类Trainer,因此processing_classcompute_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系列)。

请注意:上述数字为原指南作者在特定数据集划分与超参数下的实测结果,仅供参考;在不同环境复现时数值可能略有波动。建议读者自行尝试不同的预训练教师、学生架构与蒸馏参数(temperaturelambda),并对比"从零训练基线"以验证蒸馏收益。

进阶方向:进一步探索

  • 更换教师/学生架构:教师可替换为其他在 Hub 上微调过的 ViT 变体或更大的 CNN;学生可尝试 MobileNetV3、EfficientNet 等轻量架构,验证压缩比与精度权衡;
  • 调整蒸馏超参数:增大temperature使软目标更平滑、更强调类间关系;调节lambda平衡蒸馏损失与标签损失;
  • 结合硬标签与软目标:本示例已同时使用两者;还可引入特征层蒸馏(对齐中间特征图)、attention 蒸馏等变体;
  • 生产部署:蒸馏后的轻量学生模型更易部署。仓库提供了模型导出支持(见 exporters)与 ONNX 等格式转换能力,可将训练好的 MobileNetV2 导出用于推理服务。

总结

本文以完整可运行的代码路径,演示了基于 🤗 TransformersTrainer的图像分类知识蒸馏:通过继承Trainer并重写compute_loss(),将 ViT 教师的软目标知识以 KL 散度形式注入 MobileNetV2 学生模型,并结合真实标签损失进行联合优化。文中同时给出了数据预处理、训练配置、评估指标、结果对比等完整闭环,并补充了Trainer.compute_loss调用链、ViTForImageClassificationMobileNetV2ForImageClassification损失计算实现、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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询