《动手学深度学习》(D2L)实战:用预训练 BERT 微调自然语言推理(SNLI)的完整流程
发布时间:2026/10/1 2:39:18来源:尧图网络
文档教程人工智能深度学习NLP计算机视觉强化学习【免费下载链接】d2l-enInteractive deep learning book with multi-framework code, math, and discussions. Adopted at 500 universities from 70 countries including Stanford, MIT, Harvard, and Cambridge.项目地址https://gitcode.com/gh_mirrors/d2/d2l-en点击查看免费下载本文基于 d2l-en 开源仓库Interactive deep learning book with multi-framework code中 natural-language-inference-bert.md 一节系统讲解如何将预训练的小型 BERTbert.small在 SNLI 自然语言推理数据集上进行微调从下载并加载预训练权重、把前提-假设文本对编码为单条 BERT 输入序列到仅新增一个两层 MLP 分类头完成蕴含entailment、矛盾contradiction、中立neutral三分类。读完本文你将掌握 BERT 微调在序列级文本对分类任务上的完整代码实现、数据打包与截断细节、参数冻结/更新的实际机制并能直接在 PyTorch 与 MXNet 双框架下复现训练。为什么自然语言推理适合用 BERT 微调在本仓库的 natural-language-inference-attention.md 一节中我们已经基于注意力机制为 SNLI 数据集设计过专用的 NLI 架构而 natural-language-inference-and-dataset.md 则介绍了 SNLI 数据集的读取方式。这两条路线代表着两种不同的范式为每个下游任务手工设计专用模型工作量大、难以泛化到其他任务微调预训练模型如 BERT正如 finetuning-bert.md 所指出的BERT 的优势在于最小架构改动——对不同的下游任务只需要在 BERT 之上追加额外的全连接层即可。自然语言推理NLI本质上是**序列级文本对分类sequence-level text pair classification**问题输入是前提premise 假设hypothesis这一对文本序列输出是三分类标签。BERT 的输入表示天然支持单文本与文本对两种形态因此微调 BERT 只需要在 bert-pretraining.md 训练得到的编码器之上追加一个基于 MLP 的小型分类头整体架构如下本节接下来的实操会围绕这张图中的链路展开下载预训练 BERT → 把 SNLI 文本对打包为 BERT 输入序列 → 用 BERTClassifier 微调 → 在测试集上评估。第一步加载预训练 BERT两种预训练模型bert.base 与 bert.small在 bert-pretraining.md 中我们讨论过原始 BERT 模型参数规模高达数亿该节说明原版 BERT 有 110M/340M 两种规模。为兼顾贴近真实与便于演示两种需求仓库通过d2l.DATA_HUB提供了两个版本的预训练权重见 natural-language-inference-bert.md模型标识规模定位MXNet 校验值PyTorch 校验值bert.base接近原版 BERT Base微调需较多算力7b3820b35da691042e5d34c0971ac3edbd80d3f4225d66f04cae318b841a13d32af3acc165f253acbert.small小型版本便于演示a4e718a47137ccd1809c9107ab4f5edd317bae2cc72329e68a732bef0452e4b96a1c341c8910f81f每个预训练包内包含两个文件vocab.json定义词表集合用于重建d2l.Vocabpretrained.params预训练好的模型参数。加载函数 load_pretrained_model下面的load_pretrained_model函数负责完成下载解压 → 重建词表 → 构建 BERT 模型 → 灌入预训练参数四步PyTorch 版本def load_pretrained_model(pretrained_model, num_hiddens, ffn_num_hiddens, num_heads, num_blks, dropout, max_len, devices): data_dir d2l.download_extract(pretrained_model) # Define an empty vocabulary to load the predefined vocabulary vocab d2l.Vocab() vocab.idx_to_token json.load(open(os.path.join(data_dir, vocab.json))) vocab.token_to_idx {token: idx for idx, token in enumerate( vocab.idx_to_token)} bert d2l.BERTModel( len(vocab), num_hiddens, ffn_num_hiddensffn_num_hiddens, num_heads4, num_blks2, dropout0.2, max_lenmax_len) # Load pretrained BERT parameters bert.load_state_dict(torch.load(os.path.join(data_dir, pretrained.params))) return bert, vocabMXNet 版本逻辑一致区别仅在于使用bert.load_parameters(..., ctxdevices)将参数加载到指定设备上下文。实现要点词表重建先创建空的d2l.Vocab()再直接从vocab.json读取idx_to_token列表并据此反推token_to_idx字典——这样能保证微调阶段使用的词表与预训练时完全一致结构一致性BERTModel(len(vocab), num_hiddens, ...)的构造参数必须与预训练时完全相同否则load_state_dict/load_parameters会因形状不匹配而失败。从仓库源码 d2l/torch.py 可以看到BERTModel的实际结构class BERTModel(nn.Module): def __init__(self, vocab_size, num_hiddens, ffn_num_hiddens, num_heads, num_blks, dropout, max_len1000): super(BERTModel, self).__init__() self.encoder BERTEncoder(vocab_size, num_hiddens, ffn_num_hiddens, num_heads, num_blks, dropout, max_lenmax_len) self.hidden nn.Sequential(nn.LazyLinear(num_hiddens), nn.Tanh()) self.mlm MaskLM(vocab_size, num_hiddens) self.nsp NextSentencePred()其中encoder是核心的BERTEncoderd2l/torch.py内含 token 嵌入、segment 嵌入、可学习的位置嵌入参数pos_embedding形状(1, max_len, num_hiddens)以及若干层TransformerEncoderBlockhidden是接在编码器之上的全连接 Tanh 层用于汇聚[cls]位置的表示mlm掩码语言模型与nsp下一句预测是预训练阶段专用的两个头部微调下游任务时它们将被弃用见后文过期梯度部分。加载 bert.small为让绝大多数机器都能跑通演示本节加载小型版本devices d2l.try_all_gpus() bert, vocab load_pretrained_model( bert.small, num_hiddens256, ffn_num_hiddens512, num_heads4, num_blks2, dropout0.1, max_len512, devicesdevices)关键超参一览bert.small配置参数含义bert.small 取值bert.base练习建议num_hiddens隐藏层维度词嵌入维度256768ffn_num_hiddens前馈网络隐藏维度5123072num_heads多头注意力头数412num_blksTransformer 编码器块数212dropout丢弃率0.1—max_lenBERT 输入序列最大长度512此处为演示取 512512第二步构造用于微调的数据集 SNLIBERTDataset文本对如何打包成单条 BERT 输入对自然语言推理而言每条样本包含前提 假设一对文本。SNLIBERTDataset 的核心思想是把这一对文本打包成单条 BERT 输入序列——按照 BERT 的输入表示规范见 bert.md 中的输入表示小节序列以cls开头、前提与假设之间及序列末尾放置sep并用segment ID区分两段文本0 表示段 A1 表示段 B这一打包逻辑在仓库源码 d2l/torch.py 的get_tokens_and_segments中实现def get_tokens_and_segments(tokens_a, tokens_bNone): tokens [cls] tokens_a [sep] # 0 and 1 are marking segment A and B, respectively segments [0] * (len(tokens_a) 2) if tokens_b is not None: tokens tokens_b [sep] segments [1] * (len(tokens_b) 1) return tokens, segments截断与填充策略BERT 输入序列有固定最大长度max_len。当前提 假设拼接后超长时SNLIBERTDataset 采用逐个剔除较长序列末尾 token的策略直到满足长度约束def _truncate_pair_of_tokens(self, p_tokens, h_tokens): # Reserve slots for CLS, SEP, and SEP tokens for the BERT # input while len(p_tokens) len(h_tokens) self.max_len - 3: if len(p_tokens) len(h_tokens): p_tokens.pop() else: h_tokens.pop()注意这里的max_len - 3需要为cls、两个sep共 3 个特殊 token 预留位置因此有效文本长度上限是max_len - 3。截断完成后再用get_tokens_and_segments打包并用pad补齐到max_len同时记录valid_len真实 token 数供 Transformer 编码器做注意力掩码def _mp_worker(self, premise_hypothesis_tokens): p_tokens, h_tokens premise_hypothesis_tokens self._truncate_pair_of_tokens(p_tokens, h_tokens) tokens, segments d2l.get_tokens_and_segments(p_tokens, h_tokens) token_ids self.vocab[tokens] [self.vocab[pad]] \ * (self.max_len - len(tokens)) segments segments [0] * (self.max_len - len(segments)) valid_len len(tokens) return token_ids, segments, valid_len4 进程并行预处理SNLI 训练集规模庞大约 55 万条为加速样本生成_preprocess方法通过multiprocessing.Pool(4)启动4 个工作进程并行处理所有前提-假设token 对MXNet 与 PyTorch 实现一致def _preprocess(self, all_premise_hypothesis_tokens): pool multiprocessing.Pool(4) # Use 4 worker processes out pool.map(self._mp_worker, all_premise_hypothesis_tokens) all_token_ids [token_ids for token_ids, segments, valid_len in out] all_segments [segments for token_ids, segments, valid_len in out] valid_lens [valid_len for token_ids, segments, valid_len in out] return (torch.tensor(all_token_ids, dtypetorch.long), torch.tensor(all_segments, dtypetorch.long), torch.tensor(valid_lens))预处理结果一次性保存在三个张量中all_token_ids、all_segments、valid_lens之后__getitem__只做索引切片返回训练/测试读取阶段的开销极小。SNLI 数据读取与 DataLoader 装配数据读取依赖仓库 d2l/torch.py 中的read_snli它跳过表头过滤掉标签不在{entailment, contradiction, neutral}中的行清洗括号与多余空白并把三类标签映射为 0/1/2。随后实例化数据集并用DataLoader组批# Reduce batch_size if there is an out of memory error. In the original BERT # model, max_len 512 batch_size, max_len, num_workers 512, 128, d2l.get_dataloader_workers() data_dir d2l.download_extract(SNLI) train_set SNLIBERTDataset(d2l.read_snli(data_dir, True), max_len, vocab) test_set SNLIBERTDataset(d2l.read_snli(data_dir, False), max_len, vocab) train_iter torch.utils.data.DataLoader(train_set, batch_size, shuffleTrue, num_workersnum_workers) test_iter torch.utils.data.DataLoader(test_set, batch_size, num_workersnum_workers)实操注意点内存batch_size512是演示取值若出现显存不足OOM应调小max_len此处为节省算力取 128原版 BERT 默认max_len512BERTEncoder源码中max_len1000只是位置嵌入的创建上限更长的上下文通常带来更好的效果测试集只做组批、不做 shuffle保证评估结果可复现。第三步微调 BERT——只需一个两层 MLPBERTClassifier 结构微调的核心是BERTClassifier它复用预训练 BERT 的encoder与hidden仅新增一个输出维度为 3 的全连接层output。前向时把[cls]位置索引 0的 BERT 表示取出依次经过hidden与output得到蕴含/矛盾/中立三个类别的 logitsPyTorch 版本class BERTClassifier(nn.Module): def __init__(self, bert): super(BERTClassifier, self).__init__() self.encoder bert.encoder self.hidden bert.hidden self.output nn.LazyLinear(3) def forward(self, inputs): tokens_X, segments_X, valid_lens_x inputs encoded_X self.encoder(tokens_X, segments_X, valid_lens_x) return self.output(self.hidden(encoded_X[:, 0, :]))因为[cls]位置聚合了整条输入序列前提 假设的信息用它作为文本对整体语义的向量表示是 BERT 序列级任务的通用做法同一做法也用于 finetuning-bert.md 中的单文本分类、语义相似度回归等任务。MXNet 版本将nn.LazyLinear(3)换成nn.Dense(3)并调用net.output.initialize(ctxdevices)初始化新增层。参数的三种命运从头学、微调、冻结把预训练模型装入分类器net BERTClassifier(bert)至此net中参数的更新策略分为三类net.output新增输出层从头随机初始化并学习net.encoder、net.hidden预训练部分随下游任务继续微调fine-tunenet.mlm、net.nsp内部的 MLP 参数它们是预训练阶段计算 MLM 损失与 NSP 损失时使用的参数随bert一并被带进了net但与下游任务无关微调时不更新——即处于过期stale状态。这正是最小架构改动的具体体现除了新增分类层BERT 其余部分无需任何结构变更。过期梯度与 ignore_stale_gradTrue由于MaskLM、NextSentencePred的参数在前向中不再参与损失计算没有梯度直接调用常规step()会因参数缺少梯度而报错。仓库为此在 d2l/mxnet.py 的train_batch_ch13中显式传入了ignore_stale_gradTrue# The True flag allows parameters with stale gradients, which is useful # later (e.g., in fine-tuning BERT) trainer.step(labels.shape[0], ignore_stale_gradTrue)该标志允许带过期梯度的参数被跳过更新从而让微调流程无缝复用 image-augmentation 章节 定义的多 GPU 训练工具train_ch13。PyTorch 中Adam优化器本身不会触碰未参与本次反向传播的参数因此无需额外配置。训练配置与多 GPU 训练lr, num_epochs 1e-4, 5 trainer torch.optim.Adam(net.parameters(), lrlr) loss nn.CrossEntropyLoss(reductionnone) net(next(iter(train_iter))[0]) d2l.train_ch13(net, train_iter, test_iter, loss, trainer, num_epochs, devices)学习率1e-4微调 BERT 时通常采用远小于从头训练的学习率避免破坏预训练学到的表示轮数5 轮受限于算力的演示设置增加轮数可进一步提升精度损失交叉熵三分类net(next(iter(train_iter))[0])PyTorch 下用真实输入触发一次前向完成LazyLinear等惰性层的形状推断train_ch13d2l/torch.py内部会把网络包装为nn.DataParallel分发到devicesd2l.try_all_gpus()自动检测可用 GPU并在每轮结束后用evaluate_accuracy_gpu在test_iter上评估测试精度MXNet 侧通过d2l.split_batch_multi_inputsd2l/mxnet.py把三元组特征切分到多个设备。运行日志会实时输出训练损失、训练精度与测试精度格式形如loss x.xxx, train acc x.xxx, test acc x.xxx。受限于演示算力与bert.small规模最终精度仍有提升空间——这正是练习环节的意义。进一步改进两个练习的思考方向原文档在练习中给出了两个值得深挖的方向这里结合仓库代码补充解析练习 1切换到 bert.base 大幅提升精度将load_pretrained_model的参数替换为bert.base并把num_hiddens、ffn_num_hiddens、num_heads、num_blks分别提升到 768、3072、12、12即接近原版 BERT Base 的规模。同时增加微调轮数并适当调参目标是把测试精度做到 0.86 以上。需要注意模型规模扩大意味着显存与训练时间显著增长若出现 OOM 应先减小batch_size。练习 2按长度比例截断 vs 当前逐端剔除当前_truncate_pair_of_tokens的实现是谁长删谁的贪心策略循环中比较len(p_tokens)与len(h_tokens)总是删除较长一端的末尾 token。另一种思路是按两段文本原始长度的比例分配可用的max_len - 3预算尽量保留两侧的信息比例关系。前者实现简单、计算开销小但可能在长文本上损失较多语义后者更公平地保留两段信息但实现略复杂且对前提长而假设短等极端分布同样难以完美处理。总结围绕 SNLI 自然语言推理任务微调 BERT 的完整链路可以浓缩为三条要点数据打包前提与假设通过get_tokens_and_segments打包为以cls开头、以sep分隔/结尾的单条 BERT 输入序列segment ID 区分两段超长时按删较长端策略截断至max_len - 3不足则pad填充最小架构改动在预训练 BERT 的encoder与hidden之上仅新增一个输出维度为 3 的全连接层BERTClassifier用[cls]位置的表示完成三分类参数更新边界新增输出层从头学习预训练编码器与hidden层继续微调而仅服务于预训练损失MLM/NSP的mlm、nsp参数保持过期不更新——MXNet 侧依靠train_batch_ch13的ignore_stale_gradTrue实现这一机制。这套预训练 微调范式不仅适用于自然语言推理稍作调整更换输出层维度与输入打包方式即可迁移到 finetuning-bert.md 中介绍的情感分析、语言可接受性判断、语义文本相似度乃至 token 级标注与问答等任务。赞分享文档教程人工智能深度学习NLP计算机视觉强化学习【免费下载链接】d2l-enInteractive deep learning book with multi-framework code, math, and discussions. Adopted at 500 universities from 70 countries including Stanford, MIT, Harvard, and Cambridge.项目地址https://gitcode.com/gh_mirrors/d2/d2l-en点击查看免费下载相关推荐OneFlow深度学习框架在自然语言处理中的终极指南BERT预训练与微调完整教程OneFlow深度学习框架在自然语言处理中的终极指南BERT预训练与微调完整教程 OneFlow是一个专为深度学习设计的 用户友好、可扩展且高效 的框架在自深度学习分布式训练模型优化基于 PyTorch 与 SNLI 语料库训练自然语言推理模型SNLIClassifier 全流程实战基于 PyTorch 与 SNLI 语料库训练自然语言推理模型SNLIClassifier 全流程实战 导读 本文围绕当前仓库 legacy/snli htt示例工程人工智能深度学习动手学深度学习迁移学习与微调实战——在 ImageNet 预训练 ResNet-18 上微调热狗识别模型动手学深度学习迁移学习与微调实战——在 ImageNet 预训练 ResNet 18 上微调热狗识别模型 迁移学习transfer learning是深度人工智能深度学习机器学习教程上一篇JavaScript 调用第三方 API 实战基于 Fetch 与 Promise 的异步数据获取指南curriculum 开源课程深度解析下一篇释放存储空间AntiDupl.NET 图片去重完整上手指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网