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

自定义Attention实现与torch.nn.MultiheadAttention输出不匹配求助

问题原因

PyTorch的MultiheadAttention并非直接使用你传入的原始Q、K、V进行注意力计算,它内部会先对输入的Q、K、V分别做线性投影——也就是通过三个共享的全连接层(对应权重in_proj_weight和偏置in_proj_bias),得到真正用于注意力计算的Q'、K'、V'。而你的自定义attention函数直接用了原始的Q、K、V,没有经过这一步投影,所以权重和输出必然不匹配。

修正方案

方案1:在自定义函数中加入线性投影(对齐MultiheadAttention逻辑)

直接提取MultiheadAttention的内置投影权重,用它处理输入的Q、K、V后再执行注意力计算:

import torch
import torch.nn.functional as F
from torch.nn import MultiheadAttention

def attention_with_proj(Q, K, V, proj_weight, proj_bias):
    # 对Q、K、V做线性投影,和MultiheadAttention内部逻辑一致
    Q_proj = F.linear(Q, proj_weight, proj_bias)
    K_proj = F.linear(K, proj_weight, proj_bias)
    V_proj = F.linear(V, proj_weight, proj_bias)
    
    d_k = Q_proj.size(-1)
    scores = torch.matmul(Q_proj, K_proj.transpose(-2, -1)) / (d_k**0.5)
    attn_output_weights = F.softmax(scores, dim=-1)
    attn_output = torch.matmul(attn_output_weights, V_proj)
    return attn_output, attn_output_weights

embed_dim = 8
num_heads = 1
batch_size = 2
seq_len = 5

Q = torch.randn(batch_size, seq_len, embed_dim)
K = torch.randn(batch_size, seq_len, embed_dim)
V = torch.randn(batch_size, seq_len, embed_dim)

multihead_attn = MultiheadAttention(embed_dim=embed_dim, num_heads=num_heads, batch_first=True)
attn_output_pytorch, attn_output_weights_pytorch = multihead_attn(Q, K, V)

# 获取MultiheadAttention的投影权重和偏置
proj_weight = multihead_attn.in_proj_weight
proj_bias = multihead_attn.in_proj_bias

attn_output_custom, attn_output_weights_custom = attention_with_proj(Q, K, V, proj_weight, proj_bias)

# 此时断言会通过
assert torch.allclose(attn_output_custom, attn_output_pytorch, rtol=1e-6, atol=1e-8), "Attention output does not match."
assert torch.allclose(attn_output_weights_custom, attn_output_weights_pytorch, rtol=1e-6, atol=1e-8), "Attention weights do not match."

方案2:用投影后的张量验证注意力计算逻辑

如果只是想验证自己的softmax加权求和逻辑是否正确,可以先手动计算投影后的Q'、K'、V',再传入自定义函数:

# 手动计算MultiheadAttention内部的投影
Q_proj = F.linear(Q, multihead_attn.in_proj_weight, multihead_attn.in_proj_bias)
K_proj = F.linear(K, multihead_attn.in_proj_weight, multihead_attn.in_proj_bias)
V_proj = F.linear(V, multihead_attn.in_proj_weight, multihead_attn.in_proj_bias)

# 用投影后的张量调用原自定义attention函数
attn_output_custom, attn_output_weights_custom = attention(Q_proj, K_proj, V_proj)

# 此时断言会通过
assert torch.allclose(attn_output_custom, attn_output_pytorch, rtol=1e-6, atol=1e-8)
assert torch.allclose(attn_output_weights_custom, attn_output_weights_pytorch, rtol=1e-6, atol=1e-8)

内容的提问来源于stack exchange,提问作者user26811297

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 13:32:06