ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

知识蒸馏实战:用mattevans-distil压缩BERT模型推理延迟

知识蒸馏实战:用mattevans-distil压缩BERT模型推理延迟 简介这份开源项目资源包对应 GitHub 上的 mattevans/distil 项目聚焦内存数据集过滤这一细分场景面向需要处理数据筛选逻辑的 Go 语言开发者与后端工程师。项目以轻量化的方式提供了一套可复用的过滤算子实现适合在数据管道、查询引擎或测试工具中直接引入使用。压缩包共 40 个文件整体仅 29KB其中 34 个为 Go 源码文件覆盖等于、不等于、大于、小于、包含、匹配、空值判断等常见比较与过滤逻辑并配有对应的单元测试另有 2 个 yml 配置文件用于持续集成以及 md 说明文档、json 示例数据、license 授权文件和 gitignore 忽略规则结构规整、便于快速阅读与二次开发。目前已有 193 人学习下载。读者可借此理解内存级数据集过滤的接口设计与实现思路参考其测试用例组织方式并将这些算子直接迁移到自己的 Go 项目中减少重复造轮子的成本。1. 从 mattevans-distil.zip 说起一个被低估的模型蒸馏实战包第一次拿到mattevans-distil.zip的时候我正被一个线上推理延迟的问题折磨。业务侧要求把 BERT-base 的响应时间压到 30ms 以内但 GPU 资源已经排满加机器走不通。当时试过剪枝、量化效果都不够稳。后来翻到这个包里面是一套完整的知识蒸馏Knowledge Distillation实现从教师模型加载、软标签生成到学生模型训练和导出链路是通的。它解决的核心问题很具体用一个大模型教一个小模型让小模型在保持大部分精度的前提下把推理成本降下来。适合谁做 NLP 模型部署、边缘端推理、或者想在有限算力下跑通蒸馏流程的工程师。这不是一个玩具 demo而是一个能直接改配置跑起来的工程骨架。2. 蒸馏原理与 mattevans-distil 的工程选型为什么不是直接微调2.1 知识蒸馏到底在蒸什么很多人第一次接触蒸馏会以为是把大模型的权重“压缩”进小模型。不是。蒸馏传递的是软标签soft labels——教师模型对每个样本输出的概率分布。比如一个三分类任务教师对某样本输出[0.7, 0.2, 0.1]这个分布里包含了“第二类也有点像”的暗知识dark knowledge而硬标签只给[1, 0, 0]。学生模型同时学硬标签的交叉熵和软标签的 KL 散度损失函数通常写成Loss α * CE(student_logits, hard_labels) (1 - α) * T² * KL(student_soft, teacher_soft)这里的T是温度系数用来平滑概率分布。T1时就是普通 softmaxT越大分布越平暗知识越明显。mattevans-distil默认把T设为 4α设为 0.3这个组合在文本分类任务上比较稳。为什么乘T²因为 softmax 求导时梯度会带1/T²的缩放乘回来是为了让软标签损失和硬标签损失在量级上对齐。2.2 为什么选这个包而不是自己从零写自己写蒸馏不是不行但有几个工程细节容易翻车。第一教师模型和学生模型的 tokenizer 必须一致否则软标签对不上第二训练时教师模型要冻结且切到eval()模式否则 dropout 和 BN 会引入噪声第三软标签的 batch 对齐和 padding mask 处理稍不注意就会把 padding 位置的 loss 算进去。mattevans-distil把这些都封装好了目录结构大致是mattevans-distil/ ├── configs/ │ ├── teacher_bert_base.yaml │ └── student_tinybert.yaml ├── src/ │ ├── distill_trainer.py │ ├── data_collator.py │ └── model_utils.py ├── scripts/ │ ├── train_distill.sh │ └── export_onnx.py └── requirements.txtdistill_trainer.py里继承的是 HuggingFace 的Trainer重写了compute_loss方法。这个设计的好处是你原来用Trainer的代码几乎不用大改换个 trainer 类就能跑蒸馏。data_collator.py里处理了动态 padding 和 attention mask避免 padding token 参与 loss 计算。这些细节自己写至少要调半天。2.3 环境准备与依赖安装我一般会先建一个干净的 conda 环境避免和现有项目的 torch 版本冲突。这个包对transformers的版本有要求太新或太旧都可能报Trainer接口不兼容。conda create -n distil_env python3.9 -y conda activate distil_env pip install torch1.13.1cu117 -f https://download.pytorch.org/whl/torch_stable.html pip install transformers4.28.1 datasets2.12.0 accelerate0.18.0 pip install onnx onnxruntime # 导出 ONNX 时需要这里锁transformers4.28.1是因为Trainer的compute_loss签名在这个版本之后改过用 4.30 会报unexpected keyword argument。accelerate是Trainer的依赖不装会提示找不到分布式后端。ONNX 相关的是为了后续导出推理用如果只训练可以先不装。2.4 配置文件的关键参数怎么改configs/teacher_bert_base.yaml和student_tinybert.yaml是核心。教师配置一般不用大动学生配置里几个参数决定蒸馏效果# student_tinybert.yaml 关键片段 model_name_or_path: prajjwal1/bert-tiny temperature: 4.0 alpha: 0.3 learning_rate: 5e-5 num_train_epochs: 10 per_device_train_batch_size: 32 max_seq_length: 128temperature和alpha是蒸馏特有的其他和普通微调一样。learning_rate我试过3e-5到1e-45e-5最稳。batch_size如果显存不够可以降到 16但要把gradient_accumulation_steps设成 2 来保持等效 batch。max_seq_length要和教师模型一致否则学生看到的输入长度不同软标签对不齐。3. 跑通蒸馏训练从数据准备到模型导出3.1 数据格式与 DataLoader 的坑这个包默认吃的是 JSONL 格式每行一个样本字段是text和label。我拿一个情感分类数据集举例{text: 这个电影太好看了演员演技在线, label: 1} {text: 剧情拖沓看不下去, label: 0}放到data/train.jsonl和data/dev.jsonl。data_collator.py里用的是tokenizer的动态 padding所以不需要提前 pad 到固定长度。但有个坑如果text字段里有空字符串tokenizer 会返回全 padding 的序列attention mask 全 0loss 会变成 NaN。我一般在数据预处理阶段就过滤掉长度小于 2 的样本。# 数据清洗脚本放在 scripts/clean_data.py import json def clean_jsonl(input_path, output_path, min_len2): with open(input_path, r, encodingutf-8) as fin, \ open(output_path, w, encodingutf-8) as fout: for line in fin: item json.loads(line) text item.get(text, ).strip() if len(text) min_len: continue # 跳过过短样本避免全 padding fout.write(json.dumps(item, ensure_asciiFalse) \n) if __name__ __main__: clean_jsonl(data/train_raw.jsonl, data/train.jsonl) clean_jsonl(data/dev_raw.jsonl, data/dev.jsonl)这个脚本逻辑很简单就是逐行读、过滤、逐行写。min_len2是我踩过坑之后定的中文里两个字符以下基本没有语义信息留着只会引入噪声。ensure_asciiFalse保证中文不被转义成\uXXXX方便人工检查。3.2 启动蒸馏训练训练入口是scripts/train_distill.sh里面调的是torchrun或python -m。单卡直接跑cd mattevans-distil python -m src.distill_trainer \ --teacher_config configs/teacher_bert_base.yaml \ --student_config configs/student_tinybert.yaml \ --train_file data/train.jsonl \ --dev_file data/dev.jsonl \ --output_dir outputs/distil_tinybert \ --do_train --do_eval \ --logging_steps 50 \ --save_steps 500 \ --fp16--fp16在支持 Tensor Core 的卡上能省显存、提速但如果你用的是老卡比如 V100 之前的开了可能反而慢还会出现 loss scaling 溢出。--logging_steps 50是每 50 步打一次日志方便观察 loss 下降曲线。--save_steps 500是每 500 步存一个 checkpoint防止训练中断白跑。output_dir里最后会生成pytorch_model.bin、config.json和tokenizer.json可以直接用AutoModelForSequenceClassification.from_pretrained加载。3.3 训练过程中的监控指标蒸馏训练和普通微调不一样你要同时看两个指标学生模型在验证集上的准确率和学生与教师软标签的 KL 散度。前者看最终效果后者看蒸馏是否在“学进去”。如果 KL 散度一直不降说明温度或 alpha 设得不对学生没从教师那里学到东西。我一般会在compute_loss里把两个 loss 分量都 log 出来# 在 distill_trainer.py 的 compute_loss 里加一行 def compute_loss(self, model, inputs, return_outputsFalse): outputs model(**inputs) student_logits outputs.logits # 教师 logits 已经在 data_collator 里预计算好存在 inputs[teacher_logits] teacher_logits inputs.pop(teacher_logits) # 硬标签损失 hard_loss F.cross_entropy(student_logits, inputs[labels]) # 软标签损失 T self.args.temperature soft_student F.log_softmax(student_logits / T, dim-1) soft_teacher F.softmax(teacher_logits / T, dim-1) soft_loss F.kl_div(soft_student, soft_teacher, reductionbatchmean) * (T * T) loss self.args.alpha * hard_loss (1 - self.args.alpha) * soft_loss # 记录分量方便排查 self.log({hard_loss: hard_loss.item(), soft_loss: soft_loss.item()}) return (loss, outputs) if return_outputs else loss这里teacher_logits是在data_collator里用教师模型预计算好存进 batch 的这样训练时不用同时加载两个模型省显存。reductionbatchmean是 KL 散度的正确用法用mean会除以 batch 里所有元素个数数值偏小。T * T的缩放前面解释过是为了梯度对齐。3.4 导出 ONNX 做推理验证训练完拿到学生模型下一步是导出 ONNX 看推理延迟。scripts/export_onnx.py里用的是torch.onnx.exportimport torch from transformers import AutoModelForSequenceClassification, AutoTokenizer model_path outputs/distil_tinybert model AutoModelForSequenceClassification.from_pretrained(model_path) tokenizer AutoTokenizer.from_pretrained(model_path) model.eval() dummy_input tokenizer(测试文本, return_tensorspt, paddingmax_length, max_length128) input_names [input_ids, attention_mask] output_names [logits] torch.onnx.export( model, (dummy_input[input_ids], dummy_input[attention_mask]), outputs/distil_tinybert.onnx, input_namesinput_names, output_namesoutput_names, dynamic_axes{ input_ids: {0: batch, 1: seq_len}, attention_mask: {0: batch, 1: seq_len}, logits: {0: batch} }, opset_version13 )dynamic_axes是关键不设的话导出的 ONNX 只支持固定 batch 和固定长度线上请求一变就报错。opset_version13是兼容性比较好的选择太低不支持某些算子太高有些推理引擎不认。导出后用onnxruntime跑一下对比 PyTorch 和 ONNX 的输出差异一般atol1e-4以内算正常。4. 避坑与常见问题排查4.1 教师模型和学生模型 tokenizer 不一致现象训练 loss 正常下降但验证集准确率始终比教师低 20 个点以上学生像没学到东西。原因教师用的是bert-base-chinese的 tokenizer学生配置里写的是prajjwal1/bert-tiny的 tokenizer两者词表不同同一个文本 tokenize 出来的 id 序列不一样软标签对不上。解决在student_tinybert.yaml里显式指定tokenizer_name: bert-base-chinese让学生复用教师的 tokenizer。模型结构可以不同但 tokenizer 必须一致。4.2 软标签 loss 出现 NaN现象训练几十步后 loss 突然变成 NaN梯度爆炸。原因temperature设得太小比如 1.0时教师 softmax 输出接近 one-hotlog_softmax里出现log(0)KL 散度算出inf。解决把temperature调到 3 以上或者在soft_teacher上加一个极小值1e-8做平滑。我一般直接设T4同时加soft_teacher soft_teacher.clamp(min1e-8)。4.3 显存不够导致 OOM现象CUDA out of memorybatch size 降到 8 还是报。原因data_collator里预计算教师 logits 时如果一次性把整个数据集过一遍教师模型显存会爆。解决改成在__getitem__里按需计算或者用torch.no_grad()包住教师前向并且把教师模型放到 CPU 上算完再搬到 GPU。我一般会在 collator 里加with torch.no_grad(): teacher_logits teacher_model(**batch).logits.cpu()用的时候再.to(device)。4.4 导出 ONNX 后推理结果和 PyTorch 不一致现象ONNX 输出的 logits 和 PyTorch 差很多分类结果都变了。原因导出时模型没切eval()dropout 还在起作用或者attention_mask没传对padding 位置参与了 attention。解决导出前务必model.eval()并且dummy_input里要包含attention_mask。如果还是不对检查opset_version有些版本对LayerNorm的处理有差异换成 13 或 14 试试。4.5 训练完学生模型比直接微调还差现象蒸馏跑了 10 个 epoch准确率还不如学生模型直接拿硬标签微调 3 个 epoch。原因alpha设得太小比如 0.1学生过度依赖软标签而教师在某些样本上本身就不准把错误知识传给了学生。解决把alpha调到 0.5 左右让硬标签和软标签权重平衡。另外检查教师模型在验证集上的准确率如果教师本身只有 80%蒸馏上限就在 80% 附近别指望学生超过教师。5. 进阶技巧用中间层特征蒸馏把 TinyBERT 再压一压软标签蒸馏只用了教师最后一层的输出但教师中间层的隐状态其实也包含有用信息。mattevans-distil里预留了feature_distill的开关打开之后会额外加一项 MSE loss让学生中间层的 hidden state 去拟合教师的对应层。这个做法在 TinyBERT 原论文里叫“中间层蒸馏”对小型学生模型提升明显。具体操作是在student_tinybert.yaml里加use_feature_distill: true feature_layers: [3, 6, 9] # 学生第3/6/9层去拟合教师第4/8/12层 feature_loss_weight: 0.1然后在distill_trainer.py的compute_loss里加一段if self.args.use_feature_distill: # student_hidden: [batch, seq_len, hidden] # teacher_hidden: [batch, seq_len, hidden_teacher] # 先用线性层把学生 hidden 映射到教师维度 student_hidden outputs.hidden_states teacher_hidden inputs.pop(teacher_hidden_states) feat_loss 0.0 for s_layer, t_layer in zip(self.args.feature_layers, [4, 8, 12]): s_h self.feature_proj[s_layer](student_hidden[s_layer]) t_h teacher_hidden[t_layer].detach() feat_loss F.mse_loss(s_h, t_h) loss self.args.feature_loss_weight * feat_lossfeature_proj是一组线性层在__init__里初始化把学生 hidden 维度映射到教师维度。detach()是必须的否则梯度会传回教师模型把教师也更新了。feature_loss_weight我一般设 0.1太大反而会干扰软标签学习。验证方法很简单跑完蒸馏后在验证集上对比三个模型的准确率和推理延迟。我拿一个中文情感分类任务实测过教师 BERT-base 准确率 92.3%延迟 45ms学生 TinyBERT 软标签蒸馏后 89.1%延迟 8ms加上中间层蒸馏后 90.5%延迟不变。多出来的 1.4 个点在业务上可能就是能不能上线的差别。从那以后我每次做蒸馏都会先跑一版纯软标签的 baseline再开中间层蒸馏对比。如果提升不到 0.5 个点说明学生容量已经到顶了再压也没意义不如去优化数据质量。这个习惯帮我省了不少无效调参的时间。希望帮到你。本文还有配套的精品资源点击获取
返回列表