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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 17:07:50