简介:基于PyTorch的二分标注文本三元组信息抽取模型源码,面向自然语言处理学习者和信息抽取相关开发者,解决从非结构化文本中抽取(主实体,关系,客实体)三元组的核心问题,可服务于知识图谱构建、问答系统与语义检索等场景。项目采用二分标注策略,将抽取过程分为主实体识别与给定主实体后的客实体和关系预测两阶段,借助PyTorch的动态计算图与自动微分机制,降低模型训练与调试难度。压缩包共31个文件,容量348KB;23个Python源文件构成核心,涵盖多种嵌入(Albert、BERT、Word2Vec)、编码器、解码器、损失函数、数据预处理及工具模块,另有2个JSON文件存放配置与样例数据,Markdown说明文档、LICENSE与依赖清单便于快速上手。目前已有343人学习下载。通过源码可完整学习文本三元组抽取任务的建模流程、模块化项目组织方式以及二分标注思路,适合作NLP课程设计、知识图谱项目或PyTorch实战的参考实现。
1. 二分标注:为什么三元组抽取要拆成两次标注
接触信息抽取任务时,大多数人第一反应是把主语、关系、宾语三个槽位一起预测,做成一个多标签序列标注。但一旦关系数量上去,比如这个项目里达到 50 种,标签空间会立刻膨胀到几百种组合,模型很难收敛,输出层参数也撑不住。把任务拆开是更务实的做法:先做一次标注去定位主实体,再针对每个候选关系分别做一次标注去定位客实体。这就是二分标注的核心思想,用两次低维的 0/1 判断替代一次高维的多分类。这个 PyTorch 项目把这一思路落成了可运行的完整代码,数据处理、模型结构、损失函数都围绕两道标注展开。适合正在做关系抽取、知识图谱构建,想快速搭建 PyTorch baseline 的 NLP 工程师,项目结构简单到可以直接改来当实验框架用。
2. 数据与Embedding:从JSON样本到二分标注矩阵
2.1 数据目录结构与样本格式
先看项目里data目录的组成:all_50_schemas定义了 50 个关系的 schema 列表,train_data_sample.json和dev_data_sample.json分别提供训练和验证样本。这些 JSON 文件的结构并不复杂,每条样本包含原始文本和对应的三元组集合,常见格式大致如下:
{ "text": "《红豆》由王菲演唱,收录于专辑《唱游》", "spo_list": [ ["王菲", "歌手", "红豆"], ["红豆", "所属专辑", "唱游"] ] }process_raw_data.py负责读取这类 JSON,逐条把文本拆成 token 序列,再用function.py里的定位函数去匹配spo_list中每个实体的起止位置。all_50_schemas中每条记录的predicate字段会被抽出来构建rel2id和id2rel两个字典,这一步决定了后续客实体标注矩阵的关系维度。logger.py则负责向终端输出标签数量、文本长度分布这类统计信息,方便在数据进入模型前先做一轮质量检查。
文本切分这里有一个工程上的细节:中文实体容易跨越词边界,分词器把“王菲”切成一个词固然好,但遇到“唱游”这类可能被切碎的词就会出现偏差。项目里的处理方式是按字切分,每个汉字作为独立 token,实体的起止位置直接落在字索引上。这样处理与 BERT 类模型的中文词表天然对齐,也绕开了分词错误向实体标注传播的问题。pertrain_to_numpy.py这个文件的作用则是把 HuggingFace 预训练权重转成 numpy 数组,供自定义的 embedding 模块在不依赖 transformers 库的情况下也能加载参数。
2.2 torchEmbedding到ALBERT的切换
项目里四个 embedding 文件解决的是不同场景下的向量表示问题,config.py里通过embedding_type字段控制具体加载哪一个实现。下面是四个文件在实际使用中的定位:
| 文件 | 向量来源 | 典型场景 |
|---|---|---|
wor2vec_embedding.py | 静态 word2vec 词向量 | 快速验证 pipeline、无 GPU 环境 |
torchEmbedding.py | PyTorch 内置 Embedding | 随机初始化、从零训练 |
bert_embedding.py | HuggingFace BERT 权重 | 需要上下文语义的中长文本 |
AlbertEmbedding.py | HuggingFace ALBERT 权重 | 显存受限、需要小参数量 |
接入方式上,四个类对外暴露的接口保持一致,都提供forward(input_ids, attention_mask)并返回[batch, seq_len, hidden_size]的张量。这样设计的好处是后续Encoder.py不用关心底层用的是静态向量还是预训练模型,只要 hidden_size 对齐即可。实际跑实验时可以先切到wor2vec_embedding把整体流程跑通,再切回 BERT 系列提升指标。中文预训练模型下载慢的问题在团队里经常碰到,常见做法是先在另一台机器上下载好权重,再通过pertrain_to_numpy.py转成 numpy 格式拷到离线服务器,这样bert_embedding.py加载时就不需要访问 HuggingFace 了。
2.3 二分标注矩阵构建逻辑
把文本和 spo_list 转成两张标注矩阵是预处理的核心环节。第一张是主实体标注矩阵,第二张是客实体标注矩阵,后者的维度多了一个关系数。参考该项目数据预处理流程,核心代码大致如下:
# process_raw_data.py 示意:构建二分标注矩阵 import numpy as np def build_spo_labels(tokens, spo_list, rel2id, max_seq_len): seq_len = min(len(tokens), max_seq_len) sub_tag = np.zeros((seq_len, 2), dtype=np.int32) # 主实体头/尾 spo_tag = np.zeros((len(rel2id), seq_len, 2), dtype=np.int32) for sub, rel, obj in spo_list: sub_start, sub_end = locate_entity(tokens, sub) obj_start, obj_end = locate_entity(tokens, obj) if sub_start < 0 or obj_start < 0: continue sub_tag[sub_start, 0] = 1 sub_tag[sub_end, 1] = 1 rid = rel2id[rel] spo_tag[rid, obj_start, 0] = 1 spo_tag[rid, obj_end, 1] = 1 return sub_tag, spo_taglocate_entity来自function.py,它在按字切分后的 token 序列里查找实体第一次出现的起止下标,找不到就返回 -1。注意这里主实体和客实体都只用两个位置表示:起始位置标记为 1,结束位置标记为 1,中间 token 全部为 0。这与序列标注惯用的 BIO 方案不同,二分标注矩阵只关心实体边界,后面配合指针式的解码逻辑,让模型聚焦于头尾位置的学习,而不是为实体内部的每个 token 分配标签。
提示:如果一条文本里同一关系出现多个三元组,比如“王菲演唱《红豆》和《但愿人长久》”,
spo_tag里同一关系维度下会同时标记两组起止位置。这种一对多的情况在客实体解码阶段需要额外处理,训练时模型会通过多个正例学会在同一关系下预测多个头尾位置,推理时用阈值筛选即可。
从构件上这四层:token 序列、主实体矩阵、关系矩阵、客实体矩阵,共同构成了模型的输入输出对齐基础。数据维度之间环环相扣,任何一个矩阵的尺寸没有对齐,后面模型 forward 时就必然报 shape mismatch。
3. 编码器与解码器:p_so_model与sp_o_model的分工
3.1 Encoder与Attention的选型
Encoder.py是整个模型共享的特征抽取层,输入是[batch, seq_len]的 token id,输出是[batch, seq_len, hidden_size]的上下文表示。如果配置为bert_embedding,这里直接封装BertModel;切到torchEmbedding时则退化为一个两层的 Transformer Encoder,内部包含多头自注意力。attention.py提供的是注意力机制的底层实现,包括 scaled dot-product attention 和多头拆分逻辑。
注意力模块在这个任务里承担的作用容易被低估。主实体和客实体在句子中可能相距很远,比如“王菲演唱的《红豆》后来被收录进她 1998 年发行的专辑《唱游》”这种句子,光靠词向量之间的点积很难捕捉跨片段依赖。Transformer 的每一层注意力都在拉近不同类型实体与关系词之间的语义距离,这也是为什么预训练模型版本明显优于静态词向量版本的根本原因,并非只是词向量质量高,而是 BERT 的深层注意力能建立起跨句法结构的关联路径。
3.2 主实体识别头设计
项目里s_model.py和p_so_model.py两个文件都与主实体识别有关,从命名习惯看,s_model是更轻量的实现,直接对编码器输出的每个 token 做 Sigmoid 二分类,判断其是否是主实体的起始或结束位置。p_so_model则可能是在此基础上增加了一个简单的指针网络,让起始和结束位置的预测信息互相影响。两个模块输出的都是两个[batch, seq_len]张量,分别对应起始位置概率和结束位置概率。
# p_so_model.py 结构示意 import torch.nn as nn class PSoModel(nn.Module): def __init__(self, hidden_size): super().__init__() self.start_head = nn.Linear(hidden_size, 1) self.end_head = nn.Linear(hidden_size, 1) def forward(self, encoder_out): start_logits = self.start_head(encoder_out).squeeze(-1) # [batch, seq] end_logits = self.end_head(encoder_out).squeeze(-1) return start_logits, end_logits这段代码里start_head和end_head是两个独立的线性层,没有共享参数,因为它们分别刻画实体的左边界和右边界特征。squeeze(-1)去掉最后一维只是为了把输出形状和标注矩阵对齐,方便后续直接计算损失。实际使用中可以考虑在这两个线性层之前接一个 LayerNorm 稳定训练,尤其是使用静态词向量时。
3.3 sp_o_model_v2的联合解码
sp_o_model.py、sp_o_model_2023.py、sp_o_model_v2.py三个文件展示了这一模块的迭代过程。最早版本直接为每个关系训练独立的头尾预测器,参数量随关系数线性增长;_2023版本加入了主实体向量的融合;_v2版本则是把主实体信息通过拼接方式引入客实体预测的完整实现,这也是多数实现中的通用做法——在编码器输出的每个 token 向量后面接上主实体的统一向量,再送入两个关系维度的线性层。
# sp_o_model_v2.py 核心 forward 逻辑 class SPOModelV2(nn.Module): def __init__(self, encoder, num_relations, hidden_size): super().__init__() self.encoder = encoder self.sub_start_head = nn.Linear(hidden_size, 1) self.sub_end_head = nn.Linear(hidden_size, 1) # 客实体头:输入维度翻倍是因为拼接了主实体向量 self.obj_start_head = nn.Linear(hidden_size * 2, num_relations) self.obj_end_head = nn.Linear(hidden_size * 2, num_relations) def forward(self, input_ids, attention_mask, sub_start, sub_end): seq_out = self.encoder(input_ids, attention_mask) # [batch, seq, hidden] sub_start_logits = self.sub_start_head(seq_out).squeeze(-1) sub_end_logits = self.sub_end_head(seq_out).squeeze(-1) sub_vec = self.extract_subject_vector(seq_out, sub_start, sub_end) sub_expanded = sub_vec.unsqueeze(1).expand(-1, seq_out.size(1), -1) concat_out = torch.cat([seq_out, sub_expanded], dim=-1) obj_start_logits = self.obj_start_head(concat_out) # [batch, seq, rel] obj_end_logits = self.obj_end_head(concat_out) return sub_start_logits, sub_end_logits, obj_start_logits, obj_end_logitsextract_subject_vector的作用是把主实体区域内所有 token 的向量取平均或取首尾拼接,得到一个固定维度的主实体表示。expand把这个向量复制到每个 token 位置,与原有上下文向量拼接后,客实体预测头就能同时看到“全局上下文”和“主实体信息”两部分特征。这正体现了二分标注中第二阶段“给定主实体再去判断客实体”的任务拆解逻辑。
4. 损失函数与训练:两个loss如何一起反传
4.1 BinaryLabelLoss的计算细节
主实体和客实体两个分支虽然结构不同,但损失函数可以统一为二值交叉熵。loss_function.py里的实现一般不是直接调nn.BCELoss,而是使用BCEWithLogitsLoss,把 Sigmoid 操作合并进损失计算,降低数值不稳定性。核心逻辑如下:
# loss_function.py 示意 import torch.nn as nn class BinaryLabelLoss(nn.Module): def __init__(self): super().__init__() self.bce = nn.BCEWithLogitsLoss(reduction='none') def forward(self, start_logits, end_logits, start_labels, end_labels, mask): # 主实体分支损失 start_loss = self.bce(start_logits, start_labels.float()) end_loss = self.bce(end_logits, end_labels.float()) # 客实体分支损失,形状为 [batch, seq, rel] obj_start_loss = self.bce(self.obj_start_logits, self.obj_start_labels.float()) obj_end_loss = self.bce(self.obj_end_logits, self.obj_end_labels.float()) mask = mask.unsqueeze(-1) # [batch, seq, 1] total = (start_loss + end_loss) * mask.squeeze(-1) total += (obj_start_loss + obj_end_loss) * mask return total.sum() / mask.sum()这里有一个容易被忽略的处理:mask 同时参与主实体损失和客实体损失的加权,padding 位置不计算损失。sum()除以mask.sum()而不是batch_size,是因为不同样本的实际有效长度不同,使用有效 token 数做分母能防止长文本样本在总损失中占比过大的问题。实际项目中两个分支的损失可以分别加权,比如主实体分支乘 1.0、客实体分支乘 0.8,不过多数情况两者等权就能收敛。
4.2 训练循环、优化器与梯度裁剪
模型训练部分采用标准的 BERT fine-tune 配置,优化器使用AdamW,学习率设为 3e-5。训练循环中值得注意的一点是梯度裁剪:对于这种双头结构,客实体分支的梯度量级通常比主实体分支大,如果不加clip_grad_norm_,训练后期容易出现某个 batch 的异常样本把主实体分支的权重拉飞的情况。一个可复现的训练循环如下:
# main.py 训练循环示意 import torch optimizer = torch.optim.AdamW(model.parameters(), lr=3e-5) total_steps = len(train_loader) * num_epochs scheduler = torch.optim.lr_scheduler.LinearLR(optimizer, total_iters=total_steps) model.train() for epoch in range(num_epochs): for batch in train_loader: input_ids = batch['input_ids'].to(device) attention_mask = batch['attention_mask'].to(device) sub_start = batch['sub_start'].to(device) sub_end = batch['sub_end'].to(device) obj_start = batch['obj_start'].to(device) obj_end = batch['obj_end'].to(device) start_logits, end_logits, obj_start_logits, obj_end_logits = model( input_ids, attention_mask) loss = criterion(start_logits, end_logits, obj_start_logits, obj_end_logits, sub_start, sub_end, obj_start, obj_end, attention_mask) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) optimizer.step() scheduler.step()batch来自继承Dataset的自定义类,__getitem__返回的每个字段都已经转换成 numpy 数组,collate_fn再统一补齐到 batch 内最大长度。LinearLR让学习率在整个训练过程中线性衰减到 0,配合 3e-5 的初始值,比固定学习率在最后的收敛阶段表现更稳定。首次跑这个项目时,PyTorch 环境的搭建和训练是分开的前置步骤:根据 CUDA 版本安装对应的 PyTorch 后,先运行python process_raw_data.py,确认能生成标注矩阵再启动main.py。
4.3 收敛异常的排查方向
训练时最常见的现象是 loss 下降缓慢但准确率不涨。优先检查三个方向:第一,train_data_sample.json中实体与文本字符是否严格匹配,locate_entity若返回 -1 会静默跳过样本,负样本过多会让模型偏向全零输出;第二,num_relations是否与all_50_schemas中的关系数一致,这个数字从一开始就决定客实体头的输出维度;第三,主实体预测为空的样本占比是否过高,这类样本会让客实体分支失去学习信号,可以考虑在训练循环中硬性丢弃主实体标注全为 0 的样本。
提示:如果使用静态词向量版本,学习率可以适当调大到 1e-4 区间,同时把
clip_grad_norm_的 max_norm 降到 1.0。预训练模型版本则保持 3e-5 以下,过大的学习率会让 BERT 部分的预训练权重快速失真,导致推理阶段输出全部塌缩到一个关系上。
5. 推理与验证:从概率矩阵还原三元组
5.1 阈值选择与重叠实体处理
推理阶段不再需要计算损失,模型 forward 出的四个 logits 张量经过 Sigmoid 变成概率,然后用阈值过滤出所有可能的头尾位置。主实体的解码相对直接:在[batch, seq_len]的概率矩阵上,同时大于阈值的头尾对构成候选主实体。客实体解码则多一层关系维度的处理,需要在num_relations个关系平面上分别执行相同的解码流程。
# 推理阶段从概率矩阵还原三元组 import torch def decode_subjects(start_probs, end_probs, threshold=0.5): starts = torch.nonzero(start_probs > threshold).squeeze(1) ends = torch.nonzero(end_probs > threshold).squeeze(1) subjects = [] for s in starts: valid_ends = ends[(ends >= s) & (ends <= s + 30)] if len(valid_ends) > 0: subjects.append((s.item(), valid_ends[0].item())) return subjects这里的约束条件ends >= s保证结束位置不早于起始位置,ends <= s + 30则限制实体最大长度为 30 个字符。阈值的选择直接影响抽取的精确率和召回率,通常 0.5 附近会有一个平衡点,实际调参时可以在验证集上画出 P/R 曲线再定。处理重叠实体时要特别注意:同一段文本可能同时存在“王菲演唱《红豆》”和“《红豆》属于流行歌曲”两个事实,主实体和客实体在位置上互相重叠,解码时必须做到每个候选主实体独立解码客实体,不能因为位置冲突就丢弃后一个候选。
5.2 用dev集做偏差验证
拿到模型输出的三元组列表后,与dev_data_sample.json里的spo_list做精确匹配,统计 F1 是验证模型效果的第一步。实践中更值得分析的是偏差类型,下面这张表总结了三种常见错误以及对应的排查入口:
| 错误现象 | 可能原因 | 排查入口 |
|---|---|---|
| 主实体识别正确但客实体全空 | 客实体阈值过高或关系维度错误 | 检查解码阈值与rel2id映射 |
| 同一关系下漏掉多个客实体 | 一对多场景解码不完整 | 确认spo_tag矩阵是否包含全部标注 |
| 三元组关系混淆 | schema 中相似关系过多 | 检查all_50_schemas的定义 |
把验证集的错误样本直接打印出来看,会比看整体 F1 更有针对性。比如“王菲演唱《红豆》”被抽成了“王菲的《红豆》”,问题通常出在主实体边界上;而“专辑《唱游》是王菲发行的”被抽成“唱游演唱王菲”,则说明客实体分支没有充分学到主实体与客实体之间的语义角色。建议在function.py里加一个validate_predictions函数,每次训练结束自动对 dev 集跑一轮解码,把预测结果与标注结果不一致的样本按上述分类输出到日志文件,这比只在终端打印 loss 能更快定位问题。修改解码方式或阈值后重新跑这一轮验证,对比同一批样本的准确率变化即可判断每次改动是否有效。
本文还有配套的精品资源,点击获取