尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

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

使用 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 预训练检查点通过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_size50358词表大小hidden_size1024隐藏层维度num_hidden_layers24Transformer 层数num_attention_heads16注意力头数intermediate_size4096FFN 中间层维度hidden_actgelu激活函数hidden_dropout_prob/attention_probs_dropout_prob0.1丢弃率max_position_embeddings512最大位置编码长度initializer_range0.02权重初始化范围layer_norm_eps1e-12LayerNorm epsilonpad_token_id/bos_token_id/eos_token_id0 / 2 / 1特殊 token iduse_cacheTrue是否使用 KV 缓存加速生成is_decoderFalse是否作为解码器运行add_cross_attentionFalse是否添加交叉注意力层tie_word_embeddingsTrue是否绑定输入/输出词嵌入其中is_decoder与add_cross_attention是决定模型角色的关键开关详见 2.2 节。bos_token_id/eos_token_id支持整数eos_token_id还支持整数列表便于配置多个终止符。2.2 BertGenerationEncoder / Decoder一个主干两种角色从源码看BertGenerationEncoder与BertGenerationDecoder共享同一套主干实现二者关系如下BertGenerationEncodermodeling_bert_generation.py输出原始 hidden states不带任务头。它既可以做纯编码器只含双向自注意力也可以做解码器——当config.is_decoderTrue时前向过程会通过create_causal_mask生成因果掩码保证自回归特性BertGenerationDecodermodeling_bert_generation.py在主干之上叠加BertGenerationOnlyLMHead一个Linear(hidden_size, vocab_size)输出层并继承GenerationMixin因此天然支持generate()自回归解码其lm_head.decoder.weight与输入词嵌入通过_tied_weights_keys声明为权重绑定关系对应tie_word_embeddingsTrue。需要特别说明的是是否插入交叉注意力层由add_cross_attention控制。在BertGenerationLayer的构造逻辑中若add_cross_attentionTrue但is_decoderFalse会直接抛出ValueError因为交叉注意力只对解码器有意义。当两者同时为True时每一层在自注意力之后额外执行一次对encoder_hidden_states的交叉注意力BertGenerationCrossAttention这正是 Seq2Seq 解码器读取编码器输出的机制。此外BertGenerationPreTrainedModel声明了_supports_flash_attn、_supports_sdpa、_supports_flex_attn即该模型可选用 eager、Flash Attention、SDPA 等不同注意力后端。2.3 BertGenerationTokenizerSentencePiece 分词BertGenerationTokenizer基于 SentencePiece见 tokenization_bert_generation.py词表文件名为spiece.model默认特殊 token 为bos_tokens、eos_token/s、unk_tokenunk、pad_tokenpad、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 tokensep token (102) 作为 EOS token encoder BertGenerationEncoder.from_pretrained( google-bert/bert-large-uncased, bos_token_id101, eos_token_id102 ) # 解码器添加交叉注意力层同样使用 cls token 作为 BOS、sep token 作为 EOS decoder BertGenerationDecoder.from_pretrained( google-bert/bert-large-uncased, add_cross_attentionTrue, is_decoderTrue, bos_token_id101, eos_token_id102, ) bert2bert EncoderDecoderModel(encoderencoder, decoderdecoder) # 创建 tokenizer tokenizer BertTokenizer.from_pretrained(google-bert/bert-large-uncased) input_ids tokenizer( This is a long article to summarize, add_special_tokensFalse, return_tensorspt ).input_ids labels tokenizer(This is a short summary, return_tensorspt).input_ids # 训练前向计算 loss 并反向传播 loss bert2bert(input_idsinput_ids, decoder_input_idslabels, labelslabels).loss loss.backward()这段代码揭示了三个关键设计复用 BERT 的特殊 token 约定由于原始 BERT 没有专门的 BOS/EOS 概念文档明确建议把clstokenid 101当作 BOS、septokenid 102当作 EOS从而无需改动预训练词表即可接入 Seq2Seq 的生成流程解码器必须同时开启两个开关is_decoderTrue让主干生成因果掩码并启用 KV 缓存add_cross_attentionTrue让每一层额外插入交叉注意力子层。这一点在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_tokensFalse, return_tensorspt, ).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_cacheTrue的 KV 缓存机制可显著加速逐 token 生成。五、使用技巧与注意事项文档末尾给出了两条直接影响训练效果的经验性建议务必遵守BertGenerationEncoder 与 BertGenerationDecoder 应配合EncoderDecoderModel使用不要单独把编码器当生成模型对摘要、句子分割、句子融合和翻译任务输入无需添加特殊 token——尤其是不要在输入末尾追加 EOS token。这是因为该框架采用BOS 由编码侧 cls 充当、EOS 由解码侧负责的约定输入侧多余的特殊 token 会干扰解码器的注意力对齐。六、源码与测试佐证仓库为bert_generation提供了完整的模型与分词器测试可据此验证上述行为test_modeling_bert_generation.py覆盖了编码器前向输出形状、add_cross_attentionTrue时以解码器角色接收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_size1024若显存受限可通过自定义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),仅供参考
返回列表