news 2026/9/8 2:22:34

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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.txt:

torch>=2.0.0 transformers>=4.30.0 nltk>=3.8 rouge-score>=0.1.2 tqdm>=4.65.0 numpy>=1.24.0 datasets>=2.12.0

安装命令很简单:

conda create -n bertsum python=3.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, threshold=0.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_len=512): 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_pretrained='bert-base-uncased', hidden_size=768): 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_ids=input_ids, attention_mask=attention_mask, token_type_ids=segment_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(), lr=2e-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与warmup

BERT微调有一组被反复验证过的“黄金参数”,直接套用基本不会出大问题。我在复现过程中用的参数如下:

参数推荐值说明
学习率2e-5BERT全量微调默认值,过大容易毁掉预训练权重
训练轮数(epoch)3-5小数据时可以多加到5轮,但注意观察过拟合
Batch Size8-32按显存来,不够就用梯度累积
梯度累积步数4或8等效放大batch size,稳定梯度
Warmup步数总步数的10%防止初期loss剧烈震荡
梯度裁剪max_grad_norm=1.0防止梯度爆炸导致NaN

学习率是这里最重要的参数。BERT的预训练权重已经很“成熟”了,学习率设太大相当于把学好的参数一脚踢飞,loss会乱跳;设太小又微调不动。如果你自己改了学习率,一定要在验证集上盯ROUGE指标,不要只看训练loss。

3.4 显存不够怎么办:梯度累积与混合精度

我跑这篇论文的时候用的是一张8G显存的消费级显卡,bert-base-uncased本身就接近400M参数,全量微调时稍微加个长文档就很容易爆显存。我的解法是三层组合拳。

第一,先把batch size降到1。BERTSUM的每个样本本身已经是一整篇文档,batch size为1时batch内已经包含多个句子,效果损失很小。

第二,用梯度累积。把4个batch的梯度攒起来再统一更新,等效于batch size=4。代码上就是把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_stemmer=True) scores = scorer.score(reference_summary, prediction_summary) print(scores['rougeL'].fmeasure)

use_stemmer=True会把单词还原成词根,比如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不降或直接变成NaN

Loss一直不降,最常见的原因是学习率开太大,或者标签和输入没有对齐。BERT微调学习率超过5e-5就很容易出问题,建议回退到2e-5再试。Loss变成NaN,基本可以断定是梯度爆炸,除了降低学习率,别忘了在optimizer.step()之前加torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0),这是最稳的兜底手段。

还有个特别容易踩的坑是正负样本不均衡。如果一篇文档里绝大多数句子标签都是0,模型会学会无脑预测0,loss虽然不高但完全没用。遇到这种情况,可以把训练文档里标签全为0的样本过滤掉一部分,或者给正样本的loss加权重。

5.3 预训练模型下载慢或加载失败

初次运行代码时,transformers会自动下载bert-base-uncased权重,网络不好的话很容易卡住甚至中断。解决办法是设置镜像环境变量,让HuggingFace走国内镜像源:

export HF_ENDPOINT=https://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生成式摘要的对比实验,到时候拿数据说话。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/8 2:20:43

主从博弈框架下综合能源系统需求响应与电能交互优化调度

1. 项目概述与核心痛点分析 1.1 这个课题到底在做什么 先说人话版本:现在能源系统早就不是"发电厂→用户"的单向管道了,一个园区里可能同时存在光伏、储能、燃气轮机、电锅炉、冰蓄冷空调,还可能出现多个园区手拉手互相借电的情况…

作者头像 李华
网站建设 2026/9/8 2:19:47

HarmonyOS ArkTS层叠布局Stack深度解析:对齐、定位与避坑实战

搞了半天,终于把HarmonyOS那套ArkTS里的层叠布局(Stack)整明白了。前几天有个刚转鸿蒙开发的朋友问我,一个头像右上角的红色角标,怎么用原生组件放上去?我第一反应就是:这玩意不就是给Stack准备…

作者头像 李华
网站建设 2026/9/8 2:19:20

基于OpenCV的舞蹈镜像对比学习工具开发实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/8 2:18:23

Word论文页眉页脚设置:分节、页码与常见问题全解析

写毕业论文的时候,很多人都被 Word 页眉页脚折磨过。明明设置了页码,正文前面的摘要目录也带上了编号;明明删掉了页眉里的横线,下一页又冒出来;明明想从某一页开始插入罗马数字页码,结果整个文档全都乱了。…

作者头像 李华
网站建设 2026/9/8 2:17:35

1968道奇Charger改装:千匹马力碳纤维肌肉车技术解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/8 2:17:22

轻量级灰度发布平台实践:基于Spring Cloud Gateway与Redis的动态流量控制

1. 先想清楚再动手:这套轻量灰度平台的方案选型 先说一个我自己的线上事故。有一次发一个新版订单服务,自测、测试环境全过了,结果全量上线后不到十分钟,用户开始集中反馈下单页白屏。最后定位到是某个老浏览器不兼容新前端资源的…

作者头像 李华