ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

Bert+CRF三元组识别:从数据标注到模型训练实战 简介这套NLP实战资源以BertCRF三元组识别为主题面向希望入门信息抽取、知识图谱构建的Python学习者与开发者。项目聚焦从非结构化文本中识别主体、谓词、客体例如“马云是阿里巴巴的创始人”这类三元组可支撑问答系统、语义搜索等下游应用。包内共11个文件以6个py源码脚本为主配合Markdown说明、依赖清单和示意图压缩包约37KB轻量且结构清晰。代码覆盖数据预处理、模型搭建、训练、评估、预测与数据切分等环节并内置bert-base-chinese预训练权重方便直接复现中文命名实体识别与序列标注流程。已有122人学习。通过该项目可系统掌握Hugging Face Transformers调用Bert的方法理解CRF层如何修正标签间约束关系学会处理中文分词、填充对齐和指标度量形成从建模到部署的完整实践认识是学习NLP结构化抽取技术的实用参考。1. 从跑通到改跑BertCRF三元组识别项目到底解决什么问题从一段几百字新闻里自动抽出“马云是阿里巴巴的创始人”这种主体、谓语、客体三元组属于信息抽取里最典型的落地需求。很多人第一反应是上大模型其实bert配一层CRF就能跑得不错这也是这套BertCRF三元组识别项目的核心思路用序列标注方式给每个字分配S/P/O三类语义槽再借助CRF约束标签转移拼回结构化三元组。压缩包结构标准main.py负责训练model.py是BertCRF网络predict.py做推理split_data.py切数据data放标注语料bert-base-chinese是中文预训练权重。适合第一次想独立训练中文序列标注模型的开发也适合跑过Bert分类、想补上CRF这块拼图的人。等你真正动手会发现网络定义最省心磨人的是数据格式和标签对齐。2. 数据侧把“马云是阿里巴巴的创始人”变成BIO标签序列2.1 标注格式与读入方式谁是S谁是P谁是O先把任务翻译成模型的语言。三元组识别在这里不是做一个关系分类而是对句子做序列标注也就是给每个token赋一个标签。常见标注规范是BIOB-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_ratio0.8, dev_ratio0.1, seed42): with open(data_path, r, encodingutf-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_len128): encoding tokenizer( words, is_split_into_wordsTrue, truncationTrue, paddingmax_length, max_lengthmax_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.pyBert输出加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_tagsconfig.num_labels, batch_firstTrue) def forward(self, input_ids, attention_mask, labelsNone): outputs self.bert(input_idsinput_ids, attention_maskattention_mask) emissions self.classifier(self.dropout(outputs.last_hidden_state)) if labels is not None: loss -self.crf(emissions, labels, maskattention_mask.bool(), reductionmean) return loss decode_tags self.crf.decode(emissions, maskattention_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参数定义好网络之后剩下一半工作量在训练配置。BertCRF的损失不是逐token的交叉熵而是CRF的负对数似然它的意思是最大化目标标签序列在所有可能路径中的概率。公式层面目标路径得分减去所有路径logsumexp最后取相反数。这样学出来的东西不只是“每个token像哪个标签”还会学习标签间的转移概率比如从B-S后面大概率接I-S或O几乎不会直接跳到I-P。config.py里的核心超参数实操里最常见的一组配置如下参数取值说明bert_path./bert-base-chinese本地预训练权重路径num_labels7B/I与3种语义槽Omax_len128输入最大长度截断加paddingbatch_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_stepswarmup_steps, num_training_stepstotal_steps )这里warmup的意义在于前期用一个比较小的学习率先稳住预训练参数后面再逐渐加大最后又线性衰减收敛。哪怕只有几千条样本我也不会省掉warmup因为Bert微调非常吃这一下。保存时也别只存最后一步每个epoch结束把验证集F1最高那一份单独存出来后面predict.py加载的就是最优权重而不是最终权重。4.2 训练期常见问题按现象、原因、解决记录第一个翻车现场loss在下降但预测结果全是O。现象是训练曲线看着正常每轮loss都在降可把验证集交给模型去看输出的标签几乎全是“无关”。原因大多是标签极度不均衡无关token占绝大多数模型只要把所有位置都预测成Oloss也能压得很低。解决方法是不要在训练循环里只看token级loss每个epoch结束跑一次实体级别的精确率和召回率计算方式用seqeval这类工具。如果实体级F1一直上不来再考虑给实体标签加权重或者做负样本下采样把无关token多的句子比例压下来。第二个坑非常隐蔽max_len128的截断把三元组从中间切断了。现象是短句全对长句全错而且错误都集中在一句超过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导致训练或推理时直接报类型错误。解决方法是固定成一组经过验证的组合transformers4.20.0、torchcrf0.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_locationcpu) model.load_state_dict(state_dict) model.eval()这里有个细节如果加载时用的是cpu后续要迁移到GPU别忘了在forward前给模型和tensor都做一次cuda()迁移否则会出现设备不匹配的报错。load_state_dict时如果出现missing key先检查是不是因为保存时带了module前缀那种情况要去掉前缀再加载。保存时我也建议只存模型参数不存优化器状态文件体积小很多载入也快。如果你要把predict.py部署成一个文本处理流水线的一环还需要提前想清楚模型常驻内存还是逐个任务加载前者省时但占显存后者开销大但灵活。CRF的decode方法实际上是维特比解码它会结合emissions和转移概率一次性找到整条序列的最优标签路径而不是在每个位置单独取概率最大的标签。这也是BertCRF比Bertsoftmax在实体抽取上边界崩得更少的原因它天然考虑了标签间的顺序约束。解码时同样要把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和OBJdef 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_locationcpu)) self.model.eval() self.tokenizer BertTokenizer.from_pretrained(config.bert_path) def extract(self, sentences): encoded self.tokenizer( sentences, truncationTrue, paddingTrue, max_lengthconfig.max_len, return_tensorspt, ) 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序列和预测标签肉眼扫一遍边界。从那以后我每次上手类似的三元组识别项目都强制自己先跑通一版最小流程切分数据、编码、训练两轮、预测一条、后处理打印到屏幕上全部通了再谈调参和性能优化。这个习惯帮我筛掉过好几份标注格式有歧义的数据也让我少在模型结构上浪费无谓时间。希望帮到你。本文还有配套的精品资源点击获取
返回列表