自定义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

