PyTorch MultiheadAttention自定义计算结果与官方输出不匹配求助
问题解决:手动复现PyTorch MultiheadAttention输出不匹配
你的手动计算代码存在两个核心错误:线性层矩阵乘法顺序颠倒,以及部分细节维度处理不严谨,以下是修正方案:
错误原因
- PyTorch线性层的计算逻辑:
nn.Linear的权重形状为(out_features, in_features),实际计算是输入 @ 权重.t(),而非输入 @ 权重,你在Q/K/V投影和最终输出投影时都搞反了顺序。 - 注意力softmax维度的通用性:虽然当前场景下
dim=1可用,但更严谨的写法是dim=-1(针对最后一维,即key序列维度)。
修正后的手动计算代码
import torch import torch.nn as nn # 初始化原模型和输入 query = torch.randn(2, 4) key = torch.randn(2, 4) value = torch.randn(2, 4) model = nn.MultiheadAttention(4, 1, bias=False) # 官方模型输出 official_output, _ = model(query, key, value) official_output = official_output.squeeze() # 去掉batch维度,变为(2,4) # 手动复现步骤 # 1. 拆分投影权重 q_proj_weight = model.in_proj_weight[:4] k_proj_weight = model.in_proj_weight[4:8] v_proj_weight = model.in_proj_weight[8:12] out_proj_weight = model.out_proj.weight # 2. 计算Q/K/V投影(匹配nn.Linear的计算逻辑:输入 @ 权重.t()) Q = query @ q_proj_weight.t() K = key @ k_proj_weight.t() V = value @ v_proj_weight.t() # 3. 计算注意力分数与softmax d_k = Q.size(-1) attn_scores = Q @ K.t() / torch.sqrt(torch.tensor(d_k, dtype=torch.float32)) attn_weights = torch.softmax(attn_scores, dim=-1) # 4. 计算注意力输出与最终投影 attn_output = attn_weights @ V final_output = attn_output @ out_proj_weight.t() # 验证结果一致性(浮点数精度范围内相等) print(torch.allclose(final_output, official_output, atol=1e-6)) # 应输出True
额外说明
- 当
num_heads>1时,还需要对Q/K/V进行分头拼接的处理,但你这里num_heads=1,无需额外操作。 - PyTorch的
MultiheadAttention默认batch_first=False,输入会被解析为(seq_len, batch_size, embed_dim),你传入的2D张量会自动扩展为(2,1,4),所以手动计算时用2D张量处理后,只需挤压官方输出的batch维度即可对比。
内容的提问来源于stack exchange,提问作者apostofes
相关产品推荐
相关产品推荐

