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等于语言数量,即每个语言对拥有一套独立的可学习层 logits(num_logits=len(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,因为其后要接 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.py的global_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.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-weight | 0.0 | KL 损失退火后的最终权重;示例 0.1 |
--share-weight | 0.0 | 跨语言共享(层利用率熵)损失权重;示例 0.1 |
--soft-update | 1 | 前 N 步使用软采样(可导),之后切换硬采样;示例 500 |
--anneal-updates | 1 | KL/sparsity 权重线性退火的步数;示例 5000 |
--prior | uniform | KL 先验:uniform或agged_posterior |
--soft-select | 关闭 | 模型级开关:训练与推理全程使用软样本 |
--sampling-tau | 5.0 | Gumbel-Sigmoid 采样温度 |
4.2 参数之间的协同关系
--soft-update 500意味着前 500 步hard_select=False、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_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_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.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),仅供参考