ARTICLE DETAIL

资讯详情

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

微调BERT实现抽取式摘要:从数据预处理到训练避坑实战

微调BERT实现抽取式摘要:从数据预处理到训练避坑实战 简介面向自然语言处理开发者的BERT微调文本摘要代码包源自BertSum-master项目专注利用预训练BERT完成自动摘要生成。资源共36个文件压缩包14.99MB其中20个Python脚本覆盖数据预处理、模型构建、训练与评估搭配7个txt数据映射、2个JSON配置、说明文档与许可协议目录区分src、models等模块结构清晰。目前已有1523人学习浏览。代码完整演示了将原始文本转换为Token IDs、Segment IDs与Mask IDs经Hugging Face加载预训练BERT叠加编码器-解码器结构使用交叉熵损失和Adam优化器训练并以ROUGE指标测评生成质量。对希望深入理解微调原理、上手NLP摘要实战的开发者这是一份可直接运行的参考实现涵盖抽取式摘要全流程适合作为论文复现与实践进阶的脚手架。1. 一开始提到微调 BERT 做摘要到底在复现什么如果你最近在 GitHub 上检索过摘要生成方向的论文源码大概率见过一类仓库名字类似 “Python-微调BERT用于提取摘要的论文代码”。这类项目往往不是某个大厂官方发布的框架而是论文作者或开源贡献者把论文里的实验脚本整理出来的结果。其背后要解决的问题很具体给定一篇长文档从原文里挑出若干句子按原文顺序拼成一段摘要而不是像 GPT 这类模型那样逐字“写”一段新文本。这个任务在学术界叫抽取式摘要Extractive Summarization实际业务里常见于舆情简报、研报导读、合同要点提取等场景。为什么要微调 BERT 而不是直接拿预训练模型硬上因为 BERT 本身是一个语言理解模型它擅长判断“这段文字在讲什么”但不会自动告诉你“第几句值得进摘要”。你需要在这个理解能力之上加一个分类头让模型对每个句子输出一个 0 到 1 之间的分代表它被选入摘要的概率。微调的过程就是让 BERT 的深层语义表征和这个分类任务对齐。这篇笔记我会从一个可落地的角度拆开整套方案先是模型选型和数据准备再给出一份能跑通微调的 Python 训练脚本最后把我在复现这类论文代码时踩过的坑和验证技巧一并列出来。适合正在复现论文、准备做毕设实验或者想在公司内部快速验证 BERT 摘要效果的开发者抄作业。2. 理解 BERT 做抽取式摘要的任务设计一句话能进摘要凭什么抽取式摘要的核心不是“生成”而是“打分”。你要先搞清楚模型在训练阶段看到的是什么形状的数据以及它在推理阶段如何输出结果。这一步想不清楚后面调参全是玄学。2.1 从 [CLS] 到句子向量BERT 怎么表达一个句子BERT 的输入格式通常是在每个句子前后加上特殊标记[CLS] 放在序列开头[SEP] 放在句子结尾。微调时绝大多数论文代码的做法是取 [CLS] 位置的输出向量接入一个全连接层和一个二分类层输出该句子被选入摘要的概率。这里有一个非常重要的设计选择你是把整个文档一次性塞进 BERT让所有句子共享上下文还是一个句子一个句子单独过 BERT两种做法的差异很大。把整篇文档塞进去句子之间的注意力可以相互影响模型理论上能捕捉“这段开头已经出现过的信息后面再重复就不该进摘要”这种跨句逻辑。但 BERT 的输入长度限制是硬约束经典 BERT 是 512 token长文档必须截断或滑窗。大多数可复现代码采用的方法是以句子为单位构造输入。每个训练样本的输入是“目标句子 前后各若干句的上下文”这样既能在一定程度上保留上下文信息又不会超过 BERT 的长度限制。这种设计在论文里通常被叫做 Truncated BERT-style 输入。你需要关注代码里是否实现了滑窗拼接而不是简单地把每句独立送入。2.2 标签怎么生成从论文里的 ROUGE 到代码里的 label这个环节是新手复现时最容易懵的地方。标签不是人工一句句标注的而是用 ROUGE 指标自动算出来的。具体做法是对文档里的每个句子分别计算它和参考摘要一般是人工写的标准摘要的 ROUGE-2 / ROUGE-L F1 分数然后设定一个阈值超过阈值的句子标记为 1否则为 0。这个设计思路来自经典论文后续很多摘要研究的代码都沿用了这个套路。阈值取值范围通常在 0.3 到 0.6 之间不同数据集的最佳值有差异。我在实际体验中建议先按论文给的默认阈值跑一遍再观察训练集的标签分布。如果正样本占比低于 5%说明阈值太高模型学不到东西如果高于 30%摘要就退化成段落复述。标签生成这一步要在训练之前离线完成不要在 DataLoader 里现算否则每次跑实验都会浪费大量时间。下面是一个参考的实现思路把文档切句、算 ROUGE、生成标签三件事串在一起import nltk from rouge_score import rouge_scorer scorer rouge_scorer.RougeScorer([rouge2, rougeL], use_stemmerTrue) def generate_labels(document, reference_summary, threshold0.45): sentences nltk.sent_tokenize(document) labels [] selected_indices [] for idx, sent in enumerate(sentences): scores scorer.score(reference_summary, sent) rouge2_f scores[rouge2].fmeasure rougeL_f scores[rougeL].fmeasure # 两个指标都过阈值才标记为正样本减少噪声 if rouge2_f threshold and rougeL_f threshold: labels.append(1) selected_indices.append(idx) else: labels.append(0) return sentences, labels, selected_indices这段代码里rouge_scorer是 Google 开源rouge_score库的接口use_stemmerTrue开启词干还原避免时态和复数导致分数偏低。阈值这里同时要求 ROUGE-2 和 ROUGE-L 都过比只看 ROUGE-2 要稳一些实际使用中能减少只覆盖关键词但不连贯的噪声句子被选中。2.3 评估指标选哪个ROUGE 的局限和论文里的隐藏细节你如果只跑训练不跑评估就无法判断模型是否真的学会了抽取摘要。ROUGE 指标是抽取式摘要的通用评估口径ROUGE-1 看的是单词级别的重合度ROUGE-2 看的是二元词组ROUGE-L 看的是最长公共子序列。论文代码里通常三类都算最后报告的往往是 ROUGE-1/2/L 的 F1。有个细节容易被忽略ROUGE 计算时要不要做标准化处理。默认情况下文本会做大小写归一化和标点剥离但有些作者的代码没有开启 stemmer导致结果和论文报告有一两个百分点的偏差。复现阶段建议优先保持和论文一致的评估设置不要一边微调一边换评估细节否则很难定位是模型问题还是评估口径问题。3. 把原始文章转成 BERT 输入格式数据预处理的做法这一章要解决的是从数据集到 DataLoader 之间的“最后一公里”。标题说是提取摘要的论文代码但真正花时间的往往不是模型训练而是数据整理。3.1 选用哪类数据集CNN/DailyMail 还是自己的业务文本公开复现实验时最常用的是 CNN/DailyMail 数据集它包含约 30 万篇新闻文章和对应摘要是摘要领域的事实标准。每篇文章的正文由多个句子组成摘要通常有 3 到 4 句。这个数据集的规模很适合微调 BERT在一张 V100 或 RTX 3090 上跑两三个 epoch 就能看到明显效果。如果是自己的业务数据你需要先确认文本质量。BERT 对噪声比较敏感全角半角混用、表格转文字残留、无意义短句都会影响微调效果。我的做法是先做一轮启发式清洗去掉纯数字行、去掉重复段落、合并被换行符切断的句子。3.2 单句输入加前后文构造窗口样本训练样本的构造策略直接影响模型效果。单一句子输入的缺点是模型完全看不到上下文打分时容易把每个句子都当成独立文本选出的摘要会缺乏连贯性。我建议在代码实现时采用“目标句 前两句 后一句”的拼接方式。from transformers import BertTokenizer tokenizer BertTokenizer.from_pretrained(bert-base-uncased) MAX_LEN 256 def build_training_samples(sentences, labels, tokenizer, max_len256): input_ids_list [] attention_mask_list [] label_list [] for idx, (sent, label) in enumerate(zip(sentences, labels)): start max(0, idx - 2) end min(len(sentences), idx 2) context_sents sentences[start:end 1] text .join(context_sents) encoded tokenizer.encode_plus( text, max_lengthmax_len, truncationTrue, paddingmax_length, return_tensorspt ) input_ids_list.append(encoded[input_ids]) attention_mask_list.append(encoded[attention_mask]) label_list.append(label) return input_ids_list, attention_mask_list, label_list这里的关键参数是max_len256。如果文档句子较长两前一后的拼接可能超过 512此时需要把max_len调大但注意最大不能超过模型位置编码上限。paddingmax_length会把所有样本统一补到 256方便后面批量训练。truncationTrue会截断超长部分但截断方向默认保留头部如果上下文的关键信息在尾部效果会打折扣。一个常用的优化是调整截断策略改成truncationonly_second并对输入区分段 ID让目标句子完整保留。这个细节在小数据集上能明显影响指标。3.3 数据划分和验证集构造别让测试集掺进训练流程论文复现中最容易犯的错是把同一个文档的不同句子同时分进训练和测试集导致评估得分虚高。抽取式摘要的数据划分应该以文档为粒度而不是以句子为粒度。# 假设 documents 是一个 list每个元素是该篇文档的所有句子 doc_indices list(range(len(documents))) random.Random(42).shuffle(doc_indices) train_docs doc_indices[: int(len(doc_indices) * 0.8)] valid_docs doc_indices[int(len(doc_indices) * 0.8): int(len(doc_indices) * 0.9)] test_docs doc_indices[int(len(doc_indices) * 0.9):]以文档为粒度切分之后再对每个文档内部的句子生成样本和标签确保同一个文档的句子不会同时出现在训练集和验证集。随机种子固定成 42 方便复现结果。这里的比例可以根据数据规模调整但测试集至少要保留 100 篇以上文档否则 ROUGE 指标的置信区间会非常大一个百分点的浮动根本看不出实验差异。4. 训练一个微调 BERT 摘要模型从脚本到参数调优数据准备好之后训练脚本是整个项目的核心。这一章给出一份能直接跑的 PyTorch 训练脚本配合 Hugging Face Transformers 库操作 BERT。4.1 模型结构BERT 加分类头的组合方式这里要用到BertForSequenceClassification。它是一个二分类模型输出 logits 的维度是 2取索引 1 的 logits 作为句子的入选概率。相比之下BertForPreTraining输出的是下一句预测和掩码语言模型的 logits不适合直接用于分类。注意区分。from transformers import BertForSequenceClassification model BertForSequenceClassification.from_pretrained( bert-base-uncased, num_labels2, output_hidden_statesFalse, )num_labels2让模型在顶部自动生成一个全连接分类层。如果你想把最后一层换成自己的结构可以设置num_labels2后通过model.classifier替换常见做法是在上面再加一个 Dropout 层和线性层。4.2 核心训练循环学习率、批次大小与早停训练参数的选择对收敛速度和最终效果影响极大。BERT 微调一般使用较小的学习率比如 2e-5 到 5e-5批次大小在 8 到 32 之间。显存有限时优先减小批次而不是调低max_len因为长度对摘要效果的影响比批次大小更直接。from torch.utils.data import DataLoader, TensorDataset import torch import torch.nn as nn from transformers import AdamW def train_model(input_ids_list, attention_mask_list, label_list, epochs3): input_ids torch.cat(input_ids_list, dim0) attention_mask torch.cat(attention_mask_list, dim0) labels torch.tensor(label_list, dtypetorch.long) dataset TensorDataset(input_ids, attention_mask, labels) dataloader DataLoader(dataset, batch_size16, shuffleTrue) device torch.device(cuda if torch.cuda.is_available() else cpu) model.to(device) optimizer AdamW(model.parameters(), lr3e-5, weight_decay0.01) loss_fn nn.CrossEntropyLoss() model.train() for epoch in range(epochs): total_loss 0 for batch in dataloader: batch_input_ids batch[0].to(device) batch_attention_mask batch[1].to(device) batch_labels batch[2].to(device) outputs model( input_idsbatch_input_ids, attention_maskbatch_attention_mask, labelsbatch_labels ) loss outputs.loss total_loss loss.item() optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() avg_loss total_loss / len(dataloader) print(fepoch {epoch 1} / {epochs}, loss: {avg_loss:.4f})这里用TensorDataset把之前生成的三个列表封装成数据集batch_size16在 12GB 显存上比较稳。clip_grad_norm_设置梯度裁剪为 1.0防止长文本样本造成梯度爆炸。AdamW是 BERT 微调的标准选择配合weight_decay可以降低过拟合。4.3 评估和保存只在验证 ROUGE 上涨时保存模型只保存 loss 最低的模型不一定能得到最高的 ROUGE 指标建议每到验证阶段就跑一次评估。这一环节同时承担两个目的过滤过拟合以及确定最终要保存的参数。# 在验证集上预测每句得分然后拼成摘要 def evaluate(model, valid_sentences_list, tokenizer, device, max_len256): model.eval() summaries [] for doc_sentences in valid_sentences_list: scores [] for sent in doc_sentences: encoded tokenizer.encode_plus( sent, max_lengthmax_len, truncationTrue, paddingmax_length, return_tensorspt ) with torch.no_grad(): logits model( input_idsencoded[input_ids].to(device), attention_maskencoded[attention_mask].to(device) ).logits prob torch.softmax(logits, dim1)[0][1].item() scores.append(prob) top_k max(1, int(len(doc_sentences) * 0.3)) top_indices sorted(range(len(scores)), keylambda i: scores[i], reverseTrue)[:top_k] top_indices.sort() summary [doc_sentences[i] for i in top_indices] summaries.append( .join(summary)) return summaries这段评估逻辑里有一个值得注意的设定选取 top 30% 的句子。实际体验中这一比例在不同数据上表现差异较大建议你在验证集上做小范围扫描选择 ROUGE 最高的比例作为最终配置。top_indices.sort()保证了句子按原文顺序输出这是抽取式摘要和生成式摘要的一个重要差异。5. 微调 BERT 提取摘要的避坑实测我从翻车过程中总结的 5 个关键教训BERT 摘要复现项目看起来简单但实际操作中容易在各种细节处卡住。这些坑代码里不会写论文里也不会提只有跑挂了才会发现。5.1 256 长度截断导致关键句子被切断现象训练 loss 正常下降但验证集 ROUGE-L 一直没过 18摘要里频繁出现语义不完整的半截句。原因句子拼接后超过max_len256truncationTrue默认从尾部截断而目标句子位于文本中间偏前位置前面的上下文保留但目标句后半段被切掉了。模型看到的并不是完整句子。解决排查 DataLoader 里实际喂给 BERT 的输入把truncationonly_second加上同时把max_len提升到 384 或 512。我自己的习惯是构造样本时设置一个debug开关随机打印几个样本的 decode 文本肉眼确认目标句完整。5.2 ROUGE 标签阈值设太高正样本几乎为零现象训练 loss 在第一个 epoch 之后几乎不动模型预测全部落到 0 类。原因文档句子和参考摘要的 ROUGE-2 分数整体偏低阈值按 0.6 设正样本比例只有 1% 左右。模型学不到正类特征。解决先统计训练集标签分布阈值下调到 0.35 到 0.45让正样本比例落在 8% 到 15%。注意这里是 F1 值而不是 precision按论文里的数值设置之前先确认指标口径一致。5.3 全角半角混用导致 tokenizer 分词异常现象英文文档中出现中文逗号、中文冒号等字符BERT 的 tokenizer 会把它们切成[UNK]训练速度变慢且效果下降。原因数据清洗阶段没有统一字符编码把从 PDF 或网页抽取的文本直接拼进数据集。解决预处理阶段先用unicodedata.normalize(NFKC, text)归一化所有字符再替换标点为英文标点。这一步看起来微小但对 ROUGE 计算的 F1 影响能在 1 到 2 个百分点。5.4 训练集和验证集句子级混切ROUGE 虚高现象验证集 ROUGE-1 超过论文报告值但实际人工看摘要质量很差且多次实验指标波动剧烈。原因数据切分在句子级别而不是文档级别同一个文档的训练句子和验证句子存在上下文重叠模型在验证集上“见过”相似内容。解决改造成按文档切分同时记录文档 ID。验证时以文档 ID 为索引确保一篇文档内所有句子只出现在同一个集合。5.5 随机种子不固定两次实验 ROUGE 相差 3 个点现象同一份代码、同样的超参数换台机器跑结果就不同甚至同一台机器跑两次也有明显差异。原因数据打乱、Dropout、模型初始化的随机性没有被固定。解决训练前设置torch.manual_seed、random.seed和numpy.random.seed。同时把 DataLoader 的shuffleTrue配合固定 seed 使用并保持 PyTorch 版本一致。如果换卡跑把torch.backends.cudnn.deterministic设为True。6. 把微调后的模型用于推理ROUGE 验证和阈值调优的进阶技巧模型训练完之后真正的实战才刚刚开始。这一步要解决“模型输出哪些句子”的问题涉及两个关键控制项选择句子的阈值和数量。6.1 动态数量策略不要固定前 3 句很多论文代码为了省事直接取得分最高的 N 句话作为摘要但实际文档长短差异大短文档取 3 句可能等于全文复述长文档取 3 句又可能漏掉关键信息。我常用的做法是设置一个长度上限比例例如摘要总字数不超过原文的 25%同时句子数量不超过 5 句。def generate_summary(sentences, sentence_scores, max_ratio0.25, max_sentences5): sorted_scores sorted( [(i, score) for i, score in enumerate(sentence_scores)], keylambda x: x[1], reverseTrue ) selected [] total_len 0 for idx, score in sorted_scores: sent_len len(sentences[idx]) if total_len sent_len max_ratio * sum(len(s) for s in sentences): continue selected.append(idx) total_len sent_len if len(selected) max_sentences: break selected.sort() return [sentences[i] for i in selected]这里用原文总长度的 25% 作为摘要长度上限max_sentences兜底防止摘要过长。比例数值可以放在配置文件里在验证集上用 ROUGE 反推最优值。6.2 验证合成摘要的可读性ROUGE 之外的人工检查ROUGE 指标有一个天然缺陷它衡量的是 n-gram 重合度而不是摘要的可读性和信息完整性。模型可能选出一堆包含关键词但语义跳跃的句子ROUGE 照样给高分。我个人的习惯是微调结束之后随机抽取 20 到 30 篇文档人工对比模型输出的摘要和参考摘要。重点检查三方面选中句子是否跨段落跳跃、是否存在相同信息重复出现、是否遗漏了原文章核心观点。这个动作不用每轮实验都做但最终交付前必须做一次。6.3 一个值得尝试的方向把 Length Penalty 加进打分逻辑BERT 分类头的输出分数本质上和句子长度无关但摘要中过长的句子会侵占篇幅挤压其他有效句子的入选空间。一个轻量改进是在推理时对得分做长度惩罚final_score model_score - lambda * (len(sent) / max_len)。lambda从 0.01 到 0.1 之间扫描在验证集上通常能找到比无惩罚高 1 到 2 个 ROUGE 点的参数。这个技巧在论文代码里很少见但实际应用时很管用。我最早是在一份做新闻摘要的私有代码里看到这个做法后来在多个数据集上验证都有效。最后说一个我的习惯微调这类模型永远先在小规模数据上把整个流程跑通再放大到全量数据。数据预处理和训练脚本里的 bug 在小数据上暴露更快调试成本更低。希望这篇笔记能帮你在复现 BERT 摘要代码的路上少走一段弯路。本文还有配套的精品资源点击获取
返回列表