Appearance
多头注意力
单头注意力只能从一个角度理解词与词的关系。但语言里的关系是多维的:语法依赖、指代关系、语义相关……都想同时捕捉。多头注意力(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 是其省显存的变体。
下一步
注意力负责「跨位置混合信息」,那每个位置内部怎么加工?→ 前馈网络与残差连接。