ARTICLE DETAIL

资讯详情

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

BERT微调实战:20NewsGroups文本分类全流程与避坑指南

BERT微调实战:20NewsGroups文本分类全流程与避坑指南 简介一份面向课程实验与作业的BERT文本分类完整实践包聚焦20NewsGroups新闻数据集的多类别分类任务。压缩包共21个文件包含源码脚本、编译缓存、训练与测试文本、运行日志、配置备份、说明文档和PDF报告整体约14.42MB已有76人学习下载。源码目录提供数据读取、模型定义、训练与评估的完整实现数据目录存放切分好的新闻样本日志目录记录运行输出结构清晰便于对照复现。实践包覆盖数据预处理、模型微调与评估全流程可快速复现实验并理解BERT在新闻分类中的实际用法微调时可调整学习率、批大小等超参数借助验证集准确率评估效果。日志与检查点文件能帮助排查参数调整问题适合作为自然语言处理课程设计或入门研究的完整参考。1. BERT 在 20NewsGroups 上的分类任务一次把数据、模型和评估串起来的微调实测20NewsGroups 是文本分类绕不开的 20 类新闻邮件数据集接近两万篇英文文本用 BERT 在它上面做分类任务听起来只是把数据喂给模型实际链路却比想象长加载时要去邮件头防止泄漏分词时要处理 padding 和截断训练时要盯着 loss 防止过拟合。这篇文章按我自己的实操顺序把这套流程从数据准备写到验证评估并把踩过的坑标出来。适合刚接触 BERT 微调、想找个标准数据集练手的人也适合想确认自己流程到底哪一步导致分数虚高或翻车的熟手。2. 把 20NewsGroups 喂给 BERT数据加载、分词和 Dataset 构建2.1 用 fetch_20newsgroups 加载数据但必须 remove 掉邮件头很多人拿到这个数据集的第一反应是直接用 sklearn 的fetch_20newsgroups一把梭这没问题但如果默认参数一口气 load 下来你会在后续训练里得到一个“看起来很漂亮、实际上没学会”的模型。原因后面避坑章节会详细说先记住一个结论加载的时候要把headers、footers、quotes一起去掉。from sklearn.datasets import fetch_20newsgroups train_data fetch_20newsgroups( subsettrain, remove(headers, footers, quotes), shuffleTrue, random_state42, ) test_data fetch_20newsgroups( subsettest, remove(headers, footers, quotes), ) texts_train train_data[data] y_train train_data[target] texts_test test_data[data] y_test test_data[target] target_names train_data[target_names] print(target_names) print(len(texts_train), len(texts_test))这里target_names是个长度为 20 的列表默认按类别名排序从alt.atheism到talk.religion.misctarget就是对应这 20 个类别的 0~19 数字编号。训练集约有一万多篇测试集几千篇类别分布整体均衡所以准确率是能直接用的指标。remove(headers, footers, quotes)的作用是剥掉每封邮件头部的 From/Subject/Organization 字段、文末签名和引用内容训练集和测试集都要用同样的参数否则训练和推理时看到的文本形态不一致。2.2 先统计文本长度再决定 max_len 是多少BERT 的输入长度硬上限是 512 个 token20NewsGroups 里的文本是新闻组邮件长短差距很大。有人上来就设 512结果显存爆得莫名其妙也有人设 64模型精度明显吃亏。我一般会先抽样几百条用同一个 tokenizer 统计一下长度分布再决定截断长度。from transformers import BertTokenizer import numpy as np tokenizer BertTokenizer.from_pretrained(bert-base-uncased) sample_lens [] for text in texts_train[:2000]: ids tokenizer(text, add_special_tokensTrue)[input_ids] sample_lens.append(len(ids)) for p in [50, 90, 95, 99]: print(fpercentile {p}: {int(np.percentile(sample_lens, p))})我在这类数据集上常见的分布是中位数在 100 到 200 token 之间90 分位数可能在两三百左右。所以max_len设为 128 还是 256取决于你能接受多少信息被截断。设 128 训练更快、显存更省代价是长邮件后半段被切掉设 256 信息保留更完整但显存和训练时间几乎翻倍。如果你只是跑通流程128 够用如果追求更高准确率统计完分布后选一个能覆盖九成样本的值通常是 256。要注意这里用的是bert-base-uncased它只处理英文小写文本如果你后面换成中文数据集要换成bert-base-chinese或对应的中文预训练模型。2.3 一次性 tokenize避免每个 epoch 重复分词接下来是构造 Dataset。常见做法是自定义一个 PyTorch Dataset在__getitem__里调用 tokenizer这样写逻辑简单但每一轮训练都会重新做一次分词一万多篇文本就是三次重复劳动。更好的做法是先把所有文本一次性编码好把input_ids和attention_mask存成 Tensor训练时直接按索引取。import torch from torch.utils.data import Dataset, DataLoader train_enc tokenizer( texts_train, max_length128, paddingmax_length, truncationTrue, ) test_enc tokenizer( texts_test, max_length128, paddingmax_length, truncationTrue, ) class NewsDataset(Dataset): def __init__(self, encodings, labels): self.input_ids torch.tensor(encodings[input_ids]) self.attention_mask torch.tensor(encodings[attention_mask]) self.labels torch.tensor(labels, dtypetorch.long) def __len__(self): return len(self.labels) def __getitem__(self, idx): return { input_ids: self.input_ids[idx], attention_mask: self.attention_mask[idx], label: self.labels[idx], } train_dataset NewsDataset(train_enc, y_train) test_dataset NewsDataset(test_enc, y_test) train_loader DataLoader(train_dataset, batch_size16, shuffleTrue) test_loader DataLoader(test_dataset, batch_size32, shuffleFalse)这段代码里有几个容易被忽略的细节。paddingmax_length会把所有样本都补到 128而不是只补到本 batch 的最大长度。好处是每个 batch 张量都是固定形状训练稳定代价是短文本浪费一部分计算。如果改成paddinglongest能省一点计算但每个 batch 形状不同速度收益在这个数据规模上并不明显我倾向直接用max_length。attention_mask对应的是 padding 位置的标记1 表示真实 token0 表示 padding token模型在 self-attention 里会忽略这些位置。训练时如果把 mask 漏掉模型会把一堆无意义的 [PAD] 当成正常文本loss 会偏低但评估也受影响。truncationTrue表示超过 128 的部分从右侧截掉对新闻邮件来说通常前面主题相关性更强这个截断方向是合理的。3. BERT 为什么能处理 20 类新闻分类分类头、损失函数和微调策略3.1 预训练权重 分类头的组合逻辑BERT 本身不是一个分类器它是一个在大规模英文语料上预训练的语言模型学的是“词在上下文里的表示”。放到 20NewsGroups 任务上我们需要在 BERT 顶层接一个分类头输入一篇文章输出 20 个类别的概率分布。这就是常见的BertForSequenceClassification结构。HuggingFace 的实现里BertForSequenceClassification默认取输入序列第一个 token也就是[CLS]的最后一层隐状态经过一个线性层映射到 20 维 logits再配合 softmax 得到概率。这里的[CLS]在预训练阶段就被训练成“汇总全句信息”的角色微调时它会进一步适配当前分类数据。值得注意的是有些实验表明只拿[CLS]不一定是最优解把最后一层所有 token 的向量做 mean pooling 再进分类头在某些数据集上效果更好。如果你有时间可以在 20NewsGroups 上把两种 pooling 方式做一个消融对比这个实验成本不高结论也能直接写到论文里。微调的含义是整个 BERT 的参数都会跟着更新不只是最后那个线性层。这意味着反向传播时模型要缓存每一层的激活值显存开销明显大于“只训练分类头”。这也是为什么训练这类模型需要关注 batch size 和 max_len而不是像传统机器学习那样只关心学习率。3.2 多分类损失函数CrossEntropyLoss 和前向返回的 loss分类头输出的是 20 个 logits需要和真实标签计算损失。20NewsGroups 是单标签多分类任务标准做法是交叉熵损失。PyTorch 里CrossEntropyLoss会先对 logits 做 softmax 再计算负对数似然不需要你手动再加 softmax。如果你用BertForSequenceClassification直接把labels传进前向方法它会返回一个loss字段这个 loss 内部就是CrossEntropyLoss(logits, labels)。代码写成outputs model(input_ids..., attention_mask..., labels...)之后直接用outputs.loss做 backward 就行。验证阶段拿 logits 做argmax(dim-1)得到预测类别。这里要提醒一下训练时必须model.train()让 dropout 生效验证或推理时必须model.eval()并包在torch.no_grad()里否则 dropout 随机失活会让你每次拿到的预测结果不一样而且验证时会多算一遍梯度图白吃显存。3.3 微调超参数怎么设从学习率到 warmupBERT 微调和从头训练模型不一样它对学习率非常敏感。因为预训练权重已经是一个很好的初始点学习率太大容易灾难性遗忘太小则收敛太慢。以下是我在 20NewsGroups 这类中等规模文本分类数据集上的默认配置参数常见区间我一般用说明learning rate2e-5 ~ 5e-52e-5太大会破坏预训练权重batch size8 ~ 3216显存不够时用梯度累积max_len128 / 256128先统计长度分布再定epochs2 ~ 43一轮以上基本够太多必过拟合warmup 比例5% ~ 10%10%让学习率从 0 平稳上升weight_decay0.01 ~ 0.10.01对分类头上作用更明显warmup 这一段官方实现里叫get_linear_schedule_with_warmup前 10% 的 step 学习率线性上升后面线性衰减到 0。它的作用是避免训练初期模型权重被大步长撞出预训练好的区域这在 NLP 微调里几乎是标配。如果你看到 loss 曲线前几步剧烈抖动先检查是不是忘了加 scheduler其次再怀疑学习率过高。batch size 和学习率是联动的。显存不够把 batch 从 16 降到 8 时学习率最好也降到 1e-5 左右或者用梯度累积凑回等效的 16。20NewsGroups 文本长度不短即便 max_len128单卡 16G 也只能放几十个样本所以梯度累积是必须掌握的技巧。4. 跑通完整训练流程从 DataLoader 到最优模型保存4.1 合并数据准备、训练、评估和保存的完整脚本下面是一个可以直接跑通的完整流程把前面两章的数据处理、分词、Dataset 构造加上训练和评估都串在一起。为了控制篇幅评估这里直接用内置 test 集做最终测试如果是正式研究我建议再从 train 里切出 10% 当验证集用验证集选超参test 只碰一次。import json import torch from torch.utils.data import Dataset, DataLoader from transformers import ( BertTokenizer, BertForSequenceClassification, AdamW, get_linear_schedule_with_warmup, ) from sklearn.datasets import fetch_20newsgroups from sklearn.metrics import classification_report # ---------- 1. 加载数据 ---------- train_data fetch_20newsgroups(subsettrain, remove(headers, footers, quotes)) test_data fetch_20newsgroups(subsettest, remove(headers, footers, quotes)) texts_train, y_train train_data[data], train_data[target] texts_test, y_test test_data[data], test_data[target] target_names train_data[target_names] # ---------- 2. 一次性 tokenize ---------- tokenizer BertTokenizer.from_pretrained(bert-base-uncased) train_enc tokenizer(texts_train, max_length128, paddingmax_length, truncationTrue) test_enc tokenizer(texts_test, max_length128, paddingmax_length, truncationTrue) class NewsDataset(Dataset): def __init__(self, encodings, labels): self.input_ids torch.tensor(encodings[input_ids]) self.attention_mask torch.tensor(encodings[attention_mask]) self.labels torch.tensor(labels, dtypetorch.long) def __len__(self): return len(self.labels) def __getitem__(self, idx): return { input_ids: self.input_ids[idx], attention_mask: self.attention_mask[idx], label: self.labels[idx], } train_dataset NewsDataset(train_enc, y_train) test_dataset NewsDataset(test_enc, y_test) train_loader DataLoader(train_dataset, batch_size16, shuffleTrue) test_loader DataLoader(test_dataset, batch_size32, shuffleFalse) # ---------- 3. 模型、优化器、调度器 ---------- device torch.device(cuda if torch.cuda.is_available() else cpu) model BertForSequenceClassification.from_pretrained( bert-base-uncased, num_labelslen(target_names), ).to(device) optimizer AdamW(model.parameters(), lr2e-5, weight_decay0.01) epochs 3 total_steps len(train_loader) * epochs scheduler get_linear_schedule_with_warmup( optimizer, num_warmup_stepsint(total_steps * 0.1), num_training_stepstotal_steps, ) # ---------- 4. 训练 ---------- for epoch in range(epochs): model.train() total_loss 0.0 for step, batch in enumerate(train_loader): input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) labels batch[label].to(device) outputs model( input_idsinput_ids, attention_maskattention_mask, labelslabels, ) loss outputs.loss loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() scheduler.step() optimizer.zero_grad() total_loss loss.item() print(fepoch {epoch} avg_loss {total_loss / len(train_loader):.4f}) # ---------- 5. 最终评估 ---------- model.eval() preds, truths [], [] with torch.no_grad(): for batch in test_loader: input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) logits model(input_idsinput_ids, attention_maskattention_mask).logits preds.extend(logits.argmax(dim-1).cpu().tolist()) truths.extend(batch[label].cpu().tolist()) print(classification_report(truths, preds, target_namestarget_names, digits4)) # ---------- 6. 保存模型和映射 ---------- model.save_pretrained(./news_bert) tokenizer.save_pretrained(./news_bert) with open(./news_bert/target_names.json, w, encodingutf-8) as f: json.dump(target_names, f)4.2 脚本里的关键参数为什么要这样设第 4.1 的代码里加了torch.nn.utils.clip_grad_norm_梯度裁剪到 1.0这一步对 BERT 微调很重要。BERT 的深层网络在训练初期容易出现梯度爆炸loss 直接变成 NaN裁剪之后至少能保证训练不翻车。如果你看到 loss 突然跳到几十上百先检查是不是少了这一行。AdamW不是Adam区别在于它把权重衰减和梯度更新解耦HuggingFace 官方默认就是 AdamW。weight_decay0.01对全连接层和 attention 层都会生效能在小数据集上延缓过拟合。学习率我固定用 2e-5如果你的显存只能放下 batch_size8把学习率降到 1e-5 会更稳。训练轮数方面20NewsGroups 在 remove 掉邮件头之后BERT 一般两到三个 epoch 就能达到接近峰值的准确率。第一个 epoch 结束验证集可能还在明显上升第二个 epoch 后开始趋稳第三个 epoch 后如果不早停就会开始过拟合。我的习惯是训练时每个 epoch 存一次模型或者直接在验证集上监控准确率实现早停最后选验证集分数最高的那个 checkpoint而不是最后一个。关于用 test 集做评估代码里是直接拿fetch_20newsgroups的 test 子集算 classification_report。用交叉熵训练的分类模型classification_report里的 precision、recall、f1 会按 20 个类别分别统计。由于类别均衡macro avg 和 weighted avg 差距不大但一定要看 macro avg因为它能暴露模型在少数难分类别上的短板。模型保存这里save_pretrained会把模型权重和 config 一起存到一个目录tokenizer.save_pretrained会把词表和 tokenizer 配置也存进去之后用from_pretrained(./news_bert)就能直接加载。我额外存了target_names.json是因为 label 编号必须和类别名对齐只存权重不存映射推理时很容易错位。5. BERT 20NewsGroups 避坑记录数据泄漏、显存爆掉和过拟合5.1 不删邮件头准确率虚高的数据泄漏现象直接用默认fetch_20newsgroups加载数据训练测试集准确率轻松超过 95%甚至接近 98%看起来模型强得离谱。原因20NewsGroups 本质是新闻组邮件每篇都带着 From、Newsgroups、Organization 这些邮件头字段其中 Newsgroups 字段直接写了类别名。模型根本不需要读正文只要识别出这个字符串分类就完成了。这不是模型学到了语义而是学到了数据里的“快捷答案”。解决加载时用remove(headers, footers, quotes)训练集和测试集都这样处理。如果从原始语料自己制作数据集也要写正则把From:开头、签名块和开头的引用行剥掉。判断自己有没有踩这个坑最简单的办法是看准确率是否异常高真正的语义分类任务BERT 在这个数据集上 remove 掉头之后分数会明显下降但那是真实水平。5.2 max_len512 引发的显存爆掉现象训练跑不了几个 step显存直接 OOM或者 loss 开始大幅波动。原因BERT 反向传播需要缓存每一层 attention 的激活值max_len从 128 提高到 512显存占用近似线性放大 4 倍。20NewsGroups 很多邮件正文较长但并不是所有信息都有用强行塞满 512 只会让单 batch 变小、训练变慢、显存爆掉。解决先按照第 2.2 节的方法统计长度分布选一个能覆盖九成样本的max_len常见是 128 或 256。如果确实需要长文本信息考虑把文本做滑窗切段分别过 BERT 后做 pooling 融合但这个方案成本高20NewsGroups 上不一定值。显存还不够时把 batch_size 降到 8训练轮数和学习率都要相应调整。5.3 训练集 loss 很低验证集却在上涨过拟合来得比想象快现象第二个 epoch 开始训练集 loss 一路降到接近 0验证集准确率却不涨反跌。原因训练集约一万多篇对整个 BERT 来说是偏小的数据量。BERT-base 有 1.1 亿参数几个 epoch 就能把训练集的邮件细节背下来而不是学出可泛化的类别特征。解决epoch 控制在 2 到 4 之间配合早停验证集不再提升就停。weight_decay保留 0.01学习率不要超过 3e-5。如果过拟合还是很严重可以尝试在文本层面做轻量增强比如随机删除句子里的一些词再训练但这类增强要小心不要破坏类别关键信息。更正规的做法是把验证集切出来做模型选择用验证集上 f1 最高的模型而不是训练 loss 最低的模型。5.4 保存模型后推理错位label 映射没跟着存现象训练时准确率正常把模型保存下来换到推理脚本里输入一条新闻预测结果总是错到完全无关的类别。原因fetch_20newsgroups的target_names按字母排序label 0 到 19 对应这个固定顺序。推理时如果重新加载数据类别列表顺序变化或者只加载模型没加载target_nameslogits 的 argmax 索引就会映射到错误的类别名上。解决保存模型的同时保存 label 映射建议把target_names写成 JSON 文件放进模型目录。推理时用同一个列表恢复类别名。更稳妥的做法是把 id2label 写进 model configmodel.config.id2label {i: name for i, name in enumerate(target_names)} model.config.label2id {name: i for i, name in enumerate(target_names)} model.save_pretrained(./news_bert)这样from_pretrained时可以自动带上标签映射推理脚本里就不用额外管理 JSON 文件了。6. 把模型从“训练完”变成“验证过”最小推理和混淆矩阵实战检查6.1 用混淆矩阵看模型到底把哪些类搞混训练完只看分类报告还远远不够。20NewsGroups 里有些类别本身就高度相似比如comp.sys.ibm.pc.hardware和comp.sys.mac.hardware以及几个talk.politics.*子类模型混淆它们是正常的但从混淆矩阵里能看到是“合理的边界模糊”还是“系统性错误”。如果某个类别的预测大量跑到不相关类别去就该检查数据清洗或训练过程。from sklearn.metrics import confusion_matrix import matplotlib.pyplot as plt cm confusion_matrix(truths, preds) fig, ax plt.subplots(figsize(16, 14)) im ax.imshow(cm, cmapBlues) ax.set_xticks(range(len(target_names))) ax.set_yticks(range(len(target_names))) ax.set_xticklabels(target_names, rotation90) ax.set_yticklabels(target_names) plt.colorbar(im) plt.tight_layout() plt.savefig(confusion_matrix.png, dpi150)线上实践里我不会逐个看 20×20 的每个格子而是先找出每行中最大的非对角线项那才是模型最容易混淆的一对类别。如果这类混淆集中出现在某个主题簇里比如体育类内部说明模型学到了主题级别的区分但没学到簇内的精细差异这可以作为下一步用更大max_len或更长训练轮数来验证的起点。6.2 最小推理脚本把保存的模型跑起来训练完的模型要能真正拿来预测单条新闻才算是闭环。下面这个脚本加载第 4 章保存的目录对一条新文本做预测from transformers import BertTokenizer, BertForSequenceClassification import json model BertForSequenceClassification.from_pretrained(./news_bert) tokenizer BertTokenizer.from_pretrained(./news_bert) with open(./news_bert/target_names.json, encodingutf-8) as f: target_names json.load(f) text My MacBooks fan keeps running at full speed after the latest update. inputs tokenizer( text, max_length128, truncationTrue, paddingmax_length, return_tensorspt, ) logits model(**inputs).logits pred logits.argmax(dim-1).item() print(target_names[pred])注意推理时也要有truncation和padding并且max_length要和训练时保持一致。输入短文本时paddingmax_length会补到 128模型没问题但会浪费一点计算。如果后面把这个模型做服务化部署可以把max_length改成 256 或按业务重新训练一遍而不是直接换参数因为模型在训练时见过的长度分布会影响它对长文本的表示。以前我偷懒不删邮件头跑出过虚高的分数还高兴了两天后来用混淆矩阵才发现模型根本没学会分类只是记住了字段。现在我做任何新分类任务第一件事都是先怀疑数据里有没有捷径再谈模型调优。希望帮到你。本文还有配套的精品资源点击获取
返回列表