Deep Transformers with Latent Depth:基于 Fairseq 的多语言机器翻译自适应层深度训练实战指南
2026/9/13 14:15:27 网站建设 项目流程

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 Depth(Li 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/LatentLayersKLLoss(KL 损失)与LatentLayersSparsityLoss(稀疏性/共享损失)。

2. 核心实现:LayerSelect 与 Gumbel-Sigmoid 采样

latent depth 的关键机制体现在 modules/latent_layers.py 的LayerSelect模块中。

2.1 语言特定的层 logits

self.layer_logits = torch.nn.Parameter( torch.Tensor(num_logits, num_layers), requires_grad=True, )

layer_logits的形状为(num_logits, num_layers)。在多语言场景下,num_logits等于语言数量,即每个语言对拥有一套独立的可学习层 logitsnum_logits=len(langs),见 latent_multilingual_transformer.py 中_get_module_classnum_logits传参)。训练时通过set_lang_idx(lang_idx)指定当前样本所属语言,sample(logit_idx)便从对应的 logits 行采样,从而让每个语言对学到各自不同的层选择分布。

2.2 Gumbel-Sigmoid 采样与硬/软选择

每个采样值由 Gumbel-Sigmoid 分布生成(两个 Gumbel(0,1) 噪声之差再经过 sigmoid,因为其后要接 sigmoid):

gumbels1 = (-torch.empty_like(logits).exponential_().log()) gumbels2 = (-torch.empty_like(logits).exponential_().log()) gumbels1 = (logits + gumbels1 - gumbels2) / tau y_soft = gumbels1.sigmoid()
  • hard_select=False(软选择):直接返回y_soft,即重参数化(reparameterization)后的连续权重,层以"加权"方式参与残差连接;
  • hard_select=True(硬选择):通过 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_graph=True)保留计算图后追加反向传播。注意:论文源码中 sparsity 损失对 target-layers 项的系数也复用了share_weight(见loss/latent_depth.pyglobal_sparsity_loss分支),配置时二者需协同调整。

4. 训练配置实战:多语言 latent depth

README 给出了一个完整的 One-to-Many(O2M)训练示例: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_str="eng-aze,eng-bel,eng-ces,eng-glg,eng-por,eng-rus,eng-slk,eng-tur" databin_dir=<path 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-backend=legacy_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.1

4.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 先验:uniformagged_posterior
--soft-select关闭模型级开关:训练与推理全程使用软样本
--sampling-tau5.0Gumbel-Sigmoid 采样温度

4.2 参数之间的协同关系

  • --soft-update 500意味着前 500 步hard_select=Falsekl_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_transformer(latent_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_logits=len(langs));
  • --decoder-langtok在解码器端附加语言 token,帮助模型区分目标语言;
  • 数据侧--max-tokens 4096--update-freq 1--lr 0.0015、inverse_sqrt 学习率与 8000 步 warmup 共同构成示例的稳定训练配置。

5. 推理与评测:fairseq-generate

训练完成后,使用fairseq-generate对指定语言对和数据集划分进行翻译评测:

lang_pairs_str="eng-aze,eng-bel,eng-ces,eng-glg,eng-por,eng-rus,eng-slk,eng-tur" databin_dir=<path to binarized data> model_path=<path to checkpoint> src_lang=<source language to translate from> tgt_lang=<target language to translate to> gen_data=<name 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_dicttgt_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在每次前向时调用,训练与推理共享同一条采样路径,保证行为一致:

  1. 任务根据当前 batch 的语言对调用set_lang_idx,并依据update_num > soft_update更新hard_select标志;
  2. 编码器/解码器前向开始时执行layer_select.sample(lang_idx):从该语言对应的 logits 行做 Gumbel-Sigmoid 采样,硬选择模式下二值化;
  3. 每个 Transformer 层的residual_connection用采样值缩放非残差子层输出,实现"加权通过/跳过";
  4. 标准交叉熵 + KL 损失(每语言对)+ 稀疏性/共享损失(所有语言统一)联合优化。

这一设计使得"训练时学习到的层选择分布"能够无缝迁移到推理时的离散层选择,实现真正的自适应深度。

7. 引用与延伸阅读

若你的工作基于或参考了该方法,请按如下方式引用(摘自 README.md):

@article{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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询