尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

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

BERT微调实现多标签文本分类的Keras实战指南 简介基于Keras与Keras-bert的文本多标签分类项目包面向自然语言处理实战场景通过微调BERT完成多标签分类并以2020语言与智能技术竞赛事件抽取任务作为数据样例适合需要快速落地预训练模型的开发者和研究者。压缩包大小约1.01MB共10个文件包含4个Python脚本分别实现训练、评估、预测与FGM对抗训练2个CSV文件提供训练集和测试集2个TXT文件给出中文预训练词汇表与依赖清单另有README说明和gitignore目录结构简洁清晰便于按模块复用。目前已有1634人学习下载。借助该资源可完整体验文本多标签分类流程从数据预处理、模型微调、效果评估到单条文本预测均有代码支撑FGM对抗训练有助于提升泛化能力配套的中文词汇表与依赖文件降低了环境配置门槛非常适合课程设计、算法竞赛或工程验证前的快速试验也方便二次开发与定制。1. 多标签文本分类不是多分类为什么BERT微调是当前最稳的解法做文本分类做到一定阶段都会撞上同一个需求一条文本不再只属于一个类别。比如电商客服工单一条投诉可能同时命中“物流延迟”和“退换货流程”再比如新闻打标一篇稿件既可以归入“财经”也可以归入“宏观政策”。这种任务在工业界叫多标签分类multi-label classification它不是多分类的简单变体——多分类用softmax输出一个概率分布多标签则要为每个类别独立做“是或否”的二值判断。传统的做法有很多先把文本用TF-IDF或Word2Vec转成向量再叠一个多输出分类器但这类做法有两个天花板——一是特征表达弱同义不同形、长距离依赖都抓不住二是标签之间的相关性很难建模。于是BERT一出来文本分类基本就统一到了“预训练微调”这条路上拿一个在大规模语料上训练好的BERT模型用你的业务数据做监督学习更新权重。这个方案对多标签任务尤其友好因为BERT输出的每个token表示都携带上下文语义拿[CLS]向量往下游接一个多标签输出层效果比传统方法稳定得多。本文要讲的这套Keras实现正是沿着这条线走的用Keras-bert加载BERT预训练权重在输出层改为“多个sigmoid 二元交叉熵”对BERT做微调。这套方案适合手里有几千到几万条标注数据、想把文本多标签分类落到工程上的团队。接下来从数据准备、模型构建到踩坑记录逐步展开。2. 数据准备从原始文本到BERT能吃的输入张量数据准备这一步多标签任务比单标签任务要麻烦一些。倒不是代码复杂而是标签的形态多种多样处理不好后面整个流程都会跟着错。我一般会在写模型之前先花时间把数据的标签结构和分布摸清楚再动手写预处理脚本。2.1 多标签数据的三种常见落地格式实际项目里拿到的多标签数据常见的有三种格式。第一种是CSV文件里标签字段用逗号分隔一行一条样本例如“文本内容,标签1,标签2,标签3”。第二种是JSON行每条样本是个JSON对象里面有一个数组字段专门存放标签列表。第三种比较少见数据源直接把标签展开了每个标签一列是0/1值。在做统一处理时我习惯先把不同来源的数据统一成一个中间结构每条样本包括文本字段和一个Python列表类型的标签字段。这样无论后面接什么模型数据处理逻辑都只需要写一份。import pandas as pd # CSV格式标签列是逗号分隔的字符串 df pd.read_csv(samples.csv) df[labels] df[labels].apply(lambda x: x.split(,)) # JSON行格式每条记录自带labels数组 # records格式: [{text: ..., labels: [金融, 政策]}, ...] # 统一之后只保留两个字段 df df[[text, labels]]这段代码做的事很简单把CSV里的字符串标签切分成列表和JSON格式统一起来。后面所有流程都基于“text字段是字符串labels字段是字符串列表”这个约定。参数上要注意split的分隔符有些导出工具会用中文逗号、分号或多级分隔符需要先确认原始数据的分隔符是什么不要拿到就按英文逗号切。2.2 标签二值化与训练/验证集的正确拆分多标签任务在进入模型之前需要把标签列表转成固定维度的0/1向量这个操作在sklearn里对应MultiLabelBinarizer。它会把数据集里出现过的所有标签收集起来每个标签映射到一列样本包含该标签就在对应位置记1否则记0。这一步有两个非常容易踩坑的地方。第一个坑是fit操作必须只基于训练集不能把验证集甚至测试集的标签一起拿进来fit。理由很实际推理环境里你不可能预先知道未来数据的全部标签集合。第二个坑是验证集里出现的标签如果训练集里没见过在transform的时候会被直接忽略不会报错但模型效果评估会失真。from sklearn.preprocessing import MultiLabelBinarizer from sklearn.model_selection import train_test_split df_train, df_val train_test_split(df, test_size0.2, random_state42) mlb MultiLabelBinarizer() # 只用训练集fit y_train mlb.fit_transform(df_train[labels]) # 验证集只做transform不fit y_val mlb.transform(df_val[labels]) # 类别列表后面输出层神经元数量靠它 label_classes mlb.classes_ num_labels len(label_classes) print(f标签总数: {num_labels}, 训练集样本数: {len(y_train)}, 验证集样本数: {len(y_val)})代码里train_test_split是随机切分这在大多数场景够用。但如果数据带时间戳比如工单是按时间产生的更合适的做法是按时间先后切分避免随机切分把同一条业务流的数据同时分到训练和验证导致验证结果虚高。MultiLabelBinarizer的classes_属性是按字典序排列的这一点要记住后面模型输出向量的顺序就和它保持一致不然预测结果和标签名对不上。2.3 BERT分词与输入编码长度截断和segment_idBERT不能直接吃原始字符串它的输入是经过WordPiece分词后的token编号序列。这一步在Keras-bert里由Tokenizer类完成加载的时候需要传入BERT的词典文件vocab.txt。分词之后每个样本得到两组数组input_ids是token在词典中的编号segment_ids用来区分句子A和句子B。单条文本分类只需要一个句子segment_ids全为0即可。Transformer的输入序列长度上限通常是512但实际项目中很少有人直接用512因为序列越长显存占用和计算时间增长很快。对于绝大多数短文本分类场景128已经够了。长尾的长文本则需要做截断策略直接截断尾部可能导致关键信息丢失常见做法是保头保尾但Keras-bert的Tokenizer不直接支持这种截断需要自己处理。from keras_bert import Tokenizer vocab_path chinese_L-12_H-768_A-12/vocab.txt tokenizer Tokenizer(vocab_path) MAX_LEN 128 def encode_text(text): # 手动截断保留前MAX_LEN-2个token预留[CLS]和[SEP] indices, segments tokenizer.encode( text, max_lenMAX_LEN, truncate_to_max_lenTrue ) return indices, segments # 编码训练集 import numpy as np X_ids np.zeros((len(df_train), MAX_LEN), dtypeint) X_segs np.zeros((len(df_train), MAX_LEN), dtypeint) for i, text in enumerate(df_train[text]): ids, segs encode_text(text) X_ids[i] ids X_segs[i] segstokenizer.encode返回的indices已经包含了开头的[CLS]和结尾的[SEP]两个特殊token所以设置max_len等于128时实际有效的文本token只有126个。这个细节很多人会忽略导致最终输入长度和你预期的不一致。truncate_to_max_len参数控制是否在超长时截断而不是报错。编码完成后得到的X_ids和X_segs就是后续模型的输入shape是(样本数, 128)数据类型是int。3. 用Keras-bert搭微调模型从加载权重到多标签输出层数据准备好了接下来是模型部分。环境配置和预训练参数下载是第一个拦路虎搭模型本身反而是最公式化的一步。Keras-bert在Keras生态里用起来很顺手但它对TensorFlow版本比较敏感版本错位会带来一堆玄学报错。3.1 Keras-bert环境配置安装顺序和预训练参数下载Keras-bert是一个把BERT封装成Keras层的开源库使用方式很贴合Keras的习惯安装直接用pip。环境上要注意的是它基于Keras 2.x开发如果你用的是TensorFlow 2.x装的时候尽量让Keras-bert和tf.keras共享同一套Keras后端避免出现两套Keras实例各管各的怪问题。pip install keras2.11.0 pip install tensorflow2.11.0 pip install keras-bert安装顺序有讲究先装Keras和TensorFlow再装keras-bert。这样做的好处是keras-bert安装时检测到的依赖环境已经稳定不会反过来把Keras版本升级到3.x。Keras 3的API和Keras-bert不是完全兼容的我就见过有人装完keras-bert后Keras被自动升到3.x模型加载时报一堆“找不到层”的错误。预训练参数下载推荐BERT官方的中文预训练模型chinese_L-12_H-768_A-12。下载后是一个压缩包解压出来关键文件有四个我习惯把这些文件放到一个固定目录里文件名作用bert_config.json模型结构配置如层数、隐藏层维度、注意力头数bert_model.ckpt.data-00000-of-00001模型权重数据bert_model.ckpt.index权重索引配合data文件一起加载vocab.txt分词词典Tokenizer加载用3.2 构建模型BERT主干 多标签输出层Keras-bert提供了load_trained_model_from_checkpoint函数传入config文件和checkpoint文件路径即可加载预训练模型。值得注意的一点是checkpoint路径要传前缀就是bert_model.ckpt这一级不需要带.data或.index后缀。加载回来的模型输出是所有位置最后一层的隐层向量。文本分类任务通常只取[CLS]位置也就是序列的第一个token的向量把它接一个Dropout层和一个Dense层。Dense层的神经元数量等于标签数量激活函数用sigmoid而不是softmax因为每个标签是一个独立的二分类问题所有标签的预测概率加起来不需要等于1。from keras.layers import Dense, Dropout, Lambda from keras.models import Model from keras_bert import load_trained_model_from_checkpoint config_path chinese_L-12_H-768_A-12/bert_config.json checkpoint_path chinese_L-12_H-768_A-12/bert_model.ckpt # 加载预训练BERT bert_model load_trained_model_from_checkpoint( config_path, checkpoint_path, seq_len128, output_layer_num1 ) # 冻结参数先只训练下游层 for layer in bert_model.layers: layer.trainable False # 取CLS向量 cls_output Lambda(lambda x: x[:, 0, :], namecls_token)(bert_model.output) dropout Dropout(rate0.3)(cls_output) output Dense(num_labels, activationsigmoid, namemulti_label_output)(dropout) model Model(inputsbert_model.input, outputsoutput) model.summary()seq_len参数必须和前面数据处理时的MAX_LEN保持一致如果数据编码用的是128而这里设置了256模型输入shape对不上直接报错。output_layer_num设为1表示取最后一层输出这是微调最常用的配置。先把BERT层冻结住训练几轮等下游分类层收敛得差不多了再解开全部权重做完整微调这种做法在数据量不大时能明显减少稳定性问题。3.3 训练方案学习率、早停、checkpoint怎么设多标签分类的损失函数和编译配置很固定损失用binary_crossentropy优化器用Adam。但学习率设置是微调BERT最关键的超参数预训练模型已经收敛到了一个比较好的状态学习率太大一步就把权重冲坏了太小又训不动。常见取值范围在2e-5到5e-5之间我一般先跑2e-5如果loss下降太慢再放到5e-5。回调函数里最有用的是EarlyStopping和ModelCheckpoint。EarlyStopping监控验证集loss连续几个epoch不下降就停止训练ModelCheckpoint则负责保存验证集最优的权重防止后面过拟合把模型训废了。from tensorflow.keras.optimizers import Adam from tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint model.compile( optimizerAdam(learning_rate2e-5), lossbinary_crossentropy, metrics[accuracy] ) callbacks [ EarlyStopping( monitorval_loss, patience3, restore_best_weightsTrue ), ModelCheckpoint( best_model.h5, monitorval_loss, save_best_onlyTrue, save_weights_onlyFalse ) ] history model.fit( x[X_ids, X_segs], yy_train, validation_data([X_val_ids, X_val_segs], y_val), batch_size32, epochs10, callbackscallbacks )注意fit的输入x是一个列表包含两个numpy数组顺序对应BERT的两个输入input_ids和segment_ids。顺序不能反过来否则模型结构对不上训练过程会莫名loss不降。batch_size的选择和显存直接相关。BERT的显存占用随序列长度和batch_size线性增长12层768维的中文BERT在8GB显存上batch_size设32序列长度128是可以跑的再大就可能OOM。4. BERT微调避坑记录训练崩了、指标虚高、标签全判成一个类BERT微调我已经做过好几轮每次翻车的原因翻来覆去就那么几个。这一章把高频的问题直接列出来按“现象、原因、解决”写清楚看完能省去大量排查时间。4.1 预训练权重路径没指对模型加载完却一直在用随机初始化现象训练loss从一开始就很高而且下降极其缓慢无论怎么调学习率都没用最后的指标远低于同类任务应该有的水平。原因load_trained_model_from_checkpoint传的checkpoint_path是“bert_model.ckpt”这个前缀但实际上文件是三个。如果路径里拼写错误或目录层级不对Keras-bert在加载时并不会报一个明确的“文件未找到”错误而是可能静默地创建了一个随机初始化的BERT模型继续往下走。这个表现极具迷惑性因为训练流程完全是通的。解决在加载权重之后做一次快速的验证。构造两条明显不同语义的文本比如“今天天气很好”和“这个产品完全坏了”用加载后的模型分别拿取[CLS]向量打印出来看是否明显不同。如果两个向量几乎一样那大概率权重没加载进去。更严谨的做法是检查模型权重数值随机挑一层打印权重均值预训练权重的数值分布通常是稳定的小数值范围而随机初始化会呈现明显不同的偏置。4.2 训练用argmax评估、预测用sigmoid卡阈值准确率永远对不上现象训练时手动写评估代码用np.argmax判断预测标签验证集准确率看上去还不错但真实业务里用模型预测结果完全不可用要么每个样本都预测出一堆标签要么一个标签都出不来。原因np.argmax只取概率最大的那个位置作为1其余全是0把它当成多分类来处理了。而多标签任务的正确解读方式是模型输出层每个位置的值代表该标签独立存在的概率最终标签集合取决于阈值卡在哪里。argmax会把多标签压成单标签看起来指标不错实际上丢掉大量弱信号标签。解决评估和预测阶段统一用sigmoid输出的概率值然后对每个位置单独判断是否大于阈值。阈值默认0.5只是起点实际项目中往往需要根据业务对准确率和召回率的倾向做调整。这个我一般放在验证集上做阈值网格搜索具体方法见下一章。4.3 序列长度截断太激进长文本的关键标签被切没了现象短文本的预测效果不错但一到长文本就漏标签尤其是那些只在文档中后部出现的观点或事件标签几乎全断。进一步定位发现漏掉的内容集中在文本后半段。原因预处理时对超过128长度的文本直接从尾部截断长文本的关键信息分布在全文各处尾部截断等于把这些信息直接丢弃。BERT有512的长度上限但计算成本高很多人图省事统一截到128结果长文本各种漏信号。解决先统计训练集的文本长度分布看90分位在哪里。如果大部分文本确实在几百字以内把MAX_LEN提到256或512一般就能覆盖。如果文本特别长比如一篇文章几千字单靠BERT硬切就不合适了常见做法是先用摘要提取或关键句选择降文本压缩到BERT可接受的长度再进入分类模型。不做分析就拍脑袋定长度是这类问题反复出现的根源。4.4 多数类压倒少数类模型把多标签任务做成了“只学热门标签”现象训练结束后统计每个标签在验证集上的预测频率发现高频标签被大量预测低频标签几乎从不出现。整体F1看上去还行但拆到每个标签上一看稀有标签的召回率接近0。原因多标签数据天然存在标签不平衡有的标签出现率超过80%有的标签只有5%。binary_crossentropy对每个标签独立计算损失多数类的梯度贡献大模型自然倾向把所有样本都预测为“无稀有标签”。解决粗调阶段先不处理等模型收敛后再看每个标签的混淆矩阵。如果确认是这个问题有多种办法给稀有标签的损失加权或者在采样时保证每个batch里稀有标签样本不缺席。损失加权在Keras里可以直接给compile传class_weight参数需要先统计每个标签的正样本比例。不过这个方案要谨慎权重设太大会把正常样本的预测全带偏一般从1.5到3倍之间试。4.5 Keras和TensorFlow版本错位Keras-bert在TF2.x下的玄学报错现象环境装好之后import keras_bert本身就报错或者模型加载时出现“name get_custom_objects is not defined”这类信息。原因keras-bert编写时主要针对Keras 2.x的API而Keras 3或者高版本TensorFlow尤其是TF 2.16之后对Keras内部结构做了调整某些旧API被移除或改名。如果你通过pip install keras-bert时没有锁定版本pip很可能自动装了一个兼容性有问题的Keras版本。解决最直接的办法是把环境锁到keras 2.x和对应的TensorFlow 2.x版本。另一个思路是检查项目是否允许不用keras-bert直接改用transformers库的TFBertForSequenceClassificationAPI不同但背后的微调原理完全一致。如果必须用keras-bert的代码可以把keras降到2.11.0并确认TensorFlow版本与它兼容。这条属于环境配置问题问题本身不难难在报错信息往往不直接指向版本排查起来费时间。5. 最后一公里验证集上搜最优点位阈值让模型真正可用模型训练完、权重保存好很多项目到这里就收工了。但真实的业务埋点在这sigmoid输出的概率值怎么转成最终标签集合阈值取多少直接决定线上效果。默认阈值0.5其实只是一个平庸的起点尤其是标签分布极不平衡时。有些标签哪怕正样本很少模型给出的概率区间整体偏低用0.5卡会把所有样本都判成负例而另一些标签模型过度自信概率集中在0.9以上0.5的阈值又太宽松。所以我在每次微调之后都会做一遍阈值网格搜索。import numpy as np from sklearn.metrics import f1_score # 验证集概率输出 probs model.predict([X_val_ids, X_val_segs], batch_size32) def search_threshold_for_label(probs, y_val, label_idx): best_threshold 0.5 best_f1 0.0 for threshold in np.arange(0.2, 0.9, 0.05): pred_label (probs[:, label_idx] threshold).astype(int) real_label y_val[:, label_idx] f1 f1_score(real_label, pred_label, zero_division0) if f1 best_f1: best_f1 f1 best_threshold threshold return best_threshold, best_f1 label_thresholds {} for idx, label in enumerate(mlb.classes_): thresh, f1 search_threshold_for_label(probs, y_val, idx) label_thresholds[label] thresh print(f标签: {label}, 最优阈值: {thresh:.2f}, F1: {f1:.4f})这段代码对每个标签独立做阈值搜索在0.2到0.9之间按0.05的步长找最优值。search_threshold_for_label的参数y_val是二值化后的标签矩阵label_idx是对应标签的列索引。threshold的搜索范围和步长可以按需调整数据量大的时候0.01的步长会更细但计算成本也高0.05足够工程用了。每个标签得到独立阈值后推理阶段的逻辑就变成模型输出概率向量逐标签和对应的阈值比较大于阈值置1否则置0。这套做法做下来整体F1一般会有几个点的提升在标签分布极不平衡时尤为明显。额外提一句多标签分类的评估不要只用整体accuracy更合理的指标是micro-F1或macro-F1。涉及的情况总归是复杂的但有一点我习惯一直保持保存模型的同时把mlb.classes_、label_thresholds、词表和MAX_LEN一起归档这几个文件缺任何一个推理环境都跑不起来。这套流程多跑几遍之后会稳定很多希望你也能在自己的数据上顺利复现希望帮到你。本文还有配套的精品资源点击获取
返回列表