Keras-bert实战:用BERT微调搞定多标签文本分类
发布时间:2026/10/1 3:03:59来源:尧图网络
简介一套基于Keras与Keras-bert的文本多标签分类实践项目面向有一定NLP基础、希望快速上手BERT微调的开发者。项目以2020语言与智能技术竞赛的事件抽取数据为样例展示了多标签分类模型的完整搭建与应用思路适合作为课程设计或项目参考。压缩包共10个文件大小约1.01MB包含4个Python脚本模型训练、评估、预测及FGM对抗训练、2个CSV数据集训练集与测试集、2个TXT文件BERT中文词表与依赖清单另有README说明与gitignore配置结构清晰便于按需取用。已有1634人学习下载。借助这份资源可掌握Keras-bert加载预训练模型、微调BERT并完成多标签文本分类的核心流程同时可借鉴竞赛数据的预处理方式与对抗训练技巧节省自行整理数据和排查环境依赖的时间。1. 多标签文本分类为什么要看得上Keras-bert这条路多标签文本分类往往被当成多分类来写但真实项目里一条新闻可以同时是“财经民生”一个工单可以同时命中“网络故障”和“退款投诉”这跟二选一的单标签任务完全是两种学习问题。用Keras和Keras-bert把BERT预训练权重加载进来做微调是这类任务里启动成本最低、调参路径最清晰的方案之一你不需要自己训练词向量也不需要从头搭Transformer只要把文本编码成BERT认识的输入接一个多标签输出层用两阶段微调把参数调顺在几千条标注数据上就能拿到一个可上线的分类服务。这篇笔记写给想在项目里快速跑通文本多标签分类的工程师也写给被各种大模型微调工具绕花眼、想走一条确定路径的新手。我们这次只把Keras-bert这条路走通把参数和坑讲透。2. 把多标签任务翻译成BERT能学的东西数据预处理与编码微调BERT之前最容易被低估的是数据准备。文本多标签分类的标签不是“第几类”而是“命中哪几个类”这一点会直接影响输出层设计、损失函数和评估指标。如果一开始就按多分类的思路处理后面每一步都会跟着错。2.1 多标签不是多分类sigmoid输出、multi-hot标签和F1评估多分类任务里一个样本只能属于一个类别标签是one-hot向量输出层用softmax损失函数用categorical_crossentropy。多标签任务里一个样本可以命中多个类别标签是multi-hot向量比如“网络故障退款投诉”对应的标签向量是[1, 0, 1, 0]输出层必须用sigmoid让每个输出节点独立判断“这个标签是否命中”损失函数用binary_crossentropy。这里有一个常见的翻车点直接把多标签任务当成多分类最后softmax把所有标签概率归一化成总和1模型被迫在“网络故障”和“退款投诉”之间二选一验证集F1自然上不去。先把标签编码这一步做对from sklearn.preprocessing import MultiLabelBinarizer label_list [ [网络故障, 退款投诉], [网络故障], [账户问题, 退款投诉], ] mlb MultiLabelBinarizer() y mlb.fit_transform(label_list) print(y.shape) # (3, 3) print(mlb.classes_) # [账户问题 网络故障 退款投诉]这里MultiLabelBinarizer会把每条样本的标签列表展开成固定列数的multi-hot矩阵列顺序由classes_决定。训练时y直接喂给sigmoid输出层预测时输出的是每个标签的概率再用阈值决定是否命中。评估指标也别沿用accuracy。多标签场景下accuracy指“所有标签全部预测正确才算对”在标签数多、样本稀疏时几乎永远是0。我一般用macro F1和每个标签单独的precision/recall来评估线上对每条样本要求严格一点时才看精确匹配率。还有一点要注意如果某个标签在训练集里只出现十几次最好在验证集里单独看一眼它的F1别被整体平均掩盖掉。2.2 BERT预训练参数与词表准备keras-bert要从目录里读什么用keras-bert微调需要先拿到一份中文BERT预训练参数。常见做法是下载开源的中文BERT模型解压后目录里通常包含三类文件vocab.txt是词表bert_config.json是模型结构配置bert_model.ckpt或对应的weights文件是预训练权重。keras-bert加载时词表和权重是分开用的词表交给Tokenizer权重目录或文件路径交给load_bet_model。from keras_bert import load_vocabulary, Tokenizer, load_bet_model vocab_path chinese_bert_wwm/vocab.txt bert_ckpt_path chinese_bert_wwm/bert_model.ckpt token_dict load_vocabulary(vocab_path) tokenizer Tokenizer(token_dict) bert_model load_bet_model( bert_ckpt_path, seq_len128, trainableFalse )load_bet_model的第一个参数可以直接传.ckpt文件路径某些版本的keras-bert也支持传包含权重文件的目录。seq_len是我们自己定的最大序列长度BERT内部最多支持512个token实际分类任务通常取64到256。trainableFalse表示加载预训练权重时先把BERT层全部冻结这个开关在二阶段微调里很有用后面会展开说。加载完成后建议先打印一下模型结构和输入名确认输入顺序。keras-bert的BERT模型输入是两个张量第一个是token ids第二个是segment ids。很多初学者在这里把输入顺序搞反训练时loss不降还以为是模型问题。print(bert_model.inputs)一下看到类似[Input-Token, Input-Segment]的顺序就稳了。2.3 用Keras-bert的Tokenizer编码文本从可变长句子到定长矩阵BERT不能直接吃原始字符串需要先把文本转成token ids和segment ids。keras-bert的Tokenizer.encode(text, max_lenseq_len)会完成切词、添加[CLS]和[SEP]、截断和padding返回两个等长数组ids, segs tokenizer.encode(账户无法登录想退款, max_len32) print(ids[:8]) # [101, 1908, 3359, ... , 102, 0, 0, 0] print(segs[:8]) # [0, 0, 0, 0, 0, 0, 0, 0]ids里101是[CLS]的id102是[SEP]的id长度不足32的部分用0补齐。segments在单句分类任务里全部是0表示整段文本都属于句子A。如果做句对任务第二个句子的segment id会是1但文本分类用不到。训练时不可能一条条调用encode那样太慢。我一般写一个生成器在batch内部批量编码喂给model.fitimport numpy as np def build_data_generator(texts, labels, tokenizer, seq_len, batch_size32): n len(texts) order np.arange(n) def gen(): while True: np.random.shuffle(order) for start in range(0, n, batch_size): batch_idx order[start:start batch_size] x1 np.zeros((len(batch_idx), seq_len), dtypenp.int32) x2 np.zeros((len(batch_idx), seq_len), dtypenp.int32) for i, idx in enumerate(batch_idx): ids, segs tokenizer.encode(texts[idx], max_lenseq_len) x1[i, :len(ids)] ids x2[i, :len(segs)] segs yield [x1, x2], labels[batch_idx] return gen()这个生成器有几个关键细节。x1和x2是int32类型的定长矩阵形状是(batch_size, seq_len)每个token id不能超过BERT词表范围所以dtype用int32足够。生成器每次shuffle一次顺序然后按batch yieldfit时指定steps_per_epoch。这里有一个容易被忽略的点如果语料里存在超长文本tokenizer.encode内部已经做了截断不需要我们手动处理但截断会丢掉后半段信息。如果业务上后半段同样关键比如合同文本的条款集中在末尾就要考虑分段编码或者改用更长上下文的模型这是另一个话题了。多标签分类的数据集划分也有讲究。sklearn的train_test_split直接支持stratify但多标签的multi-hot矩阵不适合直接做分层抽样因为一条样本同时属于多个类无法简单映射到单一类别。我一般先按业务上最主要的一个标签分层或者干脆按时间顺序切分保证验证集和线上数据的分布更接近。3. 用Keras和Keras-bert搭建微调模型冻结、池化、接多标签输出层数据准备好了接下来是模型搭建。微调这个词听起来玄学本质就是把BERT已经学到的语言表示拿过来在它的基础上加一层很薄的分类头然后用你的业务数据反向传播更新参数。BERT部分可以整体微调也可以部分冻结这就是我们说的微调策略。3.1 load_bet_model返回的到底是什么输入张量、输出shape与可训练开关load_bet_model返回的是一个Keras Model对象输入是[Input-Token, Input-Segment]输出是BERT最后一层的序列特征。以BERT-Base为例hidden size是768所以bert_model.output的形状是(None, seq_len, 768)也就是每个token位置都有一个768维的向量。很多教程会直接取[CLS]位置的向量作为整句话的表示但多标签场景下我会更推荐先看序列输出自己做池化而不是直接用BERT内部的NSP池化层。原因很简单BERT的[CLS]向量在预训练时主要服务于“两句是否连续”的任务它对整句话的语义有代表性但对多标签这种“一句话里同时存在多个主题”的情况序列平均池化往往更稳。from keras.layers import Dense, Dropout, GlobalAveragePooling1D from keras.models import Model from keras.optimizers import Adam from keras.metrics import AUC seq_len 128 num_labels len(mlb.classes_) bert_model load_bet_model(bert_ckpt_path, seq_lenseq_len, trainableFalse) seq_output bert_model.output # (None, seq_len, 768) pooled GlobalAveragePooling1D()(seq_output) # (None, 768) pooled Dropout(0.3)(pooled) output Dense(num_labels, activationsigmoid)(pooled) model Model(bert_model.inputs, output) model.compile( optimizerAdam(learning_rate2e-5), lossbinary_crossentropy, metrics[AUC(nameauc)] ) model.summary()这里GlobalAveragePooling1D会对seq_len维做平均把每个token位置的768维向量压成一个768维的句向量。Dropout(0.3)是防止分类头过拟合多标签任务里标签多、正负样本不均Dropout加在池化之后、输出层之前是标准做法。输出层激活函数必须用sigmoid每个节点独立输出0到1之间的概率。为什么多标签不用softmaxsoftmax强制所有标签概率加起来等于1而多标签的语义是每个标签独立参与判定。一个工单可以既是“网络故障”又是“退款投诉”这两个概率应该可以同时高也可以同时低。sigmoid给每个标签独立的概率配合binary_crossentropy模型才能学到这种“多选多”的分布。3.2 多标签分类头怎么接平均池化、Dropout和输出层设计分类头虽然只有几层但接法直接决定模型能不能收敛。除了GlobalAveragePooling1D还可以用GlobalMaxPooling1D或者直接取CLS位置。from keras.layers import Lambda # 取CLS位置向量shape从(None, seq_len, 768)变为(None, 768) cls_output Lambda(lambda x: x[:, 0, :])(seq_output)这两种接法我都试过。取CLS在单标签情感分类上表现很好但在多标签场景里如果一句话同时涉及两个主题CLS向量容易偏向其中占主导的主题另一个主题的特征被稀释。平均池化把所有token的信息均匀揉进去对“多主题并存”的文本更友好。如果你的标签有明显的强相关结构比如“A出现时B大概率也出现”也可以尝试max pooling它能抓住最强烈的信号。没有绝对优劣我在项目里的习惯是先平均池化跑一版基线再花半天做CLS和max的对比实验。还有一个细节分类头初始化的权重范围。Dense层默认的glorot_uniform初始化对BERT输出这种已经归一化的特征通常是合适的不需要额外改动。真正要改的是Dropout比例如果训练集只有两三千条Dropout可以加到0.4如果数据过万0.2到0.3就够。多标签任务里标签之间的共现信息很重要Dropout太大会把共现关系打散。3.3 二阶段微调流程先训分类头再解冻BERT为什么一定要分两步直接加载BERT后立即全部解冻、用一个比较大的学习率训练loss很容易在第一个epoch就冲上天。这是因为BERT的预训练权重已经在语言模型任务上收敛得很好了但新加的Dense分类头是随机初始化的两者参数尺度不匹配。随机初始化的分类头在初期会输出乱七八糟的梯度如果BERT主干的权重也跟着一起大步更新就会把预训练学到的语言表示冲坏这属于典型的“灾难性遗忘”。所以二阶段微调几乎是BERT文本分类项目的标准做法。第一阶段冻结BERT全部层只训练分类头让分类头先适配BERT的输出特征第二阶段再解冻BERT层用一个很小的学习率整体微调。# 阶段一只训练分类头 for layer in bert_model.layers: layer.trainable False model.compile( optimizerAdam(learning_rate1e-3), lossbinary_crossentropy, metrics[AUC(nameauc)] ) train_gen build_data_generator(train_texts, train_labels, tokenizer, seq_len, batch_size32) val_gen build_data_generator(val_texts, val_labels, tokenizer, seq_len, batch_size32) model.fit( train_gen, steps_per_epochlen(train_texts) // 32, validation_dataval_gen, validation_stepslen(val_texts) // 32, epochs3 ) # 阶段二解冻BERT统一微调 for layer in bert_model.layers: layer.trainable True model.compile( optimizerAdam(learning_rate2e-5), lossbinary_crossentropy, metrics[AUC(nameauc)] )阶段一的学习率可以给到1e-3因为此时只有一层Dense在更新不用担心破坏BERT。阶段二切换到2e-5这是BERT微调最常用的学习率数量级。2e-5这个数字不是玄学它来自BERT原始论文的经验预训练模型的参数在微调时只能被很轻微地扰动学习率超过5e-5某些batch上就会出现loss暴涨。如果你用的是更大规模的预训练模型学习率可以再降一半。阶段一跑几个epoch就够了判断标准是训练集AUC明显起来、不再剧烈波动。阶段二通常需要更多epoch但也不要盲目跑几十轮BERT部分在少量标注数据上微调太久会过拟合。一般5到10个epoch就能看到验证集指标开始回落这时候要靠回调来截断。4. 微调BERT的超参数与回调配置训练过程怎么控制和判断模型结构搭好之后训练环节决定最终效果。BERT微调的超参数范围其实很窄不像训练词向量那样可以大幅调整关键就几个学习率、batch size、epochs、序列长度。把这些控制好再配合回调训练过程基本不会跑偏。4.1 学习率、batch size与epochsBERT微调的起步档位BERT微调的学习率建议从2e-5起步1e-5和3e-5也都是常见选择。如果训练集很小比如不到5000条用1e-5更稳如果数据量超过2万条3e-5到5e-5可以适当加速收敛。使用Adam优化器时默认的beta参数不需要动唯一建议改动的是gradient clipping后面讲坑的时候细说。batch size受显存限制常见的BERT-Base在seq_len128时batch size可以开到16到32。batch size越小梯度噪声越大BERT微调时表现为loss曲线震荡我一般用16起步显存不够就把seq_len降到64而不是强行减小batch size。seq_len降到多少直接看业务文本长度分布先统计一下语料长度如果95%的文本都在100个token以内seq_len128就足够了没必要硬上256。epochs没有标准答案。阶段一固定3到5轮阶段二看好验证集。很多初学着在阶段二跑满20个epoch结果训练loss很低、验证AUC却在第6轮开始回落这就是过拟合。多标签场景里标签一多过拟合更隐蔽因为整体AUC可能还在缓慢上升但某个低频标签的F1已经开始崩了。所以训练过程中要盯每个标签的验证指标而不是只看一个平均值。4.2 用回调控住训练过程EarlyStopping、ReduceLROnPlateau和模型存档Keras的回调机制在BERT微调里特别有用。EarlyStopping监控验证集loss连续几个epoch不下降就自动停ReduceLROnPlateau在loss平台期自动把学习率降一半给模型二次微调的机会ModelCheckpoint把每个epoch里验证集最好的权重保存下来防止最后几轮过拟合破坏了前面的最优状态。from keras.callbacks import EarlyStopping, ReduceLROnPlateau, ModelCheckpoint callbacks [ EarlyStopping( monitorval_loss, patience5, restore_best_weightsTrue ), ReduceLROnPlateau( monitorval_loss, factor0.5, patience2, min_lr1e-6 ), ModelCheckpoint( best_multi_label_model.h5, monitorval_loss, save_best_onlyTrue ) ] model.fit( train_gen, steps_per_epochlen(train_texts) // 32, validation_dataval_gen, validation_stepslen(val_texts) // 32, epochs15, callbackscallbacks )EarlyStopping的patience我一般设5也就是连续5个epoch验证loss不降就停。patience太小的风险是错过后面某个epoch的突然下降太大的风险是浪费时间。ReduceLROnPlateau的factor0.5表示学习率每次减半min_lr设在1e-6防止减到零。有一个细节阶段一里不要用ReduceLROnPlateau因为分类头本来就在快速收敛学习率减半反而拖慢节奏阶段二再用。ModelCheckpoint的monitor建议用val_loss而不是val_auc。AUC偶尔会在某些epoch跳高后又跌回来val_loss更平滑。save_best_onlyTrue保证磁盘上始终只有一份最优权重避免训练几十轮后磁盘被一堆h5文件塞满。4.3 训练结束不等于调完验证集上的多标签阈值搜索训练完的模型输出的是概率不是最终标签。多标签任务里把概率大于0.5判定为命中是最自然的想法但实际效果通常不行。原因有两个一是标签在训练集中的先验频率不同高频标签的输出概率天然偏高低频标签即使命中也只输出0.3左右二是batch size小、样本不均衡时sigmoid输出存在偏移。多标签分类的工程实践里阈值必须单独在验证集上搜索。from sklearn.metrics import f1_score # pred_val: 模型在验证集上的输出概率, shape (N, num_labels) # y_val: 验证集multi-hot标签, shape (N, num_labels) best_thresholds [] for label_idx in range(num_labels): best_f1 -1.0 best_thresh 0.5 for thr in np.arange(0.05, 0.95, 0.05): pred_binary (pred_val[:, label_idx] thr).astype(int) f1 f1_score(y_val[:, label_idx], pred_binary, zero_division0) if f1 best_f1: best_f1 f1 best_thresh thr best_thresholds.append(best_thresh)这段代码对每个标签单独搜索最优阈值搜索范围0.05到0.95、步长0.05。zero_division0是为了避免某个标签在验证集上完全没有预测为正时sklearn报警告。搜索出来的阈值通常会呈现出明显的规律高频标签阈值低因为模型倾向于保守低频标签阈值也低因为模型本来就不敢输出高分要放低门槛才能召回。阈值搜索不能拿到测试集或线上数据上去做否则就是把测试集当成训练集的一部分评估结果虚高。正确做法是阈值只在验证集上搜索一次然后固定下来再用测试集做最终评估。线上推理时阈值是写死在配置里的不能每次请求都重新搜索。5. BERT多标签微调的常见问题与排查五个实战翻车点这一章列几个我在实际项目里真实踩过、也帮别人排查过的坑。每个现象都见过不止一次写下来当个检查清单。5.1 坑一环境依赖打架keras-bert和TensorFlow版本对不上现象pip install keras-bert之后import keras_bert直接报错或者加载模型时报TypeError: str object is not callable训练时又冒出维度不匹配的诡异报错。原因keras-bert是老牌的Keras库它同时兼容原生Keras和tf.keras但pip会把keras当作依赖一并装进来。如果你的项目用的是TensorFlow 2.x而环境里又装了独立的keras包两个Keras的符号互相覆盖加载BERT模型时就会出现这种灵异报错。解决先确认项目到底走哪一套接口。我建议在TensorFlow 2.x环境下统一用tf.keras代码里from keras.layers要改成from tensorflow.keras.layersfrom keras_bert里的load_bet_model和Tokenizer保持不动因为keras-bert底层会自动对接当前Keras环境。安装依赖时把独立keras版本固定住或者干脆pip uninstall keras强制走tf.keras。配环境不是越新越好keras-bert项目更新不频繁它和某个TensorFlow版本配合良好就不要轻易升级TensorFlow。5.2 坑二验证集F1一直是0但训练loss在降现象训练过程看起来正常loss一路向下AUC也在0.8以上但验证集的F1算出来是0或者低得离谱。原因多半不是模型问题而是阈值和评估口径问题。多标签验证集上如果直接用0.5作为阈值而某个标签在验证集里的正样本比例只有5%模型输出的概率普遍在0.1到0.4之间那么阈值0.5会把所有样本都判为负F1自然是0。低频多标签场景这是常态不是bug。解决先别动模型把训练结束后的验证集概率导出来画一个直方图看看分布。然后按4.3节的阈值搜索方法跑一遍再算F1。另外一个常见原因是验证集的标签划分有误val_labels的形状和应用在y_val上是one-hot还是multi-hot不一致导致所有样本都判错。我习惯在训练前打印一下y_val.sum(axis0)看一眼每个标签在验证集里的出现次数如果某个标签是0次它在F1评估里就永远是0。5.3 坑三解冻BERT后loss直接发散甚至出现NaN现象阶段一训练正常阶段二解冻BERT后第一个epoch的loss先降后暴涨有些batch直接变成NaN。原因最常见的有两个。一是学习率太大BERT微调的解冻阶段学习率必须控制在5e-5以内2e-5是安全值很多人沿用阶段一的1e-3直接翻车。二是数据里存在异常样本比如文本本身是一串乱码、标签全零或者文本过长被截断后只剩下一堆无意义的碎片token。解冻后BERT参数被扰动这些异常样本的梯度就会把loss推上NaN。解决先把学习率降到2e-5并在Adam里加上梯度裁剪clipnorm1.0防止极端梯度冲垮参数。然后检查训练集里是否有空字符串、重复文本、全零标签的样本把明显异常的过滤掉。最后在阶段二第一个epoch用模型预测几个batch确认输出的概率区间正常再继续训练。5.4 坑四预测概率全挤在0.5附近模型像没学过一样现象训练完所有的输出概率都在0.4到0.6之间看不出哪个标签有明确的命中倾向。原因这个现象在二阶段微调里经常出现核心原因是分类头没有真正训练起来。阶段一只跑了1个epoch分类头的权重还没从随机初始化收敛过来直接进入阶段二BERT层的微调梯度淹没了分类头的学习信号。另一个可能原因是在池化时取错位置比如取到了padding区域的向量padding部分的特征全是0池化结果被稀释。解决回到阶段一把epochs增加到3到5观察训练集AUC是否明显升到0.9以上确认分类头已经学会基本的模式后再解冻。池化位置问题可以通过打印model.predict的输出和标签真值分布对比来排查。还有一个小细节验证集和训练集的seq_len必须一致不能用训好的模型换一个更长的seq_len做预测BERT的位置编码是固定的长度一变输出就会乱。5.5 坑五模型保存和加载后predict结果变了或者直接报错现象训练好的模型保存成h5重新加载后predict的结果和保存前不一样严重时直接报错找不到自定义层。原因keras-bert模型内部用了大量自定义层如果保存的是完整模型而不是weights加载时必须提供对应的custom_objects否则Keras无法反序列化。如果你在池化时用了Lambda层Lambda层保存时会记录函数的引用但换了一台机器或者换了函数名Python解释器找不到对应函数加载就失败。解决推荐的做法是只保存权重不保存完整模型结构。结构在代码里是固定的每次预测前用相同代码重建模型再load_weights训练好的权重这样彻底避开custom_objects问题。线上的时候用tf.saved_model.save保存成SavedModel格式这个格式对自定义层的依赖要小得多后面单独讲。经验是Keras的h5格式在调试阶段用用可以部署线上服务还是要换成SavedModel。6. 微调之后还没完阈值校准、单条验证与导出部署模型训练完只完成了一半后面还有三件事不能省阈值校准、单条自检、导出部署。6.1 按标签搜索阈值别用一个阈值管所有类第4章里已经贴过阈值搜索的代码这里再补一个细节搜索步长的粒度会影响结果。0.05步长足够找到合理的阈值想更精细可以先用0.05粗扫找到最优区间后在区间内再用0.01细扫。每个标签都单独搜索得到一组thresholds数组保存成json或者numpy文件线上预测时加载使用。这里还有个容易被忽略的问题如果验证集太小某个标签只有个位数的正样本搜索出来的阈值会过拟合到这几个样本上。至少要保证验证集里每个标签有30个以上正样本再谈阈值搜索。数据量不够时可以先不考虑阈值精确性把两个标签的阈值看成一个范围上线后根据线上反馈再动态调整。6.2 单条文本走一遍完整流程微调是否合格的自检方法模型部署前我会写一个不依赖服务、能直接跑的单条预测脚本把tokenizer、模型、阈值串起来自检def predict_one(text, tokenizer, model, label_names, thresholds, seq_len128): ids, segs tokenizer.encode(text, max_lenseq_len) x1 np.array([ids], dtypenp.int32) x2 np.array([segs], dtypenp.int32) prob model.predict([x1, x2], verbose0)[0] result {} for idx, name in enumerate(label_names): result[name] { probability: round(float(prob[idx]), 4), hit: bool(prob[idx] thresholds[idx]) } return result拿几类样本测一下包含多个标签的样本、完全无标签的样本、和训练集分布差异很大的新样本。多标签模型最容易出现的上线问题是“该命中的没命中”尤其是低频标签。如果单条自检时发现某个低频标签概率始终低于阈值可以回看训练集里这个标签的样本是不是文本模式太单一导致泛化不足。6.3 导出SavedModel从Keras训练到线上服务的最后一公里训练好之后导出用SavedModel格式而不是继续用h5。SavedModel是TensorFlow官方推荐的部署格式不依赖训练时的Python代码服务端加载时不需要重建模型结构。import tensorflow as tf tf.saved_model.save(model, serving_model)导出后目录里会有saved_model.pb和variables目录预测服务直接加载这个目录。有一点要注意如果模型里有自定义层SavedModel会保留自定义层类的引用服务端加载时仍然需要把对应的自定义层代码放在可见位置。keras-bert的BERT层就是这样所以生产环境里我会在服务代码里import keras_bert而不是只依赖SavedModel的序列化结果。这整套流程走下来一个文本多标签分类服务就能稳定上线了。如果遇到新标签、新语料我的习惯是动手调参之前先花几分钟统计标签分布看看哪些标签样本过少、哪些标签总是共现。多标签项目里模型结构能带来的提升是有限的真正的瓶颈往往在标签定义和数据质量上。把这一步做成固定动作之后我踩坑的次数明显少了。希望帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网