【免费下载链接】Model-Optimizer
A unified library of SOTA model optimization techniques like quantization, distillation, pruning, neural architecture search, speculative decoding, etc. It compresses deep learning models for downstream deployment frameworks like TensorRT-LLM, TensorRT, vLLM, etc. to optimize inference speed.
Model-Optimizer 的蒸馏 API(modelopt.torch.distill,简称mtd)通过一个元模型(meta-model)封装学生(Student)与教师(Teacher)两个模型,使您可以在几乎不改动原有训练脚本的情况下,仅增加一行损失计算代码即可启动知识蒸馏训练。本文以docs/source/guides/4_distillation.rst为骨架,结合仓库中modelopt/torch/distill/的源码实现与单元测试,完整讲解从模型转换、蒸馏训练、损失与平衡器配置,到 checkpoint 保存与恢复的端到端流程。
三步走:如何用 mtd 启动一条蒸馏训练管线
官方指南将整个流程浓缩为三个步骤,这也是modelopt.torch.distill的核心使用范式:
- 模型转换(convert):通过
mtd.convert()将学生模型与教师模型一起包装成一个更大的元模型(DistillationModel),抽象掉二者之间的交互细节; - 蒸馏训练(distillation training):直接用这个元模型替代原始模型跑原有训练脚本,只在损失计算处增加一行调用,即可完成知识转移;
- Checkpoint 保存与恢复:通过
mto.save()保存模型;注意mto.restore()恢复时不会重新实例化蒸馏元模型,以避免反序列化(unpickling)问题——恢复后拿到的是普通学生模型,需要时重新执行mtd.convert()即可。
该 API 的入口定义在 modelopt/torch/distill/distillation.py,其中convert()本质是调用apply_mode(model, mode, registry=DistillModeRegistry),将模型按注册的模式描述符转换为对应形态(见 mode.py)。
Convert:把学生模型转换为蒸馏元模型
使用mtd.convert()可以把任意nn.Module学生模型转换为DistillationModel。官方指南给出的最小示例:
import modelopt.torch.distill as mtd from torchvision.models import resnet50 # User-defined model (student) model = resnet50() # Configure and convert for distillation distillation_config = { # `teacher_model` is a model, model class, callable, or a tuple. # If a tuple, it must be of the form (model_cls_or_callable,) or # (model_cls_or_callable, args) or (model_cls_or_callable, args, kwargs). "teacher_model": teacher_model, "criterion": mtd.LogitsDistillationLoss(), "loss_balancer": mtd.StaticLossBalancer(), } distillation_model = mtd.convert(model, mode=[("kd_loss", distillation_config)]) # Export model in original class, with only previously-present attributes model_exported = mtd.export(distillation_model)其中mode支持字符串、Mode对象,或(mode, config)元组列表;kd_loss模式由 KnowledgeDistillationModeDescriptor 注册,其配置类为KDLossConfig。转换过程内部会先做严格校验(config._strict_validate()),再通过init_model_from_model_like()实例化教师模型,最后调用DistillationModel.modify()完成封装,详见 mode.py。
KDLossConfig 配置项说明
配置字典的合法键由 config.py 中的 KDLossConfig 定义(注意extra="forbid",传入未知键会直接报错):
| 字段 | 类型 | 说明 |
|---|---|---|
teacher_model | 模型 / 模型类 / callable / 元组 | 教师模型;元组形式为(model_cls_or_callable,)、(model_cls_or_callable, args)或(model_cls_or_callable, args, kwargs),默认None |
criterion | Loss或{(学生层名, 教师层名): Loss}字典 | 蒸馏损失。传入单个Loss实例时仅计算输出级(output-only)蒸馏;传入字典时为逐层配对蒸馏 |
loss_balancer | DistillationLossBalancer | 将多个蒸馏损失与学生原始任务损失合并为单个标量,默认None |
expose_minimal_state_dict | bool | 默认True,隐藏教师模型的state_dict以减小 checkpoint 体积;使用 FSDP 时应设为False |
criterion存在一个隐式归一化:字段校验器会把非字典形式的Loss实例统一包装成{("", ""): loss}(表示整模型输出级蒸馏),见 config.py。严格校验还规定:存在多个层损失对时必须提供 Loss Balancer,且 Loss Balancer 不允许带有可训练参数。
转换后的两个使用注意点
官方指南给出了两条易踩坑的提示:
type()失效但isinstance()有效:转换后模型类不再是原始学生类,调用type(model)不会得到预期结果;但由于模型动态地成为原类的子类(DistillationModel继承自DynamicModule),isinstance()依然成立。- 小数据量训练优先考虑
MFTLoss:当学生模型只在少量真实标签数据上训练时,建议用mtd.MFTLoss替代标准LogitsDistillationLoss。这样学生既能从教师的分布中学习,又能适应新数据,在不覆盖教师通用知识的前提下提升新数据的特化程度(详见下文 Minifinetuning 一节)。
mtd.export()则用于把蒸馏元模型还原为原始学生类。源码中它会先解包 DP/DDP 包装(并给出警告提示导出是 in-place 的,包装器需重新创建),再应用export_student模式(distillation.py)。
Distillation Concepts:核心概念与术语
为便于理解后续配置,官方文档以术语表形式给出了蒸馏的基本概念:
| 术语 | 含义 |
|---|---|
| Knowledge Distillation(知识蒸馏) | 将可学习的特征信息从教师模型迁移到学生模型 |
| Student(学生) | 待训练的模型(可以从零开始,也可以是预训练模型) |
| Teacher(教师) | 固定的、已预训练的模型,作为学生学习的目标范例 |
| Distillation loss(蒸馏损失) | 在学生与教师特征之间使用的损失函数,用于执行知识蒸馏,与学生原始任务损失相互独立 |
| Loss Balancer(损失平衡器) | 决定如何将蒸馏损失与学生原始任务损失合并为单个标量的工具 |
| Soft-label Distillation(软标签蒸馏) | 在教师与学生模型的输出 logits 之间执行知识蒸馏的具体过程 |
Knowledge Distillation(知识蒸馏)
蒸馏是一个宽泛的术语,泛指模型之间任何形式的信息压缩;本文特指基本的教师-学生知识蒸馏。其过程是在已训练模型(教师)与未训练模型(学生)之间建立一个辅助损失(或替换原始损失),期望学生学到教师已经掌握的信息(如特征图或 logits)。官方文档总结了四个典型用途:
- A. 模型尺寸缩减:更小、更高效的学生模型(可能是剪枝后的教师)达到接近甚至超过更大、更慢的教师模型的精度;
- B. 作为纯训练的替代方案:从现有模型蒸馏(再微调)通常比从头训练更快;
- C. 模块替换:将模型内某个模块替换为更高效的实现,并用蒸馏让替换后的输出与原模块输出对齐,从而无损地重新融入整体模型;
- D. 极小改动、避免灾难性遗忘:名为 Minifinetuning 的蒸馏变体,可以在很小的数据集上训练模型而不丢失原有知识。
Student(学生)与 Teacher(教师)
学生是最终希望训练并使用(或导出后部署)的模型,理想情况下满足目标架构与算力要求,但当前要么未训练、要么精度需要提升。教师则是提供已习得特征/信息用于构建损失来源的模型,通常比期望更大或更慢,但精度令人满意。在实现层面,教师模型会被冻结——DistillationModel.modify()中执行self._teacher_model.requires_grad_(False)(distillation_model.py)。
Distillation loss(蒸馏损失)
要真正"迁移"知识,需要向学生模型的原始损失函数中**添加(或替换)**一个优化目标。最简单的方式是对教师与学生两个尺寸相同的激活张量施加 MSE,前提假设是教师学到的特征质量高、应尽可能被模仿。
ModelOpt 支持为每一对层输出分别指定不同的损失函数,并提供若干预定义损失;用户也常常需要自定义损失。层配对到损失函数的映射通过配置字典的criterion键指定——顺序分别为学生、教师——损失函数本身也应接受同样顺序的输出:
# Example using pairwise-mapped criterion. # Will perform the loss on the output of ``student_model.classifier`` and ``teacher_model.layers.18`` distillation_config = { "teacher_model": teacher_model, "criterion": {("classifier", "layers.18"): mtd.LogitsDistillationLoss()}, } distillation_model = mtd.convert(student_model, mode=[("kd_loss", distillation_config)])中间层输出由DistillationModel通过 forward hook 捕获,随后用DistillationModel.compute_kd_loss()触发损失计算;若存在学生原始的非蒸馏损失,可作为参数传入。自定义损失函数往往是必要的——尤其是输出需要先经过处理才能得到 logits 或激活时;损失函数的额外参数可通过compute_kd_loss()的kwargs传入。
Loss Balancer(损失平衡器)
由于蒸馏损失可能施加于多对层,损失以字典形式返回,需要合并成标量才能反向传播。Loss Balancer(接口由DistillationLossBalancer定义)就是用来完成这一合并的。如果蒸馏损失只作用于一对层输出、且没有学生损失,则无需提供 Loss Balancer(源码在compute_kd_loss中对应断言:无 balancer 时传入 student_loss 会直接报错,见 distillation_model.py)。ModelOpt 提供了一个简单的StaticLossBalancer实现,用户也可基于上述接口编写自定义平衡器。
Soft-label Distillation(软标签蒸馏)
仅在分类模型输出 logits 上执行蒸馏的场景即软标签蒸馏。此时甚至可以完全省略学生原始的分类损失——如果教师输出被优先视为优于任何真实标签。
Minifinetuning
Minifinetuning 是一种允许模型在很小的数据集上训练而不丢失原有知识的技术。其核心是对教师的分布做算法性修正,取决于教师在新数据集上的表现:目标是保证正确 token 与错误 argmax token 之间的间隔足够大,该间隔由阈值(threshold)参数控制。ModelOpt 为此提供了预定义损失MFTDistillationLoss(即MFTLoss),可替代标准LogitsDistillationLoss使用。
源码级原理解读:DistillationModel 是如何工作的
DistillationModel是蒸馏的核心容器(distillation_model.py),它把多个教师与学生模型封装成单一模型,主要机制如下:
- 双路 forward:
forward()中先以torch.no_grad()冻结梯度地执行教师模型前向(并强制eval()),再执行学生前向;no_grad()让 PyTorch 不必为教师层保存激活,比仅冻结权重更省显存。它还提供了only_teacher_forward()与only_student_forward()上下文管理器,分别只跑教师或只跑学生前向,用于流水线并行等场景(distillation_model.py)。 - Hook 捕获中间激活:
_register_hooks()为每个 (学生层, 教师层) 对注册 forward hook,把层输出暂存到模块的_intermediate_output属性上;若前一次输出尚未被消费就又捕获到新输出,会给出警告(提示可能使用了激活检查点)。(distillation_model.py) - 损失聚合:
compute_kd_loss(student_loss=None, loss_reduction_fn=None, skip_balancer=False, **loss_fn_kwargs)逐对消费暂存的中间输出并调用对应损失(学生输出为 pred、教师输出为 target),组成{loss类名_序号: 损失值}字典;loss_reduction_fn可用于 loss-masking 场景(对非标量损失先做归约),skip_balancer=True时返回原始字典以便外部单独归约。(distillation_model.py) - 隐藏教师 state_dict:
state_dict()在expose_minimal_state_dict=True时用hide_teacher_model()临时把教师替换为空模块,从而在保存 checkpoint 时不重复存储教师权重;load_state_dict()则会智能判断 checkpoint 中是否存在教师/损失模块键,自动决定隐藏哪些部分(distillation_model.py)。
训练循环中的一行代码
蒸馏训练时,您只需把原来的loss.backward()之前的损失计算替换为:
loss = distillation_model.compute_kd_loss(student_loss=original_student_loss) loss.backward()如果配置中没有 Loss Balancer、且只存在单对蒸馏损失,直接loss = distillation_model.compute_kd_loss()即可拿到标量。单元测试 tests/unit/torch/distill/test_distill.py 中的test_distillation_model_no_balancer、test_distillation_model_multiloss_balancer、test_logits_distillation等用例分别验证了这些分支行为。
内置蒸馏损失函数
ModelOpt 在 modelopt/torch/distill/losses.py 中提供了三类预定义损失:
LogitsDistillationLoss(输出 logits 的 KL 散度)
mtd.LogitsDistillationLoss(temperature=1.0, reduction="mean")temperature:用于在计算损失前软化logits_t与logits_s的温度值;reduction:最终逐点损失的归约方式;传"none"可配合自己的归约函数(如 loss mask)使用。
其实现即经典 KD 损失:对双方 logits 除以温度后分别做log_softmax/softmax,再计算 KL 散度(假设类别 logits 维度在最后一维)。一个值得注意的实现细节:损失会乘以temperature ** 2,因为软 logits 产生的梯度幅值按1/(T^2)缩放,乘以T^2可保证调整温度超参时各 logits 的相对贡献基本不变(losses.py)。
MFTLoss(Minifinetuning 修正分布)
mtd.MFTLoss(temperature=1.0, threshold=0.2, reduction="batchmean")threshold:用于修正教师分布的阈值,保证正确与错误 argmax token 的间隔足够大,取值范围[0, 1],默认0.2。
其 forward 需要额外的labels(真实标签)参数,内部_prepare_corrected_distributions对教师分布进行修正:对 argmax 错误的 token,按(p_argmax - p_label + threshold) / (1 + p_argmax - p_label)计算混入因子把概率质量转移到正确标签;对 argmax 正确的 token(默认apply_threshold_to_all=True也一并处理),确保正确标签概率不低于1 - threshold(losses.py)。
MGDLoss(Masked Generative Distillation)
mtd.MGDLoss(num_student_channels, num_teacher_channels, alpha_mgd=1.0, lambda_mgd=0.65)针对视觉特征图(形状BxCxHxW)的掩码生成式蒸馏:当学生与教师通道数不同时自动插入1x1卷积对齐,并用量化生成器与随机掩码计算 MSE 损失(losses.py)。
Loss Balancer:把多个损失合并为标量
DistillationLossBalancer是平衡器接口(抽象方法forward(loss: dict) -> Tensor),StaticLossBalancer是其静态权重实现:
mtd.StaticLossBalancer(kd_loss_weight=0.5)kd_loss_weight为float时作用于所有蒸馏损失之和;为list时按criterion中指定的顺序对应每个蒸馏损失键逐一加权;- 若权重之和不等于 1.0,需把
student_loss传入compute_kd_loss(),差额权重将施加于学生损失上;权重和超出[0, 1]会抛出ValueError,小于 1 会给出警告。
具体聚合逻辑:蒸馏损失按权重加权求和,学生损失按1 - sum(kd_loss_weight)加权后相加(loss_balancers.py)。传入compute_kd_loss()的损失字典键格式为:学生损失键为student_loss,各蒸馏损失键为{损失类名}_{序号}(如MSELoss_0),测试test_distillation_model_multiloss_balancer验证了多损失+平衡器的组合行为。
分层蒸馏模式:layerwise_kd
除输出级kd_loss外,仓库还提供了layerwise_kd模式(LayerwiseKDConfig,见 config.py),其criterion必须是显式的层对字典,不支持输出级蒸馏。对应的LayerwiseDistillationModel(layerwise_distillation_model.py)适用于"学生 = 教师替换了部分子模块"的场景:
- 自动冻结学生除 criterion 指定层以外的所有层(仅保留待训练子模块的梯度);
- 将教师对应层的输入通过 forward pre-hook 注入学生的对应层(
student_input_bypass_fwd_hook),使被替换的子模块以教师中间特征为输入进行训练; - 若存在
lm_head,会把学生与教师的lm_head临时替换为nn.Identity()以省去不必要的计算,导出时再恢复。
在 HuggingFace 生态中的开箱即用集成:KDTrainer
对于大语言模型场景,仓库提供了开箱即用的 HF 蒸馏 Trainer(modelopt/torch/distill/plugins/huggingface.py)。其完整示例见 examples/llm_distill/main.py:
from modelopt.torch.distill.plugins.huggingface import KDTrainer class KDSFTTrainer(KDTrainer, SFTTrainer): pass trainer = KDSFTTrainer( model, # 学生模型 training_args, distill_args={"teacher_model": teacher_model}, # 预加载的教师模型 train_dataset=dset_train, eval_dataset=dset_eval, formatting_func=..., processing_class=tokenizer, ) trainer.train()该插件的使用要点:
- 仅支持 logits 级蒸馏(
criterion="logits_loss"),教师模型需先由用户预加载为nn.Module再传入distill_args.teacher_model;教师会被冻结(requires_grad_(False)); - 通过
DistillArguments可配置temperature(软化 logits 的温度)与liger_jsd_beta(启用 Liger Kernel 时 JSD 的 beta 系数,0=前向 KL,1=反向 KL); - 兼容 FSDP2、DeepSpeed ZeRO-3 与 DDP 等并行方案(FSDP1 不被支持,构造时会直接报错);
- 支持 Liger Kernel 融合的
lm_head + JSD蒸馏路径,对因果 LM 做了 logits 平移([..., :-1, :])与ignore_index掩码处理; - 评估阶段会把原始的 CE loss 作为附加指标
eval_ce_loss一并上报。
使用前建议开启mto.enable_huggingface_checkpointing()(示例中位于main.py第 76 行),以自动保存/加载 modelopt 状态。
Checkpoint:保存与恢复的正确姿势
训练完成后:
mto.save(distillation_model, "model.pth") # 保存 model = mto.restore("model.pth") # 恢复为普通学生模型由于expose_minimal_state_dict默认隐藏教师权重,保存的 checkpoint 不会重复存储教师参数,体积更小。mto.restore()不会重新实例化蒸馏元模型——这是刻意设计,以规避 unpickling 问题;如需继续蒸馏,恢复后再执行一次mtd.convert()即可。KDLossConfig还支持用model_dump()将配置转成字典(teacher_model会被原样保留而非序列化),便于记录训练配置(config.py)。测试test_distillation_save_restore、test_minimal_state_dict_mode、test_load_student_only_state分别覆盖了保存-恢复、最小 state dict 与仅学生权重加载等场景(tests/unit/torch/distill/test_distill.py)。
小结:蒸馏流程速查
mtd.convert(student, mode=[("kd_loss", config)])构建蒸馏元模型;- 训练循环中用
compute_kd_loss(student_loss=...)计算总损失并反向传播; - 需要部署时用
mtd.export(distillation_model)还原为原始学生类,或用mto.save保存 checkpoint; - 多对层损失务必配置
loss_balancer;小数据集微调优先MFTLoss;LM 蒸馏可直接使用KDTrainer插件。
完整的 API 参考可进一步查阅 modelopt/torch/distill/init.py 导出的模块,以及蒸馏相关的单元测试目录 tests/unit/torch/distill。
【免费下载链接】Model-Optimizer
A unified library of SOTA model optimization techniques like quantization, distillation, pruning, neural architecture search, speculative decoding, etc. It compresses deep learning models for downstream deployment frameworks like TensorRT-LLM, TensorRT, vLLM, etc. to optimize inference speed.
相关推荐
PaddleOCR知识蒸馏:小模型训练技巧
PaddleOCR知识蒸馏:小模型训练技巧 引言:为什么需要知识蒸馏? 在OCR(Optical Character Recognition,光学字符识别)领域
人工智能计算机视觉OCR深度学习大模型RAGModel-Optimizer 知识蒸馏实战指南:用 KDTrainer 将大模型知识迁移到小模型
Model Optimizer 知识蒸馏实战指南:用 KDTrainer 将大模型知识迁移到小模型 本文以 Model Optimizer 开源仓库中的 llm
PaddleOCR 知识蒸馏训练完全指南:DistillationModel 框架解析与检测/识别蒸馏配置实战
PaddleOCR 知识蒸馏训练完全指南:DistillationModel 框架解析与检测/识别蒸馏配置实战 知识蒸馏(Knowledge Distillat
人工智能计算机视觉OCR深度学习大模型RAG
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考