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

PyTorch中TransformerEncoder传入注意力掩码的报错及解决方法

问题根源

PyTorch的nn.TransformerEncoder默认要求输入张量的维度顺序为**[序列长度, batch大小, 嵌入维度]**,但你传入的latent是[batch大小, 序列长度, 嵌入维度],这直接导致所有形状匹配逻辑错位,引发掩码不兼容的报错。

错误细节解析

  1. 第一次报错:
    你传入的latent是[8, 320, 512],PyTorch会错误地将第一个维度(8)识别为序列长度,第二个维度(320)识别为batch大小。此时它期望2D注意力掩码的形状是[序列长度, 序列长度]即[8,8],但你传入的是[320,320],形状不匹配导致报错。

  2. 第二次报错:
    你将掩码改为[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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 00:05:00