多头注意力两种实现方式输出不一致的问题排查求助
我正在学习Sebastian Raschka的《Build a Large Language Model (From Scratch)》,尝试实现书中提到的「多头注意力简单实现」和「替代实现」两种版本,但发现两者的输出结果完全不一致(按道理来说,两种实现的最终上下文向量应该是相同的)。为了简化示例,我暂时去掉了因果注意力的逻辑,下面是我的代码和输出结果,麻烦帮忙找出我哪里写错了?
简单实现代码及输出
#Simple Implementation import torch import torch.nn as nn torch.manual_seed(1223) batch = 1 num_tokens = 2 d_out = 6 num_heads = 3 head_dim = d_out // num_heads out_proj = nn.Linear(d_out*num_heads, d_out*num_heads) print (f'{batch = } {num_tokens = } {d_out = } {num_heads = } {head_dim = }') q = torch.rand(batch,num_tokens,d_out) k = torch.rand(batch,num_tokens,d_out) v = torch.rand(batch,num_tokens,d_out) print(f'{q.shape = } {k.shape = } {v.shape = }') context_vec = [] for _ in range(num_heads) : attn_scores1 = q @ k.transpose(1, 2) context_vec1 = attn_scores1 @ v context_vec.append(context_vec1) context_vec_Final = out_proj(torch.cat (context_vec, dim = -1)) print(f'{context_vec_Final.shape = } {context_vec_Final = }')
输出结果:
context_vec_Final.shape = torch.Size([1, 2, 18])
context_vec_Final = tensor([[[ 0.4481, 0.1600, 1.1673, -0.0284, -0.4112, -0.0324, -0.8480,
-0.6424, 1.0077, 0.2438, -0.2781, 0.6076, -0.4314, 0.4139,
0.6364, -0.7987, -0.1409, 1.0451],
[ 0.6571, 0.2290, 1.5520, -0.0847, -0.6139, -0.0287, -1.0902,
-0.9133, 1.2716, 0.2759, -0.2912, 0.8003, -0.5354, 0.4829,
0.8830, -0.9920, -0.1232, 1.3950]]], grad_fn=)
替代实现代码及输出
import torch torch.manual_seed(1223) batch = 1 num_tokens = 2 d_out = 6 num_heads = 3 head_dim = d_out // num_heads print (f'{batch = } {num_tokens = } {d_out = } {num_heads = } {head_dim = }') q = torch.rand(batch,num_tokens,d_out) k = torch.rand(batch,num_tokens,d_out) v = torch.rand(batch,num_tokens,d_out) print(f'{q.shape = } {k.shape = } {v.shape = }') print(f'{q = }') #The key operation is to split the d_out dimension into num_heads and head_dim using view option #(b, num_tokens, d_out) is reshaped to dimension (b, num_tokens, num_heads, head_dim) #d_out = num_heads * head_dim q = q.view(batch,num_tokens,num_heads,head_dim) k = k.view(batch,num_tokens,num_heads,head_dim) v = v.view(batch,num_tokens,num_heads,head_dim) print(f'{q.shape = } {k.shape = } {v.shape = }') #Now transpose (b, num_tokens, num_heads, head_dim) to (b, num_heads, num_tokens, head_dim) #Logically now we have num_heads of num_tokens x head_dim... #Now we have num_heads (multiple head ) of num_tokens x head_dim q = q.transpose(1, 2) k = k.transpose(1, 2) v = v.transpose(1, 2) print(f'{q.shape = } {k.shape = } {v.shape = }') attn_scores = q @ k.transpose(2, 3) print(f'{attn_scores.shape = } {attn_scores = }') context_vec = (attn_scores @ v) print(f'{context_vec.shape = }') #Now we are transposing (b, num_heads, num_tokens, head_dim) back to (b, num_tokens, num_heads, head_dim) context_vec = context_vec.transpose(1, 2) print(f'{context_vec.shape = }') context_vec = context_vec.contiguous().view(batch, num_tokens, d_out) print(f'{attn_scores.shape = } {context_vec.shape = }') print (f'{context_vec.shape = } {context_vec = }')
输出结果:
context_vec.shape = torch.Size([1, 2, 6])
context_vec = tensor([[[0.7713, 1.1682, 0.6141, 0.8472, 0.9745, 0.6424],
[0.3327, 0.6094, 0.1885, 0.2612, 0.6739, 0.3949]]])
问题分析与修正方案
我帮你梳理了两个核心错误点,以及修正后的代码:
错误点1:简单实现的head逻辑完全错误
你在循环中每次都用完整的q/k/v(维度[1,2,6])计算注意力,而不是拆分出每个head对应的子向量。这相当于重复计算了3次完全相同的注意力结果,然后拼接成[1,2,18]的向量,完全违背了多头注意力“每个head独立处理子空间信息”的设计逻辑。
错误点2:维度与输出投影不匹配
简单实现中out_proj的输入输出维度设为d_out*num_heads(18),而替代实现根本没有添加输出投影,最终输出维度是d_out(6),两者维度都不一致,结果自然不可能相同。
修正后的简单实现
# 修正后的简单实现 import torch import torch.nn as nn torch.manual_seed(1223) batch = 1 num_tokens = 2 d_out = 6 num_heads = 3 head_dim = d_out // num_heads # 修正out_proj维度:输入输出都是d_out(拼接后是num_heads*head_dim=d_out) out_proj = nn.Linear(d_out, d_out) print(f'{batch = } {num_tokens = } {d_out = } {num_heads = } {head_dim = }') q = torch.rand(batch, num_tokens, d_out) k = torch.rand(batch, num_tokens, d_out) v = torch.rand(batch, num_tokens, d_out) print(f'{q.shape = } {k.shape = } {v.shape = }') context_vec = [] # 将q/k/v拆分为每个head的子向量 q_split = q.split(head_dim, dim=-1) k_split = k.split(head_dim, dim=-1) v_split = v.split(head_dim, dim=-1) for h in range(num_heads): q_h = q_split[h] k_h = k_split[h] v_h = v_split[h] attn_scores1 = q_h @ k_h.transpose(1, 2) context_vec1 = attn_scores1 @ v_h context_vec.append(context_vec1) # 拼接后维度为[1,2,6],和替代实现对齐 context_vec_concat = torch.cat(context_vec, dim=-1) context_vec_Final = out_proj(context_vec_concat) print(f'{context_vec_Final.shape = } {context_vec_Final = }')
修正后的替代实现
# 修正后的替代实现 import torch import torch.nn as nn torch.manual_seed(1223) batch = 1 num_tokens = 2 d_out = 6 num_heads = 3 head_dim = d_out // num_heads # 添加和简单实现一致的输出投影层 out_proj = nn.Linear(d_out, d_out) print(f'{batch = } {num_tokens = } {d_out = } {num_heads = } {head_dim = }') q = torch.rand(batch, num_tokens, d_out) k = torch.rand(batch, num_tokens, d_out) v = torch.rand(batch, num_tokens, d_out) print(f'{q.shape = } {k.shape = } {v.shape = }') # 拆分维度并转置 q = q.view(batch, num_tokens, num_heads, head_dim).transpose(1, 2) k = k.view(batch, num_tokens, num_heads, head_dim).transpose(1, 2) v = v.view(batch, num_tokens, num_heads, head_dim).transpose(1, 2) print(f'{q.shape = } {k.shape = } {v.shape = }') attn_scores = q @ k.transpose(2, 3) context_vec = attn_scores @ v # 转置回原维度并拼接 context_vec = context_vec.transpose(1, 2).contiguous().view(batch, num_tokens, d_out) # 应用输出投影 context_vec_Final = out_proj(context_vec) print(f'{context_vec_Final.shape = } {context_vec_Final = }')
现在运行这两段修正后的代码,你会发现两者的context_vec_Final结果完全一致,符合多头注意力两种实现方式的预期。
备注:内容来源于stack exchange,提问作者KJG

