BART 摘要微调实战:基于 fairseq 在 CNN-DailyMail 与 XSum 上完成从数据准备、BPE 编码到推理评估的完整流程

📅 发布时间:2026/9/14 8:26:59
BART 摘要微调实战:基于 fairseq 在 CNN-DailyMail 与 XSum 上完成从数据准备、BPE 编码到推理评估的完整流程
BART 摘要微调实战基于 fairseq 在 CNN-DailyMail 与 XSum 上完成从数据准备、BPE 编码到推理评估的完整流程【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm本篇技术指南以 unilm 仓库kosmos-2 子项目内嵌 fairseq 框架中 BART 摘要微调文档 为主体系统讲解如何用 BARTDenoising Sequence-to-Sequence Pre-training在 CNN-DailyMail 与 XSum 两大抽象式摘要abstractive summarization基准上完成端到端实战包括原始语料获取与预处理、GPT-2 BPE 编码、fairseq-preprocess 数据二值化、fairseq-train 微调以及基于束搜索的推理与 ROUGE 指标评估。读完本文你将掌握一套可直接复制运行的摘要模型微调流水线并理解每一条关键命令行参数背后的实现原理本文中涉及的源码均可从当前仓库中对应路径查证。一、背景为什么用 BART 做摘要任务BART 是一种以去噪自编码denoising autoencoding为预训练目标的双向自编码器其架构本质是一个序列到序列sequence-to-sequenceTransformer编码器为双向bidirectional结构解码器为自回归autoregressive结构。这种双向编码 自回归生成的组合使其天然适配文本摘要、生成式问答、对话回复等 NLG 任务。在 BART 项目主页 中可以查证其预训练模型族与下游任务表现模型说明参数量bart.base6 层编码器 6 层解码器140Mbart.large12 层编码器 12 层解码器400Mbart.large.mnli在 MNLI 上微调400Mbart.large.cnn在 CNN-DM 上微调400Mbart.large.xsum在 XSum 上微调400M从源码看BART 在 fairseq 中的实现是 BARTModel它直接继承自TransformerModel并通过hub_models注册了上述五个官方发布权重examples/bart/README.md 记录了 BART-large 在 CNN/Daily Mail 测试集上的 ROUGE 结果R1 44.16 / R2 21.28 / RL 40.90说明了其在摘要基准上的有效性。本文即聚焦于从零微调 BART 到 CNN-DM / XSum 摘要任务的完整流程。二、第一步下载原始语料并做非 tokenized 预处理微调的第一步是准备原始语料。文档要求的关键点在于数据必须以未分词non-tokenized、保留大小写cased的形式保存即每条样本一行、保持原始英文文本不提前做任何 tokenization 或 BPE后续的编码步骤会统一处理。2.1 CNN / Daily Mail 数据集CNN 与 Daily Mail 是两个经典的新闻摘要数据集。获取方式为按 CNN-DailyMail 官方仓库指引下载原始 CNN 和 Daily Mail 数据需要先向原数据集作者申请获取原始数据参考相关 issue 中的预处理要点或使用社区提供的预处理代码将原始数据转换为train.source/train.target/val.source/val.target/test.source/test.target这类每条样本一行、未分词、保留大小写的文件格式。2.2 XSumExtreme Summarization数据集XSum 是摘要长度更短、抽象性更强的极端摘要数据集。下载后同样需要注意保留原始数据确保不做任何 tokenization 与 BPE。其文件组织方式与 CNN-DM 一致即每个 splittrain/val/test下有 source 与 target 两个文件。三、第二步GPT-2 BPE 预处理原始文本需要先经过 GPT-2 风格的 BPEByte-Pair Encoding编码转换成空格分隔的 token id 序列供 fairseq 后续读取。这一步依赖 fairseq 自带的 multiprocessing_bpe_encoder.py 脚本RoBERTa 与 BART 共用它使用多进程并行编码以加速大规模语料处理。3.1 下载 BPE 词表文件首先需要获取 BPE 编码所需的三个文件从 fairseq 官方 gpt2_bpe 发布渠道下载wget -N .../fairseq/gpt2_bpe/encoder.json wget -N .../fairseq/gpt2_bpe/vocab.bpe wget -N .../fairseq/gpt2_bpe/dict.txtencoder.jsonGPT-2 的 BPE 编码器映射token 合并规则vocab.bpeBPE 合并列表dict.txtBART 模型的词典文件后续fairseq-preprocess会用它作为 source 与 target 的共享词典。从源码可知该脚本底层通过from fairseq.data.encoders.gpt2_bpe import get_encoder加载编码器因此这三个文件必须与 fairseq 内置的 GPT-2 BPE 实现兼容。3.2 批量执行 BPE 编码TASKcnn_dm for SPLIT in train val do for LANG in source target do python -m examples.roberta.multiprocessing_bpe_encoder \ --encoder-json encoder.json \ --vocab-bpe vocab.bpe \ --inputs $TASK/$SPLIT.$LANG \ --outputs $TASK/$SPLIT.bpe.$LANG \ --workers 60 \ --keep-empty; done done脚本关键参数说明对应 multiprocessing_bpe_encoder.py 中的 argparse 定义参数默认值说明--encoder-json必填encoder.json 路径--vocab-bpe必填vocab.bpe 路径--inputs[-]输入文件多个路径用空格分隔-表示标准输入--outputs[-]输出文件数量必须与 inputs 一致--keep-empty关闭是否保留空行不开启时空行会被过滤并计入统计--workers20并行进程数文档示例使用 60编码原理源自源码实现 multiprocessing_bpe_encoder.py每个 worker 进程内通过initializer全局加载一次get_encoder(encoder_json, vocab_bpe)避免重复加载词表主进程用Pool(workers)建立进程池pool.imap(encoder.encode_lines, zip(*inputs), 100)以每 100 行为一个 chunk 分发任务每行文本先strip()若为空且未开启--keep-empty则标记为EMPTY并过滤编码结果以空格分隔的 token id 字符串写入输出文件enc_lines.append( .join(tokens))编码期间每处理 10000 行会在 stderr 打印进度结束后会汇总输出各类被过滤行数的统计。四、第三步fairseq-preprocess 数据二值化BPE 编码后的文本文件仍是纯文本需要转换为 fairseq 的二进制格式.bin/.idx这一步由fairseq-preprocess完成fairseq-preprocess \ --source-lang source \ --target-lang target \ --trainpref ${TASK}/train.bpe \ --validpref ${TASK}/val.bpe \ --destdir ${TASK}-bin/ \ --workers 60 \ --srcdict dict.txt \ --tgtdict dict.txt;参数说明--source-lang source --target-lang target声明源语言与目标语言名称与上一步生成的文件后缀.source/.target对应--trainpref/--validpref训练集与验证集文件前缀脚本会自动补上.source与.target--destdir二值化输出目录示例中为cnn_dm-bin/后续fairseq-train直接以该目录为数据输入--srcdict/--tgtdict源端与目标端词典均指向第一步下载的dict.txt。BART 编码器与解码器共享同一个 GPT-2 词典这与微调命令中的--share-all-embeddings一脉相承--workers 60并行 worker 数。执行完成后${TASK}-bin/目录内会生成dict.source.txt、dict.target.txt以及各 split 的二进制索引文件其中dict.source.txt在后续推理阶段还需要拷贝到 checkpoint 目录。五、第四步在 CNN-DM 上微调 BART-large5.1 微调命令全貌将BART_PATH指向预训练权重即bart.large解压后的model.pt后执行TOTAL_NUM_UPDATES20000 WARMUP_UPDATES500 LR3e-05 MAX_TOKENS2048 UPDATE_FREQ4 BART_PATH/path/to/bart/model.pt CUDA_VISIBLE_DEVICES0,1,2,3,4,5,6,7 fairseq-train cnn_dm-bin \ --restore-file $BART_PATH \ --max-tokens $MAX_TOKENS \ --task translation \ --source-lang source --target-lang target \ --truncate-source \ --layernorm-embedding \ --share-all-embeddings \ --share-decoder-input-output-embed \ --reset-optimizer --reset-dataloader --reset-meters \ --required-batch-size-multiple 1 \ --arch bart_large \ --criterion label_smoothed_cross_entropy \ --label-smoothing 0.1 \ --dropout 0.1 --attention-dropout 0.1 \ --weight-decay 0.01 --optimizer adam --adam-betas (0.9, 0.999) --adam-eps 1e-08 \ --clip-norm 0.1 \ --lr-scheduler polynomial_decay --lr $LR --total-num-update $TOTAL_NUM_UPDATES --warmup-updates $WARMUP_UPDATES \ --fp16 --update-freq $UPDATE_FREQ \ --skip-invalid-size-inputs-valid-test \ --find-unused-parameters;5.2 参数逐条解读参数取值作用--task translationtranslation以标准的序列到序列翻译任务形式训练BART 微调复用该 task--truncate-source开启对超过--max-tokens上限的源文本进行截断而非报错适配长新闻原文--layernorm-embedding开启在 embedding 层后加 LayerNorm是 BART 架构的标志性配置--share-all-embeddings开启编码器与解码器共享 embeddingBART 依赖此特性--share-decoder-input-output-embed开启解码器输入与输出层共享权重--reset-optimizer --reset-dataloader --reset-meters开启从预训练权重恢复模型时重置优化器、数据加载器与统计量避免把预训练阶段的学习率/动量带入微调--arch bart_largebart_large使用 12 层编码器 12 层解码器的 BART-large 架构--criterion label_smoothed_cross_entropy—标签平滑交叉熵摘要生成任务的常用训练目标--label-smoothing0.1标签平滑系数--dropout/--attention-dropout0.1 / 0.1全连接层与注意力层的 dropout--weight-decay0.01L2 权重衰减--optimizer adam --adam-betas (0.9, 0.999) --adam-eps 1e-08—Adam 优化器超参--clip-norm0.1梯度裁剪范数阈值--lr-scheduler polynomial_decay—多项式衰减学习率调度--lr/--total-num-update/--warmup-updates3e-05 / 20000 / 500峰值学习率、总更新步数与 warmup 步数--fp16开启混合精度训练大幅降低显存占用并加速--update-freq4梯度累积步数等效扩大 batch size--skip-invalid-size-inputs-valid-test开启验证/测试阶段跳过超过max-tokens的样本--find-unused-parameters开启自动查找并忽略未被使用的参数避免 DDP 报错5.3 硬件与时长的预期文档明确给出该配置的运行前提与预期以上配置预期在1 个节点、8 张 32GB V100 GPU上运行预期训练时长约5 小时若使用4 个节点做分布式训练并将--update-freq降为 1训练时间可以进一步缩短。5.4 XSum 任务的参数差异文档给出 XSum 微调的调整建议TOTAL_NUM_UPDATES15000 UPDATE_FREQ2即相比 CNN-DMXSum 任务将总更新步数降为15000、梯度累积步数降为2其余配置保持不变。六、第五步推理生成摘要并计算 ROUGE6.1 使用 summarize.py 进行推理训练完成后checkpoint 保存在checkpoints/目录下。推理时借助 examples/bart/summarize.py 脚本cp>XSUM_KWARGS dict(beam6, lenpen1.0, max_len_b60, min_len10, no_repeat_ngram_size3) CNN_KWARGS dict(beam4, lenpen2.0, max_len_b140, min_len55, no_repeat_ngram_size3)解码参数CNN-DM默认XSum--xsum-kwargs含义beam46束搜索宽度lenpen2.01.0长度惩罚max_len_b14060输出最大长度按输入长度 b 的比例计算min_len5510输出最小长度no_repeat_ngram_size33禁止重复 n-gram 的窗口大小抑制生成退化可见CNN-DM 新闻摘要偏向较长的摘要min_len55、max_len_b140、lenpen2.0 惩罚过短输出而 XSum 的极端摘要则短得多min_len10、max_len_b60这与两个数据集的性质一致。6.3 XSum 推理XSum 推理只需追加--xsum-kwargs开关脚本即自动切换为 XSUM 解码参数cp>export CLASSPATH/path/to/stanford-corenlp-full-2016-10-31/stanford-corenlp-3.7.0.jar # Tokenize hypothesis and target files. cat test.hypo | java edu.stanford.nlp.process.PTBTokenizer -ioFileList -preserveLines test.hypo.tokenized cat test.target | java edu.stanford.nlp.process.PTBTokenizer -ioFileList -preserveLines test.hypo.target files2rouge test.hypo.tokenized test.hypo.target先安装files2rouge工具并准备 Stanford CoreNLP 的 PTBTokenizer通过CLASSPATH指定 jar 路径对假设摘要与参考答案分别做 PTB 分词保证 ROUGE 计算在统一粒度上进行files2rouge输出 ROUGE-1 / ROUGE-2 / ROUGE-L 的 F 值。参考值可对照 examples/bart/README.md 中记录的 BART-large 在 CNN/Daily Mail 测试集上的表现R1 44.16 / R2 21.28 / RL 40.90作为基准预期。七、总结完整流水线一览BART 摘要微调的全流程可归纳为五个环节语料准备获取 CNN-DM / XSum 原始数据整理为未分词、保留大小写的逐行source/target文件BPE 编码用 multiprocessing_bpe_encoder.py 配合 GPT-2 词表encoder.json、vocab.bpe将文本转为 token id 序列数据二值化fairseq-preprocess结合dict.txt生成*-bin/二进制数据集微调fairseq-train从bart.large预训练权重恢复以标签平滑交叉熵 多项式衰减调度在 8×V100 上约 5 小时完成 CNN-DM 微调XSum 使用 15000 步 /--update-freq 2推理与评估summarize.py按任务自动选择解码超参CNNbeam4/lenpen2.0XSumbeam6/lenpen1.0生成摘要后再经 PTBTokenizer files2rouge 计算 ROUGE。这套流程不仅适用于 CNN-DM 与 XSum其预训练权重恢复 序列到序列微调 束搜索推理的范式同样可迁移到其他新闻/文档摘要场景只需根据目标数据集的摘要长度特性调整解码参数尤其是max_len_b、min_len与lenpen。相关完整入口与源码均可在当前仓库中查阅BART 摘要微调文档、BART 项目主页、推理脚本、BPE 编码脚本、BART 模型实现。【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考