Transformer解码器tgt_mask与tgt_key_padding_mask配置错误求助
Transformer解码器下一个Token预测报错问题解决
问题描述
我实现了一个用于下一个token预测的Transformer解码器,传入tgt_mask避免关注未来token,传入tgt_key_padding_mask忽略padding,但持续报错。
错误日志如下:
Training Epoch 1/1: 0%| | 0/563 [00:00<?, ?it/s] tgt_emb torch.Size([16, 875, 256]) tgt_mask torch.Size([875, 875]) padding_mask torch.Size([16, 875]) /home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/functional.py:5076: UserWarning: Support for mismatched key_padding_mask and attn_mask is deprecated. Use same type for both instead. warnings.warn( Training Epoch 1/1: 0%| | 0/563 [00:01<?, ?it/s] Traceback (most recent call last): File "/scratch/harsha.vasamsetti/decoder_aug_made/main.py", line 146, in <module> output = model(input_batch) File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1518, in _wrapped_call_impl return self._call_impl(*args, **kwargs) File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1527, in _call_impl return forward_call(*args, **kwargs) File "/scratch/harsha.vasamsetti/decoder_aug_made/transformer.py", line 47, in forward output = self.transformer_decoder(tgt_emb, memory, tgt_mask=tgt_mask, tgt_key_padding_mask=padding_mask) File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1518, in _wrapped_call_impl return self._call_impl(*args, **kwargs) File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1527, in _call_impl return forward_call(*args, **kwargs) File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/transformer.py", line 460, in forward output = mod(output, memory, tgt_mask=tgt_mask, File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1518, in _wrapped_call_impl return self._call_impl(*args, **kwargs) File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1527, in _call_impl return forward_call(*args, **kwargs) File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/transformer.py", line 846, in forward x = self.norm1(x + self._sa_block(x, tgt_mask, tgt_key_padding_mask, tgt_is_causal)) File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/transformer.py", line 855, in _sa_block x = self.self_attn(x, x, x, File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1518, in _wrapped_call_impl return self._call_impl(*args, **kwargs) File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/module.py", line 1527, in _call_impl return forward_call(*args, **kwargs) File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/modules/activation.py", line 1241, in forward attn_output, attn_output_weights = F.multi_head_attention_forward( File "/home2/harsha.vasamsetti/miniconda3/envs/slices/lib/python3.9/site-packages/torch/nn/functional.py", line 5318, in multi_head_attention_forward raise RuntimeError(f"The shape of the 2D attn_mask is {attn_mask.shape}, but should be {correct_2d_size}.") RuntimeError: The shape of the 2D attn_mask is torch.Size([875, 875]), but should be (16, 16).
我查了PyTorch文档,文档说tgt_mask的形状应该是(T,T)(T是序列长度),tgt_key_padding_mask的形状应该是(batch_size,T),从错误日志里能看到我传入的形状完全符合要求,但还是报错,不知道问题出在哪。
我的实现代码如下:
import torch.nn as nn import torch import math # Define the Transformer model class class TransformerModel(nn.Module): def __init__(self, vocab_size,pad_idx, n_embd, n_head, n_layers, max_length, dropout=0.1): super().__init__() self.pad_idx = pad_idx # Add this line self.embed = nn.Embedding(vocab_size, n_embd) decoder_layer = nn.TransformerDecoderLayer(d_model=n_embd, nhead=n_head, dropout=dropout) self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_layers=n_layers) self.pos_encoder = PositionalEncoding(n_embd, dropout, max_length) self.n_embd = n_embd self.generator = nn.Linear(n_embd, vocab_size) def generate_square_subsequent_mask(self, sz): mask = (torch.triu(torch.ones(sz, sz)) == 1).transpose(0, 1) mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0)) return mask def forward(self, tgt): tgt_emb = self.pos_encoder(self.embed(tgt) * math.sqrt(self.n_embd)) # (batch_size, seq_len, emb_dim) print("tgt_emb", tgt_emb.shape) tgt_mask = self.generate_square_subsequent_mask(tgt.size(1)).to(tgt.device) # Adjusted to tgt.size(1) for seq_len print("tgt_mask", tgt_mask.shape) # Create padding mask based on EOS token used for padding if self.pad_idx is not None: padding_mask = (tgt == self.pad_idx) # (batch_size, seq_len) print("padding_mask", padding_mask.shape) else: padding_mask = None memory = torch.zeros_like(tgt_emb) # Simplified memory initialization output = self.transformer_decoder(tgt_emb, memory, tgt_mask=tgt_mask, tgt_key_padding_mask=padding_mask) return self.generator(output.transpose(0, 1)) # Adjust generator input if necessary # Positional encoding class adds information about the order of tokens class PositionalEncoding(nn.Module): def __init__(self, d_model, dropout=0.1, max_len=5000): super().__init__() self.dropout = nn.Dropout(p=dropout) position = torch.arange(max_len).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2) * -(math.log(10000.0) / d_model)) pe = torch.zeros(max_len, 1, d_model) pe[:, 0, 0::2] = torch.sin(position * div_term) pe[:, 0, 1::2] = torch.cos(position * div_term) self.register_buffer('pe', pe) def forward(self, x): x = x + self.pe[:x.size(0)] return self.dropout(x)
问题原因
PyTorch的TransformerDecoder默认期望输入序列的维度顺序是(seq_len, batch_size, emb_dim),但你的代码里tgt_emb的形状是(16, 875, 256),也就是(batch_size, seq_len, emb_dim),维度顺序搞反了。
这就导致模型把batch_size=16当成了序列长度,把seq_len=875当成了batch size,所以才会报错要求attn_mask的形状是(16,16),而不是你传入的(875,875)。
另外还有两个小问题:
- 位置编码的
forward方法里,你用x.size(0)取的是batch size,不是序列长度,加位置编码会出错 - 最后返回时
output.transpose(0,1)是把(seq_len,batch_size,emb_dim)转成(batch_size,seq_len,emb_dim),这个是对的,但前面的输入维度要先调整。
解决方法
修改代码中的几个关键位置:
- 调整输入序列的维度顺序:在生成
tgt_emb后,转置成(seq_len, batch_size, emb_dim) - 修正位置编码的索引:转置后
x的形状是(seq_len,batch_size,emb_dim),用x.size(0)获取序列长度即可 - (可选)可以直接用
tgt_is_causal=True替代手动生成的tgt_mask,PyTorch会自动生成正确的下三角掩码,更简洁。
修改后的代码如下:
import torch.nn as nn import torch import math # Define the Transformer model class class TransformerModel(nn.Module): def __init__(self, vocab_size,pad_idx, n_embd, n_head, n_layers, max_length, dropout=0.1): super().__init__() self.pad_idx = pad_idx self.embed = nn.Embedding(vocab_size, n_embd) decoder_layer = nn.TransformerDecoderLayer(d_model=n_embd, nhead=n_head, dropout=dropout) self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_layers=n_layers) self.pos_encoder = PositionalEncoding(n_embd, dropout, max_length) self.n_embd = n_embd self.generator = nn.Linear(n_embd, vocab_size) def generate_square_subsequent_mask(self, sz): mask = (torch.triu(torch.ones(sz, sz)) == 1).transpose(0, 1) mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0)) return mask def forward(self, tgt): # tgt形状: (batch_size, seq_len) tgt_emb = self.embed(tgt) * math.sqrt(self.n_embd) # (batch_size, seq_len, emb_dim) # 转置为TransformerDecoder期望的(seq_len, batch_size, emb_dim) tgt_emb = tgt_emb.transpose(0, 1) tgt_emb = self.pos_encoder(tgt_emb) print("tgt_emb", tgt_emb.shape) # 现在应该是(875,16,256) # 生成掩码,此时sz是seq_len=875 tgt_mask = self.generate_square_subsequent_mask(tgt.size(1)).to(tgt.device) print("tgt_mask", tgt_mask.shape) # (875,875) # padding_mask形状保持(batch_size, seq_len)不变,Transformer会自动处理 if self.pad_idx is not None: padding_mask = (tgt == self.pad_idx) print("padding_mask", padding_mask.shape) # (16,875) else: padding_mask = None memory = torch.zeros_like(tgt_emb) # memory形状要和tgt_emb一致: (seq_len, batch_size, emb_dim) # 传入参数,或者直接用tgt_is_causal=True替代tgt_mask output = self.transformer_decoder(tgt_emb, memory, tgt_mask=tgt_mask, tgt_key_padding_mask=padding_mask) # 转置回(batch_size, seq_len, emb_dim)给generator return self.generator(output.transpose(0, 1)) # Positional encoding class adds information about the order of tokens class PositionalEncoding(nn.Module): def __init__(self, d_model, dropout=0.1, max_len=5000): super().__init__() self.dropout = nn.Dropout(p=dropout) position = torch.arange(max_len).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2) * -(math.log(10000.0) / d_model)) pe = torch.zeros(max_len, 1, d_model) pe[:, 0, 0::2] = torch.sin(position * div_term) pe[:, 0, 1::2] = torch.cos(position * div_term) self.register_buffer('pe', pe) def forward(self, x): # x现在是(seq_len, batch_size, emb_dim),x.size(0)是序列长度 x = x + self.pe[:x.size(0)] return self.dropout(x)
如果想用tgt_is_causal=True简化代码,可以把tgt_mask的生成和传入去掉,改成:
output = self.transformer_decoder(tgt_emb, memory, tgt_is_causal=True, tgt_key_padding_mask=padding_mask)
这样PyTorch会自动生成正确的下三角掩码,避免手动生成可能的维度错误。
内容的提问来源于stack exchange,提问作者harsh
相关产品推荐
相关产品推荐

