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

PyTorch MultiheadAttention自定义计算结果与官方输出不匹配求助

问题解决:手动复现PyTorch MultiheadAttention输出不匹配

你的手动计算代码存在两个核心错误:线性层矩阵乘法顺序颠倒,以及部分细节维度处理不严谨,以下是修正方案:

错误原因

  1. PyTorch线性层的计算逻辑:nn.Linear的权重形状为(out_features, in_features),实际计算是输入 @ 权重.t(),而非输入 @ 权重,你在Q/K/V投影和最终输出投影时都搞反了顺序。
  2. 注意力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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 19:20:36