拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyTorch与BERT实现三元组抽取:二分标注与关系分类实战

简介基于PyTorch框架的文本三元组信息抽取模型源码包面向自然语言处理方向的开发者与研究者解决从非结构化文本中抽取(主实体,关系,客实体)三元组的问题可应用于知识图谱构建、智能问答、语义搜索等场景项目采用二分标注策略将抽取过程拆解为“先识别主实体、再抽对应的客实体与关系”两个阶段有效降低建模复杂度并分别建模主客体与关系。压缩包共31个文件以23个Python源文件为核心涵盖数据清洗、分词、实体标注等预处理流程以及嵌入表示、编码器、模型主体、损失函数与解码模块另有JSON配置与样例数据、说明文档、开源许可等辅助文件整体仅348KB目录按data、model、utils等划分结构紧凑便于研读。此外源码中提供多种嵌入与模型变体如Albert与BERT嵌入、sp_o_model及其改进版本可对比不同设计的效果已有343人学习。通过完整源码与示例数据读者可快速理解二分标注思想掌握基于PyTorch的三元组抽取实现流程并可在此基础上针对具体业务场景进行训练与扩展适合作为相关课题的参考实现。1. 为什么三元组抽取要自己标注而不是直接套用开源模型信息抽取里最磨人的不是模型结构而是数据。做知识图谱、做舆情分析、做合同审查最终都要落在主谓宾三元组(subject, relation, object)上。市面上的通用抽取模型在新闻语料上表现尚可换到垂直领域比如医疗术语、工业设备故障描述、法律条文准确率立刻跳水。原因在于关系类别是领域定制的通用模型没见过的关系型样本它只能乱猜。于是问题变成了没有标注数据怎么办自己标。而“自己标”这件事如果按完整序列标注来做成本极高。一个句子要标出实体边界、实体类型、关系类别每标注一条样本要盯着屏幕确认四五次。二分标注的核心思路是把“复杂标注”拆成“简单判定”先用一个序列标注模型把主语和宾语各自标出来再用一个分类模型去判断这两个实体之间是否存在目标关系以及属于哪类关系。判定是二元的标注成本大幅下降模型结构也清晰可拆。这套方案特别适合两类人。一类是做领域知识图谱的工程师标签体系自己定、语料得自己造另一类是论文复现者想在 PyTorch 下快速验证关系抽取的效果又不想从零设计复杂的联合解码。下面这套方案从标注格式设计到模型训练再到推理输出都可以直接照搬。2. 二分标注的数据格式设计与 PyTorch Dataset 构建2.1 什么是二分标注实体标注与关系判定分离二分标注Binary Tagging最早由 CasRel 等工作引入核心是把一个联合抽取任务拆成两个子任务。第一个子任务是“主语识别”和“宾语识别”分别用 BIO 标注完成第二个子任务是“(subject, relation, object)”三要素是否成立的二分类判定。更精确地说是对“给定的主语片段、给定的宾语片段、给定的关系类别”做一个二分类成立或是不成立。这样做的好处有三个数据构造简单。每条训练样本只需要一个句子和一组事实三元组不需要事先把每个 token 都标成“关系内部/外部”因为关系判定在实体对层面完成。解码时可控。先抽主语再对每个主语去匹配候选宾语最后用关系分类器过滤天然支持一个句子包含多个三元组。扩展性。新增关系类别不需要重新标注所有语料只要对已有实体对做新的关系判定即可。形式上一条样本长这样{ text: 北京市人民政府位于中国首都北京是中国的一线城市。, triples: [ {subject: 北京市, relation: 位于, object: 北京}, {subject: 北京, relation: 是首都, object: 中国} ] }这里不采用经典的 BIEO 全序列标注因为全序列标注需要把每个 token 标为「B-主语」「I-宾语」等标签集合是关系类别与实体位置的笛卡尔积类别一多就爆炸。二分标注下主语用B-SUB/I-SUB/O标记宾语用B-OBJ/I-OBJ/O标记标签集合永远只有五种而关系类别用另一个独立的分类层去预测。2.2 将句子与三元组转换为模型输入张量第一步先把文本 token 化。这里推荐用 HuggingFace 的BertTokenizer或AutoTokenizer因为预训练语言模型的字符切分粒度对中文更友好。取最大长度 128超出部分截断不足部分用[PAD]补齐。第二步是把三元组映射成标签序列。主语识别任务需要构造一个长度等于序列长度的标签数组token 落入主语片段的开头位置标为1主语片段其他位置标为1完全用一位二值标签。如果你希望更精细可以把1拆成B和I但为了与“二分”语义对齐我一般直接用二值0表示非主语1表示是主语片段的一部分。同理构造宾语标签数组。第三步是构造关系判定矩阵。假设关系类别数为R句子被 token 化为N个 token主语片段有S个候选片段宾语片段有O个候选片段。真实样本中我们要让模型学会对每个主语片段和宾语片段的组合预测一个R维向量。这部分用真实三元组构造正样本用随机采样的实体对构造负样本。以 PyTorch 的Dataset类为例import torch from torch.utils.data import Dataset from transformers import BertTokenizer class BinaryTaggingDataset(Dataset): def __init__(self, data_list, tokenizer: BertTokenizer, max_len128, rel2idNone): self.data_list data_list self.tokenizer tokenizer self.max_len max_len self.rel2id rel2id def __len__(self): return len(self.data_list) def __getitem__(self, idx): item self.data_list[idx] text item[text] triples item[triples] # 编码句子返回 input_ids 和 attention_mask encoded self.tokenizer( text, truncationTrue, paddingmax_length, max_lengthself.max_len, return_tensorspt ) input_ids encoded[input_ids].squeeze(0) attention_mask encoded[attention_mask].squeeze(0) seq_len input_ids.size(0) # 主语标签长度 seq_len sub_heads torch.zeros(seq_len, dtypetorch.long) sub_tails torch.zeros(seq_len, dtypetorch.long) # 宾语标签长度 seq_len obj_heads torch.zeros(seq_len, dtypetorch.long) obj_tails torch.zeros(seq_len, dtypetorch.long) # 关系标签默认全为 -100忽略不计 loss rel_labels torch.full((seq_len, seq_len), -100, dtypetorch.long) # 用 tokenizer 的 offset_mapping 定位实体片段 tokens self.tokenizer(text, truncationTrue, max_lengthself.max_len) offsets tokens[offset_mapping] for triple in triples: sub triple[subject] obj triple[object] rel triple[relation] if rel not in self.rel2id: continue rel_id self.rel2id[rel] # 在 token 级别定位主语 sub_start, sub_end self._find_span(text, offsets, sub) if sub_start is None: continue obj_start, obj_end self._find_span(text, offsets, obj) if obj_start is None: continue # 主语起止位置标记 sub_heads[sub_start] 1 sub_tails[sub_end] 1 obj_heads[obj_start] 1 obj_tails[obj_end] 1 # 关系矩阵主语起始 token 行宾语起始 token 列指向关系 id rel_labels[sub_start, obj_start] rel_id return { input_ids: input_ids, attention_mask: attention_mask, sub_heads: sub_heads, sub_tails: sub_tails, obj_heads: obj_heads, obj_tails: obj_tails, rel_labels: rel_labels, } staticmethod def _find_span(text, offsets, span): # 将字符串坐标转为 token 坐标返回 (start_token_idx, end_token_idx) # 如果没找到返回 (None, None) start_char text.find(span) if start_char -1: return None, None end_char start_char len(span) start_token None end_token None for i, (s, e) in enumerate(offsets): if s start_char: start_token i if start_token is not None and e end_char: end_token i break return start_token, end_token代码说明sub_heads和sub_tails分别标记主语片段的起点和终点值为 1 的位置表示该 token 是实体边界。rel_labels是一个N×N的矩阵rel_labels[i][j]表示第i个 token 是主语起点、第j个 token 是宾语起点时二者之间的关系类别。-100是 PyTorchCrossEntropyLoss默认忽略的类别编号专门用来屏蔽非实体对位置。参数说明max_len设为 128 是因为 BERT 类模型默认最大输入长度为 512中文短文本用 128 足够且能显著加快训练速度。rel2id是关系名到数字 ID 的映射字典需要在训练前从全部数据中统计得出注意把-100排除在映射之外。2.3_find_span的边界陷阱如果_find_span返回None后续所有索引操作都会抛TypeError。实际使用中句子可能与标注实体有微小差异比如全角空格、标点粘连。这里给出一个更稳健的做法直接用字符查找找不到时做一次“去除空格”后的再查找。staticmethod def _find_span(text, offsets, span): norm_text text.replace( , ) norm_span span.replace( , ) start_char norm_text.find(norm_span) if start_char -1: return None, None end_char start_char len(norm_span) start_token None end_token None for i, (s, e) in enumerate(offsets): # offset_mapping 包含特殊 token跳过 (0,0) if s 0 and e 0: continue if s start_char: start_token i if start_token is not None and e end_char: end_token i break return start_token, end_token注意这里去空格之后start_char和end_char仍然是在原始文本上的位置但offsets是原始 tokenizer 给出的因此必须用“去空格后”的字符坐标去对齐原始offsets这一步容易出错。更保险的替代方案是把整个句子按字符逐个 token 化tokenizer设置为按字切分这样offsets就与字符坐标一一对应。3. 基于 PyTorch 的三元组抽取模型结构与损失函数3.1 编码层为什么选用预训练模型做底部模型底部直接使用BertModel或BertForPreTraining的输出原因是中文语义复杂尤其是“北京位于中国”和“中国位于北京”这种带方向性的关系如果不结合上下文单靠词向量无法区分。预训练模型输出的每个 token 维度是 768BERT-base或 1024BERT-large这个向量既包含词义也包含位置信息已经足够为后续两个任务提供输入。如果显存有限可以使用bert-base-chinese大约 110M 参数在单卡 12GB 上 batch size 为 16 无压力。如果追求更高精度可以换成RoBERTa-wwm-ext对中文全词掩码的支持更好。3.2 双头结构主体定位头与关系分类头模型结构分为三个部分。第一部分是共享编码层BertModel。第二部分是主语定位头输入是sentence_outputbatch, seq_len, hidden_size通过一个线性层映射到 (batch, seq_len, 2)两条通道分别对应主语起点和主语终点。第三部分是宾语定位头加关系分类头输入是主语起点对应的编码向量、宾语起点对应的编码向量拼接后通过一个 MLP 输出关系类别概率。这里的逻辑是先找到主语再以主语为条件去找宾语和关系。这种设计让模型天然支持一个句子中多个主语每个主语独立地去找它的宾语和关系避免传统序列标注中“实体重叠”无法处理的问题。import torch import torch.nn as nn from transformers import BertModel, BertConfig class BinaryTaggingModel(nn.Module): def __init__(self, pretrained_namebert-base-chinese, num_relations10): super().__init__() self.bert BertModel.from_pretrained(pretrained_name) hidden_size self.bert.config.hidden_size # 主语起点/终点二元分类 self.sub_head_cls nn.Linear(hidden_size, 1) self.sub_tail_cls nn.Linear(hidden_size, 1) # 宾语起点/终点二元分类 self.obj_head_cls nn.Linear(hidden_size, 1) self.obj_tail_cls nn.Linear(hidden_size, 1) # 关系分类输入为 [主语起点向量; 宾语起点向量] self.rel_cls nn.Sequential( nn.Linear(hidden_size * 2, hidden_size), nn.ReLU(), nn.Dropout(0.1), nn.Linear(hidden_size, num_relations), ) def forward(self, input_ids, attention_mask): outputs self.bert(input_idsinput_ids, attention_maskattention_mask) sequence_output outputs.last_hidden_state # (batch, seq_len, hidden) sub_heads_logits self.sub_head_cls(sequence_output).squeeze(-1) # (batch, seq_len) sub_tails_logits self.sub_tail_cls(sequence_output).squeeze(-1) obj_heads_logits self.obj_head_cls(sequence_output).squeeze(-1) obj_tails_logits self.obj_tail_cls(sequence_output).squeeze(-1) # 用 argmax 提取每个 token 位置的主语起点概率和宾语起点概率 # 具体做法在编码层之上取每个 token 向量用于计算关系 # 这里省略批内 token 的提取详见训练循环 return { sub_heads: sub_heads_logits, sub_tails: sub_tails_logits, obj_heads: obj_heads_logits, obj_tails: obj_tails_logits, sequence_output: sequence_output, }模型输出中四个候选定位头的 logits 都需要与对应标签计算损失。损失函数用二元交叉熵还是带 logits 的BCEWithLogitsLoss取决于标签是二值0/1。这里给一个完整的损失计算函数它同时考虑主语定位、宾语定位和关系分类三个任务3.3 多任务损失如何把三个损失平衡在一起主体定位头的损失关注的是每个 token 是否是一个主语片段或宾语片段的起点/终点这个损失天然是类别不平衡的因为一个句子中真正的边界 token 极少。直接使用BCEWithLogitsLoss模型会学到把所有位置都预测为 0因为 0 的占比极高。解决方法是给正样本更高的权重或者使用 Focal Loss但更简单的做法是给pos_weight传入大于 1 的值让模型更重视少数类别。def compute_loss(model_output, batch, relation_weight1.0): sub_heads_pred model_output[sub_heads] sub_tails_pred model_output[sub_tails] obj_heads_pred model_output[obj_heads] obj_tails_pred model_output[obj_tails] # 标签已在 Dataset 中构造 sub_heads_label batch[sub_heads].float() sub_tails_label batch[sub_tails].float() obj_heads_label batch[obj_heads].float() obj_tails_label batch[obj_tails].float() # 计算四个二分类损失用 BCEWithLogitsLoss bce nn.BCEWithLogitsLoss() loss_sub_h bce(sub_heads_pred, sub_heads_label) loss_sub_t bce(sub_tails_pred, sub_tails_label) loss_obj_h bce(obj_heads_pred, obj_heads_label) loss_obj_t bce(obj_tails_pred, obj_tails_label) # 关系分类损失只有在实体对位置才计算 rel_logits model_output[rel_logits] # (batch, seq_len, seq_len, num_relations) rel_labels batch[rel_labels] # (batch, seq_len, seq_len) # 直接使用 CrossEntropyLoss-100 会被忽略 loss_rel nn.CrossEntropyLoss()(rel_logits.view(-1, rel_logits.size(-1)), rel_labels.view(-1)) total_loss loss_sub_h loss_sub_t loss_obj_h loss_obj_t relation_weight * loss_rel return total_lossrelation_weight是一个超参数我一般设为 1.0 或 1.5。如果你的关系类别很多超过 20 类关系分类任务的难度更大把relation_weight提到 2.0 可以让模型优先拟合关系实体定位头则靠较大的数据量慢慢收敛。注意rel_logits在forward函数中需要额外计算这里补上缺失的部分def forward(self, input_ids, attention_mask): outputs self.bert(input_idsinput_ids, attention_maskattention_mask) sequence_output outputs.last_hidden_state # (batch, seq_len, hidden) # 四个定位头 sub_heads_logits self.sub_head_cls(sequence_output).squeeze(-1) sub_tails_logits self.sub_tail_cls(sequence_output).squeeze(-1) obj_heads_logits self.obj_head_cls(sequence_output).squeeze(-1) obj_tails_logits self.obj_tail_cls(sequence_output).squeeze(-1) # 关系分类取每个 token 向量拼接主语起点向量和宾语起点向量 # 维度(batch, seq_len, seq_len, hidden*2) batch_size, seq_len, hidden sequence_output.size() # 扩展序列维度构造全包含矩阵 sub_ext sequence_output.unsqueeze(2).expand(batch_size, seq_len, seq_len, hidden) obj_ext sequence_output.unsqueeze(1).expand(batch_size, seq_len, seq_len, hidden) pair_feat torch.cat([sub_ext, obj_ext], dim-1) # (batch, seq_len, seq_len, hidden*2) rel_logits self.rel_cls(pair_feat) # (batch, seq_len, seq_len, num_relations) return { sub_heads: sub_heads_logits, sub_tails: sub_tails_logits, obj_heads: obj_heads_logits, obj_tails: obj_tails_logits, rel_logits: rel_logits, }这里expand不会真正复制数据内存占用可控。但pair_feat做cat后会产生(batch, seq_len, seq_len, hidden*2)的张量当seq_len128、hidden768、batch8时内存占用约为8*128*128*1536*2字节约 3GB。如果你的显卡只有 8GB把这个 max_len 降到 64或者只用真实实体对位置计算关系损失改为稀疏矩阵方式比直接构造全矩阵划算得多。4. 源码级训练流程从数据加载到模型收敛的完整命令4.1 训练环境准备Anaconda 配置 PyTorch 环境这一节直接给可执行的命令。先用 Anaconda 创建一个新的虚拟环境命名为tripletPython 版本选 3.9 或 3.10这两个版本与 PyTorch 和 Transformers 的兼容性最稳定。conda create -n triplet python3.9 conda activate triplet pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install transformers datasets evaluate tqdm如果机器没有 NVIDIA GPU把--index-url去掉直接pip install torch安装 CPU 版本训练时速度会慢 10 到 20 倍但跑通流程没问题。建议至少一张 8GB 显存的显卡否则seq_len128时显存会爆此时把max_len改为 64并关掉pair_feat的显式构造用torch.einsum或直接循环处理实体对。4.2 训练脚本多任务梯度回传与学习率调节写一个完整的训练循环包含优化器、学习率调度器和早停逻辑。对于这个模型AdamW 搭配线性衰减是最常用、最稳的组合。BERT 底层和头部学习率不同底层用 2e-5头部用 1e-4这样可以防止预训练参数被破坏。import torch from torch.utils.data import DataLoader from transformers import BertTokenizer, AdamW, get_linear_schedule_with_warmup from tqdm import tqdm def train_model(model, train_data, dev_data, rel2id, epochs5, batch_size16, lr2e-5, head_lr1e-4): tokenizer BertTokenizer.from_pretrained(bert-base-chinese) train_dataset BinaryTaggingDataset(train_data, tokenizer, max_len128, rel2idrel2id) train_loader DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue) dev_dataset BinaryTaggingDataset(dev_data, tokenizer, max_len128, rel2idrel2id) dev_loader DataLoader(dev_dataset, batch_sizebatch_size, shuffleFalse) no_decay [bias, LayerNorm.weight] optimizer_grouped_parameters [ { params: [p for n, p in model.named_parameters() if not any(nd in n for nd in no_decay) and rel_cls not in n], weight_decay: 0.01, lr: lr, }, { params: [p for n, p in model.named_parameters() if any(nd in n for nd in no_decay) and rel_cls not in n], weight_decay: 0.0, lr: lr, }, { params: [p for n, p in model.named_parameters() if rel_cls in n], weight_decay: 0.01, lr: head_lr, }, ] optimizer AdamW(optimizer_grouped_parameters) total_steps len(train_loader) * epochs scheduler get_linear_schedule_with_warmup( optimizer, num_warmup_stepsint(0.1 * total_steps), num_training_stepstotal_steps ) best_f1 0.0 model.train() for epoch in range(epochs): loop tqdm(train_loader, descfEpoch {epoch1}/{epochs}) total_loss 0.0 for batch in loop: input_ids batch[input_ids] attention_mask batch[attention_mask] output model(input_ids, attention_mask) loss compute_loss(output, batch) loss.backward() optimizer.step() scheduler.step() optimizer.zero_grad() total_loss loss.item() loop.set_postfix(lossloss.item()) print(fEpoch {epoch1} loss: {total_loss / len(train_loader):.4f}) # 验证阶段简化只计算准确率 dev_loss evaluate_model(model, dev_loader) print(fDev loss: {dev_loss:.4f})参数说明epochs5是一个参考值BERT 微调任务通常 3 到 5 轮就会收敛再多容易过拟合。warmup_steps10%的意思是前 10% 的训练步数里学习率从 0 线性升到设定值这能显著提升稳定性避免一开始就大步长冲乱预训练参数。4.3 训练时常见的三个坑负样本丢失、类别不平衡、梯度爆炸负样本丢失是最隐蔽的问题。rel_labels矩阵里大部分位置是-100在构造损失时被忽略。如果训练数据里每个句子只有一个三元组那么模型实际能接收到的正样本极少数关系分类器会严重偏向“没有关系”。解决办法是在构造负样本时把同一个句子中不存在三元组的实体对随机抽取一部分标记为关系 ID 为 0表示无关系并把“无关系”定义为一个独立类别。在rel2id中预留id0给 “NONE”其他关系从 1 开始。类别不平衡的处理对于sub_heads和sub_tails正负样本比例可能达到 1:100BCEWithLogitsLoss的pos_weight参数直接设为 5 到 10训练会更稳定。对rel_cls的分类任务也可以用class_weight但注意不要把-100所在位置计入权重。梯度爆炸在头几个 epoch 里偶尔出现尤其当head_lr设置太大。观察 loss 曲线如果出现 NaN先降低head_lr到 1e-5再把batch_size减半最后再考虑梯度裁剪# 在 loss.backward() 之后、optimizer.step() 之前 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)max_norm1.0是 BERT 微调常用的梯度裁剪阈值它不改变梯度方向只把梯度向量的 L2 范数限制在 1 以内防止极端样本把参数推得太远。5. 推理与解码从概率到三元组的完整后处理管线5.1 解码流程先在序列上找主语再在主语约束下找宾语和关系训练完成后模型的前向输出是一堆概率值要把它们转换成可读的三元组列表需要一个解码函数。逻辑分三步第一用阈值过滤主语起点与终点。设定threshold0.5所有sub_heads概率大于 0.5 的位置记为potential_sub_starts所有sub_tails概率大于 0.5 的位置记为potential_sub_ends。如果起点索引大于等于终点索引说明该片段非法直接丢弃。第二对每个合法主语片段提取该片段第一个 token 的编码向量并用它去和整个序列每个 token 的起点判断宾语。为了方便直接复用模型推理时的rel_logitsrel_logits[sub_start, obj_start, :]给出了该主语起点和所有宾语起点组合的关系类别概率。对每个宾语起点同样用阈值过滤其终点概率组合出宾语片段。第三过滤关系。rel_logits的最后一个维度做softmax取概率最高的关系类别。如果该关系类别是NONEID 为 0则丢弃这个实体对否则保留三元组。def decode_triples(model, tokenizer, text, rel2id, threshold0.5): model.eval() id2rel {v: k for k, v in rel2id.items()} encoded tokenizer(text, truncationTrue, max_length128, return_tensorspt) with torch.no_grad(): output model(encoded[input_ids], encoded[attention_mask]) sub_heads torch.sigmoid(output[sub_heads][0]).cpu().numpy() sub_tails torch.sigmoid(output[sub_tails][0]).cpu().numpy() obj_heads torch.sigmoid(output[obj_heads][0]).cpu().numpy() obj_tails torch.sigmoid(output[obj_tails][0]).cpu().numpy() rel_logits output[rel_logits][0].cpu().numpy() # (seq_len, seq_len, num_relations) tokens tokenizer.convert_ids_to_tokens(encoded[input_ids][0]) triples [] sub_starts [i for i in range(len(sub_heads)) if sub_heads[i] threshold] sub_ends [i for i in range(len(sub_tails)) if sub_tails[i] threshold] for s_start in sub_starts: for s_end in sub_ends: if s_start s_end: continue subject tokenizer.convert_tokens_to_string(tokens[s_start:s_end1]) subject subject.replace( , ).replace([CLS], ).replace([SEP], ) obj_starts [i for i in range(len(obj_heads)) if obj_heads[i] threshold] obj_ends [i for i in range(len(obj_tails)) if obj_tails[i] threshold] for o_start in obj_starts: # 关系类别从 rel_logits 取 argmax rel_id rel_logits[s_start, o_start].argmax() if rel_id 0: # NONE continue for o_end in obj_ends: if o_start o_end: continue object_str tokenizer.convert_tokens_to_string(tokens[o_start:o_end1]) object_str object_str.replace( , ).replace([CLS], ).replace([SEP], ) relation id2rel[rel_id] triples.append((subject, relation, object_str)) break # 每个宾语起点只取最近的终点片段 # 一个主语可能匹配多个宾语这里不 break return triples这段代码有一个关键优化点obj_ends的先遍历顺序是升序的而break在找到第一个合法片段后就跳出内层for o_end循环这样每个o_start只对应一个宾语片段。如果你希望支持多个相同开始位置的宾语很少见可以把break去掉。5.2 用真实文本测试推理效果假设我们训练好了模型输入一句话python inference.py --text 苹果公司成立于1976年总部位于美国加州库比蒂诺。预期输出三个三元组(苹果公司, 成立于, 1976年) (苹果公司, 总部位于, 美国加州库比蒂诺)如果输出为空检查三个位置。第一rel2id里是否定义了成立于和总部位于未定义的关系在训练时被跳过推理时自然也不会出现。第二阈值 0.5 是否过高如果模型的实体定位头训练不充分概率普遍在 0.3~0.5 之间波动此时调低到 0.3 再试。第三主语与宾语的起始 token 是否包含在[CLS]或[SEP]附近[CLS]的 token 位置会参与边界过滤但一般不会是实体边界。6. 从单卡训练到源码落地的三个进阶优化6.1 使用 PyTorch 混合精度训练提速如果你的显卡是 RTX 30 系或更新支持自动混合精度AMP。在训练脚本中加三行代码就能在几乎不掉精度的情况下把训练速度提升 40% 到 70%。注意GradScaler会自动处理梯度缩放不要手动乘 65536。from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for batch in train_loader: with autocast(): output model(batch[input_ids], batch[attention_mask]) loss compute_loss(output, batch) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad()使用后loss.item()的值可能略大这是因为混合精度计算时部分算子以半精度计算但梯度的数值稳定性由GradScaler保障。如果出现 loss 为 NaN检查autocast作用域内是否有一些自定义的负对数运算必要时对这些运算使用float()强制转回单精度。6.2 在 PyTorch 中嵌入源码级 Debug用分布式打印定位异常 token这个模型最头疼的问题是实体定位准确率高但关系分类差或者反之。要定位是哪一个模块出了问题最直接的方法是在decode_triples中打印每个主语和宾语片段对应的关系 logits 分布。probs torch.softmax(torch.tensor(rel_logits[s_start, o_start]), dim-1) print(fsubject{subject}, object{object_str}, probs{probs.tolist()})关注关系类别的概率分布。如果所有实体对的概率都很接近均匀分布说明关系分类头欠拟合增加relation_weight或增加负样本比例如果概率集中在某个错误的关系上说明训练数据该关系类别下正样本太少或者主语/宾语边界的 token 对齐有问题尤其要检查_find_span是否把“苹果公司”切成了“苹果公”和“司”两个片段。6.3 源码输出结构把模型、数据、推理打包成可复现的项目目录一个可复现的三元组抽取项目至少要有 6 个文件每个文件职责单一。不建议把数据加载、模型定义、训练逻辑全部写在同一个.py文件里后期改一个参数要翻数百行。project/ ├── config.py # 超参数与路径配置 ├── data/ │ ├── train.json │ ├── dev.json │ └── test.json ├── dataset.py # BinaryTaggingDataset 类 ├── model.py # BinaryTaggingModel 类 ├── train.py # 训练循环入口 ├── inference.py # 解码与推理脚本 └── requirements.txtconfig.py里把max_len、batch_size、epochs、lr、head_lr、threshold、rel2id的保存路径全部集中管理train.py运行python train.py --data_dir data/ --output_dir output/结束时用torch.save(model.state_dict(), output/model.pt)保存模型权重同时把rel2id用json.dump单独存一份。推理时只加载model.pt和rel2id.json不需要重新读取原始数据这也是工程上最常用的部署格式。6.4 最后一道验证跑一个 20 条数据的过拟合测试拿到新语料后不要立刻全量训练。先取 20 条样本在epochs50、max_len64下强行训练看模型能否把损失压到接近 0、训练集上的三元组是否完全抽对。如果 20 条数据都过拟合不了说明模型结构或数据对齐有 bug此时调参没用要回头检查_find_span、rel_labels构造和损失函数里-100的 mask 位置。这一步能省下至少半天调参时间也是所有源码调试中最具性价比的操作。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门