☰
DeepKE 少样本命名实体识别工具模块深度解析:few_shot.utils.util 函数全解与实战调用链
2026/10/5 10:08:18 网站建设 项目流程
  • 人工智能
  • NLP
  • 知识图谱
  • 深度学习

【免费下载链接】DeepKE

[EMNLP 2022] An Open Toolkit for Knowledge Graph Extraction and Construction

项目地址:https://gitcode.com/gh_mirrors/de/DeepKE
点击查看免费下载

导读

本文聚焦 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.pyCoNLL 格式数据解析与 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.pyTrainer(训练/评估/预测驱动)src/deepke/name_entity_re/few_shot/module/train.py
module/metrics.pySeq2Seq Span 指标(F1/Pre/Rec/EM)src/deepke/name_entity_re/few_shot/module/metrics.py
models/model.pyPromptBart 编码器/解码器与生成逻辑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 标签列表。解码策略分三步:

  1. 定位有效预测长度:利用 eos(id=1)在序列中的位置,通过flip + cumsum技巧计算每条样本的真实预测长度(源码第 98-101 行);
  2. 还原实体与词的配对:解码目标由三类 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] ...])校验词序单调递增;
  3. 对齐原始词并输出 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 的编排顺序如下:

  1. set_seed(cfg.seed)固定随机源(L87);
  2. ConllNERProcessor加载数据、向分词器注册<<...>>提示词 token(datasets.py);
  3. PromptBartModel构造时调用avg_token_embeddings初始化新增 token 嵌入(model.py);
  4. 训练迭代中: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的设计遵循三个原则:

  1. 职责单一:掩码、损失、种子、设备探测、嵌入初始化、解码、落盘各司其职,全部为纯函数或静态工具,不持有模型状态;
  2. 与模型解耦:所有函数只依赖torch/numpy/transformers基础 API,可被models、module、example三个层级自由引用而不会造成循环依赖;
  3. 显式契约:通过 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

项目地址:https://gitcode.com/gh_mirrors/de/DeepKE
点击查看免费下载

相关推荐

上一篇:Mac Mouse Fix终极指南:如何让普通鼠标秒变生产力神器
下一篇:DoraBox新手入门:零基础学习Web安全漏洞测试的完整路径

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询