news 2026/9/10 16:24:56

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
使用 Transformers 中的 BertGeneration 构建序列生成模型:从 BERT 预训练权重到 Bert2Bert 微调实战

使用 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 预训练检查点,通过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_size50358词表大小
hidden_size1024隐藏层维度
num_hidden_layers24Transformer 层数
num_attention_heads16注意力头数
intermediate_size4096FFN 中间层维度
hidden_act"gelu"激活函数
hidden_dropout_prob/attention_probs_dropout_prob0.1丢弃率
max_position_embeddings512最大位置编码长度
initializer_range0.02权重初始化范围
layer_norm_eps1e-12LayerNorm epsilon
pad_token_id/bos_token_id/eos_token_id0 / 2 / 1特殊 token id
use_cacheTrue是否使用 KV 缓存加速生成
is_decoderFalse是否作为解码器运行
add_cross_attentionFalse是否添加交叉注意力层
tie_word_embeddingsTrue是否绑定输入/输出词嵌入

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

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

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

  • 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=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 概念,文档明确建议把clstoken(id 101)当作 BOS、septoken(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 架构时参考复用。

【免费下载链接】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),仅供参考

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

Swin-Transformer-Unet内窥镜图像分割实战

简介&#xff1a;本资源是一套面向医学图像分析研究者与计算机视觉初学者的内窥镜图像语义分割实战代码包&#xff0c;聚焦手术场景下多组织器官的精准像素级识别任务。项目创新融合Transformer与U-Net架构&#xff0c;支持腹壁、肝脏、胆囊、胃肠道等12类解剖结构的端到端分割…

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

LeetCode三数之和问题解析与双指针解法

1. Leetcode 15三数之和问题解析三数之和是Leetcode上经典的算法题目&#xff0c;编号为15。这道题要求找出数组中所有不重复的三元组&#xff0c;使得三个数之和等于零。看似简单的问题背后隐藏着多个需要解决的难点&#xff0c;包括如何高效地遍历所有可能组合、如何避免重复…

作者头像 李华
网站建设 2026/9/10 16:21:09

商用车后轮制动器设计与CAD工程实践

1. 项目背景与需求分析CC1031载货汽车后轮制动器设计是一个典型的商用车底盘系统开发项目。作为载货汽车的核心安全部件&#xff0c;制动器设计直接关系到整车制动性能和道路行驶安全。这个项目要求完成6张CAD工程图纸、设计说明书以及三维模型&#xff0c;涵盖了从概念设计到工…

作者头像 李华
网站建设 2026/9/10 16:19:17

K8s集群安装Jenkins K8s Pod模板配置部署实操

K8s集群安装Jenkins K8s Pod模板配置部署实操 技术栈&#xff1a;Jenkins 2.440.x Kubernetes v1.32.13 Rocky Linux 8.6 Kubernetes Plugin Kaniko Helm 3.14.x 操作环境 / 对接原理 / 详细步骤 / 完整命令 / 配置文件 / 验证流程 / 排错方案 K8s集群安装Jenkins K8s …

作者头像 李华