首页
/ 使用 Transformers 中的 BertGeneration 构建序列生成模型:从 BERT 预训练权重到 Bert2Bert 微调实战

使用 Transformers 中的 BertGeneration 构建序列生成模型:从 BERT 预训练权重到 Bert2Bert 微调实战

2026-09-09 14:36:18作者:明树来

导读

本文以 Transformers 仓库中 BertGeneration 模型文档 为核心,系统讲解如何利用公开的 BERT / RoBERTa 预训练检查点,通过 BertGenerationEncoderBertGenerationDecoderEncoderDecoderModel 组装出可用于摘要、句子融合、句子分割与机器翻译等序列生成任务的 Seq2Seq 模型。读完本文,你将掌握 Bert2Bert 模型的完整搭建流程、配置项含义、训练与推理细节,以及它在仓库源码中的底层实现与测试验证依据。

一、模型背景:让预训练 BERT 承担生成任务

BertGeneration 是一类专为序列生成任务设计的 BERT 模型,其思路来自论文 Leveraging Pre-trained Checkpoints for Sequence Generation Tasks(作者 Sascha Rothe、Shashi Narayan、Aliaksei Severyn)。论文的核心观点是:大规模无监督预训练虽然已经彻底改变了 NLP,但此前业界主要把预训练检查点用于"自然语言理解"类任务;该工作则证明,公开的 BERT、GPT-2 与 RoBERTa 检查点同样可以作为编码器/解码器初始化,显著加速序列生成任务的收敛,并在机器翻译、文本摘要、句子分割(sentence splitting)和句子融合(sentence fusion)等任务上取得了当时的先进结果。

基于这一思路,Transformers 仓库将生成适配层封装为三个核心类:

  • BertGenerationConfig:模型配置;
  • BertGenerationEncoder:可充当编码器(纯自注意力)或解码器(叠加交叉注意力层)的裸 Transformer 主干;
  • BertGenerationDecoder:带语言建模头的解码器,可直接用于 CLM 微调与自回归生成;
  • BertGenerationTokenizer:基于 SentencePiece 的分词器。

它们与通用的 EncoderDecoderModel 组合,即可复用两个预训练 BERT 检查点完成端到端微调。仓库内实现位于 src/transformers/models/bert_generation/ 目录。

二、核心组件解析

2.1 BertGenerationConfig:默认参数即"大模型"配置

BertGenerationConfig 继承自 PreTrainedConfig,模型类型为 "bert-generation",其默认配置对应论文中使用的 24 层大模型规格(见 configuration_bert_generation.py):

配置项 默认值 含义
vocab_size 50358 词表大小
hidden_size 1024 隐藏层维度
num_hidden_layers 24 Transformer 层数
num_attention_heads 16 注意力头数
intermediate_size 4096 FFN 中间层维度
hidden_act "gelu" 激活函数
hidden_dropout_prob / attention_probs_dropout_prob 0.1 丢弃率
max_position_embeddings 512 最大位置编码长度
initializer_range 0.02 权重初始化范围
layer_norm_eps 1e-12 LayerNorm epsilon
pad_token_id / bos_token_id / eos_token_id 0 / 2 / 1 特殊 token id
use_cache True 是否使用 KV 缓存加速生成
is_decoder False 是否作为解码器运行
add_cross_attention False 是否添加交叉注意力层
tie_word_embeddings True 是否绑定输入/输出词嵌入

其中 is_decoderadd_cross_attention 是决定模型"角色"的关键开关(详见 2.2 节)。bos_token_id/eos_token_id 支持整数,eos_token_id 还支持整数列表,便于配置多个终止符。

2.2 BertGenerationEncoder / Decoder:一个主干,两种角色

从源码看,BertGenerationEncoderBertGenerationDecoder 共享同一套主干实现,二者关系如下:

  • BertGenerationEncodermodeling_bert_generation.py)输出原始 hidden states,不带任务头。它既可以做纯编码器(只含双向自注意力),也可以做解码器——当 config.is_decoder=True 时,前向过程会通过 create_causal_mask 生成因果掩码,保证自回归特性;
  • BertGenerationDecodermodeling_bert_generation.py)在主干之上叠加 BertGenerationOnlyLMHead(一个 Linear(hidden_size, vocab_size) 输出层),并继承 GenerationMixin,因此天然支持 generate() 自回归解码;其 lm_head.decoder.weight 与输入词嵌入通过 _tied_weights_keys 声明为权重绑定关系(对应 tie_word_embeddings=True)。

需要特别说明的是:是否插入交叉注意力层由 add_cross_attention 控制。在 BertGenerationLayer 的构造逻辑中,若 add_cross_attention=Trueis_decoder=False,会直接抛出 ValueError,因为交叉注意力只对解码器有意义。当两者同时为 True 时,每一层在自注意力之后额外执行一次对 encoder_hidden_states 的交叉注意力(BertGenerationCrossAttention),这正是 Seq2Seq 解码器读取编码器输出的机制。

此外,BertGenerationPreTrainedModel 声明了 _supports_flash_attn_supports_sdpa_supports_flex_attn,即该模型可选用 eager、Flash Attention、SDPA 等不同注意力后端。

2.3 BertGenerationTokenizer:SentencePiece 分词

BertGenerationTokenizer 基于 SentencePiece(见 tokenization_bert_generation.py),词表文件名为 spiece.model,默认特殊 token 为:bos_token="<s>"eos_token="</s>"unk_token="<unk>"pad_token="<pad>"sep_token="<::::>"。它还支持通过 sp_model_kwargs 传入 enable_samplingnbest_sizealpha 等参数启用子词正则化(subword regularization)。测试文件 test_tokenization_bert_generation.py 验证了词表转换、<s>/<unk>/<pad> 的 id 映射等行为。

三、实战一:用两个 BERT 检查点组装 Bert2Bert 模型

文档给出的核心用法是:将模型与 EncoderDecoderModel 结合,复用在 Hub 上公开的 BERT 检查点。核心代码如下(完整示例见 docs/source/ja/model_doc/bert-generation.md):

# 利用检查点构建 Bert2Bert 模型
# 编码器:使用 BERT 的 cls token (101) 作为 BOS token,sep token (102) 作为 EOS token
encoder = BertGenerationEncoder.from_pretrained(
    "google-bert/bert-large-uncased", bos_token_id=101, eos_token_id=102
)

# 解码器:添加交叉注意力层,同样使用 cls token 作为 BOS、sep token 作为 EOS
decoder = BertGenerationDecoder.from_pretrained(
    "google-bert/bert-large-uncased",
    add_cross_attention=True,
    is_decoder=True,
    bos_token_id=101,
    eos_token_id=102,
)

bert2bert = EncoderDecoderModel(encoder=encoder, decoder=decoder)

# 创建 tokenizer
tokenizer = BertTokenizer.from_pretrained("google-bert/bert-large-uncased")

input_ids = tokenizer(
    "This is a long article to summarize", add_special_tokens=False, return_tensors="pt"
).input_ids
labels = tokenizer("This is a short summary", return_tensors="pt").input_ids

# 训练:前向计算 loss 并反向传播
loss = bert2bert(input_ids=input_ids, decoder_input_ids=labels, labels=labels).loss
loss.backward()

这段代码揭示了三个关键设计:

  1. 复用 BERT 的特殊 token 约定:由于原始 BERT 没有专门的 BOS/EOS 概念,文档明确建议把 cls token(id 101)当作 BOS、sep token(id 102)当作 EOS,从而无需改动预训练词表即可接入 Seq2Seq 的生成流程;
  2. 解码器必须同时开启两个开关is_decoder=True 让主干生成因果掩码并启用 KV 缓存,add_cross_attention=True 让每一层额外插入交叉注意力子层。这一点在 BertGenerationLayer.forward 中有硬性校验——传入 encoder_hidden_states 时若没有交叉注意力层会直接报错;
  3. 端到端微调EncoderDecoderModel 前向时会把 labels 右移一位后作为解码器输入(见 modeling_encoder_decoder.py 中的 shift_tokens_right 逻辑),因此在训练时只需同时提供 decoder_input_idslabels

EncoderDecoderModel 本身是一个通用封装类(modeling_encoder_decoder.py),它通过 AutoModel.from_config 实例化编码器、AutoModelForCausalLM.from_config 实例化解码器,并在初始化时校验两侧 hidden_size 是否匹配(交叉注意力维度一致性检查)。

四、实战二:直接加载预训练好的 EncoderDecoderModel

除自行组装外,论文作者还提供了训练完成的检查点,可直接从模型 Hub 加载:

# 实例化句子融合模型
sentence_fuser = EncoderDecoderModel.from_pretrained("google/roberta2roberta_L-24_discofuse")
tokenizer = AutoTokenizer.from_pretrained("google/roberta2roberta_L-24_discofuse")

input_ids = tokenizer(
    "This is the first sentence. This is the second sentence.",
    add_special_tokens=False,
    return_tensors="pt",
).input_ids

outputs = sentence_fuser.generate(input_ids)

print(tokenizer.decode(outputs[0]))

这个例子演示了"两句话融合为一句"(sentence fusion)的推理流程:输入不加特殊 token,直接交给 generate() 做自回归解码,最后用分词器把生成的 token id 序列还原为文本。BertGenerationDecoder 继承的 GenerationMixin 提供了 generate() 的全部能力(beam search、采样、长度惩罚等),配合 use_cache=True 的 KV 缓存机制可显著加速逐 token 生成。

五、使用技巧与注意事项

文档末尾给出了两条直接影响训练效果的经验性建议,务必遵守:

  1. BertGenerationEncoder 与 BertGenerationDecoder 应配合 EncoderDecoderModel 使用,不要单独把编码器当生成模型;
  2. 对摘要、句子分割、句子融合和翻译任务,输入无需添加特殊 token——尤其是不要在输入末尾追加 EOS token。这是因为该框架采用"BOS 由编码侧 cls 充当、EOS 由解码侧负责"的约定,输入侧多余的特殊 token 会干扰解码器的注意力对齐。

六、源码与测试佐证

仓库为 bert_generation 提供了完整的模型与分词器测试,可据此验证上述行为:

  • test_modeling_bert_generation.py:覆盖了编码器前向输出形状、add_cross_attention=True 时以解码器角色接收 encoder_hidden_states 的前向、带 past_key_values 的增量解码(对比"无缓存全量前向"与"带缓存增量前向"的隐藏状态一致性)、以及带 labels 的因果语言建模 loss 计算;
  • test_tokenization_bert_generation.py:验证 SentencePiece 词表加载、token↔id 转换与特殊 token 排序。

另外,配置类文档中给出的标准用法是:从 google/bert_for_seq_generation_L-24_bbc_encoder 检查点加载配置与分词器,设置 config.is_decoder = True 后即可得到可直接前向的 BertGenerationDecoder;该检查点也是分词器测试的默认 from_pretrained_id

七、适用前提与限制

  • 本模型的直接适用场景是用 BERT/RoBERTa 预训练权重初始化 Seq2Seq 生成模型;如果只需要纯理解任务,原生 BERT 即可满足需求;
  • 默认配置为 24 层大模型(hidden_size=1024),若显存受限,可通过自定义 BertGenerationConfig 缩小 num_hidden_layershidden_size 等参数后从零初始化;
  • 所有代码示例依赖 torchsentencepiece 环境,分词器加载需要安装 sentencepiece 依赖。

综上,BertGeneration 提供了一条"复用预训练理解模型、低成本迁移到生成任务"的经典路径,其 Encoder/Decoder 双角色设计、交叉注意力开关与 EncoderDecoderModel 的组合方式,值得在自定义 Seq2Seq 架构时参考复用。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.76 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
860
1.35 K
docsdocs
暂无描述
Markdown
899
5.83 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
925
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.84 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
533
601
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.03 K
525
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
395