基于 fairseq 的语音-文本联合训练(Joint Speech Text Training)实战指南:以 MuST-C 英德与 IWSLT 2021 多语种语音翻译为例
【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm
导读
本文围绕 speech_text_joint_to_text 示例模块展开,系统讲解如何在 fairseq 中实现"语音到文本 + 文本到文本"双任务联合训练:通过引入共享编码器、引导式交叉熵损失、交叉注意力正则化与在线知识蒸馏,让语音翻译模型充分复用海量文本翻译数据。读完本文,你将掌握 MuST-C 英德(En-De)与 IWSLT 2021 多语种语音翻译两条完整流水线的数据准备、训练与评估方法,并能从源码层面理解speech_text_joint_to_text任务、dual_input_s2t_transformer模型与guided_label_smoothed_cross_entropy_with_accuracy准则的底层机制。
一、背景:为什么要做语音-文本联合训练
纯语音到文本(S2T)任务普遍受限于有标注语音数据稀缺,而纯文本翻译数据(如 WMT)体量庞大、唾手可得。speech_text_joint_to_text模块正是为了解决这一矛盾:它是 fairseq S2T 项目(详见 speech_to_text 示例)的扩展,在语音到文本任务的基础上,共训练一个文本到文本映射任务,让两条任务线共享模型参数,从而把文本语料中的翻译知识迁移到语音翻译上。
该模块的完整技术路线来自两篇论文:
- 联合训练基线:A General Multi-Task Learning Framework to Leverage Text Data for Speech to Text Tasks(Tang 等,ICASSP 2021)——提出语音与文本联合训练的基本多任务框架;
- 增强联合训练:Improving Speech Translation by Understanding and Learning from the Auxiliary Text Translation Task(Tang 等,ACL 2021)——在基线之上引入预训练模型初始化、交叉注意力正则化(CAR)与在线知识蒸馏(online KD),效果显著提升。
二、模块结构与源码地图
整个示例模块位于仓库的kosmos-2/fairseq/examples/speech_text_joint_to_text/下,其组织方式如下:
speech_text_joint_to_text/ ├── configs/ │ └── mustc_noise.list # 保留词表(噪声/掌声等标记映射,供 g2p 编码时原样保留) ├── criterions/ │ └── text_guide_cross_entropy_acc.py # 引导式标签平滑交叉熵准则(含 KD 与 CAR) ├── docs/ │ ├── ende-mustc.md # MuST-C 英德语音翻译联合训练示例 │ └── iwslt2021.md # IWSLT 2021 多语种语音翻译联合训练示例 ├── models/ │ ├── s2t_dualinputtransformer.py # 双输入 S2T Transformer(scratch 训练用) │ └── s2t_dualinputxmtransformer.py # 双输入 XM-Transformer(w2v + mBART 预训练用) ├── scripts/ │ └── g2p_encode.py # 英文→音素(phoneme)编码脚本 └── tasks/ └── speech_text_joint.py # 联合训练任务:数据加载与多模态 batch 组织四个核心组件各司其职:
| 组件 | 注册名 | 职责 |
|---|---|---|
| Task | speech_text_joint_to_text | 同时加载语音数据集与并行文本数据集,按采样比例混合成多模态 batch |
| Criterion | guided_label_smoothed_cross_entropy_with_accuracy | 对语音输入使用文本输出的概率分布做引导(在线 KD),并叠加交叉注意力正则化损失 |
| Model | dual_input_s2t_transformer | 双编码器(语音 S2T 编码器 + 文本 Transformer 编码器)+ 双解码器(共享参数) |
| 数据脚本 | g2p_encode.py | 将英文源文本转成音素序列,使文本侧与语音侧在"发音"层面更对齐 |
三、核心机制:从源码看联合训练如何工作
3.1 任务层:两种输入如何混合
SpeechTextJointToTextTask继承自SpeechToTextTask,在 tasks/speech_text_joint.py 中定义了若干关键参数:
--parallel-text-data:并行文本数据目录(即 WMT 等纯文本平行语料),为空则退化为纯语音任务;--langpairs:文本训练的语言对,逗号分隔,如en-de;--speech-sample-ratio/--text-sample-ratio:语音数据与文本数据的采样倍数,默认均为 1;--max-tokens-text/--max-positions-text:文本输入的 batch token 上限与单句最大长度(默认 400);--update-mix-data:当update-freq > 1时,在一次 update 内混合多种模态数据;--load-speech-only:推理/加载时只处理语音数据;--mask-text-ratio/--noise-token:文本源侧掩码比例(如掩码 15% 源词以增强鲁棒性),noise-token指定掩码替换符号(如▁NOISE)。
在load_dataset中,语音数据通过SpeechToTextJointDatasetCreator.from_tsv读取(tsv 中的src_text列即为音素化后的源文本),文本数据通过load_langpair_dataset读取;两者随后被包装进MultiModalityDataset的两个ModalityDatasetItem(sup_speech与text),由get_batch_iterator依据mult_ratio = [speech_sample_ratio, text_sample_ratio]采样并构造GroupedEpochBatchIterator。
3.2 模型层:双输入 Transformer 与参数共享
模型注册名为dual_input_s2t_transformer,实现在 models/s2t_dualinputtransformer.py 中,整体是一个"双编码器 + 双解码器"结构:
- 语音编码器:
S2TTransformerEncoder(含 Conv1d 子采样),可选SpeechEoSEncoder包装,在语音特征末尾追加 EOS 特征(--add-speech-eos)以对齐文本侧句边界; - 文本编码器:标准
TransformerEncoder,其嵌入层使用音素字典; - 共享层:通过
--encoder-shared-layers、--encoder-shared-layer-level与--decoder-shared-layer-level控制语音/文本编码器与解码器之间的参数共享程度(0:完全共享;1:共享全部参数但保持独立模型;2:只共享权重、不共享 bias 与 LayerNorm); - 梯度调控:
--enc-grad-mult可对两个编码器输出统一缩放梯度,--text-input-cost-ratio控制文本纯输入样本的损失权重。
模型提供dualinputs2ttransformer_s / _m / _b / _l四档架构,差异集中在嵌入维度与层数(如_s:embed 256、各 7 层;_m:embed 512、语音 10 层 + 文本 6 层 + 解码 6 层)。
3.3 损失层:引导式标签平滑交叉熵 + 在线 KD + CAR
准则注册名为guided_label_smoothed_cross_entropy_with_accuracy,实现在 criterions/text_guide_cross_entropy_acc.py 中,其关键参数包括:
| 参数 | 默认值 | 作用 |
|---|---|---|
--label-smoothing | 0.0 | 标签平滑 ε |
--guide-alpha | 0.0 | 在线 KD 权重 α:loss = α * guide_loss + (1-α) * ce_loss |
--disable-text-guide-update-num | 0 | 前 N 步只用 CE 损失(让语音解码器先站稳再被引导) |
--attentive-cost-regularization | 0.0 | 交叉注意力正则化(CAR)损失权重 β |
--attentive-cost-without-normalize | False | 计算 CAR 时不做归一化 |
在线知识蒸馏:当 batch 同时含语音与文本输入时(is_dual_input),decoder 输出被torch.chunk拆成lprobs_spch(来自语音编码路径)与lprobs_text(来自文本编码路径);文本路径的输出概率probs_teacher(detach 后)作为教师分布,指导语音路径的损失(见guide_loss_and_acc)。
交叉注意力正则化(CAR):在TransformerMultiInputDecoder.cross_attentive_loss中,利用语音与文本编码器在倒数第 N 层的中间状态(encoder_states),计算"语音序列用文本状态重建"与"语音序列用自身状态重建"之间的距离作为正则项,乘以 β 后并入总损失——这正是--attentive-cost-regularization 0.02所启用的机制。
四、示例一:MuST-C 英德(En-De)语音翻译联合训练
对应完整文档见 docs/ende-mustc.md。
4.1 数据准备
第一步:下载基础文件。官方发布了联合训练专用的 SentencePiece 模型spm.model、目标字典dict.txt、数据配置config.yaml以及音素字典src_dict.txt,请从官方 release 地址下载后放入 manifest 根目录($MANIFEST_ROOT)。
第二步:准备 MuST-C 数据集。语音部分的准备流程与 S2T 示例中的 MuST-C 说明完全一致,请遵循该流程生成 tsv manifest。
第三步:源文本音素化。将 tsv 中src_text列的英文源文本转换为音素表示:
python 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 ${must_c_en_de_src_text} \ --out-path ${must_c_en_de_src_text_pho}脚本(scripts/g2p_encode.py)基于g2p_en将英文转为 CMU 风格音素串,各选项含义:
--lower-case:统一小写;--do-filter:把连字符、破折号替换为空格;--use-word-start:每个词前加▁词首标记(与 SentencePiece 风格对齐);--no-punc:剔除标点;--reserve-word:指定保留词表文件,词表内词不参与音素化。示例模块自带的 configs/mustc_noise.list 中定义了一批噪声/语气标记(如(Applause) NOISE、(Laughter) VOICE),这些标注会被保留而非强行转音素;--parallel-process-num:可用 submitit 并行加速。
音素化完成后,用生成的音素串替换 tsv 中src_text列,并将音素字典保存到$MANIFEST_ROOT/src_dict.txt。
第四步:准备 WMT 平行文本数据。下载 WMT14 En-De 数据,按翻译示例的流程处理:英文源侧同样做音素化转换,然后生成二值化的平行数据文件,保存到$parallel_text_data。
4.2 训练
官方基线使用8 张 V100 GPU训练,共 100 个 epoch。
方案 A:从零联合训练(small 架构):
python train.py ${MANIFEST_ROOT} \ --save-dir ${save_dir} \ --num-workers 8 \ --task speech_text_joint_to_text \ --arch dualinputs2ttransformer_s \ --user-dir examples/speech_text_joint_to_text \ --max-epoch 100 --update-mix-data \ --optimizer adam --lr-scheduler inverse_sqrt \ --lr 0.001 --update-freq 4 --clip-norm 10.0 \ --criterion guided_label_smoothed_cross_entropy_with_accuracy \ --label-smoothing 0.1 --max-tokens 10000 --max-tokens-text 10000 \ --max-positions-text 400 --seed 2 --speech-encoder-layers 12 \ --text-encoder-layers 6 --encoder-shared-layers 6 --decoder-layers 6 \ --dropout 0.1 --warmup-updates 20000 \ --text-sample-ratio 0.25 --parallel-text-data ${parallel_text_data} \ --text-input-cost-ratio 0.5 --enc-grad-mult 2.0 --add-speech-eos \ --log-format json --langpairs en-de --noise-token '▁NOISE' \ --mask-text-ratio 0.0 --max-tokens-valid 20000 --ddp-backend no_c10d \ --log-interval 100 --data-buffer-size 50 --config-yaml config.yaml \ --keep-last-epochs 10方案 B:良好初始化 + 交叉注意力正则化 + 在线知识蒸馏(medium 架构)。该方案需先下载预训练模型:pretrain_encoder(多语种 ASR Transformer)与pretrain_nmt(NMT 检查点):
python train.py ${MANIFEST_ROOT} \ --save-dir ${save_dir} \ --num-workers 8 \ --task speech_text_joint_to_text \ --arch dualinputs2ttransformer_m \ --user-dir examples/speech_text_joint_to_text \ --max-epoch 100 --update-mix-data \ --optimizer adam --lr-scheduler inverse_sqrt \ --lr 0.002 --update-freq 4 --clip-norm 10.0 \ --criterion guided_label_smoothed_cross_entropy_with_accuracy \ --guide-alpha 0.8 --disable-text-guide-update-num 5000 \ --label-smoothing 0.1 --max-tokens 10000 --max-tokens-text 10000 \ --max-positions-text 400 --seed 2 --speech-encoder-layers 12 \ --text-encoder-layers 6 --encoder-shared-layers 6 --decoder-layers 6 \ --dropout 0.1 --warmup-updates 20000 --attentive-cost-regularization 0.02 \ --text-sample-ratio 0.25 --parallel-text-data ${parallel_text_data} \ --text-input-cost-ratio 0.5 --enc-grad-mult 2.0 --add-speech-eos \ --log-format json --langpairs en-de --noise-token '▁NOISE' \ --mask-text-ratio 0.0 --max-tokens-valid 20000 --ddp-backend no_c10d \ --log-interval 100 --data-buffer-size 50 --config-yaml config.yaml \ --load-pretrain-speech-encoder ${pretrain_encoder} \ --load-pretrain-decoder ${pretrain_nmt} \ --load-pretrain-text-encoder-last ${pretrain_nmt} \ --keep-last-epochs 10与方案 A 相比,方案 B 的增量体现在:
--guide-alpha 0.8:启用在线 KD,α 取 0.8(即 80% 权重给文本教师分布);--disable-text-guide-update-num 5000:前 5000 步禁用引导、只用 CE;--attentive-cost-regularization 0.02:启用 CAR,权重 0.02;--load-pretrain-speech-encoder/--load-pretrain-decoder/--load-pretrain-text-encoder-last:分别用 ASR 编码器与 NMT 检查点初始化语音编码器、解码器与文本编码器末层。从源码(DualInputEncoder.build_encoder与DualInputS2TTransformerModel.build_decoder)可见,这些参数经checkpoint_utils.load_pretrained_component_from_model按组件加载,且--load-pretrain-text-encoder-last提供了一次用预训练 MT 编码器覆盖共享层的机会。
4.3 评估
使用 fairseq 的生成脚本,以--load-speech-only仅加载语音数据:
python ./fairseq_cli/generate.py \ ${MANIFEST_ROOT} \ --task speech_text_joint_to_text \ --max-tokens 25000 \ --nbest 1 \ --results-path ${infer_results} \ --batch-size 512 \ --path ${model} \ --gen-subset tst-COMMON_st \ --config-yaml config.yaml \ --scoring sacrebleu \ --beam 5 --lenpen 1.0 \ --user-dir examples/speech_text_joint_to_text \ --load-speech-only注意--gen-subset tst-COMMON_st:这是 MuST-C 的语音测试子集(_st后缀标识),--scoring sacrebleu使用 sacreBLEU 评分。
4.4 官方结果(联合训练 + 初始化 + CAR + 在线 KD)
| 方向 | En-De | En-Es | En-Fr |
|---|---|---|---|
| BLEU | 27.4 | 31.2 | 37.6 |
官方同时发布了各方向的最终检查点(checkpoint_ave_10.pt,即最后 10 个 epoch 的平均),可在官方 release 页面获取后直接复现。
五、示例二:IWSLT 2021 多语种语音翻译联合训练
对应完整文档见 docs/iwslt2021.md,其技术方案来自 FST(FAIR Speech Translation system for the IWSLT21 Multilingual Shared Task)。
5.1 数据准备
- 下载官方发布的
spm.model、目标字典tgt_dict.txt与config.yaml; - 语音部分请遵循 speech-to-text 示例中的 mTEDx 数据准备说明,并使用
--use-audio-input选项生成原始音频 tsv 文件; - 源文本列
src_text同样需要音素化,方法与 MuST-C 示例完全一致(即 ende-mustc.md 中的g2p_encode.py流程)。
5.2 训练
该实验涉及 6 个语言(es、fr、it、pt、en),覆盖"语音到文本翻译(X→en)+ 同语言语音转写(es→es、fr→fr、pt→pt、it→it)"等方向。训练前需下载预训练mBART模型与w2v(XLSR-53 56k)模型:
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 5000该命令与 MuST-C 方案的显著差异:
- 架构换为
dualinputxmtransformer_base:实现于 models/s2t_dualinputxmtransformer.py,语音编码器以 w2v/XLSR 为基础、文本编码器与解码器以 mBART 为基础,因此出现--w2v-path、--load-pretrained-mbart-from、--stack-w2v-mbart-encoder(堆叠 w2v 与 mBART 编码器)、--drop-w2v-layers 12(丢弃 w2v 最后 12 层)、--skip-encoder-projection、--finetune-w2v-params all/--finetune-mbart-encoder-params all/--finetune-mbart-decoder-params all等微调控制参数; - 多语言子集:
--train-subset/--valid-subset显式列出 8 个训练子集与 15 个验证子集,覆盖 es/fr/it/pt 与 en 之间的多种方向; - 更长输入:
--max-source-positions 800000配合--max-tokens 500000、--max-sentences 3,适配原始音频长序列;学习率降至5e-05,标签平滑加大到0.3。
5.3 评估
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注意--infer-target-lang en:多语种解码时需要指定目标语言标记。从 tasks/speech_text_joint.py 源码可见,该参数会在setup_task中把<lang:en>语言标签映射为 decoder 的起始 token(bos_token),从而在inference_step中作为生成起点。
5.4 官方结果
| 方向 | es_en | fr_en | pt_en | it_en | fr_es | pt_es | it_es | es_es | fr_fr | pt_pt | it_it |
|---|---|---|---|---|---|---|---|---|---|---|---|
| BLEU | 31.62 | 36.93 | 35.07 | 27.12 | 38.87 | 35.57 | 34.13 | 74.59 | 74.64 | 70.84 | 69.76 |
同语言转写方向(如 es_es、fr_fr、pt_pt、it_it)BLEU 明显更高,符合"语音转写"任务本身比跨语言翻译更易的直觉;官方训练的模型检查点(checkpoint17.pt)可在 release 页面下载复现。
六、实践经验小结
- 文本数据是语音翻译的"免费午餐":通过
--text-sample-ratio 0.25控制文本样本占比、--text-input-cost-ratio 0.5控制其损失权重,可在不显著增加语音数据开销的前提下引入大量平行文本。 - 音素对齐是关键预处理:
g2p_encode.py将英文源文本转为音素表示并用▁标记词首,配合--add-speech-eos在语音侧补 EOS 特征,使两条输入模态在序列语义上更接近,这是共享编码器能有效工作的前提。 - 增强技巧按需叠加:从 scratch 训练(方案 A)→ 预训练初始化 → 在线 KD(
--guide-alpha 0.8+--disable-text-guide-update-num 5000)→ CAR(--attentive-cost-regularization 0.02),每一步都能带来稳定的翻译质量提升,官方 En-De 达到 27.4 BLEU。 - 多语种场景优先考虑预训练底座:IWSLT 2021 实验直接复用 w2v + mBART,配合全参数微调与"堆叠 + 丢弃部分 w2v 层"的策略,使多语种联合训练在数据量有限时依然取得有竞争力的结果。
参考文献
本文涉及的论文与工具引用如下(完整 BibTeX 见 speech_text_joint_to_text/README.md):
- Tang, Pino, Wang, Ma, Genzel.A General Multi-Task Learning Framework to Leverage Text Data for Speech to Text Tasks.ICASSP 2021.
- Tang, Pino, Li, Wang, Genzel.Improving Speech Translation by Understanding and Learning from the Auxiliary Text Translation Task.ACL 2021.
- Tang, Gong, Li, Wang, Pino, Schwenk, Goyal.FST: the FAIR Speech Translation System for the IWSLT21 Multilingual Shared Task.IWSLT 2021.
- Wang, Tang, Ma, Wu, Okhonko, Pino.fairseq S2T: Fast Speech-to-Text Modeling with fairseq.AACL 2020.
- Ott, Edunov, Baevski, Fan, Gross, Ng, Grangier, Auli.fairseq: A Fast, Extensible Toolkit for Sequence Modeling.NAACL-HLT 2019.
【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考