PyTorch中TransformerEncoder传入注意力掩码的报错及解决方法
PyTorch的nn.TransformerEncoder默认要求输入张量的维度顺序为**[序列长度, batch大小, 嵌入维度]**,但你传入的latent是[batch大小, 序列长度, 嵌入维度],这直接导致所有形状匹配逻辑错位,引发掩码不兼容的报错。
错误细节解析
第一次报错:
你传入的latent是[8, 320, 512],PyTorch会错误地将第一个维度(8)识别为序列长度,第二个维度(320)识别为batch大小。此时它期望2D注意力掩码的形状是[序列长度, 序列长度]即[8,8],但你传入的是[320,320],形状不匹配导致报错。第二次报错:
你将掩码改为[8,320,320]后,PyTorch依然基于错误的维度识别逻辑,认为序列长度是8、batch大小是320。若你的num_heads为8,那么3D注意力掩码需要满足[num_heads * batch大小, 序列长度, 序列长度]即[8*320,8,8] = [2560,8,8],和你传入的[8,320,320]完全不符,因此再次报错。
核心操作是调整输入张量的维度顺序,将latent转置为[序列长度, batch大小, 嵌入维度],再匹配对应形状的掩码即可。
情况1:全局统一的2D注意力掩码(所有样本共用同一掩码)
import torch import torch.nn as nn # 初始化TransformerEncoder(假设d_model=512,num_heads=8,num_layers=2) transformer_encoder = nn.TransformerEncoder( nn.TransformerEncoderLayer(d_model=512, nhead=8), num_layers=2 ) # 原始输入与掩码 latent = torch.rand(8, 320, 512) # [batch_size, seq_len, d_model] mask = torch.rand(320, 320) # [seq_len, seq_len] 全局注意力掩码 # 转置输入维度为Transformer要求的[seq_len, batch_size, d_model] latent = latent.transpose(0, 1) # 正确传入参数 output = transformer_encoder(latent, attn_mask=mask) print(output.shape) # 输出: torch.Size([320, 8, 512])
情况2:样本级的padding掩码(标记每个样本的padding位置)
如果你的掩码是用来屏蔽padding token的,应该使用src_key_padding_mask参数,形状为[batch_size, seq_len]:
# 生成padding掩码(1表示该位置是padding,会被屏蔽) src_key_padding_mask = torch.randint(0, 2, (8, 320)) # [batch_size, seq_len] # 传入padding掩码 output = transformer_encoder(latent, src_key_padding_mask=src_key_padding_mask)
情况3:每个样本独立的3D注意力掩码
若需要为每个样本的每个注意力头单独设置掩码,3D掩码的形状应为[num_heads * batch_size, seq_len, seq_len]:
num_heads = 8 # 生成3D注意力掩码 mask_3d = torch.rand(num_heads * 8, 320, 320) # [num_heads*batch_size, seq_len, seq_len] output = transformer_encoder(latent, attn_mask=mask_3d)
内容的提问来源于stack exchange,提问作者SAUMYA BHANDARY

