DrNMT 判别式重排序训练与推理实战指南:基于 unilm 仓库 fairseq 示例的端到端流程
DrNMT 判别式重排序训练与推理实战指南基于 unilm 仓库 fairseq 示例的端到端流程【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm导读本文围绕 unilm 仓库中edgelm/examples/discriminative_reranking_nmt目录下的 DrNMTDiscriminative Reranking for Neural Machine Translation示例展开系统讲解如何为一个已训练好的神经机器翻译NMT基础模型训练一个判别式重排序器并利用它从 beam search 生成的候选译文中选出更优结果从而在 BLEU / TER 指标上获得提升。读完本文你将掌握三类原始数据文件的组织格式、基于 XLMR 的句子级特征抽取与 BPE 预处理流程、DrNMT 的 fairseq 数据打包与训练配置含 KL 散度判别式训练目标、以及在 valid / test 集上进行权重调优与最终重排序评分的完整命令链。一、方法背景为什么需要判别式重排序神经机器翻译模型在推理阶段通常使用 beam search 一次性生成译文最终输出往往是束内得分最高的一条。然而束搜索得分最高的候选并不总是与人工翻译最接近。DrNMT 的思路是再训练一个独立的判别式模型对每个源句对应的 N 条候选译文hypotheses分别打分并将该得分与基础 MT 模型的得分加权融合从而挑选出最终译文。该示例目录位于 edgelm/examples/discriminative_reranking_nmt包含以下关键组件文件/目录作用README.md完整使用说明本文主体依据scripts/prep_data.py将原始文本 假设句转换为 BPE 与指标标签config/deen.yaml论文中 De-En 实验的 Hydra 训练配置tasks/discriminative_reranking_task.pyfairseq Task数据加载、beam 分组、验证期指标计算models/discriminative_reranking_model.pyBertRanker 模型基于 XLMR 编码器 分类头criterions/discriminative_reranking_criterion.pyKL 散度判别式训练损失drnmt_rerank.py推理阶段的重排序与评分脚本整体流程可概括为四步准备基础 MT 模型的译文候选 → 用 XLMR 生成句级表示并计算指标标签 → 训练判别式重排序器 → 在 valid 集上调权重、在 test 集上应用。二、数据准备构建训练重排序器所需的三类文件2.1 前提先构建基础 MT 模型重排序器需要基础 MT 模型生成候选译文因此首先需按 edgelm/examples/translation 下的说明训练一个基础 MT 模型。该目录提供了 IWSLT14 De-Enprepare-iwslt14.sh、WMT14 En-Deprepare-wmt14en2de.sh、WMT14 En-Frprepare-wmt14en2fr.sh以及多语言prepare-iwslt17-multilingual.sh等完整的数据准备与训练流程训练完成后即可用fairseq-generate/fairseq-interactive产出候选译文。2.2 三类原始文件格式对每个数据切分train / valid / test需要准备三个纯文本文件每一行是一条原始句子不经过 sentencepiece 等任何切分处理源句文件source共 L 行参考译文文件target / ground truth共 L 行与源句一一对应候选译文文件hypotheses共L × N行按每个源句的 N 条候选连续排列的顺序组织。文件内容示意如下_N_为每个源句的候选数量# 源句文件共 L 行 source_sentence_1 source_sentence_2 source_sentence_3 ... source_sentence_L # 参考译文文件共 L 行 target_sentence_1 target_sentence_2 target_sentence_3 ... target_sentence_L # 候选译文文件共 L*N 行 source_sentence_1_hypo_1 source_sentence_1_hypo_2 ... source_sentence_1_hypo_N source_sentence_2_hypo_1 ... source_sentence_2_hypo_N ... source_sentence_L_hypo_1 ... source_sentence_L_hypo_N论文中每个源句使用N50条候选。2.3 下载 XLMR 预训练模型DrNMT 的编码器基于 XLMR需要先下载其 base 版本wget https://dl.fbaipublicfiles.com/fairseq/models/xlmr.base.tar.gz tar zxvf xlmr.base.tar.gz # 解压后的文件夹应包含 dict.txt、model.pt 和 sentencepiece.bpe.model 三个文件这三个文件分别用于fairseq 数据二值化时的词典dict.txt、重排序模型的参数初始化model.pt、文本切分sentencepiece.bpe.model。2.4 用 prep_data.py 生成 BPE 数据与指标标签针对每个数据切分train、valid、test 等分别运行 scripts/prep_data.py其核心逻辑是用 sentencepiece 把源句与每条候选译文编码为 piece 序列process函数中的sp.EncodeAsPieces用sacrebleu计算每条候选与参考译文的 BLEUget_bleu或 TERget_ter的详细统计量作为训练标签按分片把结果写入split*目录。运行前需要设置四个关键变量N每个源句的候选数论文用 50SPLIT切分名即 train、valid、test若同一切分有多个数据集则用split_name、split_name1、split_name2依次命名如train、train1、valid、valid1NUM_SHARDS分片数非 train 切分必须设为 1METRIC重排序器要优化的指标支持bleu或ter。# 对每个数据切分train、valid、test 等执行 SOURCE_FILE/path/to/source_sentence_file TARGET_FILE/path/to/target_sentence_file HYPO_FILE/path/to/hypo_file XLMR_DIR/path/to/xlmr OUTPUT_DIR/path/to/output python scripts/prep_data.py \ --input-source ${SOURCE_FILE} \ --input-target ${TARGET_FILE} \ --input-hypo ${HYPO_FILE} \ --output-dir ${OUTPUT_DIR} \ --split $SPLIT \ --beam $N \ --sentencepiece-model ${XLMR_DIR}/sentencepiece.bpe.model \ --metric $METRIC \ --num-shards ${NUM_SHARDS}脚本会在${OUTPUT_DIR}/$METRIC下生成NUM_SHARDS个分片在split*/input_src、split*/input_tgt和split*/$METRIC下分别得到$SPLIT.bpe源句 BPE、候选译文 BPE与$SPLIT.$METRIC指标标签文件。从 prep_data.py 的参数定义可以看出脚本的完整可用选项与默认值参数必选默认值说明--input-source/--input-target/--input-hypo是—三类原始文本文件路径--output-dir是—输出目录--split是—切分名须以train/valid/test开头含数字后缀--beam是—每个源句的候选数等价于论文中的 N--sentencepiece-model是—XLMR 的 sentencepiece 模型--metric否bleu可选bleu或ter--num-shards否1train 切分可 1valid/test 必须为 1--n-proc否8多进程预处理并行度脚本内置了几处关键的输入校验见 prep_data.py--num-shards对 valid/test 强制为 1源句与参考译文行数必须一致候选总数必须能被--beam整除且恰好等于L × beam否则直接报错避免静默产生错位数据。2.5 二值化为 fairseq 格式用fairseq-preprocess把 BPE 文本转换为 fairseq 的 indexed dataset。第一个分片split1需要同时处理 train 与 valid 并生成词典其余分片只处理 train并通过软链接复用 split1 的 valid 文件# 若存在多个 train/valid 集合用逗号分隔 for suffix in src tgt ; do fairseq-preprocess --only-source \ --trainpref ${OUTPUT_DIR}/$METRIC/split1/input_${suffix}/train.bpe \ --validpref ${OUTPUT_DIR}/$METRIC/split1/input_${suffix}/valid.bpe \ --destdir ${OUTPUT_DIR}/$METRIC/split1/input_${suffix} \ --workers 60 \ --srcdict ${XLMR_DIR}/dict.txt done for i in seq 2 ${NUM_SHARDS}; do for suffix in src tgt ; do fairseq-preprocess --only-source \ --trainpref ${OUTPUT_DIR}/$METRIC/split${i}/input_${suffix}/train.bpe \ --destdir ${OUTPUT_DIR}/$METRIC/split${i}/input_${suffix} \ --workers 60 \ --srcdict ${XLMR_DIR}/dict.txt ln -s ${OUTPUT_DIR}/$METRIC/split1/input_${suffix}/valid* ${OUTPUT_DIR}/$METRIC/split${i}/input_${suffix}/. done ln -s ${OUTPUT_DIR}/$METRIC/split1/$METRIC/valid* ${OUTPUT_DIR}/$METRIC/split${i}/$METRIC/. done要点说明--only-source因为重排序器的输入是源句 候选译文拼接序列标签是指标得分而非文本不需要常规 MT 的 src/tgt 双通道处理--srcdict ${XLMR_DIR}/dict.txt直接复用 XLMR 词典保证 token 空间与预训练模型一致--workers 60加大预处理并行度可显著加速二值化软链接让所有分片共享同一套 valid 数据避免重复存储。三、训练判别式重排序器3.1 启动训练训练通过 Hydra 配置驱动示例目录中config/下的 deen.yaml 即论文 De-En 实验所用配置使用 16 张 GPU 与 50 条候选EXP_DIR/path/to/exp fairseq-hydra-train -m \ --config-dir config/ --config-name deen \ task.data${OUTPUT_DIR}/$METRIC/split1/ \ task.num_data_splits${NUM_SHARDS} \ model.pretrained_model${XLMR_DIR}/model.pt \ common.user_dir${FAIRSEQ_ROOT}/examples/discriminative_reranking_nmt \ checkpoint.save_dir${EXP_DIR}其中${FAIRSEQ_ROOT}在本仓库中对应根目录下的 edgelm即该示例实际位于 edgelm/examples/discriminative_reranking_nmt。硬件与超参适配指南原文档明确给出若 GPU 数量少于 16设置distributed_training.distributed_world_sizek并令optimization.update_freq[x]其中x 16/k通过梯度累积等效模拟 16 卡若候选数少于 50设置task.mt_beamN dataset.batch_sizeN dataset.required_batch_size_multipleN保证 batch 维始终按 beam 数对齐。3.2 deen.yaml 配置逐项解读结合 deen.yaml 与对应源码各配置段含义如下common / checkpointcommon: fp16: true log_format: json log_interval: 50 seed: 2 checkpoint: no_epoch_checkpoints: true best_checkpoint_metric: bleu maximize_best_checkpoint_metric: true开启混合精度训练fp16、JSON 日志模型选择以验证集 BLEU 为准best_checkpoint_metric: bleu指标越大越好maximize_best_checkpoint_metric: true。tasktask: _name: discriminative_reranking_nmt data: ??? num_data_splits: ??? include_src: true mt_beam: 50 eval_target_metric: true target_metric: bleudata与num_data_splits在命令行以???形式强制注入。include_src: true表示模型输入为源句 候选拼接mt_beam: 50对应候选数 Neval_target_metric: true使验证阶段实时计算目标指标。dataset / optimization / optimizer / lr_schedulerdataset: batch_size: 50 num_workers: 6 required_batch_size_multiple: 50 valid_subset: ??? optimization: max_epoch: 200 lr: [0.00005] update_freq: [32] optimizer: _name: adam adam_betas: (0.9,0.98) adam_eps: 1e-06 lr_scheduler: _name: polynomial_decay warmup_updates: 8000 total_num_update: 320000batch_size与required_batch_size_multiple均设为 50与mt_beam对齐训练 loss 的求和组织要求样本数能被 beam 整除见下文 criterion 源码学习率 5e-5Adam 优化器betas(0.9, 0.98), eps1e-6多项式衰减 8000 步 warmup、共 32 万次更新。criterioncriterion: _name: kl_divergence_rereanking target_dist_norm: minmax temperature: 0.5这是 DrNMT 的核心训练目标。其实现位于 criterions/discriminative_reranking_criterion.py若target_dist_norm minmax对每个源句的 beam 组内指标标签做 min-max 归一化(target - min) / (max - min eps)把标签映射到 [0, 1]用温度temperature对归一化后的标签做 softmax得到目标分布真实指标得分越高目标概率越大模型对每条候选输出一个标量 logit经 log-softmax 得到模型分布最小化两者的KL 散度等价于交叉熵减去目标分布的熵见loss -(target_dist * model_dist - target_dist * target_dist.log()).sum()。该损失不要求模型直接回归指标绝对值而是学会在束内区分优劣与重排序任务的目标天然一致。criterion 同时提供了forward_batch_size默认 32参数当 beam 很大时可把每个样本的模型前向拆成更小的批以规避显存溢出discriminative_reranking_criterion.py 中按forward_batch_size切片循环前向后再拼接。modelmodel: _name: discriminative_nmt_reranker pretrained_model: ??? classifier_dropout: 0.2模型结构与参数详见 models/discriminative_reranking_model.py加载 XLMR 的model.pt通过update_init_roberta_model_state把预训练权重重命名映射到TransformerSentenceEncoder剔除lm_head等无需迁移的参数并把layernorm_embedding改为emb_layer_norm等在 XLMR 之上加入分类头RobertaClassificationHead单输出标量tanh 激活 classifier_dropout输入序列为源句 候选译文的拼接二者用分隔符eos token隔开get_segment_labels依据分隔符生成段标签get_positions为两个 segment 分别重排位置编码模型以head表示sentence_rephead即[CLS]位置输出作为句子级特征可选sentence_rep为head/meanpool/maxpool以及joint_classificationnone/sent在束维度上再做若干层联合 Transformer 以捕捉候选间交互、freeze_embeddings、n_trans_layers_to_freeze等扩展配置均在DiscriminativeNMTRerankerConfig中有默认值与说明discriminative_reranking_model.py。distributed_trainingdistributed_training: ddp_backend: no_c10d distributed_world_size: 16论文实验使用 16 卡训练。3.3 训练期的数据处理细节源码层面task 实现 中还有几个值得注意的工程细节分片轮转load_dataset中当数据路径以数字结尾时按(epoch-1) % num_data_splits 1在每个 epoch 轮换使用不同的数据分片实现多分片训练数据全覆盖beam 分组打乱train 集的 shuffle 以 beam 为单位进行——先打乱mt_beam个起点索引再在组内平铺保证同一源句的 N 条候选始终相邻discriminative_reranking_task.py验证期指标eval_target_metric: true时验证阶段会取每个 beam 组内模型得分最高argmax的候选用其 BLEU/TER 统计量在reduce_metrics中汇总计算验证集 BLEU/TERdiscriminative_reranking_task.pyTER 取负因 TER 越低越好加载标签时对 TER 取负号np_labels -np_labelsdiscriminative_reranking_task.py使指标越大越好的优化方向对两种指标统一。四、推理与评分调权与重排序推理阶段执行 DrNMT 重排序fw 得分 reranker 得分融合。4.1 第一步在 valid 集上调优融合权重先用基础 MT 模型为每个源句生成 N 条候选记录 fw 得分再用重排序模型打分最后在 valid 集上随机搜索fw 权重 长度惩罚的最佳组合# 用基础 MT 模型生成 N 条候选fw score VALID_SOURCE_FILE/path/to/source_sentences # 每行一句已用基础 MT 模型的 sentencepiece 切分 VALID_TARGET_FILE/path/to/target_sentences # 每行一句原始文本无 sentencepiece、无 tokenization MT_MODEL/path/to/mt_model MT_DATA_PATH/path/to/mt_data cat ${VALID_SOURCE_FILE} | \ fairseq-interactive ${MT_DATA_PATH} \ --max-tokens 4000 --buffer-size 16 \ --num-workers 32 --path ${MT_MODEL} \ --beam $N --nbest $N \ --post-process sentencepiece valid-hypo.out # 调优融合权重若目标指标是 TER把 --metric bleu 换成 ter python drnmt_rerank.py \ ${OUTPUT_DIR}/$METRIC/split1/ \ --path ${EXP_DIR}/checkpoint_best.pt \ --in-text valid-hypo.out \ --results-path ${EXP_DIR} \ --gen-subset valid \ --target-text ${VALID_TARGET_FILE} \ --user-dir ${FAIRSEQ_ROOT}/examples/discriminative_reranking_nmt \ --bpe sentencepiece \ --sentencepiece-model ${XLMR_DIR}/sentencepiece.bpe.model \ --beam $N \ --batch-size $N \ --metric bleu \ --tune4.2 第二步在 test 集上应用最佳权重# 生成 test 候选fw score TEST_SOURCE_FILE/path/to/source_sentences # 每行一句已用基础 MT 模型的 sentencepiece 切分 cat ${TEST_SOURCE_FILE} | \ fairseq-interactive ${MT_DATA_PATH} \ --max-tokens 4000 --buffer-size 16 \ --num-workers 32 --path ${MT_MODEL} \ --beam $N --nbest $N \ --post-process sentencepiece test-hypo.out # 用第一步调出的 BEST_FW_WEIGHT / BEST_LENPEN 重排序并评分 # 若评估 TER把 --metric bleu 换成 ter # 添加 --target-text 可计算 BLEU/TER否则脚本只输出得分最高的候选译文 python drnmt_rerank.py \ ${OUTPUT_DIR}/$METRIC/split1/ \ --path ${EXP_DIR}/checkpoint_best.pt \ --in-text test-hypo.out \ --results-path ${EXP_DIR} \ --gen-subset test \ --user-dir ${FAIRSEQ_ROOT}/examples/discriminative_reranking_nmt \ --bpe sentencepiece \ --sentencepiece-model ${XLMR_DIR}/sentencepiece.bpe.model \ --beam $N \ --batch-size $N \ --metric bleu \ --fw-weight ${BEST_FW_WEIGHT} \ --lenpen ${BEST_LENPEN}4.3 融合打分公式与实现细节重排序的最终得分融合公式在 drnmt_rerank.py 中实现def get_score(mt_s, md_s, w1, lp, tgt_len): return mt_s / (tgt_len ** lp) * w1 md_s即最终得分 fw得分 / (译文长度 ** 长度惩罚) × fw权重 重排序模型得分。其中tgt_len ** lp是对候选长度差异的补偿fw_weightw1控制基础 MT 模型得分的贡献比例。get_best_hyps在每个 beam 组内取融合得分最高的候选作为最终输出。脚本的其他要点drnmt_rerank.pyparse_fairseq_gen解析fairseq-interactive的输出S-行提取源句D-行提取候选译文及其 fw 得分--tune模式在--lower-bound-fw-weight至--upper-bound-fw-weight默认 0.0–3.0与--lower-bound-lenpen至--upper-bound-lenpen默认 0.0–3.0范围内随机搜索--num-trials默认 1000组权重组合并用 32 进程并行评估drnmt_rerank.py评估指标用 sacrebleu 计算BLEU 取最高TER 取最低结果写入${results_path}/generate-${gen_subset}.txt每行格式为序号TAB融合得分TAB最优译文drnmt_rerank.py提供--target-text时脚本会同时打印重排序前后的 BLEU/TER 对比before reranking/after reranking with fw_weight..., lenpen...便于直接量化提升幅度drnmt_rerank.py。五、工作流程总览与实战建议完整的 DrNMT 落地流程如下训练基础 MT 模型 (examples/translation) │ fairseq-interactive --beam N --nbest N ▼ 生成 源句 / 参考译文 / N 条候选 三个原始文件 │ scripts/prep_data.pysentencepiece 编码 sacrebleu 标签 ▼ BPE 文本 指标标签split1..splitK │ fairseq-preprocess --only-source --srcdict xlmr/dict.txt ▼ fairseq 二值化数据 │ fairseq-hydra-trainconfig/deen.yaml user_dir ▼ 重排序器 checkpoint_best.pt │ drnmt_rerank.py --tunevalid 集调 fw_weight / lenpen ▼ 最佳权重 → drnmt_rerank.pytest 集重排序 BLEU/TER 评估实战建议候选数 N 的选择论文使用 50候选越多重排序器的选择空间越大但数据量与显存开销也线性增长。若资源受限可按 README 提示同步调整mt_beam、batch_size、required_batch_size_multiple训练硬件适配GPU 不足 16 卡时通过update_freq梯度累积等价放大 batchbeam 很大时可用 criterion 的forward_batch_size拆分前向验证集权重调优--tune的随机搜索范围与次数--num-trials可按需调整调出的fw_weight/lenpen必须复用到 test 集避免在 test 上二次调参造成信息泄漏指标对齐训练、验证与推理阶段的--metricbleu / ter应保持一致TER 优化时注意训练标签取负、验证取最低的细节。六、引用该示例对应的论文为原 README 附带的 BibTeXinproceedings{lee2021discriminative, title{Discriminative Reranking for Neural Machine Translation}, author{Lee, Ann and Auli, Michael and Ranzato, MarcAurelio}, booktitle{ACL}, year{2021} }如需复现论文中的 De-En 实验可直接以 config/deen.yaml 为模板若需在自己的语言对或模型上应用只需替换基础 MT 模型、XLMR 与数据路径并按其说明调整update_freq、mt_beam与batch_size即可。【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考