简介:面向中文医学文本实体关系抽取任务,这份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[f"B-{t}"] = i * 2 + 1 BIO2ID[f"I-{t}"] = i * 2 + 2上面的id分配有一个潜在技巧:同一个实体的B-标签和I-标签在id上是相邻的,比如B-Disease=1、I-Disease=2,B-Symptom=3、I-Symptom=4。这样后处理时判断“当前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_factory=list) relations: List[Relation] = field(default_factory=list) def to_bio_labels(self) -> List[str]: labels = ["O"] * len(self.text) for ent in self.entities: labels[ent.start] = f"B-{ent.type}" for i in range(ent.start + 1, ent.end): labels[i] = f"I-{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_first=True) def forward(self, input_ids, attention_mask, labels=None): outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask) seq_out = self.dropout(outputs.last_hidden_state) logits = self.classifier(seq_out) if labels is not None: loss = -self.crf(logits, labels, mask=attention_mask.bool()) return loss return self.crf.decode(logits, mask=attention_mask.bool())self.crf(logits, labels, mask)返回的是负对数似然,前面加负号就变成最小化目标。当labels=None时,CRF的decode内部走Viterbi路径,返回的是每个样本的标签id序列,形状是List[List[int]]。这里最容易被忽略的是mask=attention_mask.bool()这一步:如果不传mask,padding部分的标签会参与转移计算,解码结果会出现起始标签出现在序列中间这类非法情况。
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_ids=input_ids, attention_mask=attention_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(dim=0)) tail_vecs.append(seq_out[i, ts:te + 1].mean(dim=0)) 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_tags=len(BIO2ID)) optimizer = torch.optim.AdamW(model.parameters(), lr=2e-5) scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=200, num_training_steps=5000 ) 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, labels=labels) optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.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(f"epoch {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_size | 16~32 | 显存允许时尽量大 |
| learning_rate | 2e-5 | BERT微调通用基线 |
| max_len | 128 | 覆盖电子病历句子主体 |
| warmup_ratio | 0.1 | 前10%步数线性升温 |
| grad_clip_norm | 1.0 | 防止梯度爆炸 |
| epochs | 3~5 | 医学小语料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_tokens=True) entities.append(Entity(text=text, type=ent_type, start=i, end=j + 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 关系类型的阈值选择与否定场景过滤
关系分类模型输出每个候选实体对在所有关系类型上的概率。在医学文本里,绝大多数实体对之间没有关系,直接取概率最大的类别会把“没有关系”误判成某种关系。因此每个关系类别需要独立阈值,概率高于阈值才输出。
| 关系类型 | 接受阈值 | 建议依据 |
|---|---|---|
| treatment | 0.50 | 训练样本较多,阈值适中 |
| symptom | 0.45 | 症状关联在病历中频繁出现 |
| indication | 0.60 | 适应症边界比较模糊,收紧防错 |
| adverse | 0.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_tags=len(BIO2ID)) relation_model = RelationClassifier("./bert-base-chinese", num_relations=len(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(force=True) 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})force=True的作用是允许接口接收不带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" } ] }接口端点设计可以归纳为下表,方便对接方理解:
| 端点 | 方法 | 入参 | 返回说明 |
|---|---|---|---|
| /health | GET | 无 | 服务与模型加载状态 |
| /relation_extract | POST | JSON中的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%以上。对外汇报效果时,要统一口径,否则很容易出现前后自相矛盾的情况。
本文还有配套的精品资源,点击获取