- 人工智能
- 大模型
- 预训练
- 微调
- LoRA
- RLHF
- 强化学习
- 分布式训练
【免费下载链接】PaddleNLP
Easy-to-use and powerful LLM and SLM library with awesome model zoo.
导读
本文围绕 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-base | bart-large |
|---|---|---|
vocab_size | 50265 | 50265 |
d_model | 768 | 1024 |
num_encoder_layers/num_decoder_layers | 6 / 6 | 12 / 12 |
encoder_attention_heads/decoder_attention_heads | 12 / 12 | 16 / 16 |
encoder_ffn_dim/decoder_ffn_dim | 3072 / 3072 | 4096 / 4096 |
dropout/attention_dropout/activation_dropout | 0.1 / 0.1 / 0.1 | 0.1 / 0.1 / 0.1 |
activation_function | gelu | gelu |
max_position_embeddings | 1024 | 1024 |
init_std | 0.02 | 0.02 |
scale_embedding | False | False |
| 特殊 token id(bos/pad/eos/forced_eos/decoder_start) | 0 / 1 / 2 / 2 / 2 | 0 / 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.
相关推荐
Transformers 摘要生成实战指南:基于 T5 微调 BillSum 法律文本摘要模型
Transformers 摘要生成实战指南:基于 T5 微调 BillSum 法律文本摘要模型 摘要生成(Summarization)是 🤗 Transfor
人工智能深度学习机器学习预训练微调NLP计算机视觉语音多模态BART模型文本摘要:从原理到实战的终极完整指南
BART模型文本摘要:从原理到实战的终极完整指南 在当今信息爆炸的时代,如何从海量文本中快速提取核心信息成为迫切需求。BART模型文本摘要技术应运而生,它通过深
教程DeepSpeed加速BART文本摘要:三小时打造专业级摘要模型
DeepSpeed加速BART文本摘要:三小时打造专业级摘要模型 还在为训练文本摘要模型耗时过长而烦恼?DeepSpeed让BART模型微调变得前所未有的简单高
示例工程
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考