ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

DeepKE 少样本命名实体识别工具模块深度解析:few_shot.utils.util 函数全解与实战调用链

DeepKE 少样本命名实体识别工具模块深度解析:few_shot.utils.util 函数全解与实战调用链 人工智能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模块。该模块基于 LightNERCOLING22范式通过 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.pymodule/datasets.pyCoNLL 格式数据解析与 BIO→序列目标转换src/deepke/name_entity_re/few_shot/module/datasets.pymodule/mapping_type.py实体标签到...提示词的映射表src/deepke/name_entity_re/few_shot/module/mapping_type.pymodule/train.pyTrainer训练/评估/预测驱动src/deepke/name_entity_re/few_shot/module/train.pymodule/metrics.pySeq2Seq Span 指标F1/Pre/Rec/EMsrc/deepke/name_entity_re/few_shot/module/metrics.pymodels/model.pyPromptBart 编码器/解码器与生成逻辑src/deepke/name_entity_re/few_shot/models/model.py二、运行环境与数据准备使用前提few-shot 模块的官方运行前提见 example/ner/few-shot/README_CN.mdPython 3.8、torch 1.11、transformers 4.26.0并安装deepke包本身。数据方面需要准备如下格式的文件放在example/ner/few-shot/data目录下CoNLL2003train.txt/dev.txt/test.txt以\t分隔的词 BIO 标签MIT-movie、MIT-restaurant、ATISk-shot-train.txtk 可取 10/20/50/100/200/500与test.txtCLUENER2020中文20-shot-train.txt与test.txt。数据集与路径、映射关系的注册表集中在 run.pyDATASET_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(seed2021)全链路随机种子控制def set_seed(seed2021): torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic True np.random.seed(seed) random.seed(seed)它一次性固定 PyTorchCPU/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_lensrc_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_lentgt_tokens.size(1) - 1).eq(0) tgt_tokens tgt_tokens[:, 1:].masked_fill(mask, -100) loss F.cross_entropy(targettgt_tokens, inputpred.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_valuebos_token_id, dtypetorch.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 indexnum_tokens保证新 token 的 id 一定落在新增区域。3.6convert_preds_to_outputs(preds, raw_words, mapping, tokenizer)模型预测 → BIO 序列解码这是预测阶段最关键的转换函数util.py。它的职责是把 BART 解码器输出的 token id 序列还原为与原始句子等长的 BIO 标签列表。解码策略分三步定位有效预测长度利用 eosid1在序列中的位置通过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[i1] ...])校验词序单调递增对齐原始词并输出 BIO根据分词器对每个原始词的分词长度计算累积偏移cum_lens把词 id 映射回原始词下标最终生成B-{tag}/I-{tag}/O标签。output[start_idx] fB-{id2label[entity-2]} for _ in range(start_idx1, end_idx1): output[_] fI-{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(rtouch {}.format(path)) with open(path, w, encodingutf-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.pyif 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.pyrun.py 的编排顺序如下set_seed(cfg.seed)固定随机源L87ConllNERProcessor加载数据、向分词器注册...提示词 tokendatasets.pyPromptBartModel构造时调用avg_token_embeddings初始化新增 token 嵌入model.py训练迭代中seq_to_mask生成编码器掩码model.py→ 解码器输出 logits →get_loss计算损失run.py→ 反向传播。关键的超参few_shot.yamlnum_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 训练可追加trainfew_shot_cn覆盖默认配置官方提示全量数据微调才能达到最佳性能README_CN.md。4.2 推理链路python predict.pypredict.py 与 run.py 结构几乎一致差异在于数据加载模式为test此时ConllNERDataset.__getitem__只返回src_tokens / src_seq_len / first / raw_wordsdatasets.pymodel.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_pathtrain.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项目地址https://gitcode.com/gh_mirrors/de/DeepKE点击查看免费下载相关推荐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知识图谱深度学习上一篇Mac Mouse Fix终极指南如何让普通鼠标秒变生产力神器下一篇DoraBox新手入门零基础学习Web安全漏洞测试的完整路径创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表