尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

PaddleNLP 知识蒸馏实战:将 BERT 教师模型的任务知识蒸馏进 Bi-LSTM 学生模型

发布时间:2026/9/25 4:32:54

资讯中心
01
ARTICLE

PaddleNLP 知识蒸馏实战:将 BERT 教师模型的任务知识蒸馏进 Bi-LSTM 学生模型

PaddleNLP 知识蒸馏实战:将 BERT 教师模型的任务知识蒸馏进 Bi-LSTM 学生模型
人工智能大模型预训练微调LoRARLHF强化学习分布式训练【免费下载链接】PaddleNLPEasy-to-use and powerful LLM and SLM library with awesome model zoo.项目地址https://gitcode.com/gh_mirrors/pa/PaddleNLP点击查看免费下载本文以 PaddleNLP 中slm/examples/model_compression/distill_lstm/目录下的蒸馏实验为核心完整讲清“把特定任务下 fine-tuned BERT 的知识蒸馏到 Bi-LSTM 小模型”这一经典压缩方案包括教师/学生模型的划分、均方误差蒸馏损失与交叉熵损失的加权方式、基于 Masking 与 n-gram 采样的数据增强细节、三个训练阶段BERT fine-tuning、小模型独立训练、蒸馏的完整可运行命令、各超参数的默认值与含义以及官方给出的蒸馏实验结果。读完本文你可以按步骤在自己的环境里复现整套 BERT → Bi-LSTM 蒸馏流程并能看懂每个参数背后的源码实现。示例总览目录结构与核心文件该实验实现参考论文《Distilling Task-Specific Knowledge from BERT into Simple Neural Networks》Tang et al., 2019。目录内各文件分工如下引自 README. ├── small.py # 小模型结构以及对小模型单独训练的脚本 ├── bert_distill.py # 用教师模型BERT蒸馏学生模型的蒸馏脚本 ├── data.py # 定义了dataloader等数据读取接口 ├── utils.py # 定义了将样本转成id的转换接口 ├── args.py # 参数配置脚本 └── README.md # 文档本文件在模型蒸馏中较大的模型本例中的 BERT称为教师模型较小的模型本例中的 Bi-LSTM称为学生模型。学生模型通过蒸馏相关的损失函数学习教师模型的输出分布数据、词表与训练脚本的组织方式都围绕这一“教师—学生”结构展开。蒸馏损失MSE 为主、交叉熵可加权README 指出本实验的蒸馏损失是均方误差MSE函数输入分别为学生模型输出与教师模型输出。这一描述在 bert_distill.py 中可以得到源码级印证一次训练 step 的完整流程为教师模型前向不参与梯度# Calculate teacher models forward. with paddle.no_grad(): teacher_logits teacher.model(bert_input_ids, bert_segment_ids)学生模型前向QQP 为双句输入其余任务为单句# Calculate student models forward. if args.task_name qqp: logits model(student_input_ids_1, seq_len_1, student_input_ids_2, seq_len_2) else: logits model(student_input_ids, seq_len)联合损失由--alpha控制两个损失的权重平衡loss args.alpha * ce_loss(logits, labels) (1 - args.alpha) * mse_loss(logits, teacher_logits)结合 args.py 可以看到--alpha默认值为0.0即 README 中的三条蒸馏命令均未显式指定 alpha实际训练中损失完全来自 MSE学生模型只拟合教师模型的 logits不使用真实标签——这正是“借助教师模型的暗知识”的核心机制。如需同时利用真实标签把--alpha调大即可在 MSE 与交叉熵之间取得平衡。教师模型本身由BertForSequenceClassification.from_pretrained(teacher_dir)从第一步训练产出的检查点加载并固定处于eval()状态见 bert_distill.py 的TeacherModel类。数据增强Masking 与 n-gram 采样论文为让教师模型表达更多可供学习的信息对训练数据做了增强共提出三种方式Masking、POS-guided word replacement、n-gram sampling。README 特别说明本实验只实现了第 1 种Masking和第 3 种n-gram sampling。两者的具体实现位于 data.pyMasking以概率p_mask默认 0.1将已分词结果中的 token 替换为[MASK]。英文数据上支持--whole_word_mask参数切换为整词掩码模式n-gram sampling以概率p_ng默认 0.25从每条样本中截取一段长度为 n 的连续片段n 的取值范围默认(2, 6)中文版本为(2, 10)见 apply_data_augmentation_for_cn迭代次数每条训练样本额外增强--n_iter次默认 20因此训练集规模会膨胀为原来的约1 n_iter倍。def apply_data_augmentation( data, task_name, tokenizer, n_iter20, p_mask0.1, p_ng0.25, ngram_range(2, 6), whole_word_maskFalse, seed0 ): ... for example in data: for _ in range(n_iter): ... new_data.append({sentence: words, labels: example[labels]})数据增强在create_distill_loader中通过train_ds.map(data_aug_fn, batchedTrue)统一接入data.py并打印 “Data augmentation has been applied.” 提示。中文任务ChnSentiCorp由于 BERT 与 jieba 分词结果不一致增强时会同步产出两套 token 序列lstm_tokensjieba 分词供学生模型与bert_tokensBERT 分词供教师模型保证两个模型看到的是同一句增强文本这也是中文蒸馏能成立的关键设计。增强样本会继承原样本标签因此学生模型实际上是在“远大于原始训练集”的数据上拟合教师输出。数据、预训练模型与词表准备数据集使用 GLUE 中的 SST-2、QQP 以及中文情感分类数据集 ChnSentiCorp 的训练集作为训练语料验证集dev用于评估。数据集通过paddlenlp.datasets.load_dataset自动下载存放于paddlenlp.utils.env.DATA_HOME路径下例如 Linux 系统下 QQP 默认在~/.paddlenlp/datasets/glue/QQPChnSentiCorp 在~/.paddlenlp/datasets/chnsenticorp。源码侧对应 data.py 中的load_dataset(glue, task_name, ...)与load_dataset(task_name, ...)。教师预训练模型bert-base-uncased英文任务、bert-base-chinese与bert-wwm-ext-chinese中文任务同样自动下载到MODEL_HOME例如~/.paddlenlp/models/bert-base-uncased。学生模型词表中文任务的输入使用 jieba 分词词表与 PaddleNLP 文本分类项目使用的 senta 词表一致文件为senta_word_dict.txt可从 PaddleNLP 官方数据源下载默认存放命令见原 README 的wget说明放置在本示例目录下。为了节省显存与运行时间可以先对 ChnSentiCorp 中未出现的词做过滤再把过滤后的词表路径与词表大小分别配置到--vocab_path与--vocab_size参数中。英文任务的词表来源英文任务下学生模型直接复用教师 BERT 的BertTokenizer作为词表见 data.py非中文任务分支vocab BertTokenizer.from_pretrained(model_name)因此英文任务的--vocab_size对应 BERT 词表大小 30522中文任务则由Vocab.load_vocabulary加载 senta 词表unk_token[UNK]、pad_token[PAD]。阶段一训练教师模型BERT fine-tuning蒸馏的前提是先有一个任务特定的强教师模型。GLUE 任务可以复用本仓库 glue 示例目录 下的 run_glue.py更多说明见 glue 的 README。以 SST-2 任务为例原 README 给出的命令为cd ../../benchmark/glue export CUDA_VISIBLE_DEVICES0 export TASK_NAMESST-2 python -u ./run_glue.py \ --model_type bert \ --model_name_or_path bert-base-uncased \ --task_name $TASK_NAME \ --max_seq_length 128 \ --batch_size 128 \ --learning_rate 3e-5 \ --num_train_epochs 3 \ --logging_steps 10 \ --save_steps 10 \ --output_dir ../model_compression/distill_lstm/pretrained_models/$TASK_NAME/ \ --device gpu \注意--output_dir指向蒸馏目录下的pretrained_models/$TASK_NAME/即教师模型产物就为后续的--teacher_dir参数准备好了位置。若训练基于 ChnSentiCorp 的 BERT fine-tuning 模型原 README 建议进入文本分类示例目录 slm/examples/text_classification 下的多分类示例把预训练模型改为bert-base-chinese或bert-wwm-ext-chinese进行 fine-tuning。训练完成后将效果最好的模型保存在本示例的pretrained_models/$TASK_NAME/下模型目录应包含model_config.json、model_state.pdparams、tokenizer_config.json及vocab.txt这几个文件——这与BertForSequenceClassification.from_pretrained(teacher_dir)的加载要求一致。阶段二独立训练 Bi-LSTM 学生模型蒸馏效果的对照基线small.py 定义了 Bi-LSTM 学生模型结构并完成独立训练使用nn.CrossEntropyLoss直接拟合真实标签其作用是作为不蒸馏的对照基线。模型结构要点见 small.pynn.Embedding(vocab_size, embed_dim, padding_idx)学生模型词嵌入层默认embed_dim300nn.LSTM(embed_dim, hidden_size, num_layers, bidirectional)双向 LSTM默认 1 层、hidden_size300单句任务取双向 LSTM 末时刻的两个方向隐状态拼接后经fc2*hidden → hidden与tanh再经output_layer输出 logits双句任务QQP额外拼接两句的“和”与“绝对差”向量输入维度变为8*hidden经fc_1再输出。三个任务的独立训练命令原 README 原样保留CUDA_VISIBLE_DEVICES0 python small.py \ --task_name chnsenticorp \ --max_epoch 20 \ --vocab_size 1256608 \ --batch_size 64 \ --model_name bert-wwm-ext-chinese \ --optimizer adam \ --lr 3e-4 \ --dropout_prob 0.2 \ --vocab_path senta_word_dict.txt \ --save_steps 10000 \ --output_dir small_models/chnsenticorp/CUDA_VISIBLE_DEVICES0 python small.py \ --task_name sst-2 \ --vocab_size 30522 \ --max_epoch 10 \ --batch_size 64 \ --lr 1.0 \ --dropout_prob 0.4 \ --output_dir small_models/SST-2 \ --save_steps 10000 \ --embedding_name w2v.google_news.target.word-word.dim300.enCUDA_VISIBLE_DEVICES0 python small.py \ --task_name qqp \ --vocab_size 30522 \ --max_epoch 35 \ --batch_size 256 \ --lr 2.0 \ --dropout_prob 0.4 \ --output_dir small_models/QQP \ --save_steps 10000 \ --embedding_name w2w.google_news.target.word-word.dim300.en # 注意原README为 w2v.google_news.target.word-word.dim300.en当前版本的重要注意论文中曾使用 Google News 预训练 Word Embedding 初始化小模型的 Embedding 层但当前仓库代码中 BiLSTM.init明确抛出异常——“TokenEmbedding is deprecated in PaddleNLP since 3.0, please set embedding_name to None”。因此按当前仓库运行英文任务时请去掉--embedding_name参数Embedding 层将随机初始化后端到端训练。阶段三执行蒸馏BERT → Bi-LSTM蒸馏阶段加载教师模型并对学生模型训练数据管线会自动施加前述数据增强。三个任务的蒸馏命令原 README 原样保留CUDA_VISIBLE_DEVICES0 python bert_distill.py \ --task_name chnsenticorp \ --vocab_size 1256608 \ --max_epoch 6 \ --lr 1.0 \ --dropout_prob 0.1 \ --batch_size 64 \ --model_name bert-wwm-ext-chinese \ --teacher_dir pretrained_models/chnsenticorp/best_bert_wwm_ext_model_880 \ --vocab_path senta_word_dict.txt \ --output_dir distilled_models/chnsenticorp \ --save_steps 10000 \CUDA_VISIBLE_DEVICES0 python bert_distill.py \ --task_name sst-2 \ --vocab_size 30522 \ --max_epoch 6 \ --lr 1.0 \ --dropout_prob 0.2 \ --batch_size 128 \ --model_name bert-base-uncased \ --output_dir distilled_models/SST-2 \ --teacher_dir pretrained_models/SST-2/best_model_610 \ --save_steps 10000 \ --embedding_name w2v.google_news.target.word-word.dim300.enCUDA_VISIBLE_DEVICES0 python bert_distill.py \ --task_name qqp \ --vocab_size 30522 \ --max_epoch 6 \ --lr 1.0 \ --dropout_prob 0.2 \ --batch_size 256 \ --model_name bert-base-uncased \ --n_iter 10 \ --output_dir distilled_models/QQP \ --teacher_dir pretrained_models/QQP/best_model_17000 \ --save_steps 10000 \ --embedding_name w2v.google_news.target.word-word.dim300.en各命令的要点--teacher_dir指向阶段一保存的最佳 BERT 检查点目录命令中的best_model_610、best_model_17000、best_bert_wwm_ext_model_880为示例命名实际请以自己训练输出的最佳模型目录为准--model_name既用于加载教师的BertTokenizer蒸馏 batch 中 BERT 侧输入即由它生成中文任务下还会作为小模型训练数据的词表来源--n_iter控制每条样本的数据增强次数QQP 命令中设为 10默认 20同样注意按当前仓库代码--embedding_name参数会导致ValueError实际运行英文任务时请将其移除训练与保存逻辑与small.py一致每--save_steps步保存step_N.pdparams与step_N.pdopt每--log_freq步在 dev 集上评估并打印 loss、accQQP 另含 precision/recall/f1。关键参数说明来自 args.py完整的参数定义在 args.py核心参数与默认值如下表参数默认值说明--task_namesst-2任务名支持sst-2/qqp/chnsenticorp--optimizeradadelta优化器仅支持adam|adadeltaAdadelta 使用rho0.95--lr1.0学习率--num_layers1LSTM 层数--emb_dim/--hidden_size300/300学生模型嵌入维度 / LSTM 隐层维度--output_dim2分类数本实验均为二分类--vocab_size10000学生模型词表大小需与--vocab_path词表一致--batch_size/--max_epoch64/12批大小 / 最大训练轮数--max_seq_length128序列最大长度--n_iter20数据增强时每条样本的增强迭代次数--dropout_prob0.0LSTM dropout--init_scale0.1全连接层权重 Uniform 初始化区间[-scale, scale]--log_freq/--save_steps10/100日志频率 / 保存 checkpoint 频率按 step--model_namebert-base-uncased教师模型名其 tokenizer 会被学生数据管线复用--teacher_dir无教师模型目录蒸馏时必填--vocab_pathMODEL_HOME/bert-base-uncased/...学生模型词表路径中文任务用 senta 词表--alpha0.0交叉熵损失与 MSE 蒸馏损失的权重平衡系数--whole_word_maskFalse数据增强时使用整词掩码--init_from_ckptNone从既有 checkpoint 恢复模型与优化器--seed2021随机种子保证参数初始化与数据增强可复现--devicegpu运行设备可选gpu/cpu/xpu--embedding_nameNone预训练词嵌入名当前版本传入非空值会直接报错见上文说明README 提醒训练不同任务时需要调整对应超参数表中各命令里的 epoch、lr、batch_size、dropout 即为官方推荐的组合。蒸馏实验结果官方实验在 GLUE 的 SST-2、QQP 与 ChnSentiCorp 上进行均用各自 dev 集评价指标为准确率QQP 另含 f1。对比结果原 README 表格ModelSST-2 (dev acc)QQP (dev acc/f1)ChnSentiCorp (dev acc)ChnSentiCorp (dev acc)Teacher modelbert-base-uncasedbert-base-uncasedbert-base-chinesebert-wwm-ext-chineseBERT-base0.9300460.905813(acc)/0.873472(f1)0.9516670.955000Bi-LSTM0.8543580.856616(acc)/0.799682(f1)0.9200000.920000Distilled Bi-LSTM0.8876150.875216(acc)/0.831254(f1)0.9325000.934167可以看到用 BERT 教师模型蒸馏后的 Bi-LSTM 相对独立训练的 Bi-LSTM在 SST-2、QQP、ChnSentiCorp 上分别提升约 3.3%、1.9%、1.4%与教师模型的差距明显缩小——这正是“教师暗知识 数据增强”在小模型上的收益体现。小结与参考资料本示例用不到 600 行代码串起了任务知识蒸馏的完整闭环教师模型 fine-tuning复用 run_glue.py→ 学生模型独立训练基线small.py→ 带数据增强的蒸馏训练bert_distill.py、data.py、utils.py。其中损失函数、数据增强与词表设计均可通过args.py参数灵活调整适合作为理解“教师-学生”范式在序列分类任务上如何落地的最小完整样例。参考文献Tang R, Lu Y, Liu L, Mou L, Vechtomova O, Lin J.Distilling Task-Specific Knowledge from BERT into Simple Neural Networks. arXiv preprint arXiv:1903.12136, 2019.赞分享人工智能大模型预训练微调LoRARLHF强化学习分布式训练【免费下载链接】PaddleNLPEasy-to-use and powerful LLM and SLM library with awesome model zoo.项目地址https://gitcode.com/gh_mirrors/pa/PaddleNLP点击查看免费下载相关推荐PaddleNLP 由 BERT 到 Bi-LSTM 的知识蒸馏实战任务特定知识蒸馏完整指南PaddleNLP 由 BERT 到 Bi LSTM 的知识蒸馏实战任务特定知识蒸馏完整指南 导读 本文基于 PaddleNLP 仓库中的蒸馏示例 dist人工智能大模型预训练微调LoRARLHF强化学习分布式训练模型推理服务推理引擎模型量化模型压缩本地部署NLPPaddleNLP 知识蒸馏实战从 BERT 到 Bi-LSTM 的 Task-Specific 蒸馏指南PaddleNLP 知识蒸馏实战从 BERT 到 Bi LSTM 的 Task Specific 蒸馏指南 导读 本文基于 PaddleNLP 仓库中的 di人工智能大模型预训练微调LoRARLHF强化学习分布式训练模型推理服务推理引擎模型量化模型压缩本地部署NLP知识蒸馏实战nlp-tutorial师生模型的知识传递完整指南知识蒸馏实战nlp tutorial师生模型的知识传递完整指南 nlp tutorial是面向深度学习研究者的自然语言处理教程项目通过一系列实践案例帮助开发示例工程上一篇Thumbfast架构解析mpv播放器实时缩略图生成引擎的实现原理与实践指南下一篇sudo-rs的系统集成与systemd和日志服务交互创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

◈

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

◐

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

▲

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。