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

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),这个是对的,但前面的输入维度要先调整。

解决方法

修改代码中的几个关键位置:

  1. 调整输入序列的维度顺序:在生成tgt_emb后,转置成(seq_len, batch_size, emb_dim)
  2. 修正位置编码的索引:转置后x的形状是(seq_len,batch_size,emb_dim),用x.size(0)获取序列长度即可
  3. (可选)可以直接用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 07:30:54