
unilm FST 系统详解IWSLT21 多语言语音翻译的语音文本联合训练实现【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm本篇围绕 unilm 仓库中edgelm/examples/speech_text_joint_to_text目录下的 IWSLT21 文档iwslt2021.md展开完整还原 FAIR 语音翻译系统 FST 在多语言共享任务上的数据准备、双输入模型训练与评测流程并结合 任务实现、模型实现 与 引导损失准则 的源码解释训练脚本中每个关键参数的作用机制。读完后你将能够独立复现“wav2vec mBART 堆叠编码器 文本知识蒸馏”的联合语音文本训练理解guide-alpha、enc-grad-mult、attentive-cost-regularization等参数在源码中的落点并掌握--load-speech-only、--infer-target-lang在推理阶段的实现方式。一、FST 系统与语音文本联合训练该目录对应论文FST: the FAIR Speech Translation System for the IWSLT21 Multilingual Shared Task的公开代码是 unilm 仓库中 fairseq 扩展edgelm对 S2Tspeech-to-text框架的增强版本在标准“语音 → 文本”任务之上叠加一个与语音编码器共享解码器的“文本 → 文本”翻译任务让模型从平行文本翻译数据中直接学习语义对齐再反哺语音翻译。模块 README 中给出了该系列工作的完整引用线索Tang 等的 ICASSP 2021 多任务框架、ACL 2021 辅助文本翻译任务、IWSLT 2021 FST 系统以及 fairseq S2T 工具链论文。IWSLT21 共享任务的多语言语音翻译基于 TEDx 音频语料源语言为西班牙语es、法语fr、意大利语it、葡萄牙语pt目标语言覆盖英语en、西语、法语、意语、葡语。任务在源码中注册为speech_text_joint_to_textregister_task(speech_text_joint_to_text) class SpeechTextJointToTextTask(SpeechToTextTask): Task for joint training speech and text to text.见 tasks/speech_text_joint.py。与基类SpeechToTextTask相比它持有两套词表——src_dict承载语音旁转写的音素序列与tgt_dict目标语言文字并在setup_task中校验两者pad/eos索引一致见 setup_task 实现这正是音素源文本方案的前提。二、数据准备2.1 需要下载的基础文件按照文档训练前需从 FAIR 公开渠道获取三个 IWSLT 数据文件存放于数据目录MANIFEST_ROOT下SentencePiece 模型spm.model目标语言文字的 BPE 分词模型词典dict.txt目标语言词表tgt_dict配置文件config.yamlfairseq S2T 的数据配置指定音频 tsv 清单、词表文件名vocab_filename/src_vocab_filename等由任务在构造时解析self.data_cfg S2TJointDataConfig(Path(args.data) / args.config_yaml)见 speech_text_joint.py。2.2 音频 TSV 准备原始音频 tsv 文件的准备方式与 fairseq S2T 的 speech-to-text 示例一致关键点是带上--use-audio-input选项使每行引用原始音频文件而非特征。每一行需要包含audio音频路径/URL、tgt_text参考译文等列。2.3 音素化的源文本列src_text联合训练要求 tsv 中的src_text列是音素phoneme表示而非普通拼写文本准备方法与 MuST-C 英德示例 相同。仓库提供了音素转换脚本 scripts/g2p_encode.py配合噪声词保留表 configs/mustc_noise.list把 (Applause)、(Laughter) 等旁注映射为 NOISE/VOICE 类 tokenpython examples/speech_text_joint_to_text/scripts/g2p_encode.py \ --lower-case --do-filter --use-word-start --no-punc \ --reserve-word examples/speech_text_joint_to_text/configs/mustc_noise.list \ --data-path ${mustc_src_text} \ --out-path ${mustc_src_text_pho}生成的音素序列写回 tsv 的src_text列并据此构建音素词典src_dict。音素词表与文字词表分离是模型中双编码器共享同一个 mBART 词嵌入空间却输入不同模态的基础。三、模型dualinputxmtransformer_base 架构训练命令指定--arch dualinputxmtransformer_base对应注册模型dual_input_xm_transformers2t_dualinputxmtransformer.py注释明确其结构编码器 wav2vec 语音编码器 mBART 文本编码器解码器 mBART 文本解码器。3.1 堆叠编码器wav2vec 上叠 mBART 层--stack-w2v-mbart-encoder触发 StackedWav2VecEncoderWithAdaptorif drop_w2v_layers 0: self.w2v_encoder.w2v_model.encoder.layers ( self.w2v_encoder.w2v_model.encoder.layers[:-drop_w2v_layers] )前向流程为wav2vec 特征编码 →adaptor降采样/投影 →直接接 mBART 编码器的全部 12 层 Transformerself.mbart_encoder_layers→ 最终 LayerNorm。训练脚本中--drop-w2v-layers 12会丢弃 wav2vec 的最后 12 层只保留卷积特征提取器 前段 Transformer 层把深层语义建模交给 mBART 层从而让语音路径与文本路径共享同一套 Transformer 参数结构。文本编码器则是独立的TransformerEncoder其权重通过--load-pretrained-mbart-from从 mBART 检查点加载if getattr(args, load_pretrained_mbart_from, None): text_encoder checkpoint_utils.load_pretrained_component_from_model( componenttext_encoder, checkpointargs.load_pretrained_mbart_from )见 build_encoder。3.2 冻结策略与梯度控制build_encoder/build_decoder中w2v、mBART 编码器、mBART 解码器三组参数默认全部冻结仅当参数名匹配finetune-*-params模式时才requires_grad Truefor k, p in spch_encoder.named_parameters(): if safe_hasattr(args, finetune_w2v_params) and \ XMTransformerModel.finetune_params(args.finetune_w2v_params, k): p.requires_grad True else: p.requires_grad False本脚本使用--finetune-w2v-params all --finetune-mbart-decoder-params all --finetune-mbart-encoder-params all即三个组件全部解冻参与微调配合--w2v-path指定的预训练 xlsr 检查点与--load-pretrained-mbart-from指定的 mBART 检查点初始化。--enc-grad-mult 2.0对应 DualInputEncoder 中的GradMultiply操作双输入 batch 训练时语音与文本两路编码器输出的梯度都乘以该系数用于调节编码器侧梯度量级merge_output 实现。3.3 跨注意力正则化attentive-cost-regularization--attentive-cost-regularization 0.02一旦大于 0两处生效编码器在计算注意力损失所需的倒数第二层隐状态cross_attentive_loss_before_last_layer 0解码器 TransformerMultiInputDecoder 在训练时计算cross_attentive_loss分别对语音侧隐状态teacher与文本侧隐状态student做 L2 归一化后求相似度矩阵用 softmax 加权重构后比较两者距离得到逐位置对齐代价attn_cost最终按系数加进总损失。其语义是让文本编码器对源句的隐表示在几何上能够“解释”语音编码器学到的表示从而约束语音编码器提取出更适合跨模态对齐的特征。3.4 base 架构默认值dualinputxmtransformer_base 的默认配置编码/解码嵌入维度 1024、FFN 4096、各 12 层、16 头、pre-norm、gelu激活、嵌入做 LayerNormmbart_dropout默认 0.1脚本覆盖为 0.3mbart_attention_dropout默认 0.0脚本通过--attention-dropout 0.3传入 wav2vec 侧。--skip-encoder-projection表示跳过编码器中的投影层参数定义--normalize属于 wav2vec 编码器参数由set_default_w2v_encoder_args定义用于对输入音频特征做归一化。3.5 引导损失准则guided_label_smoothed_cross_entropy_with_accuracy--criterion guided_label_smoothed_cross_entropy_with_accuracy注册于 text_guide_cross_entropy_acc.py是联合训练的核心。对双输入 batch语音 同句音素文本模型输出按 batch 维切半lprobs_spch, lprobs_text torch.chunk(lprobs, 2)文本侧走普通 label-smoothed NLL语音侧调用 guide_loss_and_acc实现在线知识蒸馏probs_teacher lprobs_teacher.exp().masked_fill_(...) probs_teacher probs_teacher.detach() guide_loss -(probs_teacher * lprobs).sum() loss self.alpha * guide_loss (1.0 - self.alpha) * loss # loss 为语音 NLL即语音解码输出以文本解码输出detach 后的软标签为教师--guide-alpha 0.8表示损失中 80% 权重给蒸馏项、20% 给真值 NLL。--disable-text-guide-update-num 5000保证前 5000 个 update 只用 NLLmodel.num_updates self.disable_update_num时直接退化为普通 CE避免训练初期教师分布尚不可靠时误导学生。此外准则还会把attn_cost * attn_beta叠加进总损失并区分记录speech_loss、speech_nll_loss、speech_attn_loss等日志指标见 aggregate_logging_outputs便于训练中单独观察语音分支与正则项的收敛情况。注意本脚本未指定--parallel-text-data/--langpairs因此文本引导完全来自同一语音样本的音素转写dual-input batch而不是独立平行语料独立文本语料路径在任务中由 load_langpair_dataset 负责装载用于 MuST-C 示例中的 WMT 文本数据。四、训练4.1 预训练模型预训练 mBART 检查点mbart.pt用于初始化文本编码器与解码器预训练 wav2vec 检查点xlsr_53_56k.pt通过--w2v-path传入XLS-R 系列自监督语音编码器。4.2 完整训练命令文档给出的训练脚本8 语言子集训练、15 个验证子集python train.py ${MANIFEST_ROOT} \ --save-dir ${save_dir} \ --user-dir examples/speech_text_joint_to_text \ --train-subset train_es_en_tedx,train_es_es_tedx,train_fr_en_tedx,train_fr_es_tedx,train_fr_fr_tedx,train_it_it_tedx,train_pt_en_tedx,train_pt_pt_tedx \ --valid-subset valid_es_en_tedx,valid_es_es_tedx,valid_es_fr_tedx,valid_es_it_tedx,valid_es_pt_tedx,valid_fr_en_tedx,valid_fr_es_tedx,valid_fr_fr_tedx,valid_fr_pt_tedx,valid_it_en_tedx,valid_it_es_tedx,valid_it_it_tedx,valid_pt_en_tedx,valid_pt_es_tedx,valid_pt_pt_tedx \ --config-yaml config.yaml --ddp-backend no_c10d \ --num-workers 2 --task speech_text_joint_to_text \ --criterion guided_label_smoothed_cross_entropy_with_accuracy \ --label-smoothing 0.3 --guide-alpha 0.8 \ --disable-text-guide-update-num 5000 --arch dualinputxmtransformer_base \ --max-tokens 500000 --max-sentences 3 --max-tokens-valid 800000 \ --max-source-positions 800000 --enc-grad-mult 2.0 \ --attentive-cost-regularization 0.02 --optimizer adam \ --clip-norm 1.0 --log-format simple --log-interval 200 \ --keep-last-epochs 5 --seed 1 \ --w2v-path ${w2v_path} \ --load-pretrained-mbart-from ${mbart_path} \ --max-update 1000000 --update-freq 4 \ --skip-invalid-size-inputs-valid-test \ --skip-encoder-projection --save-interval 1 \ --attention-dropout 0.3 --mbart-dropout 0.3 \ --finetune-w2v-params all --finetune-mbart-decoder-params all \ --finetune-mbart-encoder-params all --stack-w2v-mbart-encoder \ --drop-w2v-layers 12 --normalize \ --lr 5e-05 --lr-scheduler inverse_sqrt --warmup-updates 50004.3 关键参数解析参数取值源码中的作用--task/--archspeech_text_joint_to_text/dualinputxmtransformer_base任务与模型注册入口见 任务 与 架构--train-subset/--valid-subset8 个训练 / 15 个验证 tedx 子集多语言配对按 tsv 清单逐一加载覆盖 es/fr/it/pt 到 en 及彼此互译--criterionguided_label_smoothed_cross_entropy_with_accuracyNLL 文本软标签蒸馏 注意力代价的复合损失--guide-alpha0.8蒸馏项与 NLL 的插值系数loss α·guide (1-α)·nll--disable-text-guide-update-num5000前 5000 步禁用蒸馏只训 NLL--label-smoothing0.3NLL 的标签平滑率 ε--enc-grad-mult2.0双输入时两路编码器梯度乘 2.0GradMultiply--attentive-cost-regularization0.02文本/语音隐状态对齐代价的权重 β--max-tokens/--max-sentences500000 / 3音频序列极长帧级 token故用极小句数、极大 token 上限控批--max-source-positions800000放开源序列长度上限以容纳长音频帧序列--w2v-path/--load-pretrained-mbart-from检查点路径分别初始化 wav2vec 编码器与 mBART 文本编码器/解码器--stack-w2v-mbart-encoder/--drop-w2v-layers开启 / 12丢弃 wav2vec 末尾 12 层其上堆叠 mBART 12 层--finetune-*-params all3 处均为 all默认全冻结all表示 w2v、mBART 编码器、解码器全解冻微调--skip-encoder-projection开关跳过编码器投影层--skip-invalid-size-inputs-valid-test开关验证时跳过长度超限样本防止长音频样本中断验证--update-freq/--max-update4 / 1000000梯度累积 4 个 mini-batch总更新步数上界--lr/--lr-scheduler/--warmup-updates5e-05 / inverse_sqrt / 5000微调大模型用小学习率 平方根倒数衰减--ddp-backendno_c10d使用非 c10d 的 DDP 后端train.py为仓库根部的 fairseq 训练入口edgelm 分支--user-dir examples/speech_text_joint_to_text使上述任务、模型、准则通过 user-dir 机制自动发现注册。五、评测与结果5.1 评测命令python ./fairseq_cli/generate.py ${MANIFEST_ROOT} \ --task speech_text_joint_to_text \ --user-dir ./examples/speech_text_joint_to_text \ --load-speech-only --gen-subset test_es_en_tedx \ --path ${model} \ --max-source-positions 800000 \ --skip-invalid-size-inputs-valid-test \ --config-yaml config.yaml \ --infer-target-lang en \ --max-tokens 800000 \ --beam 5 \ --results-path ${RESULTS_DIR} \ --scoring sacrebleu其中两个与联合任务强相关的参数在源码中可定位--load-speech-only任务级开关推理时只装载语音数据不构造文本分支任务参数定义source_dictionary在 speech-only 模式下返回None--infer-target-lang en当数据配置开启prepend_tgt_lang_tag_no_change时setup_task会解析出lang:en标签的 token id并在inference_step中作为解码起始bos_token传入生成器见 inference_step 与 LANG_TAG 常量 L34。这是多目标语言共一个解码器的条件下指定输出语言的标准做法评测其他方向时替换为对应语言即可--max-tokens 800000/--max-source-positions 800000与训练保持一致的长序列约束--scoring sacrebleu由 fairseq 生成管线直接输出 BLEU。5.2 各方向 BLEU 结果文档报告的 11 个翻译方向 BLEUTEDx 测试集directiones_enfr_enpt_enit_enfr_espt_esit_eses_esfr_frpt_ptit_itBLEU31.6236.9335.0727.1238.8735.5734.1374.5974.6470.8469.76同语对方向es_es、fr_fr 等得分显著高于跨语言方向符合语音翻译任务中“同语言 ASRMT 组合”上限的直观预期。文档同时提供训练好的checkpoint17.pt检查点供直接评测复现。六、小结与延伸阅读IWSLT21 示例展示了 unilm 仓库中语音文本联合训练最完整的落地形态音素化源文本 wav2vec/mBART 堆叠编码器 双输入 batch 的在线知识蒸馏 跨注意力对齐正则四个部件分别对应src_text数据管线、stack-w2v-mbart-encoder模型选项、guide-alpha损失项与attentive-cost-regularization正则项。若需从零训练而非继承预训练权重的基线可对照同目录的 MuST-C 英德文档dualinputs2ttransformer_s/m架构、--parallel-text-data独立文本语料、--mask-text-ratio文本掩码等选项双输入编码器/解码器的完整参数如--encoder-shared-layers、--decoder-shared-layer-level、--load-pretrain-speech-encoder定义在 s2t_dualinputtransformer.py是复现与调参时最值得通读的源码文件。【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考