ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

BERT微调实战:从零复现提取式摘要模型全流程

BERT微调实战:从零复现提取式摘要模型全流程 简介面向自然语言处理开发者与学术研究者这一项目完整实现了基于BERT的抽取式文本摘要微调流程从数据预处理、模型搭建到训练评估均有对应实现可复现论文中的摘要提取实验。压缩包共36个文件以20个Python脚本为核心覆盖分词预处理、分布式训练与模型构建同时包含7个txt数据映射文件、2个json配置、Markdown说明与license声明整包大小约14.99MB。截至目前已有1522人学习下载适合具备一定深度学习基础并希望深入理解BERT微调细节的读者。项目目录划分清晰src与models等模块分别管理训练逻辑与网络结构并附示例数据便于对照论文逐行研读。通过动手运行可以掌握Token IDs、Segment IDs、Mask IDs的构造方法Seq2Seq结构适配、交叉熵损失优化以及ROUGE指标评估等实用技能为后续开展摘要生成类研究提供可直接改造的代码基线。 说个可能得罪人的观点现在讨论大模型微调的人很多但半数以上没亲手跑通过一次完整的预训练模型微调流程。我自己的经验是想真正理解微调与其一上来就啃全量微调、LoRA、llama-factory这些大词不如先拿经典任务练手。我最近完整复现了BERTSUM这篇论文的代码用Python微调BERT做提取式摘要从数据预处理、模型改造到训练评估全部捋了一遍收获非常大。这篇博文就把我的实操路线、关键代码、参数设置和踩过的坑一次说清楚。这套内容适合三类人正在跑论文代码但总卡住的初学者、想系统理解“预训练微调”范式的开发者、以及之后打算接触LoRA、全量微调等进阶玩法但想先把地基打牢的人。别看BERT是2018年的模型它的微调范式和现在的大模型训练本质上一模一样只是规模不同而已。1. 项目定位与整体思路拆解1.1 文本摘要的两条路线抽取式和生成式做摘要任务前一定要先分清两条技术路线。生成式摘要是由模型自己组织语言重新写一段全新的文字而提取式摘要Extractive Summarization则是从原文中挑出若干“重要的句子”拼成摘要。这篇论文走的是提取式路线所以整个问题被转化成了一个句子级别的二分类问题每个句子要不要被保留进摘要。这两条路线各有各的适用场景。生成式更灵活、摘要读起来更像人写的但模型可能自己编造原文没有的信息也就是俗称的“幻觉”提取式的优势是内容完全忠实于原作绝对不出幺蛾子代价是摘要受限于原句的表达方式。对于新闻、论文这类信息密度高、要求事实准确的场景提取式至今仍然有大量应用价值。从入门学习的角度看提取式摘要更容易复现和验证。原因很简单它的训练标签是清晰的0和1效果好坏一眼能看出来也不用去处理解码、集束搜索那一套生成式模型的复杂链路。所以我强烈建议第一次跑摘要任务的读者从提取式入手。1.2 为什么选BERT来做这个任务在BERT之前提取式摘要的主流做法基本是两类一类是基于图的无监督算法比如TextRank一类是基于RNN或CNN的序列标注模型。TextRank不用训练但效果上限低RNN类模型虽然能做序列建模但上下文建模能力有限一个句子里的关键信息关联度往往抓得不够准。BERT的出现改变了这个局面。它通过在大规模语料上预训练学习到了通用的语义表示再用微调的方式适配下游任务。用在摘要上就是利用BERT强大的上下文表示能力把每个句子编码成语义向量再由一个简单的分类层判断句子是否重要。这套“预训练微调”的思路现在的大模型也完全在用所以说BERT微调是理解整个体系的最佳入门样例。1.3 BERTSUM的输入改造让BERT能同时处理多个句子标准BERT的输入格式是单个文本片段[CLS] tokens [SEP]。但提取式摘要要让模型同时“看”完整篇文档的所有句子还要知道每个句子的边界在哪。论文里对输入做了改造把文档中每个句子前面都加上一个[CLS]标记句子之间用[SEP]分隔整体输入变成这样[CLS] sentence1 [SEP] [CLS] sentence2 [SEP] [CLS] sentence3 [SEP] ...这样做的妙处在于每个[CLS]位置经过BERT编码后的hidden state就可以当作对应句子的向量表示接一个分类器直接打分结构非常干净。要注意的是因为每个句子前多了一个[CLS]、句子间多了一个[SEP]512个token的序列能装下的正文内容就变少了实际操作时经常需要对长文档做截断一般最多容纳5到6个句子。另外还有一个容易被忽略的细节segment ids也做了处理。BERT原本的segment id只区分两段文本0和1论文采用了一种interval交替的方式句子1用0、句子2用1、句子3再用0这样通过segment信息也能辅助模型感知句子边界。我之前自己写代码时直接全部填成0结果模型效果掉了一截后来补上才恢复正常。2. 环境搭建与数据准备先把最小可行版本跑通2.1 环境依赖与版本选择复现这套代码不需要太苛刻的环境但版本匹配确实是个坑。我自己用的是Python 3.9加PyTorch 2.0的组合transformers库用4.x版本整体跑得很稳。如果你用Python 3.11以上的环境部分旧版本库会编译失败建议用conda单独建一个环境省得污染主环境。下面是我的完整依赖清单可以直接保存为requirements.txttorch2.0.0 transformers4.30.0 nltk3.8 rouge-score0.1.2 tqdm4.65.0 numpy1.24.0 datasets2.12.0安装命令很简单conda create -n bertsum python3.9 conda activate bertsum pip install -r requirements.txt这里有个小建议rouge-score这个库是Google维护的ROUGE评估实现比老的pyrouge好装太多pyrouge那个工具在Linux上经常需要配置perl环境非常折磨人直接用rouge-score省心得多。2.2 数据格式与训练标签的构造论文原版使用的是CNN/DailyMail数据集规模有28万多篇新闻全量训练对普通个人电脑来说不太现实。我建议第一阶段先用验证集的一个小子集或者随便找一个几百篇的小型新闻数据跑通流程。数据格式只要做好两种字段就行article是原文highlights是参考摘要。关键问题是怎么构造训练标签。提取式摘要的监督信号不是现成的需要我们自己从原文和参考摘要的对应关系里“算”出来。常见的做法是先把原文按句子切分然后计算每个句子与参考摘要之间的ROUGE-L分数超过某个阈值比如0.4就把这个句子的标签设为1否则设为0。句子切分可以直接用nltk工具英文文本的切分效果比较稳定import nltk nltk.download(punkt) from nltk.tokenize import sent_tokenize def build_labels(article, summary, threshold0.4): sentences sent_tokenize(article) labels [] for sent in sentences: score rouge_l_score(sent, summary) # 计算句子与摘要的ROUGE-L F1 labels.append(1 if score threshold else 0) return sentences, labels这个阈值不需要太较真实践中0.3到0.5之间效果差别不大。重点是要想明白模型学习的是“哪些句子和摘要内容最接近”而不是“哪些句子本身写得最漂亮”。2.3 数据加载器的要点一篇文档就是一个样本数据加载的细节比想象中复杂。在标准分类任务里一条样本是一句话但在这里一条样本是一整篇文档文档里包含多个句子每个句子对应一个0/1标签。所以Dataset类的设计逻辑要理清先按文档切分再把文档内的句子分别tokenize最后组装成BERTSUM需要的输入格式。我简化后的Dataset类大概是这样的逻辑class ExtractiveSummarizationDataset(torch.utils.data.Dataset): def __init__(self, articles, summaries, max_len512): self.data [] for article, summary in zip(articles, summaries): sentences, labels build_labels(article, summary) tokens [] segment_ids [] label_list [] for i, sent in enumerate(sentences): sent_tokens tokenizer.tokenize(sent)[:80] # 限制每句长度 tokens.append([CLS]) tokens.extend(sent_tokens) tokens.append([SEP]) segment_ids.extend([i % 2] * (len(sent_tokens) 2)) # interval交替 label_list.append(labels[i]) # 截断到max_len self.data.append((tokens, segment_ids, label_list)) def __len__(self): return len(self.data)需要注意这里的segment_ids我是按句子索引的奇偶来做interval交替和标准BERT里token_type_ids填0/1的含义不一样但输入方式是一样的。tokenize之后记得转成input_ids再用tokenizer.build_inputs_with_special_tokens之类的工具补上attention_mask。3. 微调核心实现模型改造、训练循环与参数调优3.1 模型结构BERT加一个句子分类头模型结构本身不复杂加载一个bert-base-uncased取每个句子的[CLS]向量过一个MLP分类头输出分数。论文里用了两层全连接加ReLU和Dropout输出维度是1然后用sigmoid映射到0到1之间。下面是我改造后的模型核心代码import torch import torch.nn as nn from transformers import BertModel class BertSumExtractor(nn.Module): def __init__(self, bert_pretrainedbert-base-uncased, hidden_size768): super().__init__() self.bert BertModel.from_pretrained(bert_pretrained) self.classifier nn.Sequential( nn.Linear(hidden_size, hidden_size), nn.ReLU(), nn.Dropout(0.1), nn.Linear(hidden_size, 1) ) def forward(self, input_ids, attention_mask, segment_ids): outputs self.bert( input_idsinput_ids, attention_maskattention_mask, token_type_idssegment_ids ) sequence_output outputs.last_hidden_state # [batch, seq_len, hidden] cls_positions (input_ids tokenizer.cls_token_id).int() # 简化处理这里按固定间隔提取每个[CLS]向量 sentence_vectors sequence_output[cls_positions.bool(), :] logits self.classifier(sentence_vectors).squeeze(-1) return logits实际工程里实现“按位置提取多个[CLS]向量”要用到masked_select之类的操作比上面的伪代码复杂一些。核心逻辑就是一个序列里有N个[CLS]我们就取N个向量每个向量过一个共享权重的分类头得到N个分数。3.2 损失函数与训练循环训练目标就是句子级别的二分类损失函数直接上二元交叉熵。PyTorch里的BCEWithLogitsLoss更稳因为它在内部做了sigmoid和数值稳定处理比手动sigmoid再算BCE更好。训练循环我建议写一个标准的模板方便后面进一步改成LoRA或者全量微调from torch.cuda.amp import autocast, GradScaler model BertSumExtractor().cuda() optimizer torch.optim.AdamW(model.parameters(), lr2e-5) criterion nn.BCEWithLogitsLoss() scaler GradScaler() for epoch in range(3): for batch in train_dataloader: input_ids batch[input_ids].cuda() attention_mask batch[attention_mask].cuda() segment_ids batch[segment_ids].cuda() labels batch[labels].cuda() with autocast(): logits model(input_ids, attention_mask, segment_ids) loss criterion(logits, labels.float()) optimizer.zero_grad() scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()上面的代码里把损失、梯度裁剪和混合精度都放一起了直接跑没问题。特别提醒一句labels必须转成float再丢进BCEWithLogitsLoss我一开始用LongTensor报错半天没看出来属于特别基础但容易忽略的问题。3.3 关键训练参数学习率、epoch与warmupBERT微调有一组被反复验证过的“黄金参数”直接套用基本不会出大问题。我在复现过程中用的参数如下参数推荐值说明学习率2e-5BERT全量微调默认值过大容易毁掉预训练权重训练轮数epoch3-5小数据时可以多加到5轮但注意观察过拟合Batch Size8-32按显存来不够就用梯度累积梯度累积步数4或8等效放大batch size稳定梯度Warmup步数总步数的10%防止初期loss剧烈震荡梯度裁剪max_grad_norm1.0防止梯度爆炸导致NaN学习率是这里最重要的参数。BERT的预训练权重已经很“成熟”了学习率设太大相当于把学好的参数一脚踢飞loss会乱跳设太小又微调不动。如果你自己改了学习率一定要在验证集上盯ROUGE指标不要只看训练loss。3.4 显存不够怎么办梯度累积与混合精度我跑这篇论文的时候用的是一张8G显存的消费级显卡bert-base-uncased本身就接近400M参数全量微调时稍微加个长文档就很容易爆显存。我的解法是三层组合拳。第一先把batch size降到1。BERTSUM的每个样本本身已经是一整篇文档batch size为1时batch内已经包含多个句子效果损失很小。第二用梯度累积。把4个batch的梯度攒起来再统一更新等效于batch size4。代码上就是把loss.backward()每跑4步才调用一次optimizer.step()。第三开启混合精度。PyTorch 2.0里用torch.cuda.amp很方便显存能降低将近一半而且由于半精度矩阵乘法速度更快训练时间通常也有明显缩短。唯一的坑是混合精度下Apex或者老版本的apex安装麻烦但原生amp基本够用。如果这样还爆显存那就只能冻结BERT底层部分参数了。比如冻结bert.embeddings和前4个Encoder层只微调后面的层和分类头。虽然这相当于放弃了底层通用特征的更新但在资源受限时是性价比很高的折中方案。4. 评估与推理让模型真的输出一段好摘要4.1 ROUGE评估指标怎么算摘要任务最通用的评估指标是ROUGE它衡量的是模型生成的摘要和参考摘要之间n-gram的重合程度。ROUGE-1看的是单个词的重叠ROUGE-2看相邻两个词的重叠ROUGE-L用的是最长公共子序列来捕捉句子级结构相似度。用rouge-score库计算非常简单from rouge_score import rouge_scorer scorer rouge_scorer.RougeScorer([rouge1, rouge2, rougeL], use_stemmerTrue) scores scorer.score(reference_summary, prediction_summary) print(scores[rougeL].fmeasure)use_stemmerTrue会把单词还原成词根比如runs和running都算同一个词这个设置更贴近论文里的评测标准。需要注意ROUGE的分数在不同数据集之间不能横向比较同一个数据集上对比基线才有意义。4.2 推理时的句子选择策略Trigram Blocking与MMR模型训练好之后推理阶段要把得分高的句子组装成摘要。最朴素的做法是直接按预测分数从高到低取前3句但这样做很容易选出两句话讲同一件事的情况摘要读起来非常冗余。BERTSUM论文里用了一个非常经典且好用的技巧叫Trigram Blocking。思路是按分数从高到低逐句遍历在加入新句子之前检查新句子与已选句子有没有连续的三个词trigram是重复的如果有就跳过这个句子。这个策略简单、计算量小但能有效去除重复信息。下面是我实现的简化版本def trigram_blocking(selected_sents, candidate_sent): candidate_trigrams set() words candidate_sent.split() for i in range(len(words) - 2): candidate_trigrams.add( .join(words[i:i 3])) for sent in selected_sents: sent_words sent.split() for i in range(len(sent_words) - 2): if .join(sent_words[i:i 3]) in candidate_trigrams: return True return False除了Trigram Blocking另一种常见的方案是MMR最大边际相关公式是λ * 句子得分 - (1-λ) * max(句子与已选句子的相似度)通过惩罚和已选内容相似的句子来控制冗余。Trigram Blocking对新闻数据已经很够用MMR在句子语义相似度较高但用词不重复的场景下表现更好。4.3 一个完整的推理样例我拿一篇短新闻测试了一下微调后的模型输出效果大概是这样原文有5个句子模型给出的分数分别是0.92、0.31、0.78、0.45、0.65。按分数排序是第1句、第3句、第5句在检查trigram重复后第5句因为有和第1句重复的主题词而被跳过最终选中的摘要就是第1句和第3句。这个结果基本覆盖了新闻的核心信息而且没有明显冗余说明Trigram Blocking确实起到了作用。如果你想让摘要更长或更短可以调整最终选取句子数的上限通常新闻摘要控制在2到4句比较合适。5. 复现过程中的坑与调参避坑实录5.1 显存溢出的排查顺序如果你在训练时遇到CUDA out of memory不要一上来就换更大的显卡先按下面的顺序排查第一步把batch size降到1第二步把max_len从512降到384第三步开启混合精度第四步冻结BERT底层参数。绝大多数显存问题到这四步都能解决。顺便说一句训练时用torch.cuda.empty_cache()清理缓存作用很有限与其频繁清理不如把batch size和max_len控制好。5.2 Loss不降或直接变成NaNLoss一直不降最常见的原因是学习率开太大或者标签和输入没有对齐。BERT微调学习率超过5e-5就很容易出问题建议回退到2e-5再试。Loss变成NaN基本可以断定是梯度爆炸除了降低学习率别忘了在optimizer.step()之前加torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)这是最稳的兜底手段。还有个特别容易踩的坑是正负样本不均衡。如果一篇文档里绝大多数句子标签都是0模型会学会无脑预测0loss虽然不高但完全没用。遇到这种情况可以把训练文档里标签全为0的样本过滤掉一部分或者给正样本的loss加权重。5.3 预训练模型下载慢或加载失败初次运行代码时transformers会自动下载bert-base-uncased权重网络不好的话很容易卡住甚至中断。解决办法是设置镜像环境变量让HuggingFace走国内镜像源export HF_ENDPOINThttps://hf-mirror.com设置好之后重新运行模型参数会自动下载到本地缓存目录。如果你想把模型固定为离线加载也可以先手动下载权重放到本地文件夹然后直接用BertModel.from_pretrained(./bert-base-uncased)加载这种方式在论文复现里更可控。5.4 文本截断导致摘要信息不全BERT的最大输入长度是512个token原文超过这个长度就必须截断如果只保留开头部分很多关键信息在后面就丢了。我测试过一篇长新闻截断后模型选出来的摘要基本只覆盖导语部分细节和背景信息全没了。简单粗暴的解决方法是限制文档只取前6个句子通常新闻导语已经涵盖核心信息更进阶的做法是把一篇长文档切成多个块分别过模型后做冗余筛选再合并结果。如果你确实要处理超长文档建议考虑Longformer或者BigBird这类能处理更长序列的模型BERT的512上限在这里是绕不过去的硬约束。这套流程完整跑下来我最深的感受是微调的本质没有变不管是BERT还是现在的Llama、Qwen都是先预训练再在下游任务上适配变化的主要是模型规模和参数更新方式。你把这个BERT摘要代码弄透了再去看LoRA、全量微调、llama-factory这些工具会发现它们都是在“如何更新参数”这个环节上做优化任务训练循环的基本盘是一样的。如果你在复现这篇论文我的建议是先在小数据集上把整条链路跑通确认效果没大问题再上全量数据这样既省时间也省资源。后面我还想再做一组BART生成式摘要的对比实验到时候拿数据说话。本文还有配套的精品资源点击获取
返回列表