PyTorch TransformerEncoderLayer中attn_mask形状不匹配问题求助
PyTorch TransformerEncoderLayer 掩码形状报错的解决与说明
错误原因
你混淆了TransformerEncoderLayer中两个不同掩码参数的作用和形状要求:
- 你传入的第二个参数是
attn_mask(注意力掩码),它的作用是控制序列中每个token能否关注到其他token,因此形状必须是(seq_len, seq_len)(这里你的序列长度是16,所以需要(16,16))。 - 你误以为的“mask长度与token数量相等”对应的是
src_key_padding_mask(输入填充掩码),这个参数用于标记哪些token是无效的填充值,形状要求是(seq_len,)或(batch_size, seq_len),它是该层的第三个参数,而非第二个。
你的输入text = torch.randn(16,512)中,16是序列长度(seq_len),512是特征维度(embed_dim),所以attn_mask需要匹配序列长度的二维矩阵,而非特征维度。
修正代码
根据你的需求选择对应的掩码:
场景1:使用注意力掩码(控制token间的注意力可见性)
import torch import torch.nn as nn encoder_layers = nn.TransformerEncoderLayer(512, 8, 2048, 0.5) # 注意力掩码:形状为(seq_len, seq_len),控制i号token能否关注j号token attn_mask = torch.randint(0, 2, (16, 16)).bool() text = torch.randn(16, 512) # 格式:(seq_len, embed_dim) encoder_layers(text, attn_mask=attn_mask)
场景2:使用输入填充掩码(标记无效填充token)
如果你想屏蔽序列中的部分无效token,应该使用src_key_padding_mask参数:
import torch import torch.nn as nn encoder_layers = nn.TransformerEncoderLayer(512, 8, 2048, 0.5) # 输入填充掩码:形状为(seq_len,),标记哪些位置是无效token src_pad_mask = torch.randint(0, 2, (16,)).bool() text = torch.randn(16, 512) encoder_layers(text, src_key_padding_mask=src_pad_mask)
两类掩码的核心区别
attn_mask:二维矩阵,针对序列内token间的注意力关系进行细粒度控制,比如在Decoder中防止看到未来token,Encoder中可用于自定义注意力限制。src_key_padding_mask:一维(或二维批量)掩码,仅用于标记哪些位置是填充的无效token,这些位置会被整体排除在注意力计算之外。
内容的提问来源于stack exchange,提问作者Monsieur AZERTY
相关产品推荐
相关产品推荐

