BART 摘要微调实战:基于 fairseq 在 CNN-DailyMail 与 XSum 上完成从数据准备、BPE 编码到推理评估的完整流程
2026/9/14 8:26:41 网站建设 项目流程

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 摘要微调文档 为主体,系统讲解如何用 BART(Denoising Sequence-to-Sequence Pre-training)在 CNN-DailyMail 与 XSum 两大抽象式摘要(abstractive summarization)基准上完成端到端实战:包括原始语料获取与预处理、GPT-2 BPE 编码、fairseq-preprocess 数据二值化、fairseq-train 微调,以及基于束搜索的推理与 ROUGE 指标评估。读完本文,你将掌握一套可直接复制运行的摘要模型微调流水线,并理解每一条关键命令行参数背后的实现原理(本文中涉及的源码均可从当前仓库中对应路径查证)。

一、背景:为什么用 BART 做摘要任务

BART 是一种以去噪自编码(denoising autoencoding)为预训练目标的双向自编码器,其架构本质是一个序列到序列(sequence-to-sequence)Transformer:编码器为双向(bidirectional)结构,解码器为自回归(autoregressive)结构。这种"双向编码 + 自回归生成"的组合使其天然适配文本摘要、生成式问答、对话回复等 NLG 任务。

在 BART 项目主页 中可以查证其预训练模型族与下游任务表现:

模型说明参数量
bart.base6 层编码器 + 6 层解码器140M
bart.large12 层编码器 + 12 层解码器400M
bart.large.mnli在 MNLI 上微调400M
bart.large.cnn在 CNN-DM 上微调400M
bart.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 是两个经典的新闻摘要数据集。获取方式为:

  1. 按 CNN-DailyMail 官方仓库指引下载原始 CNN 和 Daily Mail 数据(需要先向原数据集作者申请获取原始数据);
  2. 参考相关 issue 中的预处理要点,或使用社区提供的预处理代码,将原始数据转换为train.source/train.target/val.source/val.target/test.source/test.target这类"每条样本一行、未分词、保留大小写"的文件格式。

2.2 XSum(Extreme Summarization)数据集

XSum 是摘要长度更短、抽象性更强的极端摘要数据集。下载后同样需要注意:保留原始数据,确保不做任何 tokenization 与 BPE。其文件组织方式与 CNN-DM 一致,即每个 split(train/val/test)下有 source 与 target 两个文件。

三、第二步:GPT-2 BPE 预处理

原始文本需要先经过 GPT-2 风格的 BPE(Byte-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.txt'
  • encoder.json:GPT-2 的 BPE 编码器映射(token 合并规则);
  • vocab.bpe:BPE 合并列表;
  • dict.txt:BART 模型的词典文件,后续fairseq-preprocess会用它作为 source 与 target 的共享词典。

从源码可知,该脚本底层通过from fairseq.data.encoders.gpt2_bpe import get_encoder加载编码器,因此这三个文件必须与 fairseq 内置的 GPT-2 BPE 实现兼容。

3.2 批量执行 BPE 编码

TASK=cnn_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.txtdict.target.txt以及各 split 的二进制索引文件,其中dict.source.txt在后续推理阶段还需要拷贝到 checkpoint 目录。

五、第四步:在 CNN-DM 上微调 BART-large

5.1 微调命令全貌

BART_PATH指向预训练权重(即bart.large解压后的model.pt)后,执行:

TOTAL_NUM_UPDATES=20000 WARMUP_UPDATES=500 LR=3e-05 MAX_TOKENS=2048 UPDATE_FREQ=4 BART_PATH=/path/to/bart/model.pt CUDA_VISIBLE_DEVICES=0,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开启编码器与解码器共享 embedding,BART 依赖此特性
--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-08Adam 优化器超参
--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_UPDATES=15000 UPDATE_FREQ=2

即相比 CNN-DM,XSum 任务将总更新步数降为15000、梯度累积步数降为2,其余配置保持不变。

六、第五步:推理生成摘要并计算 ROUGE

6.1 使用 summarize.py 进行推理

训练完成后,checkpoint 保存在checkpoints/目录下。推理时借助 examples/bart/summarize.py 脚本:

cp>XSUM_KWARGS = dict(beam=6, lenpen=1.0, max_len_b=60, min_len=10, no_repeat_ngram_size=3) CNN_KWARGS = dict(beam=4, lenpen=2.0, max_len_b=140, min_len=55, no_repeat_ngram_size=3)
解码参数CNN-DM(默认)XSum(--xsum-kwargs含义
beam46束搜索宽度
lenpen2.01.0长度惩罚
max_len_b14060输出最大长度(按输入长度 b 的比例计算)
min_len5510输出最小长度
no_repeat_ngram_size33禁止重复 n-gram 的窗口大小,抑制生成退化

可见:CNN-DM 新闻摘要偏向较长的摘要(min_len=55、max_len_b=140、lenpen=2.0 惩罚过短输出),而 XSum 的极端摘要则短得多(min_len=10、max_len_b=60),这与两个数据集的性质一致。

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 摘要微调的全流程可归纳为五个环节:

  1. 语料准备:获取 CNN-DM / XSum 原始数据,整理为未分词、保留大小写的逐行source/target文件;
  2. BPE 编码:用 multiprocessing_bpe_encoder.py 配合 GPT-2 词表(encoder.jsonvocab.bpe)将文本转为 token id 序列;
  3. 数据二值化fairseq-preprocess结合dict.txt生成*-bin/二进制数据集;
  4. 微调fairseq-trainbart.large预训练权重恢复,以标签平滑交叉熵 + 多项式衰减调度在 8×V100 上约 5 小时完成 CNN-DM 微调(XSum 使用 15000 步 /--update-freq 2);
  5. 推理与评估summarize.py按任务自动选择解码超参(CNN:beam=4/lenpen=2.0;XSum:beam=6/lenpen=1.0),生成摘要后再经 PTBTokenizer + files2rouge 计算 ROUGE。

这套流程不仅适用于 CNN-DM 与 XSum,其"预训练权重恢复 + 序列到序列微调 + 束搜索推理"的范式同样可迁移到其他新闻/文档摘要场景,只需根据目标数据集的摘要长度特性调整解码参数(尤其是max_len_bmin_lenlenpen)。相关完整入口与源码均可在当前仓库中查阅:BART 摘要微调文档、BART 项目主页、推理脚本、BPE 编码脚本、BART 模型实现。

【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询