基于 fairseq 的 BART 摘要微调实战:从 CNN-Dailymail 数据预处理到 Beam Search 推理
2026/9/13 8:25:02 网站建设 项目流程

基于 fairseq 的 BART 摘要微调实战:从 CNN-Dailymail 数据预处理到 Beam Search 推理

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

导读

本文基于仓库 decoding/IAD/fairseq/examples/bart/README.summarization.md 整理,完整讲解如何在 fairseq 框架下将预训练 BART 模型微调到 CNN-Dailymail 与 XSum 文本摘要任务。你将掌握一条可复现的全链路:原始数据下载与未分词清洗 → GPT-2 BPE 编码 → 数据集二值化 →fairseq-train微调参数配置 → 用BARTModel.from_pretrainedbart.sample完成 beam search 摘要生成,并理解每个命令行参数背后的源码实现依据。

背景:为什么 BART 适合做抽取式/生成式摘要

BART(Bidirectional and Auto-Regressive Transformer)是一个去噪自编码式的序列到序列预训练模型:编码器采用双向注意力理解全文,解码器采用自回归方式逐 token 生成,恰好与"读长文、写摘要"的任务形态天然匹配。在本仓库中,BART 的完整实现位于 fairseq/fairseq/models/bart/model.py,其中BARTModel直接继承自TransformerModel,并通过@register_model("bart")注册到 fairseq 模型注册表中,微调时由--arch bart_large自动构建。

从源码看,BARTModel的两个关键设计是:

  • 初始化时调用self.apply(init_bert_params),即采用 BERT 风格的随机初始化(model.py),这解释了为什么微调时可以放心使用较小的学习率;
  • 前向传播是标准的 encoder-decoder:编码器处理src_tokens,解码器以prev_output_tokens为输入做 teacher-forcing(model.py),微调与推理共用同一套TransformerModel的基础设施。

预训练模型方面,官方提供bart.base(6 层编码器/解码器,140M 参数)与bart.large(12 层,400M 参数),以及直接微调好的bart.large.cnnbart.large.xsum等变体,详见 examples/bart/README.md。在 CNN-Dailymail 测试集上,bart.large的 ROUGE-1 / ROUGE-2 / ROUGE-L 分别为 44.16 / 21.28 / 40.90,高于当时的抽取式基线 BERTSUMEXTABS(42.13 / 19.60 / 39.18)。

步骤一:下载并预处理 CNN-Dailymail 与 XSum 原始数据

CNN-Dailymail

CNN-DailyMail 是新闻摘要领域最经典的数据集。本仓库文档要求不要对原始语料做任何 tokenization 或 BPE,保持"非分词、cased"的原始形态:

# 下载原始 CNN 与 Daily Mail 数据集 # 参照 abisee/cnn-dailymail 仓库的说明进行下载与解压

处理得到的数据文件格式为:cnn_dm/train.sourcecnn_dm/train.targetcnn_dm/val.sourcecnn_dm/val.target等,每个.source文件按行存放一篇新闻正文,对应的.target文件按行存放人工撰写的摘要。后续所有脚本都建立在这个"source/target 逐行对应"的约定之上,因此这一步的格式正确性至关重要。

XSum

XSum(Extreme Summarization)任务要求生成极短的摘要(通常仅 1 句),数据处理要求与 CNN-DM 一致:保留原始数据集,确保没有做任何 tokenization 和 BPE。XSum 与 CNN-DM 的差异不仅在于数据规模,更在于生成目标长度差异巨大,这直接决定了后续微调超参(见步骤四)与推理参数(见步骤六)的不同取值。

步骤二:GPT-2 BPE 编码预处理

BART 与 GPT-2 共用同一套 BPE 词表,因此需要先下载三个词表文件,再调用 fairseq 提供的多进程 BPE 编码脚本:

wget -N 'https://dl.fbaipublicfiles.com/fairseq/gpt2_bpe/encoder.json' wget -N 'https://dl.fbaipublicfiles.com/fairseq/gpt2_bpe/vocab.bpe' wget -N 'https://dl.fbaipublicfiles.com/fairseq/gpt2_bpe/dict.txt' 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

该脚本位于 examples/roberta/multiprocessing_bpe_encoder.py,内部通过fairseq.data.encoders.gpt2_bpe.get_encoder加载 GPT-2 BPE,并用 Pythonmultiprocessing.Pool并行编码(--workers默认 20,示例中提升到 60 以加速)。

几个参数要点:

参数含义示例值
--encoder-jsonGPT-2 BPE 的 encoder 映射文件encoder.json
--vocab-bpeBPE 合并规则文件vocab.bpe
--inputs输入文件列表(可多个)cnn_dm/train.source
--outputs输出文件列表,与 inputs 一一对应cnn_dm/train.bpe.source
--workers并行进程数60
--keep-empty保留空行,不过滤(默认空行会被丢弃)-

关于 BPE 的一个易错细节:GPT-2 BPE 对前导空格敏感。从 hub_interface.py 的注释可以看到,bart.encode('Hello world')bart.encode(' world')bart.encode('world')得到的 token 序列完全不同(分别是[0, 31414, 232, 2][0, 232, 2][0, 8331, 2])。因此训练数据必须经同一套 BPE 流程处理,推理时也要走bart.encode而不要手工切词,保证词表一致。

步骤三:用 fairseq-preprocess 二值化数据集

fairseq 的训练入口读取的是二进制格式(.bin/.idx)数据,需要将 BPE 后的文本转为该格式:

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/--target-lang:指定源语言与目标语言名,这里统一命名为source/target(与后续fairseq-train --source-lang source --target-lang target保持一致);
  • --trainpref/--validpref:训练/验证集文件前缀,工具会自动拼接.bpe.source.bpe.target
  • --destdir:二值化输出目录,本示例为cnn_dm-bin/,后续微调与推理都要引用该目录;
  • --srcdict/--tgtdict:复用上一步下载的dict.txt(GPT-2 BPE 词表,词表大小 50265 级别),保证与预训练模型词表完全对齐——这是能否成功--restore-file加载预训练权重的前提。

步骤四:CNN-DM 微调与核心超参解读

官方示例命令

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;

参数分组详解

模型与任务定义

参数作用
--task translation将摘要建模为"源→目标"翻译任务(序列到序列生成)
--arch bart_large使用 12 层 encoder/decoder 的 400M 参数架构;若资源有限可换bart_base
--truncate-source超长新闻正文截断到max-positions,避免 batch 内长度爆炸
--layernorm-embedding嵌入层后加 LayerNorm(BART 架构要求)
--share-all-embeddings编码器/解码器/输出层共享 embedding 矩阵,大幅减少参数量
--share-decoder-input-output-embed解码器输入与输出投影共享权重
--restore-file $BART_PATH加载预训练权重;配合--reset-*丢弃预训练阶段的优化器状态

优化器与学习率

参数作用
--optimizer adam --adam-betas "(0.9, 0.999)" --adam-eps 1e-08Adam 优化器及其超参,注意 betas 需用引号包裹
--lr 3e-05预训练模型微调惯用小学习率,配合 BERT 式初始化(见步骤二源码分析)
--lr-scheduler polynomial_decay多项式衰减调度
--total-num-update 20000总更新步数
--warmup-updates 500前 500 步线性 warmup
--weight-decay 0.01--clip-norm 0.1权重衰减 0.01,梯度裁剪阈值 0.1

训练稳定性与显存

参数作用
--max-tokens 2048每个 batch 的 token 上限(不是样本数),2048 是 32GB V100 的常见取值
--update-freq 4梯度累积 4 步再更新一次,等效放大 batch size 4 倍
--fp16半精度混合精度训练,节省显存并加速
--criterion label_smoothed_cross_entropy --label-smoothing 0.1标签平滑 CE,缓解生成任务过拟合
--dropout 0.1 --attention-dropout 0.1常规 dropout 与 attention dropout
--reset-optimizer --reset-dataloader --reset-meters微调前重置优化器/数据加载器/统计器,防止预训练状态干扰
--skip-invalid-size-inputs-valid-test验证集跳过超长样本
--find-unused-parameters允许模型存在未使用参数(BART 微调常用,避免 DDP 报错)

硬件与耗时预期

  • 上述命令预期在1 个节点、8 张 32GB V100上运行,训练约5 小时
  • 如需缩短时间,可在4 个节点上做分布式训练并配合--update-freq 1(每节点梯度不再累积,靠数据并行扩大 batch)。

步骤五:XSum 任务的参数调整

XSum 摘要更短、数据分布不同,官方给出的微调差异仅为:

TOTAL_NUM_UPDATES=15000 UPDATE_FREQ=2

即总更新步数降为 15000、梯度累积降为 2(等效 batch 减半),其余参数(学习率、warmup、架构等)与 CNN-DM 完全一致。这说明同一份流水线可以无缝迁移到不同摘要任务,只需微调数据量与 batch 相关的超参。

步骤六:用训练好的 checkpoint 做 beam search 推理

CNN-DM 推理代码

训练完成后,checkpoint 保存在checkpoints/目录(checkpoint_best.pt),使用以下 Python 代码批量生成摘要:

import torch from fairseq.models.bart import BARTModel bart = BARTModel.from_pretrained( 'checkpoints/', checkpoint_file='checkpoint_best.pt', data_name_or_path='cnn_dm-bin' ) bart.cuda() bart.eval() bart.half() count = 1 bsz = 32 with open('cnn_dm/test.source') as source, open('cnn_dm/test.hypo', 'w') as fout: sline = source.readline().strip() slines = [sline] for sline in source: if count % bsz == 0: with torch.no_grad(): hypotheses_batch = bart.sample(slines, beam=4, lenpen=2.0, max_len_b=140, min_len=55, no_repeat_ngram_size=3) for hypothesis in hypotheses_batch: fout.write(hypothesis + '\n') fout.flush() slines = [] slines.append(sline.strip()) count += 1 if slines != []: hypotheses_batch = bart.sample(slines, beam=4, lenpen=2.0, max_len_b=140, min_len=55, no_repeat_ngram_size=3) for hypothesis in hypotheses_batch: fout.write(hypothesis + '\n') fout.flush()

代码与源码对应关系

  • BARTModel.from_pretrained(...):定义于 fairseq/models/bart/model.py,从 checkpoint 目录恢复模型、词表与配置;data_name_or_path指定二值化数据目录以加载词典;
  • bart.sample(slines, beam=4, lenpen=2.0, max_len_b=140, min_len=55, no_repeat_ngram_size=3):基于 fairseq 的 SequenceGenerator 做 beam search。参数含义:
    • beam=4:束宽 4;
    • lenpen=2.0:长度惩罚系数,>1 鼓励生成长摘要(CNN-DM 摘要较长);
    • max_len_b=140:最大生成长度 140 token;
    • min_len=55:最短生成长度 55 token;
    • no_repeat_ngram_size=3:禁止出现 3-gram 重复,抑制退化输出;
  • bart.half():推理时切换到 FP16 加速显存;bart.eval()关闭 dropout;
  • 批量逻辑:每次攒够bsz=32条样本后统一送入 GPU 解码,最后一组不足 32 条的余数在循环外处理;每条假设立即写入test.hypoflush(),方便观察进度与断点续跑。

XSum 推理参数

XSum 生成目标短,官方建议:

beam=6, lenpen=1.0, max_len_b=60, min_len=10

即更大的束宽(6)、中性长度惩罚(1.0)、更短的长度区间(60/10),与 XSum"一句话摘要"的数据特性匹配。

步骤七:ROUGE 指标评测

如需在 CNN-DM 测试集上复现论文指标,需要先对假设与参考做 PTB 分词再计算 ROUGE:

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 # Expected output: (ROUGE-2 Average_F: 0.21238)

其中files2rouge需要单独安装;评测结果与论文中的 ROUGE-2 ≈ 0.21 对应,可作为微调正确性的快速 sanity check。注意评测前必须用同一套 PTB 分词器处理假设与参考,否则 ROUGE 会因词边界不一致而失真。

常见问题与排查要点

  1. --restore-file报词表不匹配:多半是fairseq-preprocess时没有用官方dict.txt,或 BPE 编码与预训练词表不一致,回到步骤二/三核对;
  2. 训练时显存溢出(OOM):降低--max-tokens(如 1024)并适当提高--update-freq保持等效 batch;或改用bart_base架构;
  3. 生成内容大量重复:提高no_repeat_ngram_size或调低lenpen
  4. 验证集报样本超长:保留--skip-invalid-size-inputs-valid-test即可跳过;
  5. 分布式多节点训练:将--update-freq降为 1,并正确配置--distributed-world-size等分布式参数,通过节点数扩展 batch。

总结

本文以 decoding/IAD/fairseq/examples/bart/README.summarization.md 为主线,串起了 BART 摘要微调的完整数据流:原始数据(不 tokenize)→ GPT-2 BPE 多进程编码(multiprocessing_bpe_encoder.py)→ 二值化 →fairseq-train微调(bart_large+ 标签平滑 + polynomial_decay + FP16)→BARTModel.from_pretrained+bart.samplebeam search 推理 → ROUGE 评测。整套流程既可直接复现 CNN-DM(beam=4, lenpen=2.0, max_len_b=140, min_len=55),也可通过修改TOTAL_NUM_UPDATES/UPDATE_FREQ与推理参数迁移到 XSum 等短摘要任务。模型结构细节可继续研读 fairseq/models/bart/model.py 与 hub_interface.py,理解其背后的 Transformer 双向编码与自回归解码设计。

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

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

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

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

立即咨询