基于《Attention is All You Need》的Multihead Attention实现疑问
根据《Attention is All You Need》论文:
我们发现,与其使用单一的、基于dmodel维度键、值和查询的注意力函数,不如通过不同的、可学习的线性投影,分别将查询、键和值进行h次投影,得到dk、dk和dv维度的结果,这样做更为有效。
我的理解是,应该设置n_heads个不同的线性层,让每个头学习不同的投影。因此我按如下逻辑实现:
import torch import torch.nn as nn import math class Attention(nn.Module): def __init__(self, embed_size=512, out_feat=64) -> None: super().__init__() self.embed_size = embed_size self.out_feat = out_feat self.value_fc = nn.Linear(embed_size, out_feat) self.query_fc = nn.Linear(embed_size, out_feat) self.key_fc = nn.Linear(embed_size, out_feat) def forward(self, value, key, query, mask=None) -> torch.Tensor: """ value: torch.Tensor of shape (N, seq_len, embed_size) key: torch.Tensor of shape (N, seq_len, embed_size) query: torch.Tensor of shape (N, seq_len, embed_size) mask: torch.Tensor of shape (N, seq_len, seq_len) returns torch.Tensor of shape (N, seq_len, out_feat) """ value = self.value_fc(value) # N, seq_len, out_feat key = self.key_fc(key) query = self.query_fc(query) weights = torch.bmm(query, torch.transpose(key, 1, 2)) weights /= math.sqrt(self.out_feat) # N, query_len, key_len if mask is not None: weights += mask weights = torch.softmax(weights, dim=2) return torch.bmm(weights, value) class MultiHeadAttention(nn.Module): def __init__(self, embed_size=512, n_heads=8) -> None: super().__init__() assert embed_size % n_heads == 0, "Input feat. dim must be div. by n_heads" self.embed_size = embed_size self.n_heads = n_heads self.out_feat = embed_size // n_heads self.attention_layers = nn.ModuleList( [Attention(self.embed_size, self.out_feat) for _ in range(self.n_heads)] ) self.fc = nn.Linear(embed_size, embed_size) def forward(self, value, key, query, mask=None): values = [ attention(value, key, query, mask) for attention in self.attention_layers ] return self.fc(torch.cat(values, dim=2))
但我看到的所有(官方)实现都采用单一Attention层,将键、查询和值重塑为(batch_size, len, n_heads, d_model // n_heads)(计算注意力权重前先转置)。这种方式计算效率更高,但每个键、查询和值仅对应一个线性层,而非n_heads个线性层。我认为这与论文描述矛盾,虽然后者会大幅增加可学习参数,但按论文定义,我的实现才是正确的吗?
你的实现和官方高效实现在数学上完全等价,不存在谁更符合论文定义的问题,只是实现形式不同,后者是前者的高效优化版本。
参数数量一致
你的实现中,每个Attention头包含3个线性层,每个线性层参数为embed_size * out_feat + out_feat,n_heads个头部的总参数为n_heads * 3*(embed_size*out_feat + out_feat)。
官方实现通常用3个大线性层(如nn.Linear(embed_size, embed_size)),每个大层参数为embed_size*embed_size + embed_size。由于embed_size = n_heads * out_feat,代入后总参数为3*(embed_size*(n_heads*out_feat) + n_heads*out_feat) = 3*n_heads*(embed_size*out_feat + out_feat),和你的实现参数数量完全相同。
本质上,官方实现的大线性层就是把n_heads个小线性层的参数拼接在一起,一次性完成所有头部的投影,避免循环调用多个小层,大幅提升计算效率。
计算逻辑等价
官方实现的核心步骤:
- 用单个大线性层对query投影,得到
(N, seq_len, embed_size),重塑为(N, seq_len, n_heads, out_feat)后转置为(N, n_heads, seq_len, out_feat) - key和value执行相同操作
- 在头部维度上并行计算注意力权重
- 将结果转置回原维度,拼接为
(N, seq_len, embed_size)后通过最终线性层
这个过程和你每个头单独计算、再拼接的结果完全一致,只是把循环操作合并为张量维度变换和批量计算,充分利用GPU并行能力,速度更快。
契合论文描述
论文中“不同的、可学习的线性投影”的要求,两种实现都满足:官方实现的大线性层里,每个头对应的参数是独立的,和你每个头用单独线性层的参数没有区别,仅存储和计算方式不同。
综上,你的实现是正确的,但官方实现是更高效的等价方案,实际工程中优先选择后者。
内容的提问来源于stack exchange,提问作者The Exile

