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

基于《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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 14:20:24