ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

中文医学文本实体关系抽取:从BERT训练到Flask部署实战

2026/9/14 2:32:00 拓冰建站 浏览量
中文医学文本实体关系抽取:从BERT训练到Flask部署实战 简介面向中文医学文本实体关系抽取任务这份Python源码包针对期末大作业与课程设计场景适合高校人工智能、自然语言处理方向的学生参考也能帮助NLP入门者快速理清实体识别与关系抽取的工程实现思路。包内共13个文件其中12个Python脚本、1个使用说明txt压缩包仅28KB脚本按功能可分为模型定义、实体识别与关系抽取主逻辑、数据工具以及基于Flask的接口服务等模块结构清晰、分工明确可独立运行或按需修改。目前已有514人学习下载代码完整且可运行配合使用说明能快速搭建环境并跑通示例。对需要提交可演示程序或积累医学文本挖掘实战经验的学习者这份小巧的资源提供了现成框架可以在此基础上替换数据、调整模型参数并拓展功能便于二次开发与学习交流。1. 中文医学文本实体关系抽取一次跑通训练到部署中文医学文本实体关系抽取拆开看就是两件事在病历文本里找出症状、疾病、药物等实体再判断它们之间是治疗、伴随还是适应症关系。比如临床病历里“患者因持续性胸痛入院诊断为急性心肌梗死给予阿司匹林后症状缓解”至少含“胸痛”症状、“急性心肌梗死”疾病、“阿司匹林”药物三个实体以及伴随和治疗两层关系。这套基于Python的实体关系抽取源码用BERT做序列标注识别实体再用关系分类器判断实体对关系最后通过Flask把推理包装成HTTP服务。从数据构造、训练到部署都有对应脚本适合作为人工智能课程设计或期末大作业参考也适合当医疗NLP基线工程来跑一遍。2. 实体关系抽取的项目骨架const.py、data_structures.py 与 models.py 的设计先看标签体系和数据承载结构因为这两个文件决定了后续所有脚本里循环操作的对象类型。const.py定义标签体系data_structures.py定义实体和关系的承载对象models.py定义模型主干run_entity.py和run_relation.py分别负责实体识别和关系抽取的训练与推理。把这几处看明白后面读代码就不会被来回跳转的变量绕晕。2.1 实体与关系标签体系const.py 的常量约束这个项目的const.py定义实体类型、关系类型和BIO标注到id的映射。中文医学文本中实体类型通常不会太多常见的是疾病、症状、药物、治疗四类必要时再加检查、部位等类别。ENTITY_TYPES [Disease, Symptom, Drug, Treatment] RELATION_TYPES [treatment, symptom, indication, adverse] TYPE2ID {t: i for i, t in enumerate(ENTITY_TYPES)} REL2ID {r: i for i, r in enumerate(RELATION_TYPES)} BIO2ID {O: 0} for i, t in enumerate(ENTITY_TYPES): BIO2ID[fB-{t}] i * 2 1 BIO2ID[fI-{t}] i * 2 2上面的id分配有一个潜在技巧同一个实体的B-标签和I-标签在id上是相邻的比如B-Disease1、I-Disease2B-Symptom3、I-Symptom4。这样后处理时判断“当前token是否属于上一实体的延续”只需要看两个id的差值是否为1避免维护一份从字符串标签到实体类型的额外查找表解码逻辑也更紧凑。关系类型把临床推理限定在四类核心关系上treatment是“药物-疾病治疗”symptom是“症状-疾病关联”indication是“药物-适应症”adverse是“药物-不良反应”。从工程角度这四类关系已经能覆盖大多数电子病历的实体关系建模需求。如果换成专业医学标注数据集通常还会有concurrent并发、transfer转移等类型扩展时需要同步改REL2ID和关系分类层的输出维度。2.2 data_structures.py 中用 dataclass 承载实体与关系data_structures.py定义了Entity、Relation、MedicalSample三个dataclass。它们把分散的标签统一成对象后续的训练集构造、解码后处理、API返回格式都直接复用这些对象。from dataclasses import dataclass, field from typing import List dataclass class Entity: text: str # 实体文本 type: str # 实体类型Disease / Symptom / ... start: int # 起始字符下标闭区间 end: int # 结束字符下标开区间 dataclass class Relation: head: Entity # 头实体 tail: Entity # 尾实体 rel_type: str # 关系类型 dataclass class MedicalSample: text: str # 原始文本 entities: List[Entity] field(default_factorylist) relations: List[Relation] field(default_factorylist) def to_bio_labels(self) - List[str]: labels [O] * len(self.text) for ent in self.entities: labels[ent.start] fB-{ent.type} for i in range(ent.start 1, ent.end): labels[i] fI-{ent.type} return labelsto_bio_labels按原始字符长度初始化一个全O列表再把实体边界填入B和I标签。这里的start和end必须与原文的字符下标严格对齐。最容易出现的问题是训练数据来自其他标注工具时文件里保存的偏移是按词算或按字节算的中文按UTF-8编码时一个汉字占3个字节偏移不一致会导致标签错位。2.3 models.py 的双头模型实体标注与关系分类共用 BERT 编码器models.py是整个项目的核心模型文件。BERT部分负责把每个token编码成语义向量实体头做逐token分类关系头做实体对分类。共享编码器能缓解医学语料标注量少的问题实体识别学到的边界知识可以帮助关系分类关系分类的语义约束也能反过来矫正实体边界这是这套代码相比纯pipeline式两套独立模型的优势。2.3.1 实体识别头Linear CRF 的序列标注实体识别用BERT最后一层的输出过一层线性分类器再接一个CRF。CRF比单纯用softmax做逐token分类的优势在于它显式建模标签之间的转移概率例如“B-后面不能直接接另一个B-”这类约束解码时也是用Viterbi求全局最优序列。import torch.nn as nn from transformers import BertModel from torchcrf import CRF class EntityTagger(nn.Module): def __init__(self, bert_path: str, num_tags: int): super().__init__() self.bert BertModel.from_pretrained(bert_path) self.dropout nn.Dropout(0.1) self.classifier nn.Linear(self.bert.config.hidden_size, num_tags) self.crf CRF(num_tags, batch_firstTrue) def forward(self, input_ids, attention_mask, labelsNone): outputs self.bert(input_idsinput_ids, attention_maskattention_mask) seq_out self.dropout(outputs.last_hidden_state) logits self.classifier(seq_out) if labels is not None: loss -self.crf(logits, labels, maskattention_mask.bool()) return loss return self.crf.decode(logits, maskattention_mask.bool())self.crf(logits, labels, mask)返回的是负对数似然前面加负号就变成最小化目标。当labelsNone时CRF的decode内部走Viterbi路径返回的是每个样本的标签id序列形状是List[List[int]]。这里最容易被忽略的是maskattention_mask.bool()这一步如果不传maskpadding部分的标签会参与转移计算解码结果会出现起始标签出现在序列中间这类非法情况。2.3.2 关系分类头实体向量拼接后过线性层关系分类的输入不是整个序列而是实体对。实现上先用span平均池化取头实体和尾实体各自的向量拼接后过线性分类器。class RelationClassifier(nn.Module): def __init__(self, bert_path: str, num_relations: int): super().__init__() self.bert BertModel.from_pretrained(bert_path) self.fc nn.Linear(self.bert.config.hidden_size * 2, num_relations) def forward(self, input_ids, attention_mask, head_spans, tail_spans): seq_out self.bert(input_idsinput_ids, attention_maskattention_mask).last_hidden_state head_vecs, tail_vecs [], [] for i in range(seq_out.size(0)): hs, he head_spans[i] # (start, end) 字符级闭开区间 ts, te tail_spans[i] head_vecs.append(seq_out[i, hs:he 1].mean(dim0)) tail_vecs.append(seq_out[i, ts:te 1].mean(dim0)) head_vec torch.stack(head_vecs) tail_vec torch.stack(tail_vecs) return self.fc(torch.cat([head_vec, tail_vec], dim-1))这里用循环逐个样本做span池化主要是为了处理batch内不同实体span长度不一致的情况代码直观且不容易出错。追求性能时可以先把每个span转成mask矩阵用矩阵乘法和求和除以span长度来并行化但逻辑上不如循环好调试。这段模型中两个任务共享同一个BERT权重。纯单任务的实现可以各自独立训练但多任务共享编码器能有效提升医学小样本语料的效果因为两个任务处于同一编码空间时会互相约束权重的更新方向。2.4 字符级BIO标注为什么中文病历不依赖分词器中文电子病历里有大量长专名例如“急性非ST段抬高型心肌梗死”通用分词器很可能将其切碎导致后续实体边界完全错乱。字符级BIO标注把每个汉字作为独立token进行序列标注模型自己学习边界的判断。| 原文 | 患 | 者 | 因 | 胸 | 痛 | 入 | 院 | | 诊 | 断 | 为 | 急 | 性 | 心 | 肌 | 梗 | 死 | |---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---| | 标注 | O | O | O | B-Symptom | I-Symptom | O | O | O | O | O | O | O | B-Disease | I-Disease | I-Disease | I-Disease | I-Disease | I-Disease |表中的“胸痛”标注为B-Symptom、I-Symptom“急性心肌梗死”标注为B-Disease加4个I-Disease。CRF会学到相邻标签的转移规则例如I-Disease只能出现在B-Disease或其他I-Disease之后因此解码结果天然不会出现断裂的实体。3. 实体与关系抽取的代码实现run_entity.py 与 relation.py 拆解这一章偏实战。实体抽取从run_entity.py进入训练循环关系抽取在relation.py里完成候选实体对构造和关系类型的解码。理解这两条链路后整个项目的核心逻辑就掌握了。3.1 run_entity.py 的训练循环与超参数设置run_entity.py负责实体识别模型的训练与验证。整个训练循环与常规BERT微调一致区别在于loss来自CRF层输入侧要同时准备input_ids、attention_mask和BIO标签labels。from torch.utils.data import DataLoader from transformers import BertTokenizer, get_linear_schedule_with_warmup tokenizer BertTokenizer.from_pretrained(./bert-base-chinese) model EntityTagger(./bert-base-chinese, num_tagslen(BIO2ID)) optimizer torch.optim.AdamW(model.parameters(), lr2e-5) scheduler get_linear_schedule_with_warmup( optimizer, num_warmup_steps200, num_training_steps5000 ) for epoch in range(3): for step, batch in enumerate(entity_loader): input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) labels batch[labels].to(device) loss model(input_ids, attention_mask, labelslabels) optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() scheduler.step() if step % 50 0: pred_ids model(input_ids, attention_mask) acc prefix_accuracy(pred_ids, labels, attention_mask) print(fepoch {epoch} step {step} loss{loss.item():.4f} acc{acc:.4f})prefix_accuracy是自定义辅助函数一般放在utils.py里只统计attention mask为1的位置上pred_ids与labels的相等率不把padding部分计入分母。clip_grad_norm_在BERT微调中建议打开尤其是在batch_size较大或输入序列较长时可以防止尾部层梯度爆炸导致loss突然变成nan。BERT微调时最常用的超参数组合如下。如果GPU显存不足可以把batch_size降到8同时用梯度累积来模拟更大的batch超参数推荐值说明batch_size1632显存允许时尽量大learning_rate2e-5BERT微调通用基线max_len128覆盖电子病历句子主体warmup_ratio0.1前10%步数线性升温grad_clip_norm1.0防止梯度爆炸epochs35医学小语料3轮足够再多容易过拟合warmup机制在预训练模型微调中几乎是必需的。训练初期BERT权重还停留在预训练分布上直接用较大学习率会导致loss剧烈震荡前10%步数把学习率从0线性升到2e-5让模型平滑过渡到领域语料。3.2 relation.py 的关系候选构造实体对过滤与训练标签关系抽取的输入是“文本 实体对”。听起来可以拿实体识别结果直接穷举所有两两组合但实际工程里需要过滤掉明显无意义的重叠实体对。relation.py中通常包含这样一个构造函数def build_rel_pairs(text: str, entities: List[Entity]) - List[dict]: pairs [] for i in range(len(entities)): for j in range(len(entities)): if i j: continue head, tail entities[i], entities[j] if head.start tail.end or tail.start head.end: # 两个span不重叠时才是有效候选 pairs.append({ head_span: (head.start, head.end - 1), tail_span: (tail.start, tail.end - 1), head_text: head.text, tail_text: tail.text, }) return pairshead.start tail.end或反向条件用于排除重叠实体。举例来说一段文本里“胃炎”是实体而“慢性胃炎”作为更高层的复合实体也出现在标注里时两个span有重叠部分把它们都放进关系候选会造成重复且语义冲突的训练实例因为模型无法同时学习“胃炎→慢性胃炎”这一对子实体之间的关系。实际运行run_relation.py时数据的构造方式是把原始文本经过tokenizer编码将head_span和tail_span映射到token id级别。因为BERT的tokenizer会对中文按字切分偏移与原始文本基本一致但如果文本混有英文或数字子词切分会让token数量多于字符数量这时候需要对原始字符偏移做一次映射表转换否则span池化会取到错误位置的向量。3.3 从概率到结构化三元组解码与后处理3.3.1 实体边界归并实体模型拿到每个token的标签id序列后预测结果是一连串B/I/O标签需要把它们归并成实体对象。常见做法是遍历标签序列遇到B-开头就向后查找连续的I-同类型标签直到标签不匹配或序列结束。def decode_entities(pred_ids, id2type, tokenizer, input_ids): entities [] i 0 while i len(pred_ids): tag_id pred_ids[i] if tag_id % 2 1: # 奇数id都是B-标签 ent_type id2type[(tag_id - 1) // 2] j i while j 1 len(pred_ids) and pred_ids[j 1] tag_id 1: j 1 # 合并连续的I-标签 text tokenizer.decode(input_ids[i:j 1], skip_special_tokensTrue) entities.append(Entity(texttext, typeent_type, starti, endj 1)) i j i 1 return entities这段代码里用取模运算判断一个标签是否为B-。之前在设计BIO2ID时让B-为奇数、I-为偶数正好与这里的逻辑呼应。(tag_id - 1) // 2可以从B-Disease的id反推出实体类型在ENTITY_TYPES中的下标。注意tokenizer.decode(input_ids[i:j 1])恢复的是tokenizer视角的文本。如果序列里混有[CLS]、[SEP]需要将起始下标减去1或者在构造inputs时就把首尾special token排除出预测范围。3.3.2 关系类型的阈值选择与否定场景过滤关系分类模型输出每个候选实体对在所有关系类型上的概率。在医学文本里绝大多数实体对之间没有关系直接取概率最大的类别会把“没有关系”误判成某种关系。因此每个关系类别需要独立阈值概率高于阈值才输出。关系类型接受阈值建议依据treatment0.50训练样本较多阈值适中symptom0.45症状关联在病历中频繁出现indication0.60适应症边界比较模糊收紧防错adverse0.65不良反应样本少宁可少报不可多报阈值可以在验证集上做网格搜索把每个类别的阈值在[0.3, 0.7]区间以0.05步长扫一遍用验证集的关系级F1作为选择标准。另一个容易漏掉的工程细节是医学文本的否定表达。比如“无胸痛、无咳血”这类描述症状实体出现了但实际是不存在。如果处理不好会凭空多出一堆假阳性三元组。这个项目里没有单独写否定检测模块我一般会在后处理脚本里加一个简单的否定窗口规则以实体起点为中心向前找最近的一个否定词如果在3个字符内出现“无、未、不是、未见、阴性”等就把该实体标记为否定不再进入关系候选。这个规则虽然简单但在病历文本上能明显降低symptom类关系的误报。4. 部署成服务flask_server.py 与 relation_api.py 的推理链路模型训练完如果只停留在脚本里验证效果很弱。flask_server.py和relation_api.py的存在把这套代码变成了一个真正可调用的服务。如果只跑批处理不部署服务直接执行run_relation_api.py把结果写入JSON文件需要Web服务时启动flask_server.py即可。4.1 模型预加载与接口初始化relation_api.py里通常封装了一个extract_medical_triples函数用实体模型找实体、用关系模型判断实体对关系。flask_server.py在启动时完成两个模型的加载并将model.eval()和torch.no_grad()作为推理默认状态。import torch from flask import Flask, request, jsonify from relation_api import extract_medical_triples from models import EntityTagger, RelationClassifier app Flask(__name__) entity_model EntityTagger(./bert-base-chinese, num_tagslen(BIO2ID)) relation_model RelationClassifier(./bert-base-chinese, num_relationslen(REL2ID)) entity_model.load_state_dict(torch.load(./checkpoints/entity.pt)) relation_model.load_state_dict(torch.load(./checkpoints/relation.pt)) entity_model.eval() relation_model.eval() app.route(/health, methods[GET]) def health(): return jsonify({status: ok, model_loaded: True}) app.route(/relation_extract, methods[POST]) def relation_extract(): payload request.get_json(forceTrue) text payload.get(text, ) if not text: return jsonify({code: 1, message: text is required}), 400 with torch.no_grad(): triples extract_medical_triples(text, entity_model, relation_model) return jsonify({code: 0, triples: triples})forceTrue的作用是允许接口接收不带Content-Type: application/json的请求体。no_grad块确保推理过程不会构建计算图显存占用和推理速度都会明显改善。模型加载放在模块层而不是请求函数内部避免每个请求都重复读权重。实际部署时可以将extract_medical_triples的耗时打印出来配合torch.cuda.synchronize()拿到准确的GPU推理时间。4.2 请求与响应协议设计客户端调用接口时的请求体设计成下面这样curl -X POST http://127.0.0.1:5000/relation_extract \ -H Content-Type: application/json \ -d {text: 患者因胸痛入院诊断为急性心肌梗死给予阿司匹林治疗}响应结构返回三元组数组每一个元素包含头实体、关系类型、尾实体以及各自的实体类型和置信度{ code: 0, triples: [ { head: 胸痛, head_type: Symptom, relation: symptom, confidence: 0.92, tail: 急性心肌梗死, tail_type: Disease }, { head: 阿司匹林, head_type: Drug, relation: treatment, confidence: 0.87, tail: 急性心肌梗死, tail_type: Disease } ] }接口端点设计可以归纳为下表方便对接方理解端点方法入参返回说明/healthGET无服务与模型加载状态/relation_extractPOSTJSON中的text字段业务code与triples数组code字段用于区分业务成功和失败HTTP状态码则用于区分协议层面的错误。这样设计的好处是调用方可以同时检查HTTP状态和业务code避免把协议错误和业务错误混在一起。接口里加入confidence字段是一个容易被忽略的细节。很多初版实现只返回实体和关系类型但调试阶段没有置信度根本无法判断某个错误三元组是阈值问题还是模型学习的问题。加上置信度后配合日志系统可以快速定位是哪个关系类别在什么文本场景下经常误报。4.3 长文本截断与分段推理BERT的max_len限制是512个token但电子病历的入院记录动辄上千字直接截断会丢失尾部实体。常见的处理方式是先按标点符号断句再给每个句子单独推理最后合并结果。import re def split_sentences(text: str, max_chunk: int 128) - list[str]: parts re.split(r([。!?\n]), text) chunks, cur [], for part in parts: cur part if len(cur) max_chunk or part in 。!?\n: chunks.append(cur) cur if cur: chunks.append(cur) return chunks分段推理后需要做实体去重。同一个实体可能在前后两个窗口的交叉部分被识别两次但两次的偏移位置不一致。更稳妥的做法是让相邻窗口之间保留10到20个字符的overlap合并结果时以“实体文本 实体类型”为key做去重。overlap虽然会造成少量冗余计算但能避免实体正好被截断在窗口边界而漏识别。5. run_eval.py 评估口径与小样本医学调优技巧5.1 实体级F1与关系级F1的评估口径差异run_eval.py的核心职责是计算验证集上的精确率、召回率和F1。医学文本评估不能只看token级准确率因为标签分布严重偏向O全预测O都能拿到90%以上的准确率。实体级评估采用全体匹配只有当预测实体的文本、类型、起止位置与标注完全一致时才计为TP。这个标准严格适合判断模型上生产的可用性。如果只关心实体类型和文本是否对上也可以放宽为“文本类型”匹配忽略偏移。def evaluate_entity_f1(pred_entities, gold_entities): pred_set set((e.text, e.type) for e in pred_entities) gold_set set((e.text, e.type) for e in gold_entities) tp len(pred_set gold_set) precision tp / len(pred_set) if pred_set else 0.0 recall tp / len(gold_set) if gold_set else 0.0 f1 2 * precision * recall / (precision recall) return precision, recall, f1关系级评估更严格预测的三元组(head, relation, tail)必须与标注三元组完全一致。有些论文里会采用“头尾实体文本匹配关系类型匹配”的宽松标准但如果实体本身有偏移错误实体级F1已经体现了问题所以关系级用完整三元组匹配更利于定位模型短板。5.2 医学小样本场景下的三个调优技巧如果不打算在大规模医学语料上预训练只微调一个通用BERT中文模型有三件事值得做。第一用FGM对抗训练提升模型鲁棒性。FGM的原理是在embedding层添加一个沿着梯度方向的小扰动让模型对输入扰动不敏感。实现上只需要在loss.backward()之后、optimizer.step()之前把embedding的梯度归一化后叠加到embedding本身再反传一次并恢复原值。在医疗这类样本噪声大的场景中FGM带来的F1提升通常在1到2个点而成本只是训练时间翻倍。第二对易混淆实体做约束解码。比如“慢性胃炎”整体是Disease但模型有时只识别出“胃炎”。这种边界缺失问题无法通过调整阈值解决只能在解码后处理时做词典回退维护一份医学术语词典对预测实体做最长匹配延伸。这是把已有领域知识注入规则的廉价方法对实体边界切分有立竿见影的效果。第三处理类别不平衡。在四类关系中adverse不良反应的样本量往往最少模型对它的召回率最低。训练时可以给adverse类别更高的loss权重或者在采样阶段对包含该关系的样本做过采样。relation.py的损失函数里直接传一个weight张量即可不需要改模型结构。提示评估时注意区分严格F1和宽松F1。两个模型的严格F1可能只差0.3%但宽松F1可能差1.5%以上。对外汇报效果时要统一口径否则很容易出现前后自相矛盾的情况。本文还有配套的精品资源点击获取