news 2026/9/25 14:41:42

PaddleNLP 中使用 BART 进行文本摘要:模型原理、微调与生成实践

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PaddleNLP 中使用 BART 进行文本摘要:模型原理、微调与生成实践
  • 人工智能
  • 大模型
  • 预训练
  • 微调
  • LoRA
  • RLHF
  • 强化学习
  • 分布式训练

【免费下载链接】PaddleNLP

Easy-to-use and powerful LLM and SLM library with awesome model zoo.

项目地址:https://gitcode.com/gh_mirrors/pa/PaddleNLP
点击查看免费下载

导读

本文围绕 slm/examples/text_summarization/bart/README.md 所介绍的 BART 文本摘要示例,系统讲解 BART 作为 Seq2Seq 降噪自编码器的模型原理、在 PaddlePaddle 2.2 上的 PaddleNLP 实现、CNN/DailyMail 数据集上的微调流程与摘要生成方法,并深入剖析 paddlenlp/transformers/bart 目录下的源码实现。读完本文,你将掌握 BART 的结构特性、如何加载预训练模型、如何配置BartConfig完成微调,以及如何调用BartForConditionalGeneration进行摘要生成。

BART 模型简介

BART(Bidirectional and Auto-Regressive Transformer)是一种Seq2Seq 结构的降噪自编码器(denoising autoencoder)。其核心训练思想是:对原始文本施加多种噪声进行"破坏"(corrupt),再让模型学习重建(reconstruct)原文,从而在编码器-解码器框架中同时习得双向语境理解与自回归生成能力。

从结构上看,BART 使用标准的 Transformer 架构,可以看作三类预训练模型的统一泛化形式:

  • BERT 的泛化:BART 的编码器是双向的(bidirectional encoder),能够像 BERT 一样建模左右双向上下文;
  • GPT 的泛化:BART 的解码器是从左到右自回归的(left-to-right decoder),具备 GPT 式的生成能力;
  • 其他预训练结构的泛化:通过统一的双向编码 + 自回归解码组合,BART 在文本生成类任务(如摘要、翻译、对话)上表现出色。

这种"破坏-重建"的预训练目标,使得 BART 天然适合下游生成类任务,文本摘要正是其代表性应用场景之一。

本项目概览:PaddleNLP 中的 BART 摘要示例

slm/examples/text_summarization/bart/README.md所描述的项目,是BART 在 PaddlePaddle 2.2 上开源实现的文本摘要示例,覆盖了在CNN/DailyMail数据集上进行**微调(fine-tuning)与摘要生成(generation)**的完整代码路径。

CNN/DailyMail 是英文抽象式摘要领域最经典的基准数据集之一,训练样本由新闻报道原文与人工撰写的摘要句对组成,用于评估模型在长文档压缩与信息抽取方面的能力。项目以该数据集为微调目标,直接面向"输入新闻正文、输出简洁摘要"这一真实任务形态。

在当前仓库中,与 BART 直接相关的实现位于:

  • 模型与配置:paddlenlp/transformers/bart/modeling.py、paddlenlp/transformers/bart/configuration.py
  • 分词器:paddlenlp/transformers/bart/tokenizer.py
  • 单元测试:tests/transformers/bart/test_modeling.py、tests/transformers/bart/test_tokenizer.py

下面分别从预训练权重、模型配置、模型结构、分词器与生成链路几个层面展开。

预训练权重与配置:BartConfig

内置预训练模型清单

在 configuration.py 中,PaddleNLP 内置了两档 BART 预训练配置:bart-base与bart-large,对应两套预设超参数:

参数bart-basebart-large
vocab_size5026550265
d_model7681024
num_encoder_layers/num_decoder_layers6 / 612 / 12
encoder_attention_heads/decoder_attention_heads12 / 1216 / 16
encoder_ffn_dim/decoder_ffn_dim3072 / 30724096 / 4096
dropout/attention_dropout/activation_dropout0.1 / 0.1 / 0.10.1 / 0.1 / 0.1
activation_functiongelugelu
max_position_embeddings10241024
init_std0.020.02
scale_embeddingFalseFalse
特殊 token id(bos/pad/eos/forced_eos/decoder_start)0 / 1 / 2 / 2 / 20 / 1 / 2 / 2 / 2

两类配置共享 50265 词表与相同的特殊 token 定义:bos_token_id=0、pad_token_id=1、eos_token_id=2、forced_eos_token_id=2、decoder_start_token_id=2。其中decoder_start_token_id=2意味着解码器以</s>(eos)作为起始 token,forced_eos_token_id保证生成达到max_length时强制以 eos 收尾。

对应的预训练权重映射表BART_PRETRAINED_RESOURCE_FILES_MAP中登记了bart-base.pdparams与bart-large.pdparams,BartModel与BartForConditionalGeneration可通过from_pretrained("bart-base")的方式自动下载加载。

BartConfig 关键参数语义

BartConfig(见 configuration.py)继承自PretrainedConfig,model_type = "bart"。其核心参数在实例化 BART 模型时直接决定架构:

  • vocab_size:词表大小,决定 embedding 矩阵与 LM 输出头的维度,默认 50265;
  • d_model:编码器/解码器各层的隐藏维度,默认 768;
  • encoder_layers/decoder_layers:编码器与解码器的 Transformer 层数,默认均为 6(注意源码中通过attribute_map将num_encoder_layers等别名映射到encoder_layers/decoder_layers,保证与社区配置的兼容);
  • encoder_attention_heads/decoder_attention_heads:编码器与解码器的注意力头数,默认 12;
  • encoder_ffn_dim/decoder_ffn_dim:编码器与解码器前馈网络的中间维度,默认 3072;
  • activation_function:前馈网络使用的非线性激活函数,支持gelu、relu及 PaddlePaddle 支持的激活函数,默认"gelu";
  • dropout/attention_dropout/activation_dropout:全连接层、注意力概率、激活输出的 dropout 比率,默认均为 0.1;
  • max_position_embeddings:模型可接受的最大序列长度,默认 1024(预训练配置中 base 与 large 均为 1024);
  • init_std:所有权重矩阵初始化时截断正态分布的标准差,默认 0.02;
  • scale_embedding:是否按d_model的平方根缩放 embedding,默认False;
  • is_encoder_decoder=True、decoder_start_token_id=2、forced_eos_token_id=2:表明模型为编解码结构并指定解码起始与强制结束 token;
  • forced_bos_token_id:可选配置,用于部分摘要任务强制首 token 生成<s>,源码中为兼容旧版 CNN 模型保留了force_bos_token_to_be_generated的向后兼容处理。

模型结构源码解析

modeling.py 中实现了完整的 BART 模型族,顶层类定义包括:

  • BartPretrainedModel(基类,封装权重加载与初始化逻辑)
  • BartLearnedPositionalEmbedding(可学习位置编码)
  • BartEncoder/BartDecoder(双向编码器 / 自回归解码器)
  • BartModel(完整编解码主干)
  • BartForConditionalGeneration(条件生成模型,文本摘要的实际入口)
  • 以及BartForSequenceClassification、BartForQuestionAnswering等其他任务头

BartModel:编码器-解码器主干

BartModel由BartEncoder与BartDecoder组成,二者均基于标准 Transformer 层(多头自注意力 + 前馈网络 + LayerNorm + 残差连接),并共享底层的 token embedding(shared),位置编码采用BartLearnedPositionalEmbedding。前向时,编码器对input_ids进行双向编码得到encoder_output,解码器在此基础上结合decoder_input_ids做自回归解码。

BartForConditionalGeneration:摘要生成的入口

BartForConditionalGeneration(modeling.py)在主干之上叠加语言建模头:

  • 内部持有一个BartModel,并创建形状为[vocab_size, d_model]的lm_head_weight参数,配合final_logits_bias偏置构成 LM 输出头;
  • get_encoder()/get_decoder()暴露编码器与解码器,便于与生成框架协同;
  • forward接收input_ids、attention_mask、decoder_input_ids、decoder_attention_mask、encoder_output、labels等输入:labels用于计算掩码语言建模损失(索引为 -100 的 token 被忽略,不参与 loss),当labels=None时返回形状为[batch_size, sequence_length, vocab_size]的lm_logits,可直接用于解码采样;
  • 返回值为Seq2SeqLMOutput(return_dict=True时)或对应元组,use_cache开启时可返回 KV cache 加速自回归。
快速解码(FasterBART)

prepare_fast_entry展示了 BART 的快速解码接入点:通过FasterBART(位于paddlenlp.ops)在具备自定义 decoding 库的环境下启用加速解码,支持use_fp16_decoding、decoding_lib、enable_fast_encoder等开关。同时源码明确约束了快速解码的适用边界:

  • 仅支持 top-k 采样或 top-p 采样中的一种,二者不能同时启用;
  • 暂不支持repetition_penalty != 1.0;
  • 暂不支持min_length != 0;
  • 暂不支持非空的forced_bos_token_id。

在常规 CPU/GPU 环境中,直接使用标准自回归解码路径即可,上述限制仅在启用快速解码时生效。

BartTokenizer:byte-level BPE 分词

BART 的分词器BartTokenizer(tokenizer.py)基于byte-level Byte-Pair-Encoding(BPE),继承自GPTTokenizer,特殊 token 为"<s>"(bos/cls)、"</s>"(eos/sep)等。其实现要点:

  • bytes_to_unicode()构建 UTF-8 字节到 Unicode 字符的可逆映射表,这是 byte-level BPE 处理任意文本(包括未登录字符)的基础;
  • get_pairs(word)生成词内相邻符号二元组,是标准 BPE merge 过程的核心步骤;
  • 加载时需提供vocab_file(词表映射)与merges_file(merge 规则),与预训练词表配套使用;
  • PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES将bart-base/bart-large的位置编码长度标记为 1024,与配置一致。

配套的 tests/transformers/bart/test_tokenizer.py 与 tests/transformers/bart/test_modeling.py 覆盖了分词往返一致性、特殊 token 处理、模型前向与生成等关键路径,可作为自定义 BART 摘要训练的回归验证参考。

文本摘要实战:加载、微调与生成

1. 加载预训练模型与分词器

使用from_pretrained即可加载内置权重:

import paddle from paddlenlp.transformers import BartForConditionalGeneration, BartTokenizer # 加载 bart-base 预训练权重(自动下载) model = BartForConditionalGeneration.from_pretrained("bart-base") tokenizer = BartTokenizer.from_pretrained("bart-base")

"bart-base"与"bart-large"均登记在BART_PRETRAINED_INIT_CONFIGURATION与BART_PRETRAINED_RESOURCE_FILES_MAP中,可直接指定名称加载。

2. 数据准备与微调

微调目标为 CNN/DailyMail 数据集的"新闻正文 → 摘要"句对。训练时构造编解码输入:

def encode_fn(text, summary): # 编码器输入:正文 + eos src_ids = tokenizer(text, max_length=1024, truncation=True)["input_ids"] # 解码器输入:摘要前移一位并补 eos,作为 decoder_input_ids tgt_ids = tokenizer(summary, max_length=128, truncation=True)["input_ids"] return src_ids, tgt_ids

训练时将labels传入BartForConditionalGeneration.forward,模型会自动对摘要 token 计算交叉熵损失(-100位置被忽略)。注意解码器以decoder_start_token_id=2起始,因此摘要 token 序列需整体右移一位置入decoder_input_ids。

3. 摘要生成

微调完成后,使用generate接口进行自回归摘要生成,可配置的典型参数包括:

  • max_length:生成摘要的最大长度(生成结束时由forced_eos_token_id保证以 eos 收尾);
  • min_length:最小长度约束;
  • num_beams:beam search 的束宽,>1 时启用 beam search;
  • decode_strategy:"greedy_search"/"sampling"/"beam_search"等;
  • top_k/top_p:采样策略下的截断与核采样参数;
  • repetition_penalty:重复惩罚系数,抑制摘要中词语重复;
  • no_repeat_ngram_size:禁止 n-gram 重复。

示例:

inputs = tokenizer( "The quick brown fox jumps over the lazy dog .", return_tensors="pd", max_length=1024, truncation=True, ) summary_ids = model.generate( input_ids=inputs["input_ids"], max_length=64, num_beams=4, decode_strategy="beam_search", repetition_penalty=1.2, )[0] summary = tokenizer.decode(summary_ids[0], skip_special_tokens=True) print(summary)

值得注意的是,若启用快速解码路径(FasterBART),当前源码只支持单一 top-k 或 top-p 采样、repetition_penalty=1.0且min_length=0的组合,复杂解码参数请使用标准generate路径。

4. 性能与验收参考

从源码结构可以推断,bart-base(d_model=768、6 层编码器 + 6 层解码器)适合在单卡/CPU 环境快速验证流程,bart-large(d_model=1024、12+12 层)吞吐更高但显存与算力需求明显增大。项目 README 未给出具体评测数值,实际训练时建议以 CNN/DailyMail 验证集的 ROUGE 指标(ROUGE-1/2/L)作为摘要质量验收标准,并结合 tests/transformers/bart/test_modeling.py 中的生成用例确认解码链路正确。

小结

  • BART 是双向编码 + 自回归解码的 Seq2Seq 降噪自编码器,是 BERT 与 GPT 结构的统一泛化;
  • PaddleNLP 在paddlenlp/transformers/bart下完整提供了BartConfig、BartModel、BartForConditionalGeneration与BartTokenizer,内置bart-base/bart-large预训练权重;
  • slm/examples/text_summarization/bart/README.md所示的示例项目基于 PaddlePaddle 2.2,聚焦 CNN/DailyMail 数据集的微调与摘要生成;
  • 实操上通过from_pretrained加载权重、以"正文 → 摘要"句对构造编解码输入完成微调,再用generate结合 beam search、重复惩罚等参数产出摘要;
  • 启用快速解码(FasterBART)时需注意其对解码策略的参数约束,常规场景使用标准生成路径即可。
  • 人工智能
  • 大模型
  • 预训练
  • 微调
  • LoRA
  • RLHF
  • 强化学习
  • 分布式训练

【免费下载链接】PaddleNLP

Easy-to-use and powerful LLM and SLM library with awesome model zoo.

项目地址:https://gitcode.com/gh_mirrors/pa/PaddleNLP
点击查看免费下载

相关推荐

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/25 14:31:02

DeskcommCRM:以沟通为中心打造轻量高效的客户管理系统

我最早看到“DeskcommCRM”这个项目名时&#xff0c;第一反应是这名字起得有点意思&#xff1a;DeskCommCRM&#xff0c;工位、沟通、客户关系管理&#xff0c;三个词拼在一起&#xff0c;翻译过来就是“桌面通信型客户管理系统”&#xff0c;或者更直白一点——一套从一线沟通…

作者头像 李华
网站建设 2026/9/25 14:26:45

OpenWebUI 接入 TaoToken:MCPO 框架下 MCP 工具配置与 OpenAPI 验证

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/25 14:24:10

Cisco AI Assistant for Security 深度图解与 MCP 实战避坑:从架构到落地

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华