基于 fairseq 的 BART 摘要微调实战:从 CNN-Dailymail 数据预处理到 Beam Search 推理 📅 发布时间:2026/9/13 8:25:16 👁 浏览次数: 基于 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_pretrained与bart.sample完成 beam search 摘要生成并理解每个命令行参数背后的源码实现依据。背景为什么 BART 适合做抽取式/生成式摘要BARTBidirectional 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-forcingmodel.py微调与推理共用同一套TransformerModel的基础设施。预训练模型方面官方提供bart.base6 层编码器/解码器140M 参数与bart.large12 层400M 参数以及直接微调好的bart.large.cnn、bart.large.xsum等变体详见 examples/bart/README.md。在 CNN-Dailymail 测试集上bart.large的 ROUGE-1 / ROUGE-2 / ROUGE-L 分别为 44.16 / 21.28 / 40.90高于当时的抽取式基线 BERTSUMEXTABS42.13 / 19.60 / 39.18。步骤一下载并预处理 CNN-Dailymail 与 XSum 原始数据CNN-DailymailCNN-DailyMail 是新闻摘要领域最经典的数据集。本仓库文档要求不要对原始语料做任何 tokenization 或 BPE保持非分词、cased的原始形态# 下载原始 CNN 与 Daily Mail 数据集 # 参照 abisee/cnn-dailymail 仓库的说明进行下载与解压处理得到的数据文件格式为cnn_dm/train.source、cnn_dm/train.target、cnn_dm/val.source、cnn_dm/val.target等每个.source文件按行存放一篇新闻正文对应的.target文件按行存放人工撰写的摘要。后续所有脚本都建立在这个source/target 逐行对应的约定之上因此这一步的格式正确性至关重要。XSumXSumExtreme 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 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该脚本位于 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.txtGPT-2 BPE 词表词表大小 50265 级别保证与预训练模型词表完全对齐——这是能否成功--restore-file加载预训练权重的前提。步骤四CNN-DM 微调与核心超参解读官方示例命令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;参数分组详解模型与任务定义参数作用--task translation将摘要建模为源→目标翻译任务序列到序列生成--arch bart_large使用 12 层 encoder/decoder 的 400M 参数架构若资源有限可换bart_base--truncate-source超长新闻正文截断到max-positions避免 batch 内长度爆炸--layernorm-embedding嵌入层后加 LayerNormBART 架构要求--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_UPDATES15000 UPDATE_FREQ2即总更新步数降为 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_filecheckpoint_best.pt, data_name_or_pathcnn_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, beam4, lenpen2.0, max_len_b140, min_len55, no_repeat_ngram_size3) 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, beam4, lenpen2.0, max_len_b140, min_len55, no_repeat_ngram_size3) 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, beam4, lenpen2.0, max_len_b140, min_len55, no_repeat_ngram_size3)基于 fairseq 的 SequenceGenerator 做 beam search。参数含义beam4束宽 4lenpen2.0长度惩罚系数1 鼓励生成长摘要CNN-DM 摘要较长;max_len_b140最大生成长度 140 tokenmin_len55最短生成长度 55 tokenno_repeat_ngram_size3禁止出现 3-gram 重复抑制退化输出bart.half()推理时切换到 FP16 加速显存bart.eval()关闭 dropout批量逻辑每次攒够bsz32条样本后统一送入 GPU 解码最后一组不足 32 条的余数在循环外处理每条假设立即写入test.hypo并flush()方便观察进度与断点续跑。XSum 推理参数XSum 生成目标短官方建议beam6, lenpen1.0, max_len_b60, min_len10即更大的束宽6、中性长度惩罚1.0、更短的长度区间60/10与 XSum一句话摘要的数据特性匹配。步骤七ROUGE 指标评测如需在 CNN-DM 测试集上复现论文指标需要先对假设与参考做 PTB 分词再计算 ROUGEexport 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 会因词边界不一致而失真。常见问题与排查要点--restore-file报词表不匹配多半是fairseq-preprocess时没有用官方dict.txt或 BPE 编码与预训练词表不一致回到步骤二/三核对训练时显存溢出OOM降低--max-tokens如 1024并适当提高--update-freq保持等效 batch或改用bart_base架构生成内容大量重复提高no_repeat_ngram_size或调低lenpen验证集报样本超长保留--skip-invalid-size-inputs-valid-test即可跳过分布式多节点训练将--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_pretrainedbart.samplebeam search 推理 → ROUGE 评测。整套流程既可直接复现 CNN-DMbeam4, lenpen2.0, max_len_b140, min_len55也可通过修改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),仅供参考