- 人工智能
- NLP
- 知识图谱
- 深度学习
【免费下载链接】DeepKE
[EMNLP 2022] An Open Toolkit for Knowledge Graph Extraction and Construction
导读
本文聚焦 DeepKE 中少样本命名实体识别(few-shot NER)子模块的工具层 —— 由 Sphinx 文档 deepke.name_entity_re.few_shot.utils.rst 所收录的deepke.name_entity_re.few_shot.utils.util模块。该模块基于 LightNER(COLING'22)范式,通过 BART 生成式框架完成低资源场景下的实体识别。读完本文,你将掌握该工具模块中 7 个核心函数的输入输出契约、底层实现原理,以及它们如何被数据预处理、Prompt 模型、训练器与预测脚本串联成一条完整的 few-shot NER 训练/推理流水线。
一、模块定位:从 RST 文档到源码
1.1 RST 文档与源码的对应关系
docs/source/deepke.name_entity_re.few_shot.utils.rst是一份 Sphinxautomodule文档存根,其核心指令为:
.. automodule:: deepke.name_entity_re.few_shot.utils.util :members: :undoc-members: :show-inheritance:它要求 Sphinx 在生成文档时,自动从 Python 源码 util.py 中提取模块级函数及其 docstring 作为文档主体。这意味着该工具模块的可成文内容完全由源码决定,因此本文以源码级解析为主。
1.2 工具模块在 few-shot NER 包中的位置
utils子包位于 src/deepke/name_entity_re/few_shot/utils/ 下,其__init__.py通过from .util import *将所有工具函数暴露到包级别,因此下游代码既可以写from deepke.name_entity_re.few_shot.utils.util import get_loss,也可以写from ..utils import convert_preds_to_outputs(后者见 train.py)。
整个 few-shot NER 包的分工如下:
| 子模块 | 职责 | 路径 |
|---|---|---|
utils/util.py | 通用工具函数(掩码、损失、种子、解码、落盘) | src/deepke/name_entity_re/few_shot/utils/util.py |
module/datasets.py | CoNLL 格式数据解析与 BIO→序列目标转换 | src/deepke/name_entity_re/few_shot/module/datasets.py |
module/mapping_type.py | 实体标签到<<...>>提示词的映射表 | src/deepke/name_entity_re/few_shot/module/mapping_type.py |
module/train.py | Trainer(训练/评估/预测驱动) | src/deepke/name_entity_re/few_shot/module/train.py |
module/metrics.py | Seq2Seq Span 指标(F1/Pre/Rec/EM) | src/deepke/name_entity_re/few_shot/module/metrics.py |
models/model.py | PromptBart 编码器/解码器与生成逻辑 | src/deepke/name_entity_re/few_shot/models/model.py |
二、运行环境与数据准备(使用前提)
few-shot 模块的官方运行前提见 example/ner/few-shot/README_CN.md:Python 3.8、torch 1.11、transformers 4.26.0,并安装deepke包本身。数据方面,需要准备如下格式的文件(放在example/ner/few-shot/data目录下):
- CoNLL2003:
train.txt/dev.txt/test.txt,以\t分隔的词 + BIO 标签; - MIT-movie、MIT-restaurant、ATIS:
k-shot-train.txt(k 可取 10/20/50/100/200/500)与test.txt; - CLUENER2020(中文):
20-shot-train.txt与test.txt。
数据集与路径、映射关系的注册表集中在 run.py:DATASET_CLASS、DATA_PROCESS、DATA_PATH与MAPPING四个字典。其中MAPPING定义了“原始标签 → 提示词 token”的对应,例如 CoNLL2003 的{'loc': '<<location>>', 'per': '<<person>>', 'org': '<<organization>>', 'misc': '<<others>>'};中文 CLUENER2020 的完整映射见 mapping_type.py。
三、核心工具函数逐个解析
util.py 共定义 7 个模块级函数。下面按其在流水线中的作用逐一展开,并给出源码依据。
3.1set_seed(seed=2021):全链路随机种子控制
def set_seed(seed=2021): torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic = True np.random.seed(seed) random.seed(seed)它一次性固定 PyTorch(CPU/GPU)、cuDNN、NumPy 与 Pythonrandom四套随机源,保证实验可复现。在 run.py 与 predict.py 中都在构造数据集之前调用set_seed(cfg.seed)(默认seed: 1,见 few_shot.yaml)。
3.2seq_to_mask(seq_len, max_len):序列长度掩码生成
def seq_to_mask(seq_len, max_len): max_len = int(max_len) if max_len else seq_len.max().long() cast_seq = torch.arange(max_len).expand(seq_len.size(0), -1).to(seq_len) mask = cast_seq.lt(seq_len.unsqueeze(1)) return mask输入是批内每条样本的真实长度(bsz向量),输出形状为bsz × max_len的布尔掩码,位置j < seq_len[i]处为True(有效 token)。max_len不传时自动取批内最大长度。它的实际消费方是 Prompt 模型的编码器:
- 在 model.py 的
generator()中,attention_mask = seq_to_mask(src_seq_len, max_len=src_tokens.size(1)),用于遮蔽 BART 编码器的 padding; - 在 util.py 的
get_loss内部再次被复用,为损失计算构造 padding 掩码。
3.3get_loss(tgt_tokens, tgt_seq_len, pred):生成式 NER 的交叉熵损失
def get_loss(tgt_tokens, tgt_seq_len, pred): tgt_seq_len = tgt_seq_len - 1 mask = seq_to_mask(tgt_seq_len, max_len=tgt_tokens.size(1) - 1).eq(0) tgt_tokens = tgt_tokens[:, 1:].masked_fill(mask, -100) loss = F.cross_entropy(target=tgt_tokens, input=pred.transpose(1, 2)) return loss要点如下:
tgt_tokens形状为bsz × max_len,包含[sos, token, eos]完整序列;pred形状为bsz × (max_len-1) × vocab_size;- 因为用前一时刻预测后一时刻,目标需要左移一位:
tgt_tokens[:, 1:]; - 通过
seq_to_mask取反得到 padding 位置,用-100填充,使F.cross_entropy自动忽略这些位置; - 返回值是一个标量张量,由 Trainer 的
_step在训练模式下回传(train.py),随后执行loss.backward()与optimizer.step()。
在 run.py 中它被直接注册为 Trainer 的损失函数:loss = get_loss。
3.4get_model_device(model):安全获取模型所在设备
def get_model_device(model): assert isinstance(model, nn.Module) parameters = list(model.parameters()) if len(parameters) == 0: return None else: return parameters[0].device该函数先断言对象是nn.Module,再通过第一个参数张量的.device属性推断设备;对无参数模块(如纯容器)返回None而非抛异常。生成逻辑在初始化起始 token 时调用它确定设备:
device = get_model_device(decoder) # model.py _no_beam_search_generate 内部 tokens = torch.full([batch_size, 1], fill_value=bos_token_id, dtype=torch.long).to(device)见 model.py。
3.5avg_token_embeddings(tokenizer, bart_model, bart_name, num_tokens):新增提示词的嵌入初始化
当训练数据引入了<<location>>这类自定义提示词 token 时,新扩充的 embedding 是随机初始化的。该函数用“平均法”为其赋初值:先用与模型匹配的分词器(中文场景用BertTokenizer,否则用BartTokenizer,见 util.py)对<<xxx>>去除尖括号后的真实词xxx分词,再取其各子词嵌入的均值写入新 token 对应的 decoder 嵌入行:
indexes = _tokenizer.convert_tokens_to_ids(_tokenizer.tokenize(token[2:-2])) embed = bart_model.encoder.embed_tokens.weight.data[indexes[0]] for i in indexes[1:]: embed += bart_model.decoder.embed_tokens.weight.data[i] embed /= len(indexes) bart_model.decoder.embed_tokens.weight.data[index] = embed该函数在 model.py 中紧随resize_token_embeddings之后被调用:
num_tokens, _ = bart_model.encoder.embed_tokens.weight.shape bart_model.resize_token_embeddings(len(tokenizer.unique_no_split_tokens)+num_tokens) bart_model = avg_token_embeddings(tokenizer, bart_model, bart_name, num_tokens)注意其边界约束:若<<...>>被分词器错误切分成多个子词,会直接抛出RuntimeError(f"{token} wrong split");同时通过assert index>=num_tokens保证新 token 的 id 一定落在新增区域。
3.6convert_preds_to_outputs(preds, raw_words, mapping, tokenizer):模型预测 → BIO 序列解码
这是预测阶段最关键的转换函数(util.py)。它的职责是把 BART 解码器输出的 token id 序列还原为与原始句子等长的 BIO 标签列表。解码策略分三步:
- 定位有效预测长度:利用 eos(id=1)在序列中的位置,通过
flip + cumsum技巧计算每条样本的真实预测长度(源码第 98-101 行); - 还原实体与词的配对:解码目标由三类 id 组成——实体标签 id(
< word_start_index)、源词 id(>= word_start_index,其中word_start_index = len(mapping) + 2)。代码用cur_pair累积连续词 id,遇到实体 id 时把(词id序列 + 实体id)打包成 pair,并通过all([cur_pair[i] < cur_pair[i+1] ...])校验词序单调递增; - 对齐原始词并输出 BIO:根据分词器对每个原始词的分词长度计算累积偏移
cum_lens,把词 id 映射回原始词下标,最终生成B-{tag}/I-{tag}/O标签。
output[start_idx] = f'B-{id2label[entity-2]}' for _ in range(start_idx+1, end_idx+1): output[_] = f'I-{id2label[entity-2]}'其中id2label = list(mapping.keys())保证标签名与训练时注册的mapping一致。该函数被 train.py 的predict()逐批调用:
outputs = convert_preds_to_outputs(preds, raw_words, self.process.mapping, self.process.tokenizer)3.7write_predictions(path, texts, labels):以 CoNLL 格式落盘预测结果
def write_predictions(path, texts, labels): assert len(texts) == len(labels) if not os.path.exists(path): os.system(r"touch {}".format(path)) with open(path, "w", encoding="utf-8") as f: f.writelines("-DOCSTART-\tO\n\n") for i in range(len(texts)): for j in range(len(texts[i])): f.writelines("{}\t{}\n".format(texts[i][j], labels[i][j])) f.writelines("\n")它以标准 CoNLL 格式输出:文件头写入-DOCSTART-\tO,每行词\t标签,句子之间以空行分隔。调用点在 train.py:
if self.args.write_path is not None: write_predictions(self.args.write_path, texts, labels)write_path在 predict.yaml 中配置,例如"data/conll2003/predict.txt",产出结果可直接交给 CoNLL 官方评估脚本比对。
四、调用链全景:工具函数如何串起训练与推理
4.1 训练链路(python run.py)
run.py 的编排顺序如下:
set_seed(cfg.seed)固定随机源(L87);ConllNERProcessor加载数据、向分词器注册<<...>>提示词 token(datasets.py);PromptBartModel构造时调用avg_token_embeddings初始化新增 token 嵌入(model.py);- 训练迭代中:
seq_to_mask生成编码器掩码(model.py)→ 解码器输出 logits →get_loss计算损失(run.py)→ 反向传播。
关键的超参(few_shot.yaml):num_epochs: 30、batch_size: 3、learning_rate: 5e-5、eval_begin_epoch: 16、use_prompt: True、prompt_len: 10、prompt_dim: 800、freeze_plm: True、learn_weights: True。中文 few-shot 训练可追加+train=few_shot_cn覆盖默认配置;官方提示全量数据微调才能达到最佳性能(README_CN.md)。
4.2 推理链路(python predict.py)
predict.py 与 run.py 结构几乎一致,差异在于:
- 数据加载模式为
'test',此时ConllNERDataset.__getitem__只返回src_tokens / src_seq_len / first / raw_words(datasets.py); model.predict(src_tokens, src_seq_len, first)走_no_beam_search_generate/_beam_search_generate生成路径(model.py),生成过程中同样依赖get_model_device确定设备;- 解码出的 token id 依次经
convert_preds_to_outputs转成 BIO 标签,再由write_predictions写入write_path(train.py)。
4.3 评测链路:与 metrics 模块的衔接
需要澄清的是,指标计算并不直接调用util.py,而是由独立的 metrics.py 中的Seq2SeqSpanMetric完成——它内部实现了与convert_preds_to_outputs高度相似的“预测序列切分 + pair 还原 + TP/FP/FN 统计”逻辑(metrics.py),最终输出 F1、Precision、Recall 与 EM(精确匹配率)。两处逻辑相互印证了该解码约定的稳定性:word_start_index = num_labels + 2(或len(mapping) + 2)是贯穿两处的核心常量。
五、工具函数设计要点总结
从源码结构看,util.py的设计遵循三个原则:
- 职责单一:掩码、损失、种子、设备探测、嵌入初始化、解码、落盘各司其职,全部为纯函数或静态工具,不持有模型状态;
- 与模型解耦:所有函数只依赖
torch/numpy/transformers基础 API,可被models、module、example三个层级自由引用而不会造成循环依赖; - 显式契约:通过 docstring 声明输入输出形状(如
bsz × max_len),通过assert前置校验(如新增 token 数量、词序单调性、文本标签等长)保障数据一致性。
六、小结
deepke.name_entity_re.few_shot.utils.util虽名为“工具”,实则是 few-shot NER 流水线的粘合剂:set_seed保证可复现,seq_to_mask与get_loss支撑训练收敛,avg_token_embeddings让提示词 token 获得合理初值,convert_preds_to_outputs与write_predictions则完成从模型概率到可评估 BIO 标注的最后一公里。理解这 7 个函数,就等于掌握了 DeepKE 少样本 NER 从数据到评估的全链路数据契约。若需进一步深入,可依次阅读 model.py、datasets.py 与 train.py 的完整实现。
- 人工智能
- NLP
- 知识图谱
- 深度学习
【免费下载链接】DeepKE
[EMNLP 2022] An Open Toolkit for Knowledge Graph Extraction and Construction
相关推荐
DeepKE 少样本命名实体识别(few-shot NER)核心模块解析:数据、映射、指标与训练
DeepKE 少样本命名实体识别(few shot NER)核心模块解析:数据、映射、指标与训练 导读 本文围绕 DeepKE 的 deepke.name_en
人工智能NLP知识图谱深度学习DeepKE 小样本命名实体识别(Few-shot NER)模型模块深度解析:PromptBart 与 Prefix-tuning BART 实现
DeepKE 小样本命名实体识别(Few shot NER)模型模块深度解析:PromptBart 与 Prefix tuning BART 实现 导读 本文以
人工智能NLP知识图谱深度学习DeepKE 标准命名实体识别(NER)数据工具层全解析:tools.dataset 与 tools.preprocess 实战指南
DeepKE 标准命名实体识别(NER)数据工具层全解析:tools.dataset 与 tools.preprocess 实战指南 本篇技术指南围绕 Deep
人工智能NLP知识图谱深度学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考