API 实战指南:以最小代码改动完成 Teacher-Student 蒸馏训练)
【免费下载链接】Model-OptimizerA 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.项目地址https://gitcode.com/GitHub_Trending/te/Model-Optimizer点击查看免费下载Model-Optimizer 的蒸馏 APImodelopt.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, registryDistillModeRegistry)将模型按注册的模式描述符转换为对应形态见 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 定义注意extraforbid传入未知键会直接报错字段类型说明teacher_model模型 / 模型类 / callable / 元组教师模型元组形式为(model_cls_or_callable,)、(model_cls_or_callable, args)或(model_cls_or_callable, args, kwargs)默认NonecriterionLoss或{(学生层名, 教师层名): Loss}字典蒸馏损失。传入单个Loss实例时仅计算输出级output-only蒸馏传入字典时为逐层配对蒸馏loss_balancerDistillationLossBalancer将多个蒸馏损失与学生原始任务损失合并为单个标量默认Noneexpose_minimal_state_dictbool默认True隐藏教师模型的state_dict以减小 checkpoint 体积使用 FSDP 时应设为Falsecriterion存在一个隐式归一化字段校验器会把非字典形式的Loss实例统一包装成{(, ): loss}表示整模型输出级蒸馏见 config.py。严格校验还规定存在多个层损失对时必须提供 Loss Balancer且 Loss Balancer 不允许带有可训练参数。转换后的两个使用注意点官方指南给出了两条易踩坑的提示type()失效但isinstance()有效转换后模型类不再是原始学生类调用type(model)不会得到预期结果但由于模型动态地成为原类的子类DistillationModel继承自DynamicModuleisinstance()依然成立。小数据量训练优先考虑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 上执行蒸馏的场景即软标签蒸馏。此时甚至可以完全省略学生原始的分类损失——如果教师输出被优先视为优于任何真实标签。MinifinetuningMinifinetuning 是一种允许模型在很小的数据集上训练而不丢失原有知识的技术。其核心是对教师的分布做算法性修正取决于教师在新数据集上的表现目标是保证正确 token 与错误 argmax token 之间的间隔足够大该间隔由阈值threshold参数控制。ModelOpt 为此提供了预定义损失MFTDistillationLoss即MFTLoss可替代标准LogitsDistillationLoss使用。源码级原理解读DistillationModel 是如何工作的DistillationModel是蒸馏的核心容器distillation_model.py它把多个教师与学生模型封装成单一模型主要机制如下双路 forwardforward()中先以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_lossNone, loss_reduction_fnNone, skip_balancerFalse, **loss_fn_kwargs)逐对消费暂存的中间输出并调用对应损失学生输出为 pred、教师输出为 target组成{loss类名_序号: 损失值}字典loss_reduction_fn可用于 loss-masking 场景对非标量损失先做归约skip_balancerTrue时返回原始字典以便外部单独归约。distillation_model.py隐藏教师 state_dictstate_dict()在expose_minimal_state_dictTrue时用hide_teacher_model()临时把教师替换为空模块从而在保存 checkpoint 时不重复存储教师权重load_state_dict()则会智能判断 checkpoint 中是否存在教师/损失模块键自动决定隐藏哪些部分distillation_model.py。训练循环中的一行代码蒸馏训练时您只需把原来的loss.backward()之前的损失计算替换为loss distillation_model.compute_kd_loss(student_lossoriginal_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(temperature1.0, reductionmean)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。MFTLossMinifinetuning 修正分布mtd.MFTLoss(temperature1.0, threshold0.2, reductionbatchmean)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_allTrue也一并处理确保正确标签概率不低于1 - thresholdlosses.py。MGDLossMasked Generative Distillationmtd.MGDLoss(num_student_channels, num_teacher_channels, alpha_mgd1.0, lambda_mgd0.65)针对视觉特征图形状BxCxHxW的掩码生成式蒸馏当学生与教师通道数不同时自动插入1x1卷积对齐并用量化生成器与随机掩码计算 MSE 损失losses.py。Loss Balancer把多个损失合并为标量DistillationLossBalancer是平衡器接口抽象方法forward(loss: dict) - TensorStaticLossBalancer是其静态权重实现mtd.StaticLossBalancer(kd_loss_weight0.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必须是显式的层对字典不支持输出级蒸馏。对应的LayerwiseDistillationModellayerwise_distillation_model.py适用于学生 教师替换了部分子模块的场景自动冻结学生除 criterion 指定层以外的所有层仅保留待训练子模块的梯度将教师对应层的输入通过 forward pre-hook 注入学生的对应层student_input_bypass_fwd_hook使被替换的子模块以教师中间特征为输入进行训练若存在lm_head会把学生与教师的lm_head临时替换为nn.Identity()以省去不必要的计算导出时再恢复。在 HuggingFace 生态中的开箱即用集成KDTrainer对于大语言模型场景仓库提供了开箱即用的 HF 蒸馏 Trainermodelopt/torch/distill/plugins/huggingface.py。其完整示例见 examples/llm_distill/main.pyfrom modelopt.torch.distill.plugins.huggingface import KDTrainer class KDSFTTrainer(KDTrainer, SFTTrainer): pass trainer KDSFTTrainer( model, # 学生模型 training_args, distill_args{teacher_model: teacher_model}, # 预加载的教师模型 train_datasetdset_train, eval_datasetdset_eval, formatting_func..., processing_classtokenizer, ) trainer.train()该插件的使用要点仅支持 logits 级蒸馏criterionlogits_loss教师模型需先由用户预加载为nn.Module再传入distill_args.teacher_model教师会被冻结requires_grad_(False)通过DistillArguments可配置temperature软化 logits 的温度与liger_jsd_beta启用 Liger Kernel 时 JSD 的 beta 系数0前向 KL1反向 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小数据集微调优先MFTLossLM 蒸馏可直接使用KDTrainer插件。完整的 API 参考可进一步查阅 modelopt/torch/distill/init.py 导出的模块以及蒸馏相关的单元测试目录 tests/unit/torch/distill。赞分享【免费下载链接】Model-OptimizerA 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.项目地址https://gitcode.com/GitHub_Trending/te/Model-Optimizer点击查看免费下载相关推荐PaddleOCR知识蒸馏小模型训练技巧PaddleOCR知识蒸馏小模型训练技巧 引言为什么需要知识蒸馏 在OCROptical Character Recognition光学字符识别领域人工智能计算机视觉OCR深度学习大模型RAGModel-Optimizer 知识蒸馏实战指南用 KDTrainer 将大模型知识迁移到小模型Model Optimizer 知识蒸馏实战指南用 KDTrainer 将大模型知识迁移到小模型 本文以 Model Optimizer 开源仓库中的 llmPaddleOCR 知识蒸馏训练完全指南DistillationModel 框架解析与检测/识别蒸馏配置实战PaddleOCR 知识蒸馏训练完全指南DistillationModel 框架解析与检测/识别蒸馏配置实战 知识蒸馏Knowledge Distillat人工智能计算机视觉OCR深度学习大模型RAG上一篇Bundlephobia微服务改造将单体应用拆分为独立服务的思路下一篇APT-Hunter核心功能详解从事件日志中挖掘APT攻击痕迹的完整教程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考