自定义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
相关产品推荐
相关产品推荐

