BERT+BiLSTM+CRF中文命名实体识别实战:从数据预处理到模型部署
2026/9/23 16:23:16 网站建设 项目流程

简介:面向中文命名实体识别(NER)的Python项目源码,以BERT+BiLSTM+CRF为核心框架,同时提供BiLSTM+CRF、IDCNN+CRF等多种对比实现,覆盖数据预处理、模型训练与评估全流程。压缩包共58个文件,以16个Python源码和19个pyc编译文件为主,另有模型说明、数据预处理说明(txt/md)及可视化图片(png)等配套文档,整体大小13.75MB,目录涵盖预训练模型、中文语料数据、模型定义、数据处理与训练入口等模块,结构清晰,便于按需查阅。压缩包内代码均经过测试运行成功,可直接在MSRA、人民日报等中文标注语料上开展实验;模型定义、数据预处理和训练脚本相互分离,配套的README与词表处理脚本能帮助读者快速理解BERT、BiLSTM、CRF三者的协作原理,以及命名实体识别任务从字向量、上下文编码到标签约束的完整实现细节。该套项目既适合计算机、人工智能、大数据等相关专业学生作为毕业设计或课程设计的起点,也可供刚入门NER的开发者对照练习,目前已有1209人浏览学习,参考价值较高。

1. 从标题看本质:这是一个能跑通的中文NER基准方案

看到这个标题,第一反应别是「又是三件套缝合」。BERT+BILSTM+CRF 这个组合在中文命名实体识别里能成为标配,不是因为新,而是因为它稳——BERT 负责把字变成含语境的向量,BiLSTM 负责前后文双向建模,CRF 负责保证标签序列合法,三个模块各管一段,任何一个单独拿出来都有明显短板,但串起来就是目前性价比最高的非生成式NER方案。

这个 zip 适合三类人:刚入 NLP 方向的学生,需要一个人物/地点/组织识别基线做对比实验;做知识图谱或信息抽取的工程师,想快速在内部数据上迭代一版实体抽取;以及被预训练大模型价格劝退、想在 CPU 或单卡上跑推理的小团队。项目说明和模型打包在一起,意味着你不用从零搭环境,重点是理解数据格式、改配置、跑通训练、看指标,再替换成自己的业务数据。

接下来按我自己的复现习惯,从架构拆解、数据预处理、训练调参到避坑和部署,完整走一遍。

2. 拆开 BERT+BILSTM+CRF:三个组件各管什么,为什么这样串

2.1 BERT 层:中文实体识别为什么离不开字向量

中文 NER 的第一道坎是分词。分词错了,实体边界就跟着错,比如「武汉市长江大桥」切错成「武汉/市长/江大桥」直接改变实体类型。BERT 用字级输入绕开了这个风险——每个汉字独立成一个 token,通过多层 Transformer 的自注意力机制,让每个字向量都携带上下文字义,「长」在「长江」和「市长」里得到不同表示,这就是动态词向量的价值。

选择 BERT 而不是 Word2Vec 或 ELMo,关键在三点。第一,BERT 的 12 层 Transformer 在预训练阶段见过海量中文语料,学到的先验知识在小数据集上特别管用,业内叫 few-shot 友好;第二,字向量天然规避了分词错误传播,这在中文场景几乎是决定性的;第三,BERT 的输出是 768 维向量(base 版本),和下游 BiLSTM 的输入维度直接对接,不用做复杂的维度转换。

这一层的常见做法是直接用「bert-base-chinese」预训练权重,嫌大可以换「bert-base-chinese」蒸馏出来的小模型,但精度会掉 2~3 个点。如果是垂直领域(医学、法律、金融),更好的做法是在领域语料上做二次预训练,但这不是这个 zip 能解决的,需要自己额外准备语料和显卡。

注意:BERT 层在训练时可以冻结(freeze)也可以微调(fine-tune)。小数据集上冻结能防过拟合,但会损失领域适配能力;我的习惯是在通用数据上冻结,换到自己业务数据后就解冻,学习率调到 2e-5。

2.2 BiLSTM 层:为什么单向 LSTM 不够,双向就够

BERT 输出的是字向量序列,但实体识别本质是一个序列标注任务——每个字要打上「B-PER(人名开始)/I-PER(人名内部)/O(非实体)」这类标签。判断一个字是不是实体的一部分,既需要看左边(它前面是什么),也需要看右边(它后面接什么),单向 LSTM 只能编码一个方向的信息,天然吃亏。

BiLSTM 的做法是正向跑一遍、反向跑一遍,把两个方向的隐藏状态拼接起来,作为每个字的最终上下文表示。比如「张三在北京工作」里「三」这个字,正向看到「张」知道它是人名开头,反向看到「在」知道它后面跟的是地点介词,两个方向的信息合在一起,「三」是 I-PER 的概率就很高了。

这里有一个常被忽略的细节:BiLSTM 的输出维度是 hidden_size × 2(正反向拼接)。如果你设置 hidden_size=128,那进入 CRF 层的特征维度就是 256。这个细节在你手写模型而不是直接套框架时特别容易错,我见过很多人在这里维度对不上报错,然后怀疑是 GPU 驱动问题,其实就是一个乘 2 的事。

BiLSTM 层还有一个作用——压缩 BERT 的输出。BERT 每个 token 输出 768 维,直接接分类器参数太多,容易过拟合。BiLSTM 把 768 维压到 2 × hidden_size(通常是 256 维),既降了维度,又加了序列建模能力,算是一举两得。

2.3 CRF 层:给标签序列加规则约束,把明显错误拦下来

如果只用 Softmax 分类器,每个位置独立预测标签,会出现「B-PER 紧跟 I-ORG」「B-LOC 后面直接跟 I-PER」这类非法序列。CRF 层做的事就是给相邻标签转移加一个可学习的约束矩阵,让模型学会「B 后面必须跟 I 或 O」「I 不能出现在句首」这类硬规则。

具体来说,CRF 的损失函数由两部分构成:发射分数(来自 BiLSTM 的每个位置的标签概率)和转移分数(来自 CRF 层的标签转移矩阵)。训练时最大化正确标签序列的得分,推理时用 Viterbi 解码找全局最优路径——不是每个位置取最大值,而是整条序列联合最优。

转移矩阵是可学习的,这意味着模型不仅知道「I 不能跟在 O 后面」这种语法规则,还能学到「PER 后面不太可能出现 LOC」这类语义偏好。这种约束在小数据集上特别有用,能把规则外的明显错误直接挡掉。CRF 的另一个优势是推理时天然保证输出序列合法,不需要额外的后处理规则。

2.4 三明治反例:什么情况下这个组合会崩

这套组合不是万能的。嵌套实体(如「联合国安理会」里「联合国」和「安理会」都是组织)CRF 只能选一条路径,天然处理不了;长实体(超过 32 个字)因为 BERT 的 position embedding 限制,容易截断;推理速度上,BERT 编码 + CRF 解码在 CPU 上单条句子约 50~100ms,达不到高并发在线服务的延迟要求。

所以落地时先问自己三个问题:我的实体是不是嵌套的?实体平均长度多少?延迟要求多高?如果三个问题里有两个踩中,纯 BERT+BILSTM+CRF 就不够用,需要换 GlobalPointer 或机器阅读理解的方案。如果不是,这个组合就是性价比最高的选择。

3. 数据准备与标签体系:中文 NER 数据长什么样,怎么喂给模型

3.1 输入格式:CONLL 格式的每一列是什么

这个 zip 里的数据大概率是 CONLL 格式——每行一个字,空格分隔三列:字本身、该字在句子中的位置标识、标签。句子之间用空行隔开。大概长这样:

李 B-PER 小 I-PER 明 I-PER 在 O 北 B-LOC 京 I-LOC 工 O 作 O

第一列是汉字本身,第二列是分词边界标识(B 表示词首,I 表示词内),第三列是实体标签。注意第二列和第三列是不同的东西——第二列是词边界,第三列是实体类型。有些数据集只有两列(字 + 标签),那就需要自己从 B/I 前缀里恢复词边界。

中文 NER 标签体系最常用的是 BIO 和 BIOES 两种。BIO 只有三种前缀:B(Begin)、I(Inside)、O(Outside);BIOES 多了 E(End)和 S(Single),对实体边界描述更精确,CRF 学起来更容易。如果数据集是 BIO,我一般会转成 BIOES 再用,效果通常好 1~2 个点。转换脚本长这样:

def bio_to_bioes(labels): new_labels = [] for i, label in enumerate(labels): if label == 'O': new_labels.append('O') continue prefix, entity = label.split('-') if prefix == 'B': if i + 1 < len(labels) and labels[i + 1].endswith('-' + entity) and labels[i + 1].startswith('I'): new_labels.append('B-' + entity) else: new_labels.append('S-' + entity) # 单个字成实体 elif prefix == 'I': if i + 1 < len(labels) and labels[i + 1].endswith('-' + entity) and labels[i + 1].startswith('I'): new_labels.append('I-' + entity) else: new_labels.append('E-' + entity) # 实体结束 return new_labels

这段代码的核心逻辑是:遇到 B 先看后面还有没有跟着的 I,有就是真 B,没有就改成 S(单字实体);遇到 I 先看后面还有没有 I,有就保持 I,没有就改成 E(实体结尾)。注意比较时用 startswith 而不是直接相等,因为要同时匹配「PER」「LOC」「ORG」等不同实体类型。

转换后记得用一个断言检查一下合法性——B 的后一个不能是 O 或不同实体的 I,I 的前一个必须是 B 或同实体的 I:

def check_labels(labels): valid_entities = {'PER', 'LOC', 'ORG'} for i, label in enumerate(labels): if label == 'O': continue prefix, entity = label.split('-') assert entity in valid_entities, f"非法实体类型: {entity}" if prefix == 'I': prev = labels[i - 1] if i > 0 else 'O' assert prev.startswith('B-') and prev.endswith('-' + entity) or \ prev.startswith('I-') and prev.endswith('-' + entity), \ f"第 {i} 个标签 {label} 的前一个标签非法: {prev}"

3.2 标签映射:从字符串到数字的唯一通道

模型不能直接吃字符串,需要把每个标签映射成一个整数。这里的关键是映射必须全局唯一且稳定——训练和推理用同一套映射表。常见做法是把标签表硬编码在配置文件里,或者从训练集自动生成,但自动生成时要注意测试集里可能没有的标签类型。

label_list = ['O', 'B-PER', 'I-PER', 'B-LOC', 'I-LOC', 'B-ORG', 'I-ORG'] label2id = {label: idx for idx, label in enumerate(label_list)} id2label = {idx: label for label, idx in label2id.items()} # 训练和推理共用 import json with open('label2id.json', 'w', encoding='utf-8') as f: json.dump(label2id, f, ensure_ascii=False, indent=2)

我把 label2id 存成 json 文件,推理时直接加载,不依赖训练脚本。这个习惯救过我一次——训练脚本重写后忘了重新生成映射表,推理结果全乱了。标签映射表必须作为项目的一个独立文件管理,不能散落在各个脚本里。

3.3 数据集划分:随机抽样为什么会翻车

NER 数据集划分和图片分类不一样,图片可以随意 shuffle,但 NER 的句子之间可能有上下文依赖(比如一篇文章里「他」指代前文的人名)。如果按句子随机划分,同一文档的句子会同时出现在训练集和验证集,导致验证集指标虚高。

常见做法是分两个层级划分:先按文档划分,文档内句子不跨集合;如果数据量小,再按句子随机划分但要求来自同一文档的句子在同一集合内。划分比例一般是 8:1:1(训练:验证:测试),验证集用来调参和早停,测试集只在最终评估时碰一次。

from sklearn.model_selection import train_test_split # sentences 是列表,每个元素是一篇文章的句子列表 doc_train, doc_temp = train_test_split(documents, test_size=0.2, random_state=42) doc_val, doc_test = train_test_split(doc_temp, test_size=0.5, random_state=42)

划分完成后务必检查三个集合的实体类型分布,特别是低频实体。如果「ORG」在训练集出现 500 次、验证集才 20 次,模型在验证集上 ORG 的 F1 大概率很低,但这不是模型的问题而是划分不均。发现分布差异大就重新设置 random_state 再划一次,直到各集实体分布相对一致。

3.4 字符规范化的隐藏坑

中文 NER 数据里最常见的烂数据是字符不统一——全角半角混合,中文标点里混着英文逗号。「,」和「,」看起来都是逗号,但模型把它们当两个不同的 token。这个坑在训练时不容易发现(loss 照样降),但推理时遇到没见过的那一种标点就会预测不稳。

我的做法是加一个字符归一化函数,在数据读取阶段就把全角标点转成半角,数字字母统一小写:

def normalize_text(text): text = text.strip() text = text.lower() # 全角转半角 half_char = '' for char in text: code = ord(char) if code == 0x3000: code = 0x20 elif 0xFF01 <= code <= 0xFF5E: code -= 0xFEE0 half_char += chr(code) return half_char

这个函数放在数据读取和预处理的入口,保证进入模型的所有字符都是同一套规范。注意这个方法只能处理 Unicode 范围内的转换,特殊符号(如 emoji)需要单独处理——我的做法是直接过滤掉训练集中出现次数低于 5 次的字符,因为这些低频字符大概率是脏数据。

4. 模型构建与训练:从配置文件到跑通第一个 epoch

4.1 配置文件:哪些参数影响最大

这个项目的核心参数集中在配置文件里。我把最重要的几个参数按影响程度列出来:

参数推荐值影响
max_seq_len128超过这个长度会被截断,长实体直接丢
batch_size16(显存小就 8)太小梯度震荡大,太大 OOM
learning_rate2e-5(BERT层) / 1e-3(下游层)分层设学习率是常识,统一设必炸
warmup_ratio0.1前 10% 步数学习率从 0 线性涨到目标值
num_epochs10看验证集早停,不固定
hidden_size128BiLSTM 隐层维度,太大过拟合太小欠拟合

这里最核心的分层学习率——BERT 预训练参数用的是小学习率(2e-5 到 5e-5),因为它是「已经学好的」参数,大学习率会把预训练知识冲掉(业内叫 catastrophic forgetting);BiLSTM 和 CRF 是随机初始化的参数,需要大学习率(1e-3)才能快速收敛。我把两种学习率写到同一个配置里:

bert: learning_rate: 2e-5 weight_decay: 0.01 bilstm: hidden_size: 128 num_layers: 1 dropout: 0.5 crf: learning_rate: 1e-3 training: max_seq_len: 128 batch_size: 16 epochs: 10 warmup_ratio: 0.1

4.2 训练入口:一条命令跑通全流程

假设 zip 里的源码已经封装好了训练脚本,通常长这样:

python train.py \ --config configs/default.yaml \ --data_dir data/person_location_org \ --output_dir checkpoints/ \ --gpu_id 0

如果没有现成脚本,需要自己拼模型的话,核心训练步骤是建一个继承 torch.nn.Module 的类,把 BERT、BiLSTM、CRF 三个组件装进去。关键代码如下,注意每一处注释:

import torch import torch.nn as nn from transformers import BertModel, BertTokenizer from torchcrf import CRF class BertBilstmCrf(nn.Module): def __init__(self, config): super().__init__() # 加载预训练 BERT,返回 token 级别的特征 self.bert = BertModel.from_pretrained(config['bert']['pretrained']) # 双向 LSTM,输入维度 768,输出维度 hidden_size * 2 self.bilstm = nn.LSTM( input_size=config['bert']['hidden_size'], # 768 hidden_size=config['bilstm']['hidden_size'], # 128 num_layers=config['bilstm']['num_layers'], bidirectional=True, batch_first=True, dropout=config['bilstm']['dropout'] if config['bilstm']['num_layers'] > 1 else 0 ) # 将 BiLSTM 的输出映射到标签空间 self.fc = nn.Linear( config['bilstm']['hidden_size'] * 2, # 正反向拼接,维度翻倍 config['num_labels'] ) # CRF 层,接收发射分数和标签数量 self.crf = CRF(config['num_labels'], batch_first=True) def forward(self, input_ids, attention_mask, labels=None): # 1. BERT 编码 outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask) sequence_output = outputs.last_hidden_state # [batch, seq_len, 768] # 2. BiLSTM 编码 lstm_out, _ = self.bilstm(sequence_output) # [batch, seq_len, 256] # 3. 线性层得到发射分数 emissions = self.fc(lstm_out) # [batch, seq_len, num_labels] if labels is not None: # 训练模式:计算 CRF 损失(降低对数似然,取负) loss = -self.crf(emissions, labels, mask=attention_mask.bool()) return loss # 推理模式:Viterbi 解码 predictions = self.crf.decode(emissions, mask=attention_mask.bool()) return predictions

这段代码里有几个关键点。第一,BiLSTM 的 dropout 参数只在层数大于 1 时才生效,单层 LSTM 传 dropout 不会报错但会静默忽略——这是 PyTorch 的行为,不会警告;第二,CRF 的 mask 参数用得是 attention_mask,因为 batch 内句子长度不一致时 padding 位置不应该参与 CRF 解码;第三,forward 在训练和推理模式下返回的不一样——训练返回的是 loss,推理返回的是解码后的标签序列,这个「一个 forward 两套逻辑」的写法在 NER 项目里很常见。

4.3 训练循环的骨架代码

模型搭好之后,训练循环要处理三个关键点:梯度裁剪(防止 BiLSTM 梯度爆炸)、warmup 学习率调度(防止 BERT 参数在初期被冲坏)、早停(防止过拟合)。核心代码如下:

optimizer = torch.optim.AdamW([ {'params': model.bert.parameters(), 'lr': 2e-5}, {'params': model.bilstm.parameters(), 'lr': 1e-3}, {'params': model.fc.parameters(), 'lr': 1e-3}, {'params': model.crf.parameters(), 'lr': 1e-3}, ], weight_decay=0.01) # warmup 和线性衰减 total_steps = len(train_loader) * num_epochs scheduler = torch.optim.lr_scheduler.LambdaLR( optimizer, lr_lambda=lambda step: min((step + 1) / warmup_steps, 1.0) * (1 - (step - warmup_steps) / (total_steps - warmup_steps)) if step > warmup_steps else (step + 1) / warmup_steps ) for epoch in range(num_epochs): model.train() for batch in train_loader: input_ids, attention_mask, labels = [t.to(device) for t in batch] loss = model(input_ids, attention_mask, labels) optimizer.zero_grad() loss.backward() # 梯度裁剪:CLIP 到 5.0 防止爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) optimizer.step() scheduler.step()

梯度裁剪的 max_norm 设 5.0 是我调过多次的经验值。设太小模型收敛慢,设太大梯度照样爆。如果你发现 loss 在训练到一半时突然变成 nan,先看两个地方:一是学习率是不是太大,二是梯度裁剪值是不是太高。BERT 层的梯度大小通常在 0.1~1.0 之间,BiLSTM 层的梯度可能到 10,所以 5.0 的裁剪值主要约束的是 BiLSTM。

4.4 评估脚本:F1 而不是 accuracy

NER 的评估不能看准确率。因为一个句子里「O」标签通常占 80% 以上,模型全预测 O 准确率也有 80%,但这种模型毫无用处。NER 的黄金指标是实体级别的 Precision、Recall、F1——必须整段实体预测正确才算对,「人」和「人民」这种部分匹配算错。

python evaluate.py \ --model_dir checkpoints/bert_bilstm_crf \ --test_file data/test.txt \ --output predictions.txt

评估脚本的核心逻辑是:先把每个 token 的标签还原到原始句子,然后把连续相同实体类型的 token 拼接成一个实体,和标准答案逐字对齐。比如预测出「张(B-PER) 三(I-PER)」作为一个实体「张三」,标准答案是「张(B-PER) 三(I-PER)」,完全匹配才算一个 TP。

判断模型好坏不能只看 F1 总分,要分实体类型看。中文 NER 里 PER(人名)通常最容易(边界清晰),LOC(地名)次之,ORG(机构名)最难(长度长、组合复杂)。如果你的 ORG F1 比其他类型低 10 个点以上,不要急着调模型,先看数据里 ORG 样本够不够多,标注一致性有没有问题。

5. 避坑与排查:中文 NER 训练中的五个高频翻车点

5.1 BERT 模型下载不了或下载极慢

现象:运行脚本时卡在「Downloading model」超过 10 分钟,甚至直接超时报错。

原因:transformers 库默认从 HuggingFace 下载模型权重,国内网络经常连不上或速度极慢。

解决:提前把模型权重下载到本地。先找一台网络通畅的机器(或者用镜像站),下载 bert-base-chinese 的 config.json、pytorch_model.bin、vocab.txt 三个文件,放在项目根目录的pretrained/bert-base-chinese/下。然后用from_pretrained('pretrained/bert-base-chinese')指定本地路径,不用默认的模型名。

我一般会在拿到项目后第一时间检查有没有 pretrained 目录,没有的话先去解决权重下载,不然训练脚本永远跑不起来。如果你是在内网环境做项目,这一步必须先解决,否则后面全是白等。

5.2 BiLSTM 输出维度拼接出错

现象:forward 时线性层报维度不匹配错误,mat1 and mat2 shapes cannot be multiplied

原因:BiLSTM 的 hidden_size 是 128,但双向输出是 256,线性层的输入维度写成了 128。

解决:线性层输入维度必须写hidden_size * 2。在模型定义处加一个断言,把维度错误在初始化时就炸出来:

assert self.fc.in_features == config['bilstm']['hidden_size'] * 2, \ f"线性层输入维度应为 {config['bilstm']['hidden_size'] * 2},实际为 {self.fc.in_features}"

很多人在维度上翻车后第一反应是去改隐藏层大小,但这是错误方向。正确做法是接受「双向 = 双倍维度」这个事实,把线性层维度和 CRF 输入维度都翻倍。维度错误是小问题,但排查成本高,一个断言能省你一小时。

5.3 BERT 微调导致过拟合,F1 越训越低

现象:训练集 loss 持续下降,但验证集 F1 在某个 epoch 后开始下滑,而且下滑幅度明显。

原因:BERT 层学习率设置过高或没有早停。BERT 预训练参数在领域数据集上微调时,如果学习率大于 5e-5,会把预训练学到的通用语言知识「洗掉」。

解决:把 BERT 层学习率降到 2e-5,开启早停——连续 3 个 epoch 验证集 F1 不提升就停止训练,并保存最佳模型。最佳模型取验证集 F1 最高的那个 checkpoint,不取最后一个。

# 训练脚本里加早停参数 --patience 3 \ --save_best_on_val True

这里面还有个玄学问题:验证集 F1 的波动范围通常在 1~2 个点内,如果连续 3 个 epoch 只波动不提升,可以再多看两个 epoch。但如果连续 5 个 epoch 都没超过历史最佳,基本可以确定模型已经到极限了,继续训练只会过拟合。

5.4 标签类别不平衡导致 ORG 永远识别不出来

现象:训练完看分类型指标,PER F1 有 85,LOC 有 80,ORG 只有 40。

原因:数据里 ORG 标注数量太少,模型见过几百个「张三」但只见过几十个「阿里巴巴」,学不到机构名的特征。

解决:这不是模型问题,是数据问题。三个方向去补——增加 ORG 标注数据(最有效);对 ORG 样本做简单数据增强(实体替换:把「阿里巴巴」替换成「腾讯」「华为」,标签不变);或者降低 ORG 的分类阈值(但会引入误报,不推荐)。

我的习惯是先看数据增强能不能把 ORG 样本量翻倍,翻了之后看 F1 提升幅度。提升超过 5 个点则继续增强,不足 3 个点则可能是标注规范问题——去看看原始数据里 ORG 是不是标注得不一致(有些标了「阿里巴巴集团」,有些标「阿里巴巴」,模型无法对齐)。

5.5 推理阶段 CRF 解码报错或预测全为 O

现象:模型训练没问题,但推理时出现index out of range错误或所有预测都是 O。

原因:推理时用的 label2id 和训练时不一致,或者输入文本的长度和 max_seq_len 对不上。具体来说,如果推理脚本用的是另一个 label_list(比如少了某个实体类型),标签索引错位,CRF 解码时就可能越界;输入超过了配置的 max_seq_len 时,超出部分被截掉,如果实体刚好在截断边界,就预测成了 O。

解决:推理脚本统一从label2id.json加载映射;对输入文本先做长度检查,超过 max_seq_len 的长文本先用分句工具拆成多个短句再逐个预测,最后拼回;加一个兜底逻辑:

if len(text) > config['max_seq_len'] - 2: # 预留 [CLS] 和 [SEP] text = text[:config['max_seq_len'] - 2]

这个截断策略很简单粗暴。如果业务场景里确实有长实体(机构全称往往超过 20 个字),更好的做法是做滑窗或分句,而不是加大 max_seq_len——加大会显著增加显存占用和推理时间,不划算。

6. 模型导出与推理:把一个训练好的 Pytorch 模型用起来

模型训练完只是第一步,真正落到业务里要解决两个问题:模型怎么存、推理怎么快。我的做法是把训练好的模型包成一个独立的推理类,和训练代码彻底解耦。

先说模型导出。训练过程中每个 epoch 结束都会保存 checkpoint,但 checkpoint 里有优化器状态、epoch 信息这些推理用不到的东西。我一般会单独写一个导出脚本,把最优 epoch 的模型权重单独保存:

python export_model.py \ --checkpoint checkpoints/best.pt \ --output exports/ner_model.bin

导出脚本的逻辑很简单——加载 checkpoint、去掉优化器状态、只保存 model.state_dict()。还有一个关键点是把 label2id.json 一并复制到 exports 目录,确保推理时的标签空间和训练时完全一致。这一步做好了,后续代码升级、换机器部署都不会出现「推理结果对不上」的问题。

然后是推理类。我的设计是把分词、编码、预测、解码整个流程封装成一个类,业务方只需要调用predict(text)接口:

class NERPredictor: def __init__(self, model_dir): self.tokenizer = BertTokenizer.from_pretrained(model_dir) self.model = BertBilstmCrf.from_pretrained(model_dir) self.model.eval() self.label2id = json.load(open(f'{model_dir}/label2id.json')) self.id2label = {int(k): v for k, v in self.label2id.items()} def predict(self, text, max_len=128): # 1. 文本长度检查 text = text[:max_len - 2] # 2. 编码 inputs = self.tokenizer(text, return_tensors='pt', truncation=True, max_length=max_len) # 3. 推理,不需要梯度 with torch.no_grad(): predictions = self.model(inputs['input_ids'], inputs['attention_mask']) # 4. 解码:把 token id 映射回文本 tokens = self.tokenizer.convert_ids_to_tokens(inputs['input_ids'][0]) entities = [] current_entity = None for token, pred_id in zip(tokens[1:-1], predictions[0][1:-1]): label = self.id2label[pred_id] if label == 'O': if current_entity: entities.append(current_entity) current_entity = None elif label.startswith('B-'): if current_entity: entities.append(current_entity) current_entity = {'text': token.replace('##', ''), 'type': label[2:], 'start': 0} elif label.startswith('I-'): if current_entity and current_entity['type'] == label[2:]: current_entity['text'] += token.replace('##', '') current_entity['start'] += 1 if current_entity: entities.append(current_entity) return entities

这个类里有几个细节值得注意。第一,tokens[1:-1]去掉了 BERT 自动拼上的 [CLS] 和 [SEP] 标记,相应预测结果也去掉首尾两位;第二,token.replace('##', '')是处理 BERT 词片分词的还原逻辑,中文里不常见(中文基本是单字),但英文场景必须有;第三,实体边界用 start/end 字段记录,方便业务方在原文里高亮定位。

推理性能方面,如果单条文本平均长度在 100 字左右,BERT-BiLSTM-CRF 在 CPU 上的推理延迟约 80ms,GPU 上约 10ms。如果要做高并发在线服务,常见手法是上 GPU 并用 batch 推理,把多条请求合并成一个 batch,GPU 利用率会显著提升。还有一个优化技巧是把输入 padding 到 8 的整数倍——GPU 对对齐的内存访问更快,这个白嫖的性能提升有时候能达到 15%。

落到最终的业务接入,我一般建议保留两部分代码:一部分是训练和评估用的(带日志和可视化),一部分是优化过的最小化推理服务(只留 predict 接口)。训练代码和推理代码混在一起是项目后期最痛的维护负担——线上出了一次预测错误,你不想在训练脚本里翻半天找预处理在哪一行。

这个项目里最有价值的不是那三个模型,而是数据转换、标签对齐、评估脚本这些「周边代码」。数据转换错了模型白训,标签对不齐评估白跑,评估方法错了调参方向全错。我自己每次拿到新数据集,第一件事不是跑模型,是先把数据可视化出来——随便挑 50 条打上标签的句子打出来逐字看一遍,确认标注规范和代码假设一致。这个习惯帮我挡掉了至少三次「训练很顺利但结果没法用」的翻车,希望帮到你。

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询