Deep Transformers with Latent Depth:基于 Fairseq 的多语言机器翻译自适应层深度训练实战指南

📅 发布时间:2026/9/13 14:15:38
Deep Transformers with Latent Depth:基于 Fairseq 的多语言机器翻译自适应层深度训练实战指南
Deep Transformers with Latent Depth基于 Fairseq 的多语言机器翻译自适应层深度训练实战指南【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm本篇技术指南围绕 unilm 仓库中 edgelm/examples/latent_depth/README.md 所介绍的Deep Transformers with Latent DepthLi et al., 2020, arXiv:2009.13102展开讲解如何在 Fairseq 中通过概率框架自动学习 Transformer 每层的选与不选并将其用于共享编码器/解码器的多语言机器翻译One-to-Many, O2M训练与推理。读完本文你将掌握 latent depth 的完整训练命令、每个超参数的语义与推荐值、底层 Gumbel-Sigmoid 采样与 KL/稀疏性损失的作用原理以及如何用训练好的模型进行解码评测。1. 方法背景为什么需要潜在深度标准的深层 Transformer 为所有输入固定使用全部层数推理时计算量恒定。Li et al. (2020) 提出的 latent depth 框架则让网络自动学习每一层是否被使用通过为层选择layer selection学习后验分布模型可以在不同样本、不同语言对上跳过不必要的层从而在保持精度的同时降低推理成本。在多语言场景下这一框架被扩展为训练一个共享的 Transformer 网络为每个语言对学习不同的层选择后验分布。例如在 One-to-Many 翻译中英语eng到不同目标语言的翻译其最有效层数可能各不相同——部分语言对只需要较浅的网络即可达到较好效果而另一些则需要更深的结构。latent depth 让这种差异由数据驱动地自动涌现而不是靠人工为每个语言对挑选层数。从源码结构看这一实现位于 edgelm/examples/latent_depth/latent_depth_src包含四个核心子模块task/实为 multilingual_translation_latent_depth.py注册multilingual_translation_latent_depth任务负责解析 latent depth 相关超参数并在训练/验证/推理阶段为每个语言对注入正确的采样索引models/注册latent_multilingual_transformer模型提供支持 latent depth 的编码器/解码器与模型结构参数modules/LayerSelect模块核心的层采样实现loss/LatentLayersKLLossKL 损失与LatentLayersSparsityLoss稀疏性/共享损失。2. 核心实现LayerSelect 与 Gumbel-Sigmoid 采样latent depth 的关键机制体现在 modules/latent_layers.py 的LayerSelect模块中。2.1 语言特定的层 logitsself.layer_logits torch.nn.Parameter( torch.Tensor(num_logits, num_layers), requires_gradTrue, )layer_logits的形状为(num_logits, num_layers)。在多语言场景下num_logits等于语言数量即每个语言对拥有一套独立的可学习层 logitsnum_logitslen(langs)见 latent_multilingual_transformer.py 中_get_module_class的num_logits传参。训练时通过set_lang_idx(lang_idx)指定当前样本所属语言sample(logit_idx)便从对应的 logits 行采样从而让每个语言对学到各自不同的层选择分布。2.2 Gumbel-Sigmoid 采样与硬/软选择每个采样值由 Gumbel-Sigmoid 分布生成两个 Gumbel(0,1) 噪声之差再经过 sigmoid因为其后要接 sigmoidgumbels1 (-torch.empty_like(logits).exponential_().log()) gumbels2 (-torch.empty_like(logits).exponential_().log()) gumbels1 (logits gumbels1 - gumbels2) / tau y_soft gumbels1.sigmoid()当hard_selectFalse软选择直接返回y_soft即重参数化reparameterization后的连续权重层以加权方式参与残差连接当hard_selectTrue硬选择通过 straight-through estimator 将y_soft二值化y_soft 0.5置 1否则置 0同时保留可导的软路径y_hard - y_soft.detach() y_soft保证梯度能够回传到 logits。hard_select的初始值由soft_select参数决定self.hard_select not soft_select默认硬选择温度tau默认 5.0。而训练过程中任务会根据 update 数动态切换硬/软选择model.models[lang_pair].decoder.layer_select.hard_select ( update_num self.args.soft_update )即--soft-update步数之内使用软采样可导、便于梯度传播之后切换为硬采样离散化、贴近真实推理行为。具体逻辑见 multilingual_translation_latent_depth.py。2.3 层如何被跳过LayerSelect被注入到每一个 Transformer 层中层的残差连接被改写为def residual_connection(self, x, residual): return residual x * self.layer_select(self.idx)这是 latent depth 的关键改动非残差子层self-attention、FFN的输出x乘以上一步采样得到的该层权重再与残差相加。若该层采样值接近 0则该层的实际贡献趋近于零等价于被跳过见 latent_transformer.py 与解码器层的对应实现。同时编码器/解码器在前向开始时统一调用一次layer_select.sample(lang_idx)一次性为所有层采样见 latent_transformer.py 和 latent_transformer.py这有利于分布式训练时的计算效率。3. 损失设计KL 正则 稀疏性/共享正则latent depth 的训练目标除了标准的交叉熵损失外还叠加了两类辅助损失均定义在 loss/latent_depth.py。3.1 LatentLayersKLLoss把层选择拉向先验KL 损失约束采样分布不要过度集中或过度发散--prior uniform以 0.5 为基准的均匀先验kl_loss (samples * (log(samples) - log(0.5))).sum(-1)--prior agged_posterior使用聚合后验aggregated posterior作为先验即先对当前 batch 中所有语言的采样做归一化统计再约束各语言对的分布与之接近。KL 损失还会按层数归一化并乘以一个随 update 退火anneal的权重kl_weight min( self.args.sparsity_weight, (update_num - self.args.soft_update) * self.args.sparsity_weight / self.args.anneal_updates, )即权重从 0 开始经过--anneal-updates步线性增长到--sparsity-weight避免训练初期就施加过强约束。3.2 LatentLayersSparsityLoss目标层数 跨语言共享稀疏性损失在update_num soft_update anneal_updates之后生效is_valid判断--target-layers 0时关闭包含两部分目标层数约束统计所有语言平均每层被选中的概率得到layer_utilization计算期望被选层数expeted_layers sum(layer_utilization)再对其与--target-layers做 L2 损失(expected - target) ** 2鼓励模型实际使用的有效层数接近目标值跨语言共享约束share loss对layer_utilization计算熵的负值-sum(v * log(v))当--share-weight 0时鼓励不同语言对的层选择模式趋于一致从而更好地共享单一网络。这两个损失都在train_step中对所有语言对的采样统一计算见 multilingual_translation_latent_depth.py并通过loss.backward(retain_graphTrue)保留计算图后追加反向传播。注意论文源码中 sparsity 损失对 target-layers 项的系数也复用了share_weight见loss/latent_depth.py的global_sparsity_loss分支配置时二者需协同调整。4. 训练配置实战多语言 latent depthREADME 给出了一个完整的 One-to-ManyO2M训练示例8 个英语 → 其他语言方向eng-aze,eng-bel,eng-ces,eng-glg,eng-por,eng-rus,eng-slk,eng-tur数据使用与 Balancing Training for Multilingual NMT (Wang et al., 2020) 相同的 TED8 数据集需先经过 numberize 与 binarize 预处理。lang_pairs_streng-aze,eng-bel,eng-ces,eng-glg,eng-por,eng-rus,eng-slk,eng-tur databin_dirpath to binarized data fairseq-train ${databin_dir} \ --user-dir examples/latent_depth/latent_depth_src \ --lang-pairs ${lang_pairs_str} \ --arch multilingual_transformer_iwslt_de_en \ --task multilingual_translation_latent_depth \ --criterion label_smoothed_cross_entropy --label-smoothing 0.1 \ --share-encoders \ --share-decoders \ --decoder-langtok \ --share-decoder-input-output-embed \ --dropout 0.3 --attention-dropout 0.3 \ --optimizer adam --adam-eps 1e-06 --adam-betas (0.9, 0.98) \ --lr-scheduler inverse_sqrt --stop-min-lr 1e-9 --warmup-init-lr 1e-7 --warmup-updates 8000 \ --max-tokens 4096 --update-freq 1 \ --lr 0.0015 \ --clip-norm 1.0 \ --seed 2 \ --ddp-backendlegacy_ddp \ --encoder-layers 12 \ --decoder-layers 24 \ --decoder-latent-layer \ --sparsity-weight 0.1 \ --anneal-updates 5000 \ --soft-update 500 \ --target-layers 12 \ --share-weight 0.14.1 关键参数速查表以下参数中latent depth 专属参数由任务multilingual_translation_latent_depth.py与模型latent_multilingual_transformer.py注册参数默认值含义与建议--task multilingual_translation_latent_depth—启用 latent depth 的多语翻译任务必须与--user-dir配合加载插件--encoder-latent-layer关闭在编码器中启用层选择README 示例仅启用解码器--decoder-latent-layer关闭在解码器中启用层选择本示例开启--target-layers-1不约束期望的有效层数示例为 12即希望 24 层解码器实际只使用约 12 层--sparsity-weight0.0KL 损失退火后的最终权重示例 0.1--share-weight0.0跨语言共享层利用率熵损失权重示例 0.1--soft-update1前 N 步使用软采样可导之后切换硬采样示例 500--anneal-updates1KL/sparsity 权重线性退火的步数示例 5000--prioruniformKL 先验uniform或agged_posterior--soft-select关闭模型级开关训练与推理全程使用软样本--sampling-tau5.0Gumbel-Sigmoid 采样温度4.2 参数之间的协同关系--soft-update 500意味着前 500 步hard_selectFalse、kl_weight为 0网络以常规方式预热--anneal-updates 5000表示此后约 5000 步内KL 权重从 0 线性升至--sparsity-weight 0.1--target-layers 12与--share-weight 0.1对应的稀疏性/共享损失只在update_num soft_update anneal_updates之后才真正生效且其权重同样经历退火见 loss/latent_depth.py 与 loss/latent_depth.py。4.3 模型结构要点--arch multilingual_transformer_iwslt_de_en是注册于标准多语 Transformer 之上的架构名latent depth 模型本身注册为latent_multilingual_transformerlatent_multilingual_transformer.py其默认结构为编码器 embed/FFN 512/1024、4 头注意力、12 层解码器 512/1024、4 头、24 层并默认开启共享编码器/解码器及其嵌入latent_multilingual_transformer.py。README 命令通过--encoder-layers 12 --decoder-layers 24显式覆盖启用 latent layer 时任务会强制校验共享设置编码器启用 latent layer 必须--share-encoders解码器启用必须--share-decoders否则直接断言报错multilingual_translation_latent_depth.py--share-encoders与--share-decoders使所有语言对共用一套编码器/解码器参数这正是一个共享网络、每个语言对一套层选择 logits的前提num_logitslen(langs)--decoder-langtok在解码器端附加语言 token帮助模型区分目标语言数据侧--max-tokens 4096、--update-freq 1、--lr 0.0015、inverse_sqrt 学习率与 8000 步 warmup 共同构成示例的稳定训练配置。5. 推理与评测fairseq-generate训练完成后使用fairseq-generate对指定语言对和数据集划分进行翻译评测lang_pairs_streng-aze,eng-bel,eng-ces,eng-glg,eng-por,eng-rus,eng-slk,eng-tur databin_dirpath to binarized data model_pathpath to checkpoint src_langsource language to translate from tgt_langtarget language to translate to gen_dataname of data split, e.g. valid, test, etc fairseq-generate ${databin_dir} \ --path ${model_path} \ --task multilingual_translation_latent_depth \ --decoder-latent-layer \ --lang-pairs ${lang_pairs_str} \ -s ${src_lang} -t ${tgt_lang} \ --gen-subset $gen_data \ --scoring sacrebleu \ --remove-bpe sentencepiece \ --lenpen 1.0 \ --beam 5 \ --decoder-langtok \ --max-tokens 4096推理阶段任务会在inference_step中根据--source-lang/--target-lang自动设置编码器/解码器的语言索引src_lang_idx_dict、tgt_lang_idx_dict见 multilingual_translation_latent_depth.py从而为该语言对选择训练好的层选择 logits。此时hard_select保持为训练后期状态硬选择模型按学习到的该用哪些层执行实际推理。注意如果模型训练时同时启用了--encoder-latent-layer则推理命令中同样需要补上--encoder-latent-layer参数验证阶段valid loss同样会为每个语言对设置语言索引multilingual_translation_latent_depth.py因此可以用fairseq-validate按语言对评估。6. 从源码理解训练-推理的一致性LayerSelect.sample在每次前向时调用训练与推理共享同一条采样路径保证行为一致任务根据当前 batch 的语言对调用set_lang_idx并依据update_num soft_update更新hard_select标志编码器/解码器前向开始时执行layer_select.sample(lang_idx)从该语言对应的 logits 行做 Gumbel-Sigmoid 采样硬选择模式下二值化每个 Transformer 层的residual_connection用采样值缩放非残差子层输出实现加权通过/跳过标准交叉熵 KL 损失每语言对 稀疏性/共享损失所有语言统一联合优化。这一设计使得训练时学习到的层选择分布能够无缝迁移到推理时的离散层选择实现真正的自适应深度。7. 引用与延伸阅读若你的工作基于或参考了该方法请按如下方式引用摘自 README.mdarticle{li2020deep, title{Deep Transformers with Latent Depth}, author{Li, Xian and Stickland, Asa Cooper and Tang, Yuqing and Kong, Xiang}, journal{arXiv preprint arXiv:2009.13102}, year{2020} }本文介绍的插件目录 edgelm/examples/latent_depth 位于 unilm 仓库的 edgelm 示例集中与其配套的 Fairseq 框架代码位于 edgelm/fairseq。多语翻译任务基类MultilingualTranslationTask与共享模型基类MultilingualTransformerModel的实现分别对应 edgelm/fairseq/tasks/multilingual_translation.py 与 edgelm/fairseq/models/multilingual_transformer.py读者可结合这些文件进一步理解 latent depth 插件对标准多语翻译流程的扩展点。【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考