You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

自定义PyTorch编码器+BERT自回归解码器的目标掩码配置问题

问题描述

我正在实现编码器-解码器架构,采用自定义PyTorch编码器搭配HuggingFace的BERT作为解码器(无法使用HuggingFace的EncoderDecoder类),任务为机器翻译。需要让BERT以自回归解码器方式工作,输入仅需padding掩码,训练时必须对未来目标token进行掩码以避免“作弊”。

已将配置参数is_decoder和use_cross_attention设为True,打印模型摘要显示其已与编码器正确关联,但不清楚在forward方法中应传入什么参数来正确实现目标掩码。

根据BertModel的forward方法声明:

def forward(input_ids: typing.Optional[torch.Tensor] = None, 
    attention_mask: typing.Optional[torch.Tensor] = None, 
    token_type_ids: typing.Optional[torch.Tensor] = None, 
    position_ids: typing.Optional[torch.Tensor] = None, 
    head_mask: typing.Optional[torch.Tensor] = None, 
    inputs_embeds: typing.Optional[torch.Tensor] = None, 
    encoder_hidden_states: typing.Optional[torch.Tensor] = None, 
    encoder_attention_mask: typing.Optional[torch.Tensor] = None, 
    past_key_values: typing.Optional[typing.List[torch.FloatTensor]] = None, 
    use_cache: typing.Optional[bool] = None, 
    output_attentions: typing.Optional[bool] = None, 
    output_hidden_states: typing.Optional[bool] = None, 
    return_dict: typing.Optional[bool] = None ) → transformers.modeling_outputs.BaseModelOutputWithPoolingAndCrossAttentions or tuple(torch.FloatTensor)

文档说明attention_mask和encoder_attention_mask均用于避免关注padding(作为源掩码),想知道当BERT处于is_decoder模式时,二者中是否有一个可作为目标掩码,或是目标掩码由HuggingFace自动处理,亦或是需求无法实现?

尝试向其中一个或两个参数传入随机值均能得到无异常输出,但无法知晓内部逻辑,也无法判断输出是否正确。由于训练成本极高,无法逐一尝试,需要明确内部机制,同时希望获得基于自定义编码器的BERT解码器通用示例。


解决方案

核心机制说明

当BERT被配置为解码器(is_decoder=True)时:

  • attention_mask参数:承担两个核心作用:一是掩码目标序列中的padding token,二是在自注意力层中自动生成未来token掩码。HuggingFace的BERT实现会根据is_decoder的状态,将传入的目标padding掩码与下三角未来掩码组合,生成最终的自注意力掩码,完全满足训练时防止“作弊”的需求。
  • encoder_attention_mask参数:仅用于掩码编码器输出中的padding token(即源序列的padding掩码),与目标序列无关。

简单来说:你只需要将目标序列的padding掩码传给attention_mask,BERT会自动处理未来token的掩码逻辑,无需手动构造三角掩码。

内部逻辑验证

BERT的自注意力层在is_decoder=True时,会调用_prepare_decoder_attention_mask方法,将传入的attention_mask(目标padding掩码)转换为包含未来掩码的注意力矩阵——这个矩阵会阻止解码器关注当前token之后的所有token,同时忽略padding位置,确保训练的自回归特性。

自定义编码器+BERT解码器示例

import torch
import torch.nn as nn
from transformers import BertConfig, BertModel, BertTokenizer

# 自定义PyTorch编码器
class CustomEncoder(nn.Module):
    def __init__(self, vocab_size, embed_dim, num_layers):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, embed_dim)
        self.transformer_layers = nn.ModuleList([
            nn.TransformerEncoderLayer(
                d_model=embed_dim,
                nhead=8,
                dim_feedforward=embed_dim*4,
                dropout=0.1,
                batch_first=True
            ) for _ in range(num_layers)
        ])
    
    def forward(self, src_input_ids, src_attention_mask):
        # 生成源序列嵌入
        src_embeds = self.embedding(src_input_ids)
        # 逐层编码
        for layer in self.transformer_layers:
            src_embeds = layer(src_embeds, src_key_padding_mask=~src_attention_mask.bool())
        return src_embeds

# 初始化BERT解码器(适配编码器维度)
def init_bert_decoder(embed_dim):
    config = BertConfig.from_pretrained('bert-base-uncased')
    config.is_decoder = True
    config.use_cross_attention = True
    # 确保解码器隐藏维度与编码器一致
    config.hidden_size = embed_dim
    decoder = BertModel.from_pretrained('bert-base-uncased', config=config)
    return decoder

# 训练步示例
def train_step(src_input_ids, src_attention_mask, tgt_input_ids, tgt_attention_mask, encoder, decoder, tokenizer, loss_fn, optimizer):
    # 编码器前向传播
    encoder_outputs = encoder(src_input_ids, src_attention_mask)
    
    # 解码器前向传播
    decoder_outputs = decoder(
        input_ids=tgt_input_ids,
        attention_mask=tgt_attention_mask,  # 传入目标padding掩码,BERT自动添加未来掩码
        encoder_hidden_states=encoder_outputs,
        encoder_attention_mask=src_attention_mask,  # 传入源padding掩码
        return_dict=True
    )
    
    # 计算翻译损失(移位目标:用前N个token预测后N个token)
    logits = decoder_outputs.last_hidden_state
    shift_logits = logits[:, :-1, :].reshape(-1, logits.size(-1))
    shift_labels = tgt_input_ids[:, 1:].reshape(-1)
    
    loss = loss_fn(shift_logits, shift_labels)
    
    # 反向传播与优化
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    return loss.item()

# 示例运行
if __name__ == "__main__":
    embed_dim = 768
    src_vocab_size = 10000
    tgt_tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
    
    # 初始化组件
    encoder = CustomEncoder(src_vocab_size, embed_dim, num_layers=3)
    decoder = init_bert_decoder(embed_dim)
    loss_fn = nn.CrossEntropyLoss(ignore_index=tgt_tokenizer.pad_token_id)
    optimizer = torch.optim.Adam(list(encoder.parameters()) + list(decoder.parameters()), lr=1e-4)
    
    # 模拟训练输入
    batch_size = 2
    src_seq_len = 10
    tgt_seq_len = 12
    
    src_input_ids = torch.randint(0, src_vocab_size, (batch_size, src_seq_len))
    src_attention_mask = torch.ones((batch_size, src_seq_len))
    src_attention_mask[:, -2:] = 0  # 模拟源序列padding
    
    tgt_inputs = tgt_tokenizer(
        ["hello world this is a test", "another example sentence"],
        return_tensors="pt",
        padding=True,
        max_length=tgt_seq_len,
        truncation=True
    )
    tgt_input_ids = tgt_inputs["input_ids"]
    tgt_attention_mask = tgt_inputs["attention_mask"]
    
    # 执行训练步
    loss = train_step(src_input_ids, src_attention_mask, tgt_input_ids, tgt_attention_mask, encoder, decoder, tgt_tokenizer, loss_fn, optimizer)
    print(f"Training loss: {loss:.4f}")

关键注意事项

  • 确保编码器输出的隐藏层维度与BERT解码器的hidden_size一致,若不一致需添加线性层做维度转换。
  • 训练时,目标序列需采用移位目标策略:用前N个token预测后N个token,同时通过ignore_index忽略padding位置的损失。
  • 推理阶段启用use_cache=True可缓存之前生成的key/value对,加速自回归生成过程。

内容的提问来源于stack exchange,提问作者MIRCO RAMO

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.05 12:45:24