使用 Transformers 中的 BertGeneration 构建序列生成模型:从 BERT 预训练权重到 Bert2Bert 微调实战
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
导读
本文以 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 架构时参考复用。
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考