SimulST on MuST-C:基于 EdgeLM (fairseq) 的 wait-k 端到端同时语音翻译训练与评测实战
【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm
本文以仓库中edgelm/examples/speech_to_text/docs/simulst_mustc_example.md教程为主体,系统讲解如何在 MuST-C 英德数据集上完成"数据准备 → 离线 ASR 预训练 → wait-k / 单调多头注意力的同时语音翻译(SimulST)训练 → 基于 SimulEval 的延迟与质量评测"的完整流水线,并结合 convtransformer_simul_trans.py、fixed_pre_decision.py、label_smoothed_cross_entropy_latency_augmented.py 等源码,深入解释--simul-type、--fixed-pre-decision-ratio等关键参数的底层实现,以及 READ/WRITE 决策、BLEU 与 AL/DAL/AP 延迟指标的计算原理。
1. 任务背景:从 SimulMT 到 SimulST
同时语音翻译(Simultaneous Speech Translation, SimulST)要求模型在音频流式到达的过程中边听边译,而不是听完整句再翻译。本教程实现的方法来自论文 "SimulMT to SimulST: Adapting Simultaneous Text Translation to End-to-End Simultaneous Speech Translation"(AACL 2020),核心思路是:先把文本领域的同时翻译策略(wait-k、单调多头注意力)适配到语音上,语音侧采用 ConvTransformer 编码器 + 带"读出策略"的单调解码器。
数据集选用MuST-C(多语言语音到文本翻译语料库,基于英文 TED 演讲,含 8 种语言译文)。从源码看,prep_mustc_data.py 中的MUSTC类明确定义了支持的划分与语言:
SPLITS = ["train", "dev", "tst-COMMON", "tst-HE"] LANGUAGES = ["de", "es", "fr", "it", "nl", "pt", "ro", "ru"]__init__会读取${MUSTC_ROOT}/en-{lang}/data/{split}/下的txt/{split}.yaml(音频分段信息:wav 文件名、offset、duration、speaker_id)以及txt/{split}.en、txt/{split}.{lang}两个逐行对齐的文本文件,按wav → 分段分组后用 soundfile 切片,得到(waveform, sample_rate, src_utt, tgt_utt, spk_id, utt_id)形式的样本。
整体训练-评测链路为:
- 下载 MuST-C 数据并按
en-{lang}目录组织; - 运行
prep_mustc_data.py分别生成 ASR 与 ST 两套 manifest、特征、词表与配置; - 训练一个离线 ASR 模型(
convtransformer_espnet),得到可复用的 encoder 权重; - 用
--load-pretrained-encoder-from加载该权重,训练同时翻译模型; - 用 SimulEval 框架 + 仓库自带的 fairseq_simul_st_agent.py 在线评测 BLEU 与延迟指标。
2. 数据准备(Data Preparation)
安装额外依赖后,在fairseq(本仓库中对应 edgelm 目录)下对 ASR 与 ST 两个任务各跑一次数据准备脚本:
# Additional Python packages for S2T data processing/model training pip install pandas torchaudio sentencepiece # Generate TSV manifests, features, vocabulary, # global cepstral and mean estimation, # and configuration for each language cd fairseq python examples/speech_to_text/prep_mustc_data.py \ --data-root ${MUSTC_ROOT} --task asr \ --vocab-type unigram --vocab-size 10000 \ --cmvn-type global python examples/speech_to_text/prep_mustc_data.py \ --data-root ${MUSTC_ROOT} --task st \ --vocab-type unigram --vocab-size 10000 \ --cmvn-type globalMuST-C 原始数据需从官方站点下载,并解压到${MUSTC_ROOT}/en-{target_lang}(例如${MUSTC_ROOT}/en-de)。
2.1 脚本参数详解
结合 prep_mustc_data.py 的argparse定义,各参数含义与默认值如下:
| 参数 | 默认值 | 说明 |
|---|---|---|
--data-root/-d | 必填 | MuST-C 根目录,其下应为en-{lang}/data/{split}结构 |
--task | 无 | asr或st,决定 manifest 中tgt_text取源句英文还是目标语言译文 |
--vocab-type | unigram(必填) | 词表类型,可选bpe/unigram/char |
--vocab-size | 8000 | sentencepiece 词表大小,教程中设为 10000 |
--cmvn-type | utterance | 倒谱均值方差归一化类型,可选global/utterance;教程使用global |
--gcmvn-max-num | 150000 | 估计全局 CMVN 统计量时最多使用的句子数 |
--joint | 关闭 | 8 语言联合训练模式(本文单语向英德翻译不需要) |
--use-audio-input | 关闭 | 用原始波形(flac)代替 fbank 特征 |
2.2 脚本内部流程(源码印证)
process(args)对每个存在en-{lang}目录的语言依次执行以下操作:
- 特征提取:遍历
MUSTCDataset,逐句调用extract_fbank_features生成fbank80/{utt_id}.npy(80 维 log-mel fbank,来自 data_utils.py); - 全局 CMVN 估计:当
split == 'train'且--cmvn-type global时,缓存训练集特征,调用cal_gcmvn_stats估计全局均值/标准差并保存为gcmvn.npz。这也是配置文件里global_cmvn.stats_npz_path的来源; - 打包与 manifest:把特征目录压成
fbank80.zip,读取 zip manifest 后生成 5 列的 TSV(id, audio, n_frames, tgt_text, speaker)。注意tgt_text的取值逻辑:src_utt if args.task == "asr" else tgt_utt——即 ASR 任务的"标签"是英文转写,ST 任务的"标签"是目标语言译文; - 词表生成:把训练集
tgt_text写入临时文件,按spm_{vocab_type}{size}_{task}命名生成 sentencepiece 模型与词典(如spm_unigram10000_st.model/.txt)。ST 词表会额外注入<lang:{lang}>特殊符号(见process_joint中special_symbols的构造,用于多语向场景); - 配置生成:调用
gen_config_yaml生成config_asr.yaml/config_st.yaml,其中包含 sentencepiece 路径、gcmvn 路径、specaugment 策略(fbank 模式为lb策略)等。
处理完成后,${MUSTC_ROOT}/en-de/目录下将得到类似产物:fbank80.zip、train_asr.tsv、dev_asr.tsv、train_st.tsv、dev_st.tsv、spm_unigram10000_asr.model、spm_unigram10000_st.model、gcmvn.npz、config_asr.yaml、config_st.yaml等。
3. ASR 预训练
同时语音翻译需要先有一个预训练好的离线 ASR 模型,其 encoder 将在 ST 阶段被整体复用(--load-pretrained-encoder-from)。假设保存目录为${ASR_SAVE_DIR},教程中的命令默认在 1 张 GPU 上训练(若用 8 卡可去掉--update-freq 8):
fairseq-train ${MUSTC_ROOT}/en-de \ --config-yaml config_asr.yaml --train-subset train_asr --valid-subset dev_asr \ --save-dir ${ASR_SAVE_DIR} --num-workers 4 --max-tokens 40000 --max-update 100000 \ --task speech_to_text --criterion label_smoothed_cross_entropy --report-accuracy \ --arch convtransformer_espnet --optimizer adam --lr 0.0005 --lr-scheduler inverse_sqrt \ --warmup-updates 10000 --clip-norm 10.0 --seed 1 --update-freq 8关键参数说明:
--config-yaml config_asr.yaml:这里传的是相对于 data 目录的 basename,即第 2 步生成的config_asr.yaml。数据目录为${MUSTC_ROOT}/en-de,其中已包含train_asr.tsv/dev_asr.tsv,因此--train-subset/--valid-subset只需给出 TSV 前缀;--arch convtransformer_espnet:ESPnet 风格的 ConvTransformer 语音编码器架构,是后续 SimulST 模型的编码器基础;--criterion label_smoothed_cross_entropy --report-accuracy:标准标签平滑交叉熵损失,并额外报告 top-1 准确率;--lr 0.0005、inverse_sqrt调度、--warmup-updates 10000、--max-update 100000、梯度裁剪--clip-norm 10.0;--update-freq 8:累积 8 个 batch 再更新一次,等效于放大 batch size,适配单卡显存。
教程说明可以从论文作者的公共发布地址下载预训练好的 ASR checkpoint(must_c_v1_en_de_pretrained_asr),跳过本步骤;下载后其checkpoint_best.pt即为 ST 阶段--load-pretrained-encoder-from的输入。
4. 同时语音翻译训练
ST 模型架构注册在 convtransformer_simul_trans.py:SimulConvTransformerModel继承自ConvTransformerModel(复用 ASR 的 ConvTransformer 编码器),把解码器替换为TransformerMonotonicDecoder(定义在 transformer_monotonic_attention.py)。该解码器的每一层 encoder attention 是一个单调注意力头,前向传播时会同时输出一个action:0 = READ(还没读够,需要更多输入),1 = WRITE(可以吐出下一个词元)。架构convtransformer_simul_trans_espnet就是在convtransformer_espnet(args)基础上注册的别名。
解码策略(READ/WRITE 的判定规则)由--simul-type选择,对应fixed_pre_decision.py中通过register_monotonic_attention注册的三类单调注意力:
--simul-type取值 | 底层类 | 策略含义 |
|---|---|---|
waitk_fixed_pre_decision | WaitKAttention+ fixed pre-decision | wait-k 策略(--waitk-lagging控制 lagging) |
hard_aligned_fixed_pre_decision | MonotonicAttention+ fixed pre-decision | 硬对齐单调注意力 |
infinite_lookback_fixed_pre_decision | MonotonicInfiniteLookbackAttention+ fixed pre-decision | 可无限回看的单调多头注意力(MMA) |
"fixed pre-decision" 指在固定分块的边界上做出 READ/WRITE 决策:FixedStrideMonotonicAttention会用pre_decision_ratio对 key 序列做池化(默认average,即AvgPool1d(kernel_size=ratio, stride=ratio, ceil_mode=True);last则取每块最后一个位置),在池化后的粗粒度序列上算p_choose,再用insert_zeros上采样回原始分辨率(块内其余位置概率置零),从而保证"每ratio个编码器步才产生一次读出决策"。相关超参数在 fixed_pre_decision.py 中注册:
--fixed-pre-decision-ratio(必填):多少个编码器状态步触发一次同时决策;源码断言ratio > 1;--fixed-pre-decision-type:average(默认)或last池化;--fixed-pre-decision-pad-threshold:默认0.3,池化块中 pad 占比超过该阈值则整块视为 pad。
另外,infinite_lookback变体在推理时会把池化长度向下取整(math.floor),避免"提前看到最后一块不完整分块"的偏差——源码注释中明确写到 "The floor instead of ceil is used for inference"。
4.1 Wait-k + 固定预决策
以固定预决策比例 7(每 7 个编码器状态做一次 READ/WRITE 决策)+ wait-3 策略为例,假设 ST 模型保存目录为${ST_SAVE_DIR}:
fairseq-train ${MUSTC_ROOT}/en-de \ --config-yaml config_st.yaml --train-subset train_st --valid-subset dev_st \ --save-dir ${ST_SAVE_DIR} --num-workers 8 \ --optimizer adam --lr 0.0001 --lr-scheduler inverse_sqrt --clip-norm 10.0 \ --criterion label_smoothed_cross_entropy \ --warmup-updates 4000 --max-update 100000 --max-tokens 40000 --seed 2 \ --load-pretrained-encoder-from ${ASR_SAVE_DIR}/checkpoint_best.pt \ --task speech_to_text \ --arch convtransformer_simul_trans_espnet \ --simul-type waitk_fixed_pre_decision \ --waitk-lagging 3 \ --fixed-pre-decision-ratio 7 \ --update-freq 8与 ASR 阶段相比的要点:
--load-pretrained-encoder-from ${ASR_SAVE_DIR}/checkpoint_best.pt:整体加载离线 ASR 的编码器权重,这是 SimulST 方法"从 ASR 迁移到 ST"的关键;--simul-type waitk_fixed_pre_decision:wait-k 策略在 fixed pre-decision 的池化序列上执行,--waitk-lagging 3表示 wait-3(读到第 k 个输入后先产出滞后 3 步的词元);--lr 0.0001(比 ASR 阶段小一个量级)、--warmup-updates 4000、--seed 2;- 该阶段使用普通
label_smoothed_cross_entropy损失,不显式惩罚延迟。
4.2 单调多头注意力(MMA)+ 固定预决策
第二种策略使用可无限回看的单调多头注意力,并切换为延迟增广损失:
fairseq-train ${MUSTC_ROOT}/en-de \ --config-yaml config_st.yaml --train-subset train_st --valid-subset dev_st \ --save-dir ${ST_SAVE_DIR} --num-workers 8 \ --optimizer adam --lr 0.0001 --lr-scheduler inverse_sqrt --clip-norm 10.0 \ --warmup-updates 4000 --max-update 100000 --max-tokens 40000 --seed 2 \ --load-pretrained-encoder-from ${ASR_SAVE_DIR}/${CHECKPOINT_FILENAME} \ --task speech_to_text \ --criterion latency_augmented_label_smoothed_cross_entropy \ --latency-weight-avg 0.1 \ --arch convtransformer_simul_trans_espnet \ --simul-type infinite_lookback_fixed_pre_decision \ --fixed-pre-decision-ratio 7 \ --update-freq 8其损失函数latency_augmented_label_smoothed_cross_entropy实现于 label_smoothed_cross_entropy_latency_augmented.py,值得注意的实现细节:
- 延迟可微化:
compute_latency_loss从每层单调注意力头取出软对齐分布alpha(net_output[1].attn_list),把源端位置索引steps = arange(1, 1+src_len)与alpha加权求和得到每个 (batch, 层×头, 目标步) 的expected_delays,再用 SimulEval 的LATENCY_METRICS(average_lagging/average_proportion/differentiable_average_lagging,默认类型为differentiable_average_lagging)计算期望延迟; - 多头聚合:
--latency-gather-method支持average/weighted_average(对多头延迟做 softmax 加权,默认)/max三种聚合方式; - 权重与门控:总延迟损失为
avg_loss + var_loss,其中avg_loss = latency_avg_weight * expected_latency(即教程里的--latency-weight-avg 0.1);还支持latency_update_after在训练前若干步内关闭延迟项; - 硬依赖:该 criterion 在
__init__中断言LATENCY_METRICS is not None,即必须先安装 SimulEval(pip install simuleval),否则无法训练 MMA 变体; - 训练日志会额外输出
latency、delays_var、latency_loss三个标量(reduce_metrics中按句子数平均),可在训练过程中直接监控延迟下降趋势。
MMA 变体不指定--waitk-lagging,而是由注意力本身学习"何时该读、何时该写",配合延迟损失显式压低平均滞后。
5. 推理与评测(SimulEval)
评测框架使用 SimulEval。注意本仓库自带了评测 agent:fairseq_simul_st_agent.py(文档命令中写作${FAIRSEQ}/examples/speech_to_text/simultaneous_translation/agents/fairseq_simul_st_agent.py):
git clone https://github.com/facebookresearch/SimulEval.git cd SimulEval pip install -e . simuleval \ --agent ${FAIRSEQ}/examples/speech_to_text/simultaneous_translation/agents/fairseq_simul_st_agent.py --source ${SRC_LIST_OF_AUDIO} --target ${TGT_FILE} --data-bin ${MUSTC_ROOT}/en-de \ --config config_st.yaml \ --model-path ${ST_SAVE_DIR}/${CHECKPOINT_FILENAME} \ --output ${OUTPUT} \ --scores5.1 输入文件格式
${SRC_LIST_OF_AUDIO}:每行一个 wav 文件绝对路径的列表,例如:
/home/user/data/audio-1.wav /home/user/data/audio-2.wav${TGT_FILE}:每行一条对应音频的参考译文:
Translation_1 Translation_2若评测集就是 MuST-C 官方切分,无需手工准备上述文件——仓库提供 seg_mustc_data.py 直接从原始 MuST-C 切出逐段 wav 与文本:
python ${FAIRSEQ}/examples/speech_to_text/seg_mustc_data.py \ --data-root ${MUSTC_ROOT} --lang de \ --split ${SPLIT} --task st \ --output ${EVAL_DATA}该脚本参数(--data-root/--task(asr|st)/--lang/--output/--split,--split取值即MUSTC.SPLITS:dev、tst-COMMON、tst-HE、train)会复用MUSTCDataset,输出到${EVAL_DATA}下:逐段音频{utt_id}.wav、参考文本{split}.{lang}、音频路径清单{split}.wav_list(其中文本取tgt列,即--task st时为目标语言译文)。
5.2 配置与数据目录的对应关系
- 如果数据是自己从原始 MuST-C 准备的,
--data-bin与--config必须与训练章节保持一致; - 若只做评测,可使用官方发布目录
must_c_v1.0_en_de_databin.tgz,其中包含:spm_unigram10000_st.model:sentencepiece 模型;spm_unigram10000_st.txt:对应词典;gcmvn.npz:全局倒谱均值/方差统计;config_st.yaml:配置样例(见下)。若使用下载的数据目录,需要把sentencepiece_model与stats_npz_path改为本机绝对路径:
bpe_tokenizer: bpe: sentencepiece sentencepiece_model: ABS_PATH_TO_SENTENCEPIECE_MODEL global_cmvn: stats_npz_path: ABS_PATH_TO_GCMVN_FILE input_channels: 1 input_feat_per_channel: 80 sampling_alpha: 1.0 specaugment: freq_mask_F: 27 freq_mask_N: 1 time_mask_N: 1 time_mask_T: 100 time_mask_p: 1.0 time_wrap_W: 0 transforms: '*': - global_cmvn _train: - global_cmvn - specaugment vocab_filename: spm_unigram10000_st.txt注意一个容易踩坑的细节:一旦设置了--data-bin,--config传的是 config yaml 的basename而非完整路径(agent 内部以os.path.join(args.data_bin, args.config)打开该文件并从中读取global_cmvn.stats_npz_path)。
5.3 Agent 在线推理机制(源码印证)
fairseq_simul_st_agent.py 中的FairseqSimulSTAgent实现了 SimulEval 的SpeechAgent协议,核心机制:
- 在线特征提取:
OnlineFeatureExtractor按默认 25 ms 窗长 / 10 ms 移步(--window-size 25/--shift-size 10,16 kHz 采样率)用kaldi.fbank流式生成 80 维 fbank,并应用与训练一致的全局 CMVN 变换(np.subtract/np.divide); - 决策步长对齐:
speech_segment_size默认 40 ms(4 倍池化比 × 10 ms 移步),若解码器 attention 层带有pre_decision_ratio,则speech_segment_size *= pre_decision_ratio——即每读入40ms × fixed_pre_decision_ratio的音频才调用一次policy(),与训练时的"每 ratio 个编码器步做一次决策"严格对齐; - READ/WRITE 循环:
policy()中把当前已读特征喂给 encoder(update_model_encoder增量更新encoder_states),随后执行 decoder 一步前向,从返回的outputs.action得到READ_ACTION(继续要音频)或WRITE_ACTION(predict()取 argmax 词元输出); - 子词→词的去分词:
units_to_segment用 sentencepiece 的\u2581(BOW 前缀)判断词边界,把子词拼成完整词再发给 SimulEval 服务端,并受--max-len(默认 200)截断与--force-finish(源音频未读完时是否强制结束)控制; - 模型加载:
load_model_vocab通过 checkpoint 中的cfg重建 task 与模型(load_pretrained_encoder_from置空以保证strict=True加载),并设置torch.set_grad_enabled(False)。
agent 暴露的常用参数包括:--model-path(必填)、--data-bin(必填)、--config、--global-stats、--tgt-splitter-type/--tgt-splitter-path、--max-len、--force-finish以及特征窗口参数。
5.4 参考结果与指标解读
官方发布的convtransformer_wait5_pre7checkpoint(wait-5、预决策 280 ms 的模型)在tst-COMMON上的评测结果为:
{ "Quality": { "BLEU": 13.94974229366959 }, "Latency": { "AL": 1751.8031870037803, "AL_CA": 2338.5911762796536, "AP": 0.7931395378788959, "AP_CA": 0.9405103863210942, "DAL": 1987.7811616943081, "DAL_CA": 2425.2751560926167 } }指标含义(均基于去分词文本计算):
- Quality / BLEU:去分词 BLEU,因此必须保证发给 SimulEval 服务端的是完整词而非子词——这正是 agent 中
units_to_segment的职责; - AL(Average Lagging):平均滞后时间(ms);
AL_CA为 Context-Adaptive 变体,把首词输出前的"上下文填充时间"剔除; - AP(Average Proportion):平均产出比例,衡量输出进度对输入进度的跟随程度;
AP_CA为其上下文自适应版本; - DAL(Differentiable Average Lagging):可微平均滞后,即训练阶段
latency_augmented_label_smoothed_cross_entropy中用于反传的那个延迟量的在线实测值,可与训练日志中的latency相互印证。
带上--output ${OUTPUT}后,详细日志(逐句 READ/WRITE 事件、词元与延迟明细)和分数会一并保存到${OUTPUT}目录,便于复盘策略行为。
6. 完整流程小结与注意事项
| 阶段 | 关键命令/文件 | 产物 |
|---|---|---|
| 数据准备 | prep_mustc_data.py(--task asr/--task st) | *.zip、{split}_{task}.tsv、spm_*.model/.txt、gcmvn.npz、config_*.yaml |
| ASR 预训练 | fairseq-train+--arch convtransformer_espnet | ${ASR_SAVE_DIR}/checkpoint_best.pt |
| SimulST 训练 | fairseq-train+--arch convtransformer_simul_trans_espnet+--simul-type {waitk,infinite_lookback}_fixed_pre_decision | ${ST_SAVE_DIR}/checkpoint_*.pt |
| 评测数据切分 | seg_mustc_data.py | ${EVAL_DATA}/{split}.wav_list、{split}.{lang} |
| 在线评测 | simuleval+ fairseq_simul_st_agent.py | BLEU、AL/AP/DAL(含_CA) |
实操中需要特别注意的前置条件:
- 训练/评测命令均以"在
fairseq(EdgeLM)仓库根目录下运行"为前提,examples/speech_to_text/prep_mustc_data.py等路径都是相对该根目录; - MMA 变体的训练与延迟评测依赖 SimulEval(criterion 初始化时即断言其存在),需先
pip install simuleval; --data-bin与--config必须同源于同一份数据准备结果,且--config传 basename;- 评测 BLEU 基于去分词文本,若自行修改 agent,务必保留 BOW 前缀的词元拼接逻辑;
- 官方预训练 ASR checkpoint 与
must_c_v1_en_deST 模型(wait-5、280 ms 预决策)均提供下载,可跳过训练直接验证评测链路是否跑通,再回头做自定义--waitk-lagging/--fixed-pre-decision-ratio的实验。
掌握以上内容后,你可以完整复现 MuST-C 英德方向 SimulST 的 wait-k 与 MMA 两类基线,通过--waitk-lagging、--fixed-pre-decision-ratio、--latency-weight-avg等超参搜索"延迟-质量"折中,并基于 SimulEval 日志对 READ/WRITE 决策行为做逐句分析。
【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考