资讯动态

Keras-bert微调BERT实现文本多标签分类实战指南

发布时间:2026/10/9 3:01:37 来源:尧图企业网站定制
简介针对文本多标签分类任务这是一份基于Keras与Keras-bert微调BERT的完整项目实现面向NLP学习者或需要快速搭建多标签分类基线的算法工程师。项目以2020语言与智能技术竞赛事件抽取数据集为样例代码覆盖数据加载、模型训练、FGM对抗训练、评估与预测等完整流程有助于理解多标签场景下的数据组织方式与模型调优要点。包体共10个文件包含4个Python脚本、2个CSV数据文件、2个TXT配置与词典文件以及README说明整体仅1.01MB结构清晰便于按模块阅读与复用。资源已有1634人学习下载适合希望借助赛事真实数据演练BERT微调、Keras-bert使用和多标签评价指标的中级NLP开发者。1. 文本多标签分类加 BERT 微调这个 Keras 项目在解决什么事我常被问到的一句话是“帮我做一个多标签分类输入一句话输出几个可能的标签。” 这类需求多出现在工单分类、文章打标、投诉质检里。单标签做多了的人很容易下意识写 Softmax推到线上才发现一段文本同时命中两个标签时模型完全没办法。这个项目的做法是用 Keras 作为训练框架Keras-bert 作为 BERT 的加载入口在预训练 BERT 的最后一层接一个多标签输出头做微调一次前向就拿到多个标签概率。项目本身不算复杂但把坑填过一遍还是能省掉不少试错成本。适合已经会 Keras 基本流程、想上手 BERT 微调但还没踩过数据预处理和参数的人。2. 为什么多标签文本分类必须微调 BERTsigmoid、F1 与数据边界2.1 多标签不是“多个单分类”输出层和损失函数要一起换单标签文本分类到最后是 softmax 加上 categorical_crossentropy模型必须在类别里赌一个总共概率加起来等于 1。多标签分类不一样一条“人工智能-项目实践-文本分类”的标题按业务规则应该同时命中“人工智能”“项目实践”“文本分类”三个标签。如果你用 softmax 输出模型会因为概率互斥去压制其中一个标签这在训练数据的真实分布里根本学不出来。所以我把输出层换成 sigmoid每个标签独立算一个概率。损失函数换成 binary_crossentropy它在这里的动作是对每一个标签节点把真实标签和预测概率放进去算二值交叉熵最后对全部标签取平均。这种“多个二分类”的观点很粗糙但实践中非常稳定。拆成多个二分类模型会浪费 BERT 里各标签共享的语义特征所以我通常只保留一个 BERT 主体再接一个 Dense 输出层。from keras.layers import Dense, Dropout from keras_bert import load_trained_model_from_checkpoint seq_len 128 num_labels 5 bert load_trained_model_from_checkpoint( config_file, checkpoint_file, seq_lenseq_len, output_layer_num1, ) cls_output bert.output[:, 0, :] dropout Dropout(0.3, namecls_dropout)(cls_output) logits Dense(num_labels, activationsigmoid, namelabel_logits)(dropout)这里cls_output取的是 BERT 序列输出的第一个 token 向量。加载出来的 bert 模型输出维度是(batch, seq_len, hidden)[:, 0, :]把长度维度压掉留下 CLS 向量。如果你对池化方式没把握也可以改为GlobalMaxPooling1D()或GlobalAveragePooling1D()但我的项目里 CLS Dropout 是最省心的一版。实在不行再换池化不要一开始就在池化上纠结。那个 0.3 的 Dropout 也是微调阶段的经验值。标签之间有一定关联比如“人工智能”和“机器学习”经常一起出现Dropout 太低会让分类头把这些共现关系背下来结果换一批文本就崩。2.2 微调边界要画清楚哪些层更新哪些层要拆掉Keras-bert 把预训练 checkpoint 还原成 Keras 模型后默认所有层都可训练。全参数微调的好处是能利用 BERT 已经学到的句法、指代和上下文信息让它为我们的多标签任务重新组织特征。缺点是标注量不够时很容易过拟合。项目里我会先看训练集规模标注在 1 万条以上、标签不超过几十个全量微调标注在几千条附近冻结底部 8 层只让上部 transformer 层和分类头参与更新。冻结可以用一段循环控制for i, layer in enumerate(bert.layers): if i 8: layer.trainable False这段代码放在load_trained_model_from_checkpoint之后、组装新模型之前。注意bert.layers里前几层是 embedding、位置编码和 transformer 的前半部分。冻结它们之后这些层的前向计算仍然会执行只是反向传播时梯度不再更新。用 Keras 观察训练时被冻结层的权重数值不会变这个可以用layer.get_weights()在开头和结尾对比确认。还有一个边界容易漏load_trained_model_from_checkpoint返回的模型里可能带着预训练阶段的 pooler 和分类层。我们要做的是把bert.output当作特征入口然后自己重新接分类头。如果你直接把旧模型的输出接 Dense得到的结果可能是个连维度都对不上的怪胎。最稳妥的办法是读完权重后只保留bert.input和bert.output再用函数式 API 新建模型如第 3 章所示。有过一次惨痛教训我当时图省事直接拿bert这个模型继续add(Dense())结果 Keras 报错说模型内部已经存在输出层。后来老实改成函数式 API才把流程理顺。所以别把预训练加载出来的对象当普通 Sequential它是一个完整的 Keras Model改结构要走Model(inputs..., outputs...)。2.3 多标签数据预处理标签向量和长文本截断先定死文本进 BERT 前要先变成 token id 和 segment id这个在后面章节会展开。这里先讲数据的两个前置问题标签向量怎么编长文本怎么截。多标签分类的数据集一般长这样每一行是一条文本tags 字段用空格或逗号分隔多个标签。我们需要把这个字段拆开再映射成固定的向量。from sklearn.preprocessing import MultiLabelBinarizer tag_list [ [人工智能, 项目实践, 文本分类], [大模型微调, 微调平台], ] mlb MultiLabelBinarizer() Y mlb.fit_transform(tag_list)MultiLabelBinarizer会把所有类别推进classes_属性里然后给每一条样本生成一个 0/1 向量。注意这个向量里的 0 代表“这个标签不适用”不能把它当成负样本去参与损失计算之外的操作比如采样或过采样时要按真实文本分布来。除非标签噪声特别大否则不要轻易把 0 改成 1那会扭曲任务定义。文本截断也是一个大坑。BERT 的输入序列最大长度通常默认 512如果你的业务文本动不动上千字直接掐头去尾会丢掉关键信息。我一般先统计训练集的长度分布选一个能覆盖 90% 样本的seq_len比如 128 或 256而不是一上来就拉到 512。拉长序列会线性增加计算量我们在第 4 章会专门说 batch size 怎么配合。预处理做完下一步就要进入真正的模型组装和数据 pipeline。3. 用 Keras-bert 在本地跑通 BERT 微调数据 pipeline、模型组装与训练闭环3.1 环境选择和依赖安装先把版本坑占住直接pip install keras-bert很简单真正难的是和 TensorFlow/Keras 的版本组合。Keras-bert 最初面向 Keras 2 和 TensorFlow 1.15 设计如果你跑的是 TF 2.x需要确认你没有混装独立的 Keras 包。常见做法是pip install keras-bert tensorflow1.15 keras2.3.1如果你已经是 TF 2.x可以只装 keras-bert然后检查 Keras 是用的tf.keras还是独立的 Keras。我遇到过最拧巴的情况是模型里一部分层来自keras一部分来自tf.keras训练时 loss 计算直接爆出类型不匹配。后来统一成from keras import ...才安定下来。另外微调 BERT 建议准备一张至少 6GB 显存的显卡。CPU 也能跑但一个 epoch 可能要几小时。如果想在 CPU 上先验证代码能跑通把seq_len调到 64batch_size调到 4先把流程走通再上 GPU这是最省时间的做法。3.2 数据 pipeline从原始文本到两个定长数组Keras-bert 自带一个Tokenizer用来把中文文本转成 token id 序列。首先要加载 vocab 文件。中文 BERT 的 vocab 通常是字符级别一行一个 token。加载方式很简单from keras_bert import Tokenizer token_dict {} with open(vocab_path, r, encodingutf-8) as f: for line in f: token line.strip() token_dict[token] len(token_dict) tokenizer Tokenizer(token_dict)然后对每条文本做 encode。Tokenizer.encode会返回两个数组indices是 token idsegments是句子 id。多标签分类通常只有一条句子segments全部是 0。max_len表示补齐或截断后的最终长度一般和load_trained_model_from_checkpoint里的seq_len保持一致。indices, segments tokenizer.encode( firsttext, max_lenseq_len, paddingpost, truncatingpost, )paddingpost是在尾部补 0truncatingpost是从尾部截断。对大部分业务文本保留开头信息比保留结尾更合理所以我会把 truncating 放在后面。如果你发现专业知识常出现在文本中后段可以改成truncatingpre但一般不建议。接着把全部样本转成 numpy 数组作为模型的输入数据import numpy as np X1 np.zeros((len(texts), seq_len), dtypeint32) X2 np.zeros((len(texts), seq_len), dtypeint32) for i, text in enumerate(texts): ids, segs tokenizer.encode(firsttext, max_lenseq_len) X1[i] ids X2[i] segs注意到数组 dtype 用的是int32。Keras 的 Embedding 层默认接受int32如果你不小心用了int64在 TF 1.x 里通常会频繁报错。这种细枝末节最烦人提前统一能省下不少排错时间。提示如果数据集很大不要一次性把所有样本转成数组堆内存。用一个按 batch 生成的 generator 更好但为了把流程说清楚这里先用全量数组演示。3.3 组装模型并启动训练用 fit 完成一次完整微调现在把模型组合起来。先加载 BERT再在 CLS 输出上接 Dropout 和 Dense最后用Model包装from keras.models import Model from keras.optimizers import Adam bert load_trained_model_from_checkpoint( config_file, checkpoint_file, seq_lenseq_len, output_layer_num1, ) cls_output bert.output[:, 0, :] cls_output Dropout(0.3)(cls_output) logits Dense(num_labels, activationsigmoid)(cls_output) model Model(inputsbert.inputs, outputslogits) model.compile( optimizerAdam(2e-5), lossbinary_crossentropy, metrics[accuracy], )这里bert.inputs是一个列表包含 token id 输入和 segment id 输入。fit时要对应传入[X1, X2]。Adam(2e-5)是 BERT 微调的常见起点详细说明在第 4 章。训练代码非常直接model.fit( [X1, X2], Y, validation_split0.1, batch_size16, epochs3, shuffleTrue, )为什么batch_size16BERT 微调不像普通 CNN它每个样本要过 12 层 Transformer显存开销大batch 太大直接 OOM。16 是一个安全值如果显存还有富余可以抬到 32但通常没必要。epochs 先用 3配合验证集表现再决定要不要继续。训练完成后保存模型和标签编码器model.save(bert_multilabel.h5) # mlb 也要留存预测时候需要 classes_ 还原标签名 import pickle with open(mlb.pkl, wb) as f: pickle.dump(mlb, f)很多新手只保存模型忘记保存MultiLabelBinarizer上线时预测结果全是 0/1根本不知道哪个位置对应哪个标签。这个后悔药提前吃把mlb.classes_和模型一起存。4. 微调 BERT 的四个关键参数学习率、batch_size、epochs 与阈值4.1 学习率为什么 BERT 微调常用 2e-5 而不是 0.001BERT 是预训练模型权重已经很接近一个比较好的局部最优解。如果学习率调到 0.001 甚至更高一步更新就可能把预训练学到的语义分布冲散。常见的微调学习率范围是 2e-5 到 5e-5我习惯先跑2e-5观察训练 loss 走势。如果下降太慢下一个 trial 用 3e-5如果验证集 loss 在第一个 epoch 就起飞立刻降回 1e-5。BERT 微调还有一个细节是 warmup前几步用一个很小的学习率让模型先适应任务数据再逐步升到目标学习率。Keras-bert 自带AdamWarmup如果你不想自己实现学习率调度可以这样用from keras_bert import AdamWarmup optimizer AdamWarmup( decay0.01, warmup_shift-50, lr2e-5, ) model.compile(optimizeroptimizer, lossbinary_crossentropy)warmup_shift参数控制学习率从什么时候开始起跳。这个参数在不同版本里名字可能不一样使用前先看库的文档。如果不确定最稳妥的方案是用 Keras 自带的LearningRateScheduler回调在每个 epoch 开始时手动设置学习率。顺便提一句现在很多新项目转向 LoRA 微调、adapter 微调用小规模参数达到接近全量微调的效果。但在这个多标签任务规模下直接全参数微调更省心也不需要额外封装。只有当模型太大、单张显卡放不下时再考虑 LoRA 这类低成本方案。4.2 batch_size 与显存、梯度更新如何取舍batch_size同时影响显存和梯度质量。BERT 的中间激活值很大一个 batch 的显存开销可以占到总显存的一半。常见设置是 8、16、32。我通常先在训练集上抽样一个 mini batch 试跑一次前向如果 OOM 就减半直到不炸为止。batch 太小的副作用是梯度噪声大训练不稳定。解决办法是梯度累积把几个小 batch 的梯度累加后再更新一次参数效果近似大 batch。在 Keras 里做梯度累积比较绕可以继承Optimizer重写get_updates但项目初期不建议浪费时间。先把 batch 调到能放进显存的最大值不够稳再加梯度累积。多标签分类里标签是否平衡也和 batch_size 有关。如果某个标签出现频率只有 1%一个 batch 16 条里可能一条都没有模型这个 batch 的梯度对那个标签完全没有反馈。这时候可以适当提高 batch_size或者用class_weight后面避坑章节会讲。4.3 epochs 与早停什么时候该让模型停下BERT 微调通常用 3 到 5 个 epoch 就能看到效果续训太多轮反而过拟合。判断过拟合最简单的方式是观察验证 loss前几个 epoch 稳步下降后面开始回升说明模型开始背训练集了。用 EarlyStopping 回调是最稳妥的做法from keras.callbacks import EarlyStopping, ModelCheckpoint early_stop EarlyStopping( monitorval_loss, patience2, restore_best_weightsTrue, ) checkpoint ModelCheckpoint( best_bert_multilabel.h5, monitorval_loss, save_best_onlyTrue, ) model.fit( [X1, X2], Y, validation_split0.1, batch_size16, epochs10, callbacks[early_stop, checkpoint], )patience2表示连续两个 epoch 没有改善就停止训练。restore_best_weightsTrue很关键它会自动回滚到验证 loss 最低的权重相当于买了后悔药。如果你只设置了 EarlyStopping 而没开restore_best_weights最终保存的可能是过拟合后的模型而不是最优模型。epochs 的确定不要只看训练集 loss。多标签分类中某些稀有标签可能在训练集前几个 epoch 里被模型忽略验证集 F1 却显示整体不错。这个现象很坑所以我在验证时会把每个标签的 F1 单独打印出来看到稀有标签一直没有改善就考虑加训练轮数或者做数据增强。4.4 阈值不是 0.5用验证集扫描出每个标签的决策边界sigmoid 输出的概率不一定是校准过的概率。模型对某些标签天然保守输出只有 0.3 也算命中对另一些标签又过度自信0.6 都可能不对。固定用 0.5 切所有标签是单标签思维的多标签翻车点。我的做法是在验证集上对每个标签分别扫描阈值from sklearn.metrics import f1_score best_thresholds [] for label_idx in range(num_labels): best_score 0.0 best_thr 0.5 for thr in np.arange(0.2, 0.8, 0.05): y_pred (pred_val[:, label_idx] thr).astype(int) score f1_score(Y_val[:, label_idx], y_pred, zero_division0) if score best_score: best_score score best_thr thr best_thresholds.append(best_thr)pred_val是验证集预测概率。扫描范围 0.2 到 0.8 足够覆盖多数情况。如果某个标签在 0.5 附近分数很低到 0.3 反而最高说明模型对它的置信度整体偏低可能训练样本太少或特征太隐晦。这个结论能指导你做数据分析而不是只调模型。5. 避坑指南Keras-bert 微调多标签分类的 5 个常见翻车现场5.1 翻车现场一加载 checkpoint 后输出层维度对不上现象把bert.output直接接到 Dense 层Keras 报错Shape must be rank 2 but is rank 3。原因bert.output是三维张量(batch, seq_len, hidden)Dense 不能直接处理序列维度。解决先取 CLS 向量或用池化层压缩维度。CLS 向量就是bert.output[:, 0, :]。取完后维度变成(batch, hidden)。如果是用GlobalMaxPooling1D()需要在调用时明确data_formatchannels_last否则可能把 batch 维度池化掉训练时直接炸。这个错误几乎每个人都遇过看到报错先检查池化。5.2 翻车现场二train loss 不降或直接 NaN现象训练第一个 epoch loss 就在几百甚至 NaN后续也降不下来。原因最常见是学习率过大。BERT 预训练权重本身很敏感用一个在 CNN 上习惯的 0.001 学习率几轮更新后数值直接爆炸。另一个原因是文本里有无法映射的 tokentokenizer.encode时某些字符不在 vocab 中导致 token id 为 OOV 的未知值模型学不到有效信息。解决把学习率降到 2e-5 级别同时检查 vocab 是否覆盖了所有业务字符。中文标点、特殊符号很容易漏。构建 tokenizer 前可以先做一次全量文本的字符统计把不在 vocab 里的字符提前清洗或替换成[UNK]。另外loss 变成 NaN 时去查输入数据里有没有np.nan填充的序列这种隐蔽错误也会导致梯度异常。5.3 翻车现场三验证集 acc 很高F1 却很难看现象训练日志里 validation accuracy 到了 0.95线上效果却一塌糊涂多标签命中率低。原因多标签数据往往严重不平衡。例如 100 个标签里最热的标签出现频率 80%最冷的只有 0.1%。模型只要把所有样本的冷门标签都预测为 0accuracy 也能很高因为大多数位置本来就是 0。解决不要再盯着 accuracy改成多标签的 F1并且按标签分组观察。重点关注每个标签的 recall 而不是整体 accuracy。如果某个冷门标签始终预测不出来考虑增加该类样本、使用class_weight或做简单的文本复制增强。模型不会因为一个标签样本少就自动学会它需要人为干预。5.4 翻车现场四batch_size 降到底还是 OOM现象batch_size2都报显存不足模型根本没法训练。原因seq_len太长是最大元凶。BERT 显存占用和序列长度基本成正比如果你把seq_len设成 512即使 batch 只有 2也可能把 6GB 显存吃满。另一个原因是加载模型时没有释放之前的模型图多个 Graph 同时存在。解决先统计训练集长度把seq_len压到覆盖 90% 样本的长度。多标签分类里很多短文本不需要 512 的长度128 一般够用。另外在 notebook 里反复运行模型组装代码时用keras.backend.clear_session()清理旧图再建立新模型能释放不少显存碎片。如果显存还是不够可以把batch_size降到 4然后使用混合精度训练。Keras 里可以通过tf.keras.mixed_precision设置 float16 计算但和 keras-bert 的兼容性要看版本项目紧急时不如直接换更大显存的机器。5.5 翻车现场五class_weight 传进 fit 后训练更不稳定现象按单标签的思路给class_weight传一个字典训练 loss 忽高忽低甚至直接报错。原因多标签任务的每个样本可能对应多个标签Keras 的class_weight原理是按样本类别加权通常会期望每个样本只属于一个类别。多标签环境下同一个样本对不同标签可能要不同权重这个逻辑没法用简单字典表达。硬传上去模型会对某些输出节点施加过大的梯度更新直接破坏 sigmoid 输出的平衡。解决不用class_weight改用sample_weight。先为每个样本计算一个权重例如样本包含冷门标签就给更高权重把权重数组传给model.fit(..., sample_weightweights)。如果嫌麻烦更简单的做法是过采样包含冷门标签的原始文本让模型天然看到更多冷门样本。我在实际项目里优先做过采样因为它不改变损失函数语义训练也更稳定。6. 从训练到上线用 F1、阈值扫描和样本复盘验证多标签模型训练结束后不要只看训练日志里最后一行 loss。我习惯先把验证集预测结果保存下来用第 4.4 节的阈值扫描拿到一组最佳阈值然后生成完整的多标签分类报告from sklearn.metrics import classification_report pred_val model.predict([X1_val, X2_val]) y_pred (pred_val best_thresholds).astype(int) print(classification_report(Y_val, y_pred, target_namesmlb.classes_, zero_division0))classification_report会打印每个标签的 precision、recall 和 F1。这个表是判断模型能不能上线的核心依据。我会先看 macro F1再看最少样本的那几个标签。如果冷门标签 F1 全是 0线上也会是 0因为模型没见过足够特征这个不是玄学是数据分布决定的。光看指标还不够我会随机抽取一批验证样本把真实标签和预测标签并排打出来一条条看。模型常犯的毛病只有看具体样本才能发现比如“人工智能”标签经常被预测在“算法”上可能是两个标签在训练集里经常共现模型把共现当成了因果。这种问题要么增加数据要么人工规则兜底指望调参解决不现实。最后一步是用独立的测试集跑一次最终评估而不是沿用验证集。因为验证集已经被用来调阈值了模型和阈值都可能对验证集“过度适应”。拿另一个没见过的测试集跑一遍得到的 F1 才是线上预期值。我自己的习惯是先跑完验证调参再用测试集只测一次不做任何回环避免把测试集也变成训练的一部分。这个流程坚持下来后多标签分类项目很少会翻大车。希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑