尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

深入解析 Generalized Aggressive Decoding(GAD):基于非自回归草稿模型的无损自回归翻译加速实践

深入解析 Generalized Aggressive Decoding(GAD):基于非自回归草稿模型的无损自回归翻译加速实践 深入解析 Generalized Aggressive DecodingGAD基于非自回归草稿模型的无损自回归翻译加速实践【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm导读本文以 unilm 仓库中 decoding/GAD/readme.md 为核心系统讲解 Generalized Aggressive DecodingGAD广义激进解码的原理、模型下载、环境配置、数据预处理、NAT 草稿模型训练、GAD/GAD 推理以及 compound split BLEU 评估的完整实战流程。GAD 是一种通过一次性生成整块目标 token 自回归验证来无损加速机器翻译解码的算法文章将结合仓库中 inference.py、BlockNAT.py 等源码带读者理解 GAD 的 block 解码、top-beta 验证、tau 容差等核心机制掌握可复现、可运行的完整实现方案。一、GAD 是什么从逐 token 生成到整块草稿 一次验证1.1 问题背景自回归解码的速度瓶颈传统的自回归Autoregressive, AR机器翻译模型在解码阶段逐 token 生成输出——每一步只预测下一个 token并将该 token 拼接到输入序列后再次前向传播。这意味着生成长度为 T 的句子需要 T 次串行的模型前向计算GPU 利用率低下推理延迟与序列长度成正比。1.2 GAD 的核心思路一次解码整块一次验证整块GAD论文Lossless Speedup of Autoregressive Translation with Generalized Aggressive Decoding的核心思想是草稿生成Drafter使用一个非自回归Non-Autoregressive, NAT模型作为草稿模型一次性并行生成整块block目标 token验证Verifier使用原始自回归模型对整块 token 做一次并行前向验证一次性接受前缀中所有与 AR 模型 top-beta 候选一致的 token。由于验证一次即可接受多个 token而验证与生成都只消耗一次模型前向GAD 在不改变输出分布、不损失翻译质量的前提下显著降低了模型调用次数。这与仓库中Lossless Speedup无损加速的定位完全一致。1.3 本仓库的实现定位本仓库unilm 项目下的 decoding/GAD 目录是论文作者发布的 GAD 官方代码代码原始出处为 GAD 项目主页基于 GLATGlancing Transformer代码框架改造而来。整个目录自带了完整的 fairseq 工具链fairseq、GAD 专属的模型/任务/损失插件block_plugins、训练与推理入口脚本train.py、inference.py以及预训练模型下载入口是一个开箱即用的完整实验仓库。二、GAD 的两种变体Vanilla GAD 与 GAD在进入实操之前需要先理解 GAD 系列中两个关键变体因为推理脚本中的beta与tau两个超参数直接对应它们的区别变体betatop-beta 候选数tau容差特性Vanilla GADbeta1tau0仅当草稿 token 恰好等于 AR 模型 argmax 时才接受最严格、最保守GADbeta1如 5tau0如 3.0只要草稿 token 落在 AR 模型 top-beta 候选集合内即接受配合 tau 容忍 logit 差距接受率更高、加速比更大从 inference.py 的forward_decoder源码可以看到验证阶段的具体实现AR 模型对整块输入做一次前向后取出 logits 的topk(scores, beta)top-beta 候选随后对每个候选判断最佳分数与当前候选分数的差距是否超过 tautopk_scores, indexes torch.topk(decoder_out_tuple[0], beta, dim-1) ... for k, s in enumerate(topk_scores_list[i][j]): if topk_scores_list[i][j][0] - s tau: indexes_list[i][j][k] -1即差距超过tau的候选被标记为-1不参与接受判断。当beta1, tau0时只有 argmax 候选保留退化为严格的 Vanilla GAD当beta1, tau0时即为更激进的 GAD。注意GAD 放宽验证条件后理论上存在极小的分布偏移严格意义上的无损以 Vanilla GADbeta1为准——这一点在推理章节中还会结合--strategy gad的参数组合再强调。三、模型下载AT 验证器与 NAT 草稿器GAD 推理需要两个模型配合at-verifier-base自回归AT模型承担最终验证与质量保证nat-drafter-base (k25)非自回归NAT草稿模型block size 为 25负责一次性生成整块候选 token。仓库提供了四组语言方向的预训练模型可在 decoding/GAD/readme.md 中查看下载地址语言对模型wmt14.en-deat-verifier-base、nat-drafter-base (k25)wmt14.de-enat-verifier-base、nat-drafter-base (k25)wmt16.en-roat-verifier-base、nat-drafter-base (k25)wmt16.ro-enat-verifier-base、nat-drafter-base (k25)下载后建议统一放置到checkpoints/目录下。后续训练、推理命令中的checkpoint_path、AR_checkpoint_path即指向这两个文件。四、环境要求与安装4.1 依赖要求Python 3.7PyTorch 1.5.0仓库的 setup.py 中还列出了更细的依赖cython、numpy、regex、sacrebleu1.4.12、tqdm、hydra-core1.1、omegaconf2.1等。安装时会通过 Cython/C 编译 fairseq 的扩展模块如fairseq.libbleu、fairseq.libnat、fairseq.ngram_repeat_block_cuda等因此环境中需具备可用的 C 编译器与 CUDA 工具链。4.2 安装步骤conda create -n gad python3.7 cd GAD pip install --editable .--editable可编辑/开发模式安装会让fairseq包与源码保持同步便于修改block_plugins等插件后即时生效。安装完成后fairseq-train、fairseq-preprocess、fairseq-score等命令行工具即可直接使用见 setup.py 中注册的 console_scripts 入口。五、数据预处理复用仓库 BPE 码表与词典仓库在 decoding/GAD/data 目录下发布了各语言对的 BPE 码表与词典例如wmt14.en-de/bpe.32000BPE 码表、dict.en.txt、dict.de.txtwmt16.en-ro/dict.en.txt、dict.ro.txt其中wmt16.en-ro目录还提供了 get_data.sh可一键下载并整理 wmt16 原始数据含 BPE 切分的 train/valid/test 文件。将你的平行语料按 train/valid/test 划分并完成 BPE 编码后使用 fairseq-preprocess 制作二值化数据textPATH_YOUR_DATA srcsource_language tgttarget_language model_pathPATH_TO_MODEL_DICT_DIR fairseq-preprocess --source-lang ${src} --target-lang ${tgt} \ --trainpref $text/train --validpref $text/valid --testpref $text/test \ --destdir PATH_TO_BIN_DIR --workers 60 \ --srcdict ${model_path}/dict.${src}.txt \ --tgtdict ${model_path}/dict.${tgt}.txt关键点--srcdict/--tgtdict必须指向仓库发布的词典保证训练数据与预训练模型的词表完全一致--workers 60控制并行处理线程数可按机器核数调整--destdir为二值化输出目录即后续训练命令中的${bin_path}。六、训练 NAT 草稿模型BlockNAT6.1 训练命令GAD 的草稿模型是一个基于 GLAT 思想的 Block 级 NAT 模型通过--arch block注册对应 block_plugins/models/BlockNAT.py 中的BlockNAT类。训练命令如下仓库已提供完整参数的 train.shpython train.py ${bin_path} --arch block --noise block_mask --share-all-embeddings \ --criterion glat_loss --label-smoothing 0.1 --lr ${lr} --warmup-init-lr 1e-7 \ --stop-min-lr 1e-9 --lr-scheduler inverse_sqrt --warmup-updates ${warmup} \ --optimizer adam --adam-betas (0.9, 0.999) --adam-eps 1e-6 \ --task translation_lev_modified --max-tokens ${max_tokens} --weight-decay 0.01 \ --dropout ${dropout} --encoder-layers 6 --encoder-embed-dim 512 --decoder-layers 6 \ --decoder-embed-dim 512 --fp16 --max-source-positions 1000 \ --max-target-positions 1000 --max-update ${update} --seed ${seed} --clip-norm 5 \ --save-dir ./checkpoints --src-embedding-copy --log-interval 1000 \ --user-dir block_plugins --block-size ${size} --total-up ${update} \ --update-freq ${update_freq} --decoder-learned-pos --encoder-learned-pos \ --apply-bert-init --activation-fn gelu仓库 train.sh 中给出的一组可直接参考的参数取值参数值说明lr0.0005初始学习率dropout0.1dropout 比例warmup10000warmup 步数sizeblock-size25草稿块大小seed1随机种子max_tokens4096单批最大 token 数update_freq4梯度累积步数updatemax-update / total-up300000总训练步数6.2 关键机制解析结合源码1)--task translation_lev_modified块级掩码噪声注入该任务定义在 block_plugins/tasks/translation_lev_modified.py 中与标准 Levenshtein Transformer 任务translation_lev的区别在于支持block_mask噪声模式。inject_noise中_block_mask的实现逻辑translation_lev_modified.py为随机选择一个截断位置cutoff_length至少掩码 1 个 token将目标序列从截断位置起、长度为block_size的连续片段全部替换为unk形成草稿模型的输入prev_target同时保留完整的 padded 目标作为监督信号。也就是说训练时模型学到的是根据左侧已生成的前缀一次性预测接下来连续 block_size 个 token的能力——这正是推理阶段 NAT 草稿器一次生成整块的基础。2)--criterion glat_lossGLAT 课程式学习该损失定义在 block_plugins/criterions/glat_loss.py 中。GLATGlancing Transformer的核心是课程学习式的部分掩码策略在 translation_lev_modified.py 的train_step中随着训练步数推进context_p上下文保留比例从start_p线性衰减到start_p - minus_ptrain_ratio max(0, min(1, update_num / self.cfg.total_up)) sample[glat] {context_p: self.cfg.start_p - self.cfg.minus_p * train_ratio}其中start_p0.5、minus_p0.2默认值见 translation_lev_modified.py。训练早期模型看到更多已完成的上下文拟合更容易后期逐步过渡到更难的全块预测从而让模型能力平滑增强。在 BlockNAT.py 的forward中GLAT 的实现还会先用当前模型在无梯度条件下预测一遍掩码位置统计预测正确数与上下文比例动态决定本步实际掩码哪些位置实现模仿式的逐步引导训练并输出glat_accu、glat_context_p两个监控指标。3) 模型架构与训练技巧--arch block对应 BlockNAT.py 中的 6 层编码器 6 层解码器、embed dim 512、FFN dim 2048decoder_embed_dim*4的默认架构--apply-bert-init表示对编码器/解码器应用 BERT 风格初始化init_bert_params--src-embedding-copy表示将源语言词嵌入复制给解码器作为初始输入见 BlockNAT.py 的参数定义--decoder-learned-pos --encoder-learned-pos使用可学习位置编码--fp16开启混合精度训练--clip-norm 5设置梯度裁剪。七、推理GAD / GAD 解码7.1 推理命令与参数GAD 推理入口为 inference.py完整命令如下见 inference.sh其中设置beta5即为 GAD将beta1即为 vanilla GADpython inference.py ${data_dir} --path ${checkpoint_path} --user-dir block_plugins \ --task translation_lev_modified --remove-bpe --max-sentences 20 \ --source-lang ${src} --target-lang ${tgt} --iter-decode-max-iter 0 \ --iter-decode-eos-penalty 0 --iter-decode-with-beam 1 --gen-subset test \ --AR-path ${AR_checkpoint_path} --input-path ${input_path} \ --output-path ${output_path} --block-size ${block_size} --beta ${beta} --tau ${tau} \ --batch ${batch} --beam ${beam} --strategy ${strategy}inference.sh 中给出的参考取值strategygad、batch32、beam5、beta5、tau3.0、block_size25、srcen、tgtde。7.2 三种解码策略--strategy在 inference.py 的入口逻辑中--strategy支持三个取值用于横向对比与效果验证strategy含义实现函数fairseqfairseq 原始 beam search 解码作为标准参考fairseq_generateAR简化的自回归贪心解码基线baseline_generategadGAD / GAD 解码gad_generate其中fairseq策略直接复用 fairseq 的SequenceGeneratorbeam5AR策略则是逐 token 贪心解码见 inference.py 的baseline_generate两者均可作为 GAD 加速效果与质量一致性的对照基线。7.3 GAD 解码循环的源码级拆解gad_generateinference.py是 GAD 的核心循环每个 step 执行一次gad_forwardinference.py其流程为NAT 草稿将当前前缀含未填充的unk块一次性喂给 NAT 草稿模型model.decoder得到整块的候选 tokenoutput_tokens[i, start_pos:start_pos block_size]AR 验证拼接eos与整块候选后喂给 AR 验证器AR_model通过forward_decoder得到每个位置上的 top-beta 候选集合含 tau 容差过滤分叉点寻找bifurcation从块起始位置逐 token 检查——只要草稿 token 出现在 AR 候选集合中就被接受一旦遇到不在候选集合中的 token立即停止接受该位置记为分叉点bifurcation回退与续接被接受的 token 保留分叉点位置替换为 AR 模型在该位置的候选 tokenAR_verify_tokens[...][0]其后重新填入block_size个unk等待下一轮草稿终止条件若发现eos或达到max_len则截断句子并标记该样本完成start_pos_list[i] -1当所有样本完成时循环退出。因此每轮循环只消耗一次 NAT 前向 一次 AR 前向却可能接受数十个 token这就是 GAD 相对传统自回归解码的加速来源。而分叉点只替换一个 token这一设计保证了生成结果与 AR 模型逐步贪心解码的分布一致在 beta1 时严格无损。7.4 延迟测试脚本仓库还提供了面向论文延迟评测的 inference_paper.py。根据 readme 的说明论文中报告 GAD 推理延迟使用的是batch 1单样本实现因此若需复现论文中的加速比数据应使用该脚本而非默认的inference.py。八、评估计算 compound split BLEUGAD 官方评测使用compound split BLEU对德语等复合词做拆分后计算的 BLEU以更公平地反映翻译质量。仓库提供了两个评估入口8.1 一键脚本ref.sh直接运行即可对推理输出计算 compound split BLEU./ref.shref.sh 内部流程为从标准 fairseq 生成文件如output/beam5_result_en_de.out中提取参考译文将 GAD 输出如output/block.out与参考同时做连字符拆分perl -ple s{(\S)-(\S)}{$1 ##AT##-##AT## $2}g再调用fairseq-score计算 BLEU。脚本中默认的参考文件是仓库自带的 data/test.de.compound.ref。8.2 通用脚本compound_split_bleu.sh仓库还提供了更通用的 scripts/compound_split_bleu.sh用法为传入任意fairseq-generate的标准输出文件bash scripts/compound_split_bleu.sh YOUR_GENERATE_OUTPUT该脚本会从标准输出中分别抽取假设^H行与参考^T行做同样的连字符拆分后调用fairseq-score计算 compound split BLEU。九、完整复现流程速查将以上各环节串联一次完整的 GAD 实验如下安装环境创建 conda 环境Python 3.7进入 GAD 目录执行pip install --editable .准备数据按 get_data.sh 或自行 BPE 编码整理数据用仓库词典执行fairseq-preprocess训练草稿模型按 train.sh 的参数执行python train.py得到 NAT drafter checkpoint下载验证模型从 readme 模型表中下载对应语言对的 at-verifier-base推理按 inference.sh 修改checkpoint_path、AR_checkpoint_path、input_path、output_path与语言方向设置strategygad、beta、tau、block_size后执行python inference.py需要对照时分别用strategyfairseqbeam search与strategyAR贪心跑基线评测运行./ref.sh计算 compound split BLEU对比 GAD 与基线输出的质量是否一致并结合 inference_paper.pybatch 1统计解码延迟与加速比。十、注意事项与限制代码来源readme 明确说明本实现基于 GLAT 框架其上游为 fairseq 的 Levenshtein Transformer 任务体系translation_lev_modified任务即是在此基础上的 block 化改造无损性边界vanilla GADbeta1在设计上严格保持与 AR 模型一致的输出分布GADbeta1, tau0通过扩大候选集提升接受率适合追求更高加速比且对极轻微分布偏移不敏感的场景实验时应同时报告两者与基线fairseq/AR的 BLEU 差异运行环境命令中的 Python/PyTorch 版本约束Python 3.7、PyTorch 1.5.0来自 readmesetup.py 中的依赖版本如hydra-core1.1、omegaconf2.1也是该代码库可正常工作的前提在较新环境中安装时需注意版本兼容性延迟复现论文中的延迟数据基于 batch 1 实现若使用默认inference.py的 batch 模式加速比数值会有所不同本文所有命令与参数均以仓库实际文件为准模型下载需访问 readme 中提供的存储服务请在网络可用环境下进行。【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表