Skip to content

多头注意力

单头注意力只能从一个角度理解词与词的关系。但语言里的关系是多维的:语法依赖、指代关系、语义相关……都想同时捕捉。多头注意力(Multi-Head Attention) 的做法是:把向量切成若干份,每份(每个「头」)独立做一次注意力,再拼回来。

1. 一个直观类比

同一句话,不同"专家"同时看:
  头1 可能学会了关注"主谓关系"
  头2 可能学会了关注"代词指代"
  头3 可能学会了关注"相邻短语"
  头4 可能学会了关注"标点/句界"
最后把这些不同视角的结论合并 → 更立体的理解

就像开会时请了好几位不同背景的顾问,各看各的门道,最后汇总意见。

2. 它是怎么「切」和「拼」的

关键:多头并不是把参数量放大 h 倍,而是把 d_model 维度均分给 h 个头,每个头在低维子空间里做注意力。

d_model = 512,  h = 8 个头
 → 每个头维度 d_head = 512 / 8 = 64

流程:
 X ──W_Q──► Q[512] ─┐
                    ├─ 切成 8 段,每段 64 维 → 8 个头各算注意力
 X ──W_K──► K[512] ─┘        ↓
                     每个头输出 64 维

                     8 个头拼接 → 512 维

                     再 × W_O 融合 → 最终输出 512 维

3. 公式

[ \text{MultiHead}(Q,K,V) = \text{Concat}(head_1, \dots, head_h),W_O ] [ head_i = \text{Attention}(QW_Q^{(i)},; KW_K^{(i)},; VW_V^{(i)}) ]

其中每个头有自己的投影矩阵 (W_Q^{(i)}, W_K^{(i)}, W_V^{(i)} \in \mathbb{R}^{d_{model}\times d_{head}})。

4. PyTorch 完整实现

python
import torch
import torch.nn as nn
import math

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, num_heads=8):
        super().__init__()
        assert d_model % num_heads == 0
        self.h = num_heads
        self.d_k = d_model // num_heads          # 每个头的维度
        self.W_q = nn.Linear(d_model, d_model)   # 一次线性变换=所有头的投影打包
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)
        self.W_o = nn.Linear(d_model, d_model)   # 拼接后的融合矩阵

    def split(self, x):
        # [B, T, d_model] -> [B, h, T, d_k]
        B, T, _ = x.shape
        return x.view(B, T, self.h, self.d_k).transpose(1, 2)

    def forward(self, q, k, v, mask=None):
        B, T, _ = q.shape
        Q, K, V = self.split(self.W_q(q)), self.split(self.W_k(k)), self.split(self.W_v(v))
        # 缩放点积注意力(批量, 每个头独立)
        scores = Q @ K.transpose(-2, -1) / math.sqrt(self.d_k)   # [B,h,T,T]
        if mask is not None:
            scores = scores.masked_fill(mask == 0, float("-inf"))
        attn = torch.softmax(scores, dim=-1)
        out = attn @ V                                            # [B,h,T,d_k]
        # 拼接: [B,h,T,d_k] -> [B,T,d_model]
        out = out.transpose(1, 2).contiguous().view(B, T, -1)
        return self.W_o(out)

mha = MultiHeadAttention(512, 8)
x = torch.randn(2, 10, 512)            # batch=2, seq=10
mask = torch.tril(torch.ones(10, 10)).unsqueeze(0).unsqueeze(0)  # 因果掩码
print(mha(x, x, x, mask).shape)        # torch.Size([2, 10, 512])

自注意力里 q=k=v=x;交叉注意力里 q 来自解码器、k/v 来自编码器——多头机制完全复用,只是换了 Q/K/V 的来源。

5. 多头到底好在哪

对比单头多头
视角数量1 种关系h 种关系并行
子空间在完整 d_model 上一个 softmax各头在 d_head 子空间更专注
参数量一份 Q/K/V总量不变(维度被均分)
表达力强,实测效果显著更好

一个容易被忽略的点:多头几乎不增加参数量(还是那几个 d_model×d_model 的投影矩阵),却带来了多视角归纳能力,性价比极高。

6. 头数怎么选

  • 必须满足 d_model % num_heads == 0(能均分)。
  • 常见组合:d_model=512, h=8(原始论文 / GPT-2 small)、d_model=768, h=12(BERT base)。
  • 头太少 → 视角单一;头太多 → 每个头维度太小(d_k 很小),单头表达力不足。
  • 现代大模型为省 KV 显存,还会用 MQA/GQA(多 query 头共享少量 KV 头),属于「多头」的 Efficiency 变体,详见 高效注意力

7. 在整层中的位置

输入 ──► [多头自注意力] ──► Add&Norm ──► [前馈网络FFN] ──► Add&Norm ──► 输出
              ▲ 本页                              ▲ 下一页

小结

  • 多头 = 把 d_model 切成 h 份,各头独立做注意力再拼接融合。
  • 不显著增加参数,却获得多子空间、多关系的捕捉能力。
  • 自注意力与交叉注意力共用同一套多头机制,只是 Q/K/V 来源不同。
  • 头数需整除 d_model;MQA/GQA 是其省显存的变体。

下一步

注意力负责「跨位置混合信息」,那每个位置内部怎么加工?→ 前馈网络与残差连接