在 Fairseq 中实现 Transformer Pointer-Generator:OOV 词汇复制的完整实战指南
【免费下载链接】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)机制,使模型能够在生成时直接"复制"输入序列中的词语,从而有效处理词汇表外的 OOV(Out-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-N>(N 为该词在输入序列中的位置); - 模型:将这些位置标记统一映射到
<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.forward(transformer_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_layer,transformer_pg.py):- 正常词表 logits 经 softmax 后乘以
p_gens,并在末尾拼接num_oov_types个零填充,得到"生成部分"在扩展词表上的分布; - 注意力权重乘以
(1 - p_gens),然后通过scatter_add_按源 token ID 散射到扩展词表对应位置,得到"复制部分"分布; - 两者相加得到最终在
num_types(= 词表 + 位置标记数)上的分布。
- 正常词表 logits 经 softmax 后乘以
由于输出已是归一化分布,get_normalized_probs(transformer_pg.py)不再重复 softmax,仅在返回 log 概率时做clamp(1e-10, 1.0)保护。
4. 关键命令行参数
add_args(transformer_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,强制设为 P(1.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_size=10000 position_markers=1000 export LC_ALL=C 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 个特殊 token(<s>、<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 Summarization(XSum)数据集上的完整流程。数据从 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_updates=20000 warmup_updates=500 lr=0.001 max_tokens=4096 update_freq=4 pointer_layer=-2 CUDA_VISIBLE_DEVICES=0,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_size=32 beam_size=6 max_length=60 length_penalty=1.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 'NF<1024'过滤超长源序列,保证源/目标一一对应):
./postprocess.py \ --source <(awk 'NF<1024' 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>标记,同时也有真<unk>):
middlesbrough striker <unk> de <unk-1> has joined spanish side <unk> on a season-long loan .
后处理后的最终摘要(<unk-1>被替换为源文档第 1 个位置的词roon):
middlesbrough 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.py
- XSum 完整示例: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),仅供参考