news 2026/10/5 13:53:15

Bert+CRF三元组识别:从数据标注到模型训练实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Bert+CRF三元组识别:从数据标注到模型训练实战

简介:这套NLP实战资源以Bert+CRF三元组识别为主题,面向希望入门信息抽取、知识图谱构建的Python学习者与开发者。项目聚焦从非结构化文本中识别主体、谓词、客体,例如“马云是阿里巴巴的创始人”这类三元组,可支撑问答系统、语义搜索等下游应用。包内共11个文件,以6个py源码脚本为主,配合Markdown说明、依赖清单和示意图,压缩包约37KB,轻量且结构清晰。代码覆盖数据预处理、模型搭建、训练、评估、预测与数据切分等环节,并内置bert-base-chinese预训练权重,方便直接复现中文命名实体识别与序列标注流程。已有122人学习。通过该项目,可系统掌握Hugging Face Transformers调用Bert的方法,理解CRF层如何修正标签间约束关系,学会处理中文分词、填充对齐和指标度量,形成从建模到部署的完整实践认识,是学习NLP结构化抽取技术的实用参考。

1. 从跑通到改跑:Bert+CRF三元组识别项目到底解决什么问题

从一段几百字新闻里自动抽出“马云是阿里巴巴的创始人”这种主体、谓语、客体三元组,属于信息抽取里最典型的落地需求。很多人第一反应是上大模型,其实bert配一层CRF就能跑得不错,这也是这套Bert+CRF三元组识别项目的核心思路:用序列标注方式给每个字分配S/P/O三类语义槽,再借助CRF约束标签转移,拼回结构化三元组。压缩包结构标准,main.py负责训练,model.py是Bert+CRF网络,predict.py做推理,split_data.py切数据,data放标注语料,bert-base-chinese是中文预训练权重。适合第一次想独立训练中文序列标注模型的开发,也适合跑过Bert分类、想补上CRF这块拼图的人。等你真正动手会发现,网络定义最省心,磨人的是数据格式和标签对齐。

2. 数据侧:把“马云是阿里巴巴的创始人”变成BIO标签序列

2.1 标注格式与读入方式:谁是S,谁是P,谁是O

先把任务翻译成模型的语言。三元组识别在这里不是做一个关系分类,而是对句子做序列标注,也就是给每个token赋一个标签。常见标注规范是BIO:B-S/I-S表示主体片段,B-P/I-P表示谓词片段,B-OBJ/I-OBJ表示客体片段,O表示与目标任务无关的普通字。这里有个命名细节我吃过亏:客体(Object)不要直接用“O”当标签名,否则会和“无关”标签O完全撞车,训练期会直接体现为loss混乱。项目里更常见的做法是把客体写成OBJ,或者用T0/T1/T2这种码表,标签集合才算干净。

原始语料最常见的落地格式是TSV或JSON。一行样本用制表符分隔两个字段:句子和标签序列,句子按空格切词,标签也按空格对齐。比如“马云是阿里巴巴的创始人”这一行写成:

文本: 马云 是 阿里巴巴 的 创始人 标签: B-S B-P B-OBJ I-OBJ I-OBJ

压缩包data目录里装的就是这类标注文件,格式足够简单,打开就能直接看结构,数据替换也很方便。因为bert-base-chinese基本是字符级切词,中文里一个汉字绝大多数情况对应一个token,省掉了大量对齐麻烦。如果以后换英文或多语种数据,就要重新处理WordPiece切分导致的标签扩散问题。

再看切分。split_data.py就是把全量标注按8:1:1切成训练集、验证集、测试集,核心逻辑如下:

import random def split_dataset(data_path, train_ratio=0.8, dev_ratio=0.1, seed=42): with open(data_path, "r", encoding="utf-8") as f: lines = [line.strip() for line in f if line.strip()] random.seed(seed) random.shuffle(lines) total = len(lines) train_end = int(total * train_ratio) dev_end = train_end + int(total * dev_ratio) train_lines = lines[:train_end] dev_lines = lines[train_end:dev_end] test_lines = lines[dev_end:] return train_lines, dev_lines, test_lines

这段逻辑本身不难,但有两个细节别图省事。train_ratio和dev_ratio分别控制前80%和后10%的划分,剩下10%留作测试;seed必须固定,否则每次切分结果不一致,你调了三天的参数,回头发现验证集换过一轮,之前所有对比都作废。第二是切分前一定要shuffle,很多真实数据集是按来源排列的,比如同一家公司的新闻都排在一起,不洗牌会让验证集分布和训练集差异特别大。我一般会把归一化后的统计写在日志里,比如句子平均长度、标签分布,至少先确认数据不是一股脑地偏到某个类别上。

2.2 编码对齐:word_ids、label_ids与padding

数据读进来之后还没法进模型,utils.py负责把句子和标签转换成模型能吃的张量。这部分是新手最容易翻车的区域,因为Bert自带的tokenizer对中文人名、数字、符号处理时,可能把一个词拆成多个subword,而原始标签还是按词粒度给的,此时要对齐。Hugging Face的tokenizer保留了word_ids方法,也就是每个token对应的原词下标,遍历一遍就能把标签复制到所有subword上:

def encode_example(tokenizer, words, labels, label2id, max_len=128): encoding = tokenizer( words, is_split_into_words=True, truncation=True, padding="max_length", max_length=max_len, ) word_ids = encoding.word_ids() aligned_label_ids = [] previous_word_idx = None for word_idx in word_ids: if word_idx is None: aligned_label_ids.append(-100) elif word_idx != previous_word_idx: aligned_label_ids.append(label2id[labels[word_idx]]) else: aligned_label_ids.append(label2id[labels[word_idx]]) previous_word_idx = word_idx return { "input_ids": encoding["input_ids"], "attention_mask": encoding["attention_mask"], "labels": aligned_label_ids, }

这段代码里,-100是PyTorch交叉熵损失里的约定,表示“该位置不参与loss计算”,专门用来屏蔽[CLS]、[SEP]和PAD。这里要注意标签是否越界,如果某个词的标签在label2id里查不到,多半是标注数据里混入了无关标记,我建议在编码前先做一遍标签集合校验,把唯一标签全部打印出来核对。

这样编码完成后,每条样本变成固定长度的input_ids、attention_mask和labels三个张量。训练时按batch堆叠,labels里那些-100的位置不影响Bert模块,但CRF层需要知道哪些位置是真实token。常见做法是把attention_mask转成bool,在loss和解码时都传进去,这样PAD位置的转移矩阵就不会被CRF当成有效路径来计算。数据流走到这里就通了,模型侧的事情交给model.py。

3. 模型侧:Bert编码上下文,CRF把标签转移约束起来

3.1 model.py:Bert输出加Linear层,再接CRF解码

这一层是整个项目的核心。Bert部分做的事情是拿预训练好的中文模型,对每个token生成一个包含上下文信息的向量表示,很多人只关心它的CLS向量做分类,在这里用到的反而是每个位置的全量输出。具体来说,把输入句子经过Bert得到768维的last_hidden_state,再过一层Linear,把每个位置映射到7维的标签得分上,这7个维度分别对应B-S/I-S/B-P/I-P/B-OBJ/I-OBJ/O。这个7维得分就是CRF层说的emissions,表示每个位置对不同标签的原始打分。

model.py里网络结构大致是:

import torch import torch.nn as nn from transformers import BertModel from torchcrf import CRF class BertCRF(nn.Module): def __init__(self, config): super().__init__() self.bert = BertModel.from_pretrained(config.bert_path) self.dropout = nn.Dropout(config.dropout) self.classifier = nn.Linear(self.bert.config.hidden_size, config.num_labels) self.crf = CRF(num_tags=config.num_labels, batch_first=True) def forward(self, input_ids, attention_mask, labels=None): outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask) emissions = self.classifier(self.dropout(outputs.last_hidden_state)) if labels is not None: loss = -self.crf(emissions, labels, mask=attention_mask.bool(), reduction="mean") return loss decode_tags = self.crf.decode(emissions, mask=attention_mask.bool()) return decode_tags

这个结构里最有必要解释的是CRF层要的mask。注意这里用的是attention_mask.bool(),而不是labels里那组-100,原因是CRF只关心“真实的token位置”,而attention_mask本身对PAD是0,对真实token是1,天然可当mask用。labels里的-100是给loss用的,二者用途不同。很多人在复现阶段踩的第一个坑,就是直接把带-100的labels丢给CRF,结果标签类型不匹配直接报错。

backbone选择上,这个项目用的是bert-base-chinese,也就是压缩包里那个目录。它约1.1亿参数,对中文任务来说是起步配置。如果只是做演示,也可以换成更小的中文预训练模型来降低显存占用,但边界token表示会弱一些。我自己的习惯是先在bert-base-chinese上跑通,再根据验证集行为决定要不要压缩。

3.2 损失函数、优化器与config.py参数

定义好网络之后,剩下一半工作量在训练配置。Bert+CRF的损失不是逐token的交叉熵,而是CRF的负对数似然,它的意思是:最大化目标标签序列在所有可能路径中的概率。公式层面,目标路径得分减去所有路径logsumexp,最后取相反数。这样学出来的东西不只是“每个token像哪个标签”,还会学习标签间的转移概率,比如从B-S后面大概率接I-S或O,几乎不会直接跳到I-P。

config.py里的核心超参数,实操里最常见的一组配置如下:

参数取值说明
bert_path./bert-base-chinese本地预训练权重路径
num_labels7B/I与3种语义槽+O
max_len128输入最大长度,截断加padding
batch_size16小显存也能跑
lr_bert2e-5Bert层学习率
lr_crf1e-3CRF层学习率,可以调大点
epochs5小数据集5轮左右足够

这里最值得说的是“两个学习率”。Bert部分微调学习率通常2e-5,用太大会把预训练权重破坏;CRF层是从零训练的随机参数,转移矩阵收敛速度不要求太高,用1e-3反而更合适。所以常见训练脚本里会给CRF单独配一组参数,对Bert层用AdamW,对CRF层也用AdamW,但学习率分开。如果不想麻烦,统一用2e-5也能训,只是转移矩阵收敛慢一点,体现在验证集上就是实体边界时好时坏。

另一个细节是dropout。Bert输出到Linear层之间加了个0.1左右的dropout,这是为了降低预训练特征在少量标注数据上的过拟合。很多初学者的序列标注模型预测全是O,很大程度上是直接把Bert输出接Linear,没有dropout也没冻结策略,在小数据集上几轮就把标签多样性学没了。我一般建议config里的dropout值不小于0.1,数据量越小这个值越不太敢往下调。

4. 训练避坑:PAD填充、loss玄学与标签集合杂病

4.1 main.py训练循环:优化器warmup与保存

模型结构定了,训练循环其实比较模板化。main.py的职责是加载配置、加载数据、实例化模型、跑epoch循环,并在每个epoch末尾用验证集算一次指标。区别只在于优化器和学习率调度。这部分的典型写法:

from transformers import AdamW, get_linear_schedule_with_warmup total_steps = len(train_loader) * config.epochs warmup_steps = int(total_steps * 0.1) optimizer = AdamW([ {"params": model.bert.parameters(), "lr": config.lr_bert}, {"params": model.crf.parameters(), "lr": config.lr_crf}, ]) scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=warmup_steps, num_training_steps=total_steps )

这里warmup的意义在于前期用一个比较小的学习率先稳住预训练参数,后面再逐渐加大,最后又线性衰减收敛。哪怕只有几千条样本,我也不会省掉warmup,因为Bert微调非常吃这一下。保存时也别只存最后一步,每个epoch结束把验证集F1最高那一份单独存出来,后面predict.py加载的就是最优权重而不是最终权重。

4.2 训练期常见问题,按现象、原因、解决记录

第一个翻车现场:loss在下降,但预测结果全是O。现象是训练曲线看着正常,每轮loss都在降,可把验证集交给模型去看,输出的标签几乎全是“无关”。原因大多是标签极度不均衡,无关token占绝大多数,模型只要把所有位置都预测成O,loss也能压得很低。解决方法是不要在训练循环里只看token级loss,每个epoch结束跑一次实体级别的精确率和召回率,计算方式用seqeval这类工具。如果实体级F1一直上不来,再考虑给实体标签加权重,或者做负样本下采样,把无关token多的句子比例压下来。

第二个坑非常隐蔽:max_len=128的截断把三元组从中间切断了。现象是短句全对,长句全错,而且错误都集中在一句超过100字的样本上。原因是三元组里的主体和客体分布在句子首尾两端,截断后一头被切没了。解决方法是先统计一遍训练集的句子长度分布,如果中位数已经接近90,就把max_len拉到192或256,代价只是batch_size下调到8。还有一个常见误操作是截断时直接丢尾部,但往往宾语才是三元组里最长的那一部分,我一般会检查被截断样本中标签是否落在最后几个token上,确认要不要改截断策略。

第三个问题就是前面提过的标签O冲突。现象是training loss卡在某个值附近不降,或者验证集里客体一直被识别成无关。原因分析下来往往就是标注规范里用O表示无关,又用O表示Object,同一个标签id同时表达两个语义。解决很简单:把客体统一改成OBJ或者T2,全项目搜索替换后再重新切分数据。这个错误在人工标注格式里最容易出现,因为大家习惯写O做无关标签,顺手又把Object缩写成了O。

第四个问题常出现在重新搭建环境的时候。压缩包里依赖文件名保存的是requests.txt,不是常见的requirements.txt,复制粘贴时容易混淆,里面往往只列了包名没锁版本,装出来的transformers版本不同,CRF的mask接口从int变成bool,导致训练或推理时直接报类型错误。解决方法是固定成一组经过验证的组合:transformers==4.20.0、torchcrf==0.3.1,装好后从split_data到predict全流程跑一遍,确认没报错再调参。

4.3 不要只看loss,要盯实体级精确率、召回率与F1

这个项目的评估指标不是准确率,而是实体级别的精确率、召回率和F1。token级的正确率会被无关标签刷到95%以上,但三元组可能一个都没抽对。所以main.py里评估要按span去匹配,比如预测出来的主体标签连续片段和真实片段一致,才算一个正例。推荐seqeval库,它本身就支持BIO格式的span评估。我在实跑这段时有一条亲测有效的流程:先跑2个epoch看loss能不能降到接近0,再用验证集算F1,如果loss降了但F1不动,基本就是数据标签问题而非模型问题,优先回到2.2节做标签集合校验。

5. predict.py 实战:加载权重、维特维解码与后处理

5.1 从checkpoint恢复成可推理模型

训练结束,predict.py要把之前保存的最优权重加载回来做新数据预测。因为网络结构里有CRF层,加载权重不能像普通分类任务那样只load_state_dict,还得保证emissions经过CRF解码而不是做softmax。这个项目的做法是重新构建一个BertCRF(config)实例,然后从保存的checkpoint中把model_dict灌进去:

model = BertCRF(config) state_dict = torch.load(config.ckpt_path, map_location="cpu") model.load_state_dict(state_dict) model.eval()

这里有个细节:如果加载时用的是cpu,后续要迁移到GPU,别忘了在forward前给模型和tensor都做一次cuda()迁移,否则会出现设备不匹配的报错。load_state_dict时如果出现missing key,先检查是不是因为保存时带了module前缀,那种情况要去掉前缀再加载。保存时我也建议只存模型参数不存优化器状态,文件体积小很多,载入也快。如果你要把predict.py部署成一个文本处理流水线的一环,还需要提前想清楚模型常驻内存还是逐个任务加载,前者省时但占显存,后者开销大但灵活。

CRF的decode方法实际上是维特比解码,它会结合emissions和转移概率,一次性找到整条序列的最优标签路径,而不是在每个位置单独取概率最大的标签。这也是Bert+CRF比Bert+softmax在实体抽取上边界崩得更少的原因,它天然考虑了标签间的顺序约束。解码时同样要把attention_mask传进去,否则PAD位置会参与路径打分,把尾巴上那些“无关”标签算进去。

提示:保存checkpoint时最好只存model state_dict,不要混优化器state,否则predict.py加载慢,还可能因为优化器版本差异产生兼容问题。

5.2 后处理:把标签序列拼回三元组字符串

模型输出的是一串标签id,比如[O, B-S, I-S, B-P, B-OBJ, I-OBJ, O],这还不是我们能直接入库的三元组。后处理要做的是把连续相同类型的标签span拼接起来,组成(主体, 谓词, 客体)结构。一个句子里可能出现多个三元组,就需要按序扫描标签序列,遇到B-S就开始收集,直到离开I-S停下,再类似处理P和OBJ:

def parse_triples(words, tags): triples = [] i = 0 while i < len(tags): if tags[i] == "B-S": s = [words[i]] i += 1 while i < len(tags) and tags[i] == "I-S": s.append(words[i]) i += 1 subject = "".join(s) pred, obj = "", "" if i < len(tags) and tags[i] == "B-P": p = [words[i]] i += 1 while i < len(tags) and tags[i] == "I-P": p.append(words[i]) i += 1 pred = "".join(p) if i < len(tags) and tags[i] == "B-OBJ": o = [words[i]] i += 1 while i < len(tags) and tags[i] == "I-OBJ": o.append(words[i]) i += 1 obj = "".join(o) if subject and pred and obj: triples.append((subject, pred, obj)) else: i += 1 return triples

这段拼接逻辑要注意的是别用查找B-P的方式来重新定位,因为主语可能出现多次,顺序扫描能保证谓词和客体紧跟在主语后面。拼接时“的”这类字通常被标在OBJ内部,比如“阿里巴巴的创始人”,拼起来是一个完整客体,不需要额外去除。但如果数据里把“的”标成无关标签,后处理时词与词之间直接拼接就会得到“阿里巴巴创始人”,这种结果是标注风格不一致造成的,我一般会在数据清洗阶段提前定好规则,要么“的”都进客体,要么都排除。

predict.py真正部署前我还会做一次随机抽样检查:拿10条训练集句子和10条验证集句子,把预测结果和标注结果并排打印出来,逐条核对边界。序列标注模型最怕的不是整体跑偏,而是边界差一个字,比如主体“阿里巴巴集团”被识别成“阿里巴巴”,这种错误在精确率指标上很刺眼,但在代码逻辑里完全看不出来,只能靠人工抽样。

6. 把三元组识别包成可复用函数:一次加载、批处理与验证技巧

6.1 封装与批量调用

模型能单句预测,但正常使用场景更想一次传几十句进来。常见做法是抽出一个TripleExtractor类,把模型、tokenizer和后处理都收进去,输出直接就是三元组列表:

class TripleExtractor: def __init__(self, config): self.model = BertCRF(config) self.model.load_state_dict(torch.load(config.ckpt_path, map_location="cpu")) self.model.eval() self.tokenizer = BertTokenizer.from_pretrained(config.bert_path) def extract(self, sentences): encoded = self.tokenizer( sentences, truncation=True, padding=True, max_length=config.max_len, return_tensors="pt", ) with torch.no_grad(): tag_ids = self.model(encoded["input_ids"], encoded["attention_mask"]) return [ parse_triples(sent.split(), id2tag[t]) for sent, t in zip(sentences, tag_ids) ]

批处理比for循环逐句调用快三到四倍,尤其在GPU上,因为单条短句的GPU利用率很低。封装的时候还要留意tokenizer会自动加[CLS]和[SEP],预测结果里的标签数组要和输入句子对齐,这是最容易在封装阶段翻车的地方。之后验证封装逻辑有没有破坏对齐,我会单独挑一条包含超长实体的句子打印token序列和预测标签,肉眼扫一遍边界。

从那以后我每次上手类似的三元组识别项目,都强制自己先跑通一版最小流程:切分数据、编码、训练两轮、预测一条、后处理打印到屏幕上,全部通了再谈调参和性能优化。这个习惯帮我筛掉过好几份标注格式有歧义的数据,也让我少在模型结构上浪费无谓时间。希望帮到你。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/5 13:53:05

携程福利组合拳刷屏:混合办公与生育补贴背后的员工体验设计逻辑

昨天热搜上一出现“携程放大招”这几个字&#xff0c;我第一反应是又是什么营销噱头&#xff0c;结果点进去翻完评论区&#xff0c;画风完全一致——一排的“这才是企业该有的样子”“慕了慕了”。作为在职场里摸爬滚打过十几年的人&#xff0c;我太知道这种集体羡慕有多难得了…

作者头像 李华
网站建设 2026/10/5 13:51:42

扫码点餐系统源码部署实战:从解压到上线全流程避坑指南

简介&#xff1a;一套基于SpringBoot与uniapp(vue3)的扫码点餐系统完整源码&#xff0c;面向Java毕业设计、课程大作业及小程序开发学习者。后端采用Spring Security OAuth2实现安全认证&#xff0c;前端可发布为微信小程序或H5&#xff0c;覆盖多门店、外卖与自取等典型餐饮场…

作者头像 李华
网站建设 2026/10/5 13:49:37

Android与iOS平台测试的差异解析:从架构到发布审核

做APP测试这行当久了&#xff0c;你会发现「Android和iOS平台测试的区别」不只是换个手机跑一遍那么简单。同样是点一个按钮、发一个请求&#xff0c;两边的行为可能千差万别——Android后台杀得一干二净&#xff0c;iOS还能把你从崩溃现场拉回来&#xff1b;Android上一个Toas…

作者头像 李华
网站建设 2026/10/5 13:49:14

SolidWorks拉伸切除失败的五大根因与排查技巧

做了这么多年SolidWorks相关的工作&#xff0c;也带过不少新人刷练习题&#xff0c;我越来越觉得&#xff0c;练习题的含金量往往不在于题目本身有多复杂&#xff0c;而在于它能不能逼你把某个隐藏的坑踩一遍。就拿这几天好几个朋友都在问的问题来讲——为什么有时拉伸切除会执…

作者头像 李华
网站建设 2026/10/5 13:47:46

OpenShell 智能体执行框架:从架构设计到安全落地的工程实践

1. 从零认识 OpenShell&#xff1a;它到底解决什么问题第一次听到 OpenShell 这个名字&#xff0c;很多人会下意识以为它又是一个新的命令行工具或者某种终端美化方案。实际上&#xff0c;OpenShell 的定位要更底层、也更有意思——它是一套面向智能体&#xff08;Agent&#x…

作者头像 李华
网站建设 2026/10/5 13:44:42

V免签支付系统:ThinkPHP+安卓监控端搭建免签约收款回调方案

简介&#xff1a;面向急需接入免签约收款能力的中小商家与PHP开发者&#xff0c;这套基于Thinkphp内核的免签支付系统提供了安卓监控端与后端服务完整源码&#xff0c;可直接对接支付宝和微信支付&#xff0c;实现支付结果回调、收款实时监控及数据统计&#xff0c;省去与支付机…

作者头像 李华