ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

在 Fairseq 中实现 Transformer Pointer-Generator:OOV 词汇复制的完整实战指南

在 Fairseq 中实现 Transformer Pointer-Generator:OOV 词汇复制的完整实战指南 在 Fairseq 中实现 Transformer Pointer-GeneratorOOV 词汇复制的完整实战指南【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm导读本文围绕 Fairseq 中的transformer_pointer_generator模型展开讲解如何在 Transformer 架构中引入指针生成pointer-generator机制使模型能够在生成时直接复制输入序列中的词语从而有效处理词汇表外的 OOVOut-of-Vocabulary词尤其适合小词表下的文本摘要、翻译等序列生成任务。读完本文你将掌握该模型的设计原理、源码级实现细节以及从词表构建、数据预处理、模型训练到生成后处理的完整落地流程。该实现位于本仓库的 decoding/IAD/fairseq/examples/pointer_generator 目录下。背景从 RNN 到 Transformer 的指针生成机制指针生成网络Pointer-Generator Network最初由 See et al.2017在论文Get To The Point: Summarization with Pointer-Generator Networks中提出用于 RNN 编码器-解码器注意力模型。其核心思想是在每个解码步模型不是直接产生一个词表上的分布而是将从词表生成的概率分布与从输入复制的注意力分布进行插值混合。Transformer 同样可以借鉴这一思路复用 Transformer 中众多的注意力分布之一作为指针分布。具体来说将模型对输入词的注意力分布与正常的词表输出分布进行插值最终分布 p_gen × 词表生成分布 (1 - p_gen) × 输入注意力分布这样即使某个词不在词表中只要它出现在输入序列里模型依然有机会通过指向它来输出。这对小词表场景特别有价值——例如 README.xsum.md 中 XSum 摘要任务仅使用 10000 词的词表大量专有名词人名、地名、俱乐部名都依赖复制机制才能正确出现在摘要中。Fairseq 的独特实现不侵入模型之外的任何代码与 See et al. 的实现不同Fairseq 的这一版本采用了截然不同的工程策略。See 的原始实现需要把词语身份信息贯穿整个模型内部传递而 Fairseq 版本将指针机制完整封装在模型文件内部避免对代码库其余部分如 SequenceGenerator、任务、数据集做任何改动。实现 OOV 复制的方式是在数据预处理阶段替换 OOV 词在生成后处理阶段恢复原词。其思路是预处理把输入中每个不在词表里的词替换为位置标记unk-NN 为该词在输入序列中的位置模型将这些位置标记统一映射到unk的词嵌入并在输出层把注意力分布写入扩展词表对应的unk-N位置上后处理生成结果中若出现unk-N则用原始输入中第 N 个位置的词替换回去。这一设计使得整个机制自包含在 pointer_generator_src/transformer_pg.py 单个模型文件中通过--user-dir注册即可使用无需修改 Fairseq 核心代码。源码级剖析transformer_pg.py 的核心机制模型文件 transformer_pg.py 通过register_model(transformer_pointer_generator)注册了TransformerPointerGeneratorModel并派生出自定义的编码器与解码器。以下是几个关键实现点1. 扩展词表与共享词嵌入Embedding 子类build_model中首先强制要求源、目标共享词典joined dictionary否则直接报错Pointer-generator requires a joined dictionary。随后通过自定义的Embedding子类transformer_pg.py构建词嵌入词表末尾的source_position_markers个位置标记虽然占用词典索引但全部映射到unk的嵌入。其forward中利用torch.where将所有索引大于等于num_embeddings的输入替换为unk_idx。启动时模型会打印类似日志fairseq.models.transformer_pg | dictionary indices from 10000 to 10999 will be mapped to 3即索引 10000–10999 这 1000 个位置标记共享unk的词嵌入。2. 编码器把源 token 透传给解码器普通 Transformer 编码器不会把源 token ID 传给解码器而指针生成需要源 token ID 来做注意力分布的散射scatter。因此TransformerPointerGeneratorEncoder.forwardtransformer_pg.py在父类输出的基础上额外返回src_tokens: [src_tokens]。源码注释明确说明虽然更优雅的做法是把源 token 同时传给解码器的forward但那需要改动SequenceGenerator于是选择在编码器输出里携带。3. 解码器生成概率 p_gen 的预测与分布混合TransformerPointerGeneratorDecoder的核心在两点p_gen 预测transformer_pg.py用一个线性层project_p_gens接收当前解码输入嵌入 解码器输出特征拼接向量输出经 sigmoid 得到每个位置的生成概率 p_gen偏置初始化为 0。输出层混合output_layertransformer_pg.py正常词表 logits 经 softmax 后乘以p_gens并在末尾拼接num_oov_types个零填充得到生成部分在扩展词表上的分布注意力权重乘以(1 - p_gens)然后通过scatter_add_按源 token ID 散射到扩展词表对应位置得到复制部分分布两者相加得到最终在num_types 词表 位置标记数上的分布。由于输出已是归一化分布get_normalized_probstransformer_pg.py不再重复 softmax仅在返回 log 概率时做clamp(1e-10, 1.0)保护。4. 关键命令行参数add_argstransformer_pg.py定义了模型专属参数参数说明默认值--alignment-heads N用于指向的注意力头数量架构默认 1--alignment-layer I用于指向的解码器层号0 表示最底层支持负数如 -1 表示倒数第一层架构默认 -1解码后自动换算为decoder_layers alignment_layer--source-position-markers N词典末尾额外添加的 OOV 位置标记数量全部映射到unk嵌入max_source_positions--force-generation P不预测 p_gen强制设为 P1.0 表示纯生成0.0 表示纯指向None其中--alignment-layer/--alignment-heads的用法与transformer_align模型一致选取某个解码器层的若干注意力头做平均得到指向用的对齐分布。此外模型还预置了transformer_pointer_generator、_iwslt_de_en、_wmt_en_de、_vaswani_wmt_en_de_big等多套架构变体transformer_pg.py。使用流程四个步骤落地指针生成第 1 步构建词表并追加源位置标记指针机制在小词表下最有效前提是能恢复被复制的 OOV 词身份。为此需要把unk-0、unk-1、unk-2…… 等特殊标记追加到词表末尾。下面示例构建一个包含 10000 个最常用词 1000 个位置标记的词表vocab_size10000 position_markers1000 export LC_ALLC cat train.src train.tgt | tr -s [:space:] \n | sort | uniq -c | sort -k1,1bnr -k2 | head -n $((vocab_size - 4)) | awk { print $2 $1 } dict.pg.txt python3 -c [print(unk-{} 0.format(n)) for n in range($position_markers)] dict.pg.txt注意head -n $((vocab_size - 4))预留出 4 个特殊 tokens、pad、/s、unk的位置。生成的dict.pg.txt形如the 4954867 . 4157552 , 3439668 ... unk-0 0 unk-1 0 unk-2 0 unk-3 0 unk-4 0 ...第 2 步用 preprocess.py 替换 OOV 词核心思想文本中任何unk词若出现在输入第 1 个位置则替换为unk-0第 2 个位置则替换为unk-1依此类推。这由目录下的 preprocess.py 完成其replace_oovs函数逐序列处理源序列中不在词表里的 token用其首次出现的位置编号生成unk-N同一 OOV 词重复出现时复用同一个位置标记通过word_to_pos字典记忆目标序列中若出现源序列里的 OOV 词同样替换为对应位置的unk-N不在源序列里的词保持原样。用法./preprocess.py --source train.document --target train.summary --vocab (cut -d -f1 dict.pg.txt) --source-out train.pg.src --target-out train.pg.tgt其中--source/--target为源/目标文本文件--vocab为只含词条不带频次的词表文件--source-out/--target-out为输出文件--target、--target-out均可选纯源端预处理时可省略。第 3 步训练模型用fairseq-preprocess二值化数据后通过fairseq-train训练。位置标记数量通过--source-position-markers传给模型指向所用的注意力分布通过--alignment-heads和--alignment-layer选择用法与transformer_align相同。核心训练命令完整示例见下文 XSum 一节fairseq-train bin \ --user-dir examples/pointer_generator/pointer_generator_src \ --task translation \ --source-lang src --target-lang tgt \ --arch transformer_pointer_generator \ --alignment-layer -2 \ --alignment-heads 1 \ --source-position-markers 1000 \ ...注意使用模型文件目录必须通过--user-dir examples/pointer_generator/pointer_generator_src指定训练、验证和生成时都要带上。第 4 步生成文本并后处理生成时输入文本要与训练数据做同样的预处理把 OOV 词替换为unk-N。若这些标记被复制到输出用 postprocess.py 从未处理过的原始输入中恢复真实词语任何unk-N都应替换为原始输入序列中第 N 个位置的词。./postprocess.py --source test.document --target generate.hyp --target-out generate.hyp.processed该脚本用正则^unk-([0-9])$匹配标记并把位置超出源序列长度的情况判定为错误抛出OOVIndexError这通常意味着源/目标序列错位或指向机制关注到了序列末尾之后的位置。端到端实战XSum 极简摘要训练示例README.xsum.md 给出了在 Extreme SummarizationXSum数据集上的完整流程。数据从 XSum 原始发布处获取后应有{train,validation,test}.{document,summary}六个文件。随后依次执行1. 构建词表与上文命令一致把train.src train.tgt换成train.document train.summary生成含 1 万高频词 1 千位置标记的dict.pg.txt。2. 预处理数据./preprocess.py --source train.document --target train.summary --vocab (cut -d -f1 dict.pg.txt) --source-out train.pg.src --target-out train.pg.tgt ./preprocess.py --source validation.document --target validation.summary --vocab (cut -d -f1 dict.pg.txt) --source-out valid.pg.src --target-out valid.pg.tgt ./preprocess.py --source test.document --vocab (cut -d -f1 dict.pg.txt) --source-out test.pg.src3. 二值化使用--joined-dictionary与模型对共享词典的要求一致fairseq-preprocess \ --source-lang src \ --target-lang tgt \ --trainpref train.pg \ --validpref valid.pg \ --destdir bin \ --workers 60 \ --srcdict dict.pg.txt \ --joined-dictionary4. 训练total_updates20000 warmup_updates500 lr0.001 max_tokens4096 update_freq4 pointer_layer-2 CUDA_VISIBLE_DEVICES0,1,2,3,4,5,6,7 fairseq-train bin \ --user-dir examples/pointer_generator/pointer_generator_src \ --max-tokens $max_tokens \ --task translation \ --source-lang src --target-lang tgt \ --truncate-source \ --layernorm-embedding \ --share-all-embeddings \ --encoder-normalize-before \ --decoder-normalize-before \ --required-batch-size-multiple 1 \ --arch transformer_pointer_generator \ --alignment-layer $pointer_layer \ --alignment-heads 1 \ --source-position-markers 1000 \ --criterion label_smoothed_cross_entropy \ --label-smoothing 0.1 \ --dropout 0.1 --attention-dropout 0.1 \ --weight-decay 0.01 --optimizer adam --adam-betas (0.9, 0.999) --adam-eps 1e-08 \ --clip-norm 0.1 \ --lr-scheduler inverse_sqrt --lr $lr --max-update $total_updates --warmup-updates $warmup_updates \ --update-freq $update_freq \ --skip-invalid-size-inputs-valid-test这里指定了词典含 1000 个源位置标记并选用解码器倒数第二层-2的 1 个注意力头做指向。训练日志会确认词典中 10000 之后的索引被映射到unk嵌入README.xsum.md 中记录了当时训练产生的日志8 卡 V100 上约 5.5 小时完成 2 万步此处时间仅代表该文档记录的环境结果实际耗时取决于硬件与数据规模fairseq.tasks.translation | [src] dictionary: 11000 types fairseq.tasks.translation | [tgt] dictionary: 11000 types fairseq.models.transformer_pg | dictionary indices from 10000 to 10999 will be mapped to 35. 生成batch_size32 beam_size6 max_length60 length_penalty1.0 fairseq-interactive bin \ --user-dir examples/pointer_generator/pointer_generator_src \ --batch-size $batch_size \ --task translation \ --source-lang src --target-lang tgt \ --path checkpoints/checkpoint_last.pt \ --input test.pg.src \ --buffer-size 200 \ --max-len-a 0 \ --max-len-b $max_length \ --lenpen $length_penalty \ --beam $beam_size \ --skip-invalid-size-inputs-valid-test | tee generate.out grep ^H generate.out | cut -f 3- generate.hyp6. 后处理恢复 OOV 词由于生成时跳过了过长输入后处理同样用awk NF1024过滤超长源序列保证源/目标一一对应./postprocess.py \ --source (awk NF1024 test.document) \ --target generate.hyp \ --target-out generate.hyp.processed一个直观的示例源自 README.xsum.md——原始源文档de roon moved to teesside in june 2016 for an initial # 8.8 m fee ...预处理后的源文档人名roon、teesside、数字8.8等 OOV 词被替换为位置标记de unk-1 moved to unk-4 in june 2016 for an initial # unk-12 m fee ...生成的原始摘要模型复制出了unk-1标记同时也有真unkmiddlesbrough striker unk de unk-1 has joined spanish side unk on a season-long loan .后处理后的最终摘要unk-1被替换为源文档第 1 个位置的词roonmiddlesbrough striker unk de roon has joined spanish side unk on a season-long loan .可以看到模型成功复制了词表外的专有名词roon而无法恢复的真 OOV 词仍以unk呈现。测试验证仓库中的自动化回归用例本仓库的 tests/test_binaries.py 中提供了test_transformer_pointer_generator端到端测试它使用小规模 dummy 数据经过数据预处理、以transformer_pointer_generator架构2 层编码器/解码器、8 维嵌入、--source-position-markers 0训练并在验证与生成阶段都通过--user-dir examples/pointer_generator/pointer_generator_src加载模型。该测试印证了该模型可通过--user-dir方式无缝接入 Fairseq 的标准训练/生成流程也验证了位置标记数量为 0 时模型依然可以正常训练与推理此时退化为无扩展词表的纯生成模式。小结transformer_pointer_generator用预处理替换 共享unk嵌入 注意力散射这一自包含方案把 See et al. 的指针生成思想完整移植到 Transformer并在不改动 Fairseq 其余代码的前提下解决了 OOV 复制问题。无论是小词表的摘要任务还是其他需要从源文本中抽取实体的生成任务这套词表扩展 位置标记 生成/指向概率插值的工程模式都值得参考。深入阅读源码可继续查看核心模型pointer_generator_src/transformer_pg.py数据预处理脚本preprocess.py输出后处理脚本postprocess.pyXSum 完整示例README.xsum.md自动化测试tests/test_binaries.py【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表