资讯动态

AutoGluon AutoMM 模型蒸馏实战指南:从 GLUE 到 PAWS-X 的知识蒸馏训练与调参

发布时间:2026/9/15 22:28:09 来源:尧图企业网站定制
AutoGluon AutoMM 模型蒸馏实战指南从 GLUE 到 PAWS-X 的知识蒸馏训练与调参【免费下载链接】autogluonFast and Accurate ML in 3 Lines of Code项目地址: https://gitcode.com/GitHub_Trending/au/autogluon导读本指南基于 AutoGluon 多模态预测模块AutoMM的官方蒸馏示例系统讲解如何使用MultiModalPredictor完成文本模型的知识蒸馏Knowledge Distillation。你将掌握蒸馏训练脚本的完整命令行参数、GLUE 与 PAWS-X 两类任务的实验设置、六大蒸馏损失函数的原理与权重调参方法并结合仓库源码理解fit(teacher_predictor...)的底层实现链路。完成阅读后你可以直接用给出的命令复现实验结果也能为自定义任务设计蒸馏配置。1. 示例概览一行fit()开启蒸馏仓库中的蒸馏示例位于 examples/automm/distillation/包含两个可运行脚本automm_distillation_glue.py在 GLUE 七项 NLP 任务上演示蒸馏automm_distillation_pawsx.py在跨语言的 PAWS-X 数据集上演示多语言蒸馏。两者的核心调用方式完全一致先用一个较大的“教师模型”训练并保存再在MultiModalPredictor.fit()中通过teacher_predictor...参数把教师模型的知识蒸馏进较小的“学生模型”。例如 GLUE 示例中学生模型的蒸馏调用student_predictor MultiModalPredictor(labellabel, eval_metricaccuracy) student_predictor.fit( train_df, hyperparameters{ env.num_gpus: args.num_gpu, optim.max_epochs: args.max_epochs, model.hf_text.checkpoint_name: args.student_model, distiller.temperature: args.temperature, distiller.hard_label_weight: args.hard_label_weight, distiller.soft_label_weight: args.soft_label_weight, distiller.softmax_regression_weight: args.softmax_regression_weight, distiller.output_feature_loss_weight: args.output_feature_loss_weight, distiller.rkd_distance_loss_weight: args.rkd_distance_loss_weight, distiller.rkd_angle_loss_weight: args.rkd_angle_loss_weight, distiller.soft_label_loss_type: args.soft_label_loss_type, distiller.softmax_regression_loss_type: args.softmax_regression_loss_type, distiller.output_feature_loss_type: args.output_feature_loss_type, }, teacher_predictorteacher_predictor, seedargs.seed, )从源码看teacher_predictor参数在 predictor.py 中接受MultiModalPredictor对象或其保存路径字符串内部统一转换为teacher_learner后传入self._learner.fit()。真正承载蒸馏训练循环的是 DistillerLitModule一个 PyTorch Lightning 模块它由 base.py 在学生训练时自动创建并传入student_model、teacher_model以及全部蒸馏超参数。2. 运行方式与命令行参数详解两个脚本均通过 argparse 提供统一的命令行入口格式为python automm_distillation_task_name.py --flag value完整参数清单如下表默认值取自两个脚本的argparse定义参数适用脚本作用默认值glue_taskGLUE指定在哪个 GLUE 任务上运行实验见第 3 节数据集qnlipawsx_teacher_tasksPAWS-X教师模型使用哪些语言的数据训练[en,de,es,fr,ja,ko,zh]pawsx_student_tasksPAWS-X学生模型使用哪些语言的数据训练[en,de,es,fr,ja,ko,zh]teacher_model两者教师模型HuggingFace checkpoint 名称GLUEgoogle/bert_uncased_L-12_H-768_A-12PAWS-Xmicrosoft/mdeberta-v3-basestudent_model两者学生模型 checkpoint 名称GLUEgoogle/bert_uncased_L-6_H-768_A-12PAWS-Xnreimers/mMiniLMv2-L6-H384-distilled-from-XLMR-Largeseed两者随机种子123precisionPAWS-XGPU 运算精度脚本默认bf16README 复现示例中使用16bf16max_epochs两者学生模型最大训练轮数GLUE1000由命令行显式指定PAWS-X10time_limit两者每个模型的最大训练时间秒Nonenum_gpu两者训练使用的 GPU 数量-1表示使用全部-1temperature两者软交叉熵soft cross entropy的温度系数5.0hard_label_weight两者硬标签损失ground truth logits权重0.1soft_label_weight两者软标签损失教师模型输出权重1.0softmax_regression_weight两者softmax 回归损失权重GLUE0.1PAWS-X0output_feature_loss_weight两者输出特征损失权重0.01rkd_distance_loss_weight两者RKD 距离损失权重0.0rkd_angle_loss_weight两者RKD 角度损失权重0.0soft_label_loss_type两者软标签损失函数类型空字符串自动选择softmax_regression_loss_type两者softmax 回归损失函数类型mseoutput_feature_loss_type两者输出特征损失函数类型支持mean_square_error与cosine_distancemsefinetuned_model_cache_folder/save_path两者无蒸馏训练的模型缓存目录./AutogluonModels/cache_finetunedresume两者为True时从缓存加载已训练好的无蒸馏模型跳过重复训练Falsetrain_nodistillGLUE是否训练无蒸馏的对照学生模型Trueaug_scalePAWS-X文本平凡增强的缩放系数0.0值得注意的一个细节是脚本中的soft_label_loss_type默认值为空字符串。从 default.yaml 看蒸馏器默认配置为soft_label_loss_type: 、softmax_regression_loss_type: mse、output_feature_loss_type: mseAutoMM 会根据任务类型如回归任务 stsb自动选择合适的软标签损失函数。2.1 缓存与断点续训机制两个脚本都采用“先训练教师再训练无蒸馏学生对照最后蒸馏学生”的三段式流程并利用--resume与save_path实现缓存复用教师模型缓存路径为{save_path}/{task}-{teacher_model}模型名中的/被替换为-无蒸馏学生模型缓存路径为{save_path}/{task}-{student_model}每次尝试MultiModalPredictor.load(path)若加载成功且开启--resume则直接复用否则重新训练并save()到缓存。这意味着跑多组蒸馏权重对比实验时只需在第一次执行时训练教师与对照模型后续实验可以全部命中缓存大幅节省算力。3. 数据集说明GLUE 与 PAWS-X3.1 GLUE 任务示例借用了 GLUE 基准 [1] 中的 7 个 NLP 任务任务名沿用原论文缩写mnli(m/mm)、qqp、qnli、sst2、stsb、mrpc、rte。数据通过 HuggingFacedatasetsAPI 的load_dataset(glue, task)加载若本地不存在会自动从 HuggingFace 下载因此运行环境需要联网。脚本中的任务配置见 automm_distillation_glue.py为GLUE_METRICS { mnli: {val: accuracy, eval: [accuracy]}, qqp: {val: accuracy, eval: [accuracy, f1]}, qnli: {val: accuracy, eval: [accuracy]}, sst2: {val: accuracy, eval: [accuracy]}, stsb: {val: pearsonr, eval: [pearsonr, spearmanr]}, mrpc: {val: accuracy, eval: [accuracy]}, rte: {val: accuracy, eval: [accuracy]}, }其中val是训练期间用于模型选择/早停的验证指标eval是最终评测时报告的指标列表qqp同时报告 accuracy 与 F1stsb是回归任务报告 Pearson/Spearman 相关系数因此会自动选择对应的软标签损失函数。GLUE 的测试集标签不公开因此示例用训练集做训练/验证用官方验证集做测试。此外脚本还支持mnlim与mnlimm两个变体参数来分别选择 MNLI 的 matched/mismatched 验证集见 automm_distillation_glue.py。3.2 PAWS-X 数据集PAWS-X [3] 是 PAWS 中 Wikipedia 部分扩展到六种语言西班牙语、法语、德语、中文、日语、韩语的跨语言版本每个目标语言包含 23,659 对人类翻译的释义判断样本训练集则是对原始 PAWS 英文训练集49,401 对的机器翻译结果与 XNLI [4] 的多语料构建思路一致。示例用它测试多语言场景下的蒸馏效果通过--pawsx_teacher_tasks与--pawsx_student_tasks分别指定教师和学生使用的语言子集合法值来自PAWS_TASKS [en, de, es, fr, ja, ko, zh]。脚本会为每个语言加载 train/validation/test 三个划分并pd.concat拼接成训练与验证数据评测时分别报告每个语言测试集上的 accuracy教师与学生的推理时间也分别统计用于计算加速比。4. 蒸馏损失函数六项加权的组成与原理学生模型的总体损失是六项损失的加权和hard_label_loss、soft_label_loss、softmax_regression_loss、output_feature_loss、rkd_distance_loss、rkd_angle_loss。对应的加权求和逻辑实现在 DistillerLitModule._compute_loss 中只有权重大于 0 的损失项才会被加入总损失这样你可以通过将权重设为 0 来单独开启/关闭某一项。4.1 Hard Label Loss硬标签损失学生模型直接以真实标签为监督使用标准的交叉熵等分类损失。从 lit_distiller.py 可以看出该损失会遍历学生模型每个输出头按各头权重加权求和。4.2 Soft Label Loss软标签损失教师模型输出的 logits 经温度缩放后作为软标签指导学生。实现中lit_distiller.py学生与教师的 logits 都会先除以temperature若损失函数是nn.CrossEntropyLoss教师 logits 会先经过F.softmax变成概率分布再参与计算。较高的温度让概率分布更平滑能传递更多的类间相似性信息。4.3 Output Feature Loss输出特征损失Output Feature Loss 是教师与学生模型输出特征embedding之间的某种距离度量欧氏距离、余弦距离等对应 lit_distiller.py 中的_compute_output_feature_loss。由于师生模型输出维度可能不同学生特征会先经过一个output_feature_adaptor如nn.Linear映射到教师特征维度再计算损失损失函数由output_feature_loss_type指定支持mean_square_error与cosine_distance。README 的 GLUE 实验正是通过调节这一项的权重来展示其对蒸馏效果的关键作用。4.4 Softmax Regression Losssoftmax 回归损失Softmax Regression Loss 来源于论文《Knowledge Distillation via Softmax Regression Representation Learning》[5]。从实现看lit_distiller.py它的做法是把学生特征经output_feature_adaptor映射后送入教师的分类头self.teacher_model.head得到学生视角的 logits再与教师 logits 计算损失从而把“教师分类器”的判别能力也作为监督信号。这种损失在 PAWS-X 实验中显示出明显效果。4.5 RKD Distance/Angle Loss关系知识蒸馏损失RKD 损失源于《Relational Knowledge Distillation》[6]关注的是样本之间的结构化关系而非单个样本的输出实现在 rkd_loss.py 的RKDLoss中Distance Loss计算批内样本两两欧氏距离pdist并对师生两边的距离矩阵做归一化后用smooth_l1_loss对齐Angle Loss计算样本两两连线的夹角余弦矩阵同样用smooth_l1_loss对齐教师与学生的角度关系。RKDLoss内部的默认权重为距离 25、角度 50但在示例中通过distiller.rkd_distance_loss_weight与distiller.rkd_angle_loss_weight在总损失层面显式控制且_compute_loss只有在两者任一大于 0 时才计算 RKD 项。5. 实验复现与性能对照5.1 GLUE 上的输出特征损失消融README 用 QNLI 任务展示 output feature loss 的重要性其中max_epoch12脚本实际参数名为max_epochsglue_taskqnli teacher_modelgoogle/bert_uncased_L-12_H-768_A-12 student_modelgoogle/bert_uncased_L-6_H-768_A-12 seed123 max_epoch12 metricaccuracy temperature5 hard_label_weight0.1 soft_label_weight1 python3 automm_distillation_glue.py --teacher_model ${teacher_model} \ --student_model ${student_model} \ --seed ${seed} \ --max_epoch ${max_epoch} \ --hard_label_weight ${hard_label_weight} \ --soft_label_weight ${soft_label_weight} \ --glue_task ${glue_task}复现最佳结果的命令教师为 BERT-Large 12 层、学生为 6 层同结构模型python3 automm_distillation.py --teacher_model google/bert_uncased_L-12_H-768_A-12 \ --student_model google/bert_uncased_L-6_H-768_A-12 \ --seed 123 \ --max_epoch 8 \ --hard_label_weight 0.5 \ --soft_label_weight 5实验对比结果如下数据来自 README 表格Distillation Ratio [2] 表示学生模型在教师与无蒸馏基线之间恢复的性能比例Speed Up 为教师推理时间与学生推理时间之比output_feature_loss_weightTeacher Model AccPretrained Model AccStudent Model AccDistillation RatioSpeed Up00.917260.894010.897130.133.52x0.010.917260.894010.902980.393.52x0.10.917260.894010.89365-0.023.52x从表格可以读出两个关键结论其一加入 0.01 的输出特征损失权重能将蒸馏比例从 0.13 提升到 0.39说明特征对齐能有效帮助小模型逼近大模型其二权重并非越大越好0.1 的权重反而使学生模型低于无蒸馏基线说明过强的特征约束会干扰分类损失的学习蒸馏权重需要精细调参。5.2 PAWS-X 上的 softmax 回归与 RKD 损失消融PAWS-X 实验使用的公共设置如下# pawsx_teacher_tasks [en,de,es,fr,ja,ko,zh] # pawsx_student_tasks [en,de,es,fr,ja,ko,zh] teacher_modelmicrosoft/mdeberta-v3-base student_modelnreimers/mMiniLMv2-L6-H384-distilled-from-XLMR-Large seed123 precision16 max_epochs10 num_gpu-1 temperature5 hard_label_weight0.1 soft_label_weight1 softmax_regression_loss_typemse output_feature_loss_typemse复现最佳模型的命令python3 automm_distillation_pawsx.py --precision 16 \ --output_feature_loss_weight 0.01 \ --softmax_regression_weight 0.1 \ --rkd_distance_loss_weight 1 \ --rkd_angle_loss_weight 2对照实验矩阵结果取各语言测试集 accuracy 的平均数据来自 READMETeacher ModelStudent ModelSoftmax Regression WeightRKD Distance WeightRKD Angle WeightTeacher PerformanceNo Distill PerformanceStudent PerformanceDistillation RatioSpeed Upmdeberta-v3-basemMiniLMv2-L6-H3840000.912930.884930.888570.13010~3.3xmdeberta-v3-basemMiniLMv2-L6-H3840.01000.912930.884930.885500.02041~3.3xmdeberta-v3-basemMiniLMv2-L6-H3840.1000.912930.884930.890930.21429~3.3xmdeberta-v3-basemMiniLMv2-L6-H3840.10.10.20.912930.884930.892640.27551~3.3xmdeberta-v3-basemMiniLMv2-L6-H3840.1120.912930.884930.891000.21684~3.3xmdeberta-v3-basemMiniLMv2-L6-H3840.110200.912930.884930.889710.17092~3.3x该实验的要点softmax 回归损失从 0 提升到 0.1 时蒸馏比例从 0.13 提升到 0.21在此基础上叠加 RKD 距离/角度损失0.1/0.2进一步把蒸馏比例提升到 0.27551达到表格中的最佳效果。而继续增大 RKD 权重10/20反而略有下降再次印证“中等权重最优”的调参规律。整体上学生模型可获得约 3.3 倍的推理加速。5.3 关于实验数据的说明以上所有数字均直接引用自仓库 README 的实验记录使用的是论文引用的模型与数据集。读者在自己机器上复现时由于环境GPU 型号、库版本、随机种子差异结果可能略有浮动建议把表格数值当作相对趋势参考并结合第 4 节的损失原理指导自己的权重搜索。6. 蒸馏的底层实现链路深入源码可以完整还原fit(teacher_predictor...)的调用链MultiModalPredictor.fit()在 predictor.py 中把teacher_predictor统一转换为teacher_learner字符串路径或_learner对象传给self._learner.fit()学生 Learner 的 get_litmodule_per_run 检测到self._teacher_learner非空时构造DistillerLitModule(student_model..., teacher_model...)而非普通LitModuleDistillerLitModule 的_shared_step中学生模型正常前向并保留梯度而教师模型切换为eval()并在torch.no_grad()下前向教师不参与梯度更新各项损失按权重加权求和后回传优化器只更新学生模型、输出特征适配器等可训练参数。默认的蒸馏器配置温度 5.0、硬标签权重 0.1、软标签权重 1.0、softmax 回归权重 0、输出特征损失权重 0.01、RKD 权重为 0记录在 configs/distiller/default.yaml与两个示例脚本的 argparse 默认值保持一致你可以把它作为自定义蒸馏任务的起点。7. 调参建议与实践要点结合源码实现与 README 的消融数据给出以下实践建议师生模型结构对齐README 明确建议学生模型与教师模型使用相同 backbone 家族如同为 BERT 或同为 ELECTRA这能保证 logits 与特征空间的语义可比性是软标签与特征损失生效的前提从默认权重出发温度 5、硬标签权重 0.1、软标签权重 1 是经过验证的起点output_feature_loss_weight建议在 0.01 量级尝试softmax 回归与 RKD 权重可从 0 逐步加到 0.1/1 量级切忌一开始就开大权重善用缓存复跑先用一次完整运行生成教师与无蒸馏基线的缓存--save_path之后用--resume复跑不同蒸馏权重组合只重训学生模型节省大量算力按任务选损失回归类任务如 GLUE 的 stsb会自动切换软标签损失函数output_feature_loss_type在 MSE 与 cosine distance 之间按特征尺度选择控制蒸馏比例脚本会在结尾打印 Distillation Ratio (学生 - 无蒸馏基线) / (教师 - 无蒸馏基线)若该值接近 0 或为负说明蒸馏配置无效甚至有害应回调权重该公式的实现见两个示例脚本的结尾打印逻辑。参考论文[1] Wang A, Singh A, Michael J, et al. GLUE: A multi-task benchmark and analysis platform for natural language understanding. arXiv preprint arXiv:1804.07461, 2018.[2] H He, X Shi, J Mueller, S Zha, M Li, G Karypis. Towards Automated Distillation: A Systematic Study of Knowledge Distillation in Natural Language Processing. International Conference on Automated Machine Learning: Late-Breaking Workshop, 2022.[3] Yang Y, Zhang Y, Tar C, et al. PAWS-X: A cross-lingual adversarial dataset for paraphrase identification. arXiv preprint arXiv:1908.11828, 2019.[4] Conneau A, Lample G, Rinott R, et al. XNLI: Evaluating cross-lingual sentence representations. arXiv preprint arXiv:1809.05053, 2018.[5] Yang J, Martinez B, Bulat A, et al. Knowledge distillation via softmax regression representation learning. ICLR 2021.[6] Park W, Kim D, Lu Y, et al. Relational knowledge distillation. CVPR 2019: 3967-3976.【免费下载链接】autogluonFast and Accurate ML in 3 Lines of Code项目地址: https://gitcode.com/GitHub_Trending/au/autogluon创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价