使用 Transformers 中的 BertGeneration 构建序列生成模型:从 BERT 预训练权重到 Bert2Bert 微调实战
导读
本文以 Transformers 仓库中 BertGeneration 模型文档 为核心,系统讲解如何利用公开的 BERT / RoBERTa 预训练检查点,通过 BertGenerationEncoder、BertGenerationDecoder 与 EncoderDecoderModel 组装出可用于摘要、句子融合、句子分割与机器翻译等序列生成任务的 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_decoder 与 add_cross_attention 是决定模型"角色"的关键开关(详见 2.2 节)。bos_token_id/eos_token_id 支持整数,eos_token_id 还支持整数列表,便于配置多个终止符。
2.2 BertGenerationEncoder / Decoder:一个主干,两种角色
从源码看,BertGenerationEncoder 与 BertGenerationDecoder 共享同一套主干实现,二者关系如下:
BertGenerationEncoder(modeling_bert_generation.py)输出原始 hidden states,不带任务头。它既可以做纯编码器(只含双向自注意力),也可以做解码器——当config.is_decoder=True时,前向过程会通过create_causal_mask生成因果掩码,保证自回归特性;BertGenerationDecoder(modeling_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=True 但 is_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_sampling、nbest_size、alpha 等参数启用子词正则化(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()
这段代码揭示了三个关键设计:
- 复用 BERT 的特殊 token 约定:由于原始 BERT 没有专门的 BOS/EOS 概念,文档明确建议把
clstoken(id 101)当作 BOS、septoken(id 102)当作 EOS,从而无需改动预训练词表即可接入 Seq2Seq 的生成流程; - 解码器必须同时开启两个开关:
is_decoder=True让主干生成因果掩码并启用 KV 缓存,add_cross_attention=True让每一层额外插入交叉注意力子层。这一点在BertGenerationLayer.forward中有硬性校验——传入encoder_hidden_states时若没有交叉注意力层会直接报错; - 端到端微调:
EncoderDecoderModel前向时会把labels右移一位后作为解码器输入(见 modeling_encoder_decoder.py 中的shift_tokens_right逻辑),因此在训练时只需同时提供decoder_input_ids与labels。
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 生成。
五、使用技巧与注意事项
文档末尾给出了两条直接影响训练效果的经验性建议,务必遵守:
- BertGenerationEncoder 与 BertGenerationDecoder 应配合
EncoderDecoderModel使用,不要单独把编码器当生成模型; - 对摘要、句子分割、句子融合和翻译任务,输入无需添加特殊 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_layers、hidden_size等参数后从零初始化; - 所有代码示例依赖
torch与sentencepiece环境,分词器加载需要安装sentencepiece依赖。
综上,BertGeneration 提供了一条"复用预训练理解模型、低成本迁移到生成任务"的经典路径,其 Encoder/Decoder 双角色设计、交叉注意力开关与 EncoderDecoderModel 的组合方式,值得在自定义 Seq2Seq 架构时参考复用。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0631
MiniCPM5-2BMiniCPM5-2B 是一款面向端侧、本地部署和资源受限场景的 2B 稠密 Transformer,能够达到同尺寸开源模型 SOTA 水平。Markdown00
video-shotcraftAI宣传片skill,使用 Remotion 制作电影级产品视频:提供106 张镜头配方卡和可复用的视频魔板。适用于 Claude Code 与 Codex以及所有其他智能体Markdown00
HivisionIDPhotos⚡️HivisionIDPhotos: a lightweight and efficient AI ID photos tools. 一个轻量级的AI证件照制作算法。Python09
DragonOSDragonOS is an operating system developed from scratch using Rust, with Linux compatibility. It is designed for **Serverless** scenarios. 使用Rust从0自研内核,具有Linux兼容性的操作系统,面向云计算Serverless场景而设计。Rust00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00