Attention 机制的进化:从 Scaled Dot-Product 到 Multi-Head Attention
一、算法概览
Attention 机制本质上是一种"软寻址 + 加权聚合"操作:给定一个查询(Query),在一组键值对(Key-Value)中,根据 Query 与每个 Key 的匹配程度,对所有 Value 做加权求和。从最初的 Bahdanau Attention(2015)到 Transformer 的 Scaled Dot-Product Attention(2017),再到今天大模型普遍采用的 Multi-Head Attention 和 Causal Attention,Attention 的进化史就是一部"让模型看得更全、更准、更快"的优化史。
二、历史背景与问题起源
Attention 的诞生源于 Seq2Seq 模型的一个根本瓶颈:信息瓶颈(Information Bottleneck)。
在传统 Seq2Seq 中,Encoder 把整个输入序列压缩成一个固定长度的上下文向量,Decoder 只能从这个向量中解码。当输入序列很长时,这个"压缩包"必然丢失大量信息。
Bahdanau 等人(2015)的洞见是:为什么只给 Decoder 一个全局摘要?让它自己决定每次该看输入的哪些部分!
于是诞生了 Additive Attention:
e_ij = v^T · tanh(W_a · s_i-1 + U_a · h_j)
α_ij = softmax(e_ij)
c_i = Σ_j α_ij · h_j
Decoder 每一步都动态计算对 Encoder 所有隐藏状态的注意力权重,然后加权求和得到上下文向量。信息瓶颈被打破了。
但这里有个效率问题:tanh(W·s + U·h) 这个加法操作引入了额外的参数矩阵,计算量大且不利于并行。于是 Transformer 提出了更优雅的方案。
三、核心原理深入
3.1 Scaled Dot-Product Attention:化繁为简
Transformer 的作者们做了一个大胆的简化:扔掉 tanh,扔掉加法,直接用点积衡量相似度。
Attention(Q, K, V) = softmax( Q·K^T / √d_k ) · V
为什么点积就够了?直觉上:
- Query 向量代表"我想要什么信息"
- Key 向量代表"我能提供什么信息"
- 它们的点积
Q·K^T天然就是余弦相似度的缩放版——点积越大,方向越一致,匹配度越高
相比 Additive Attention,Dot-Product Attention 的优势是纯粹的矩阵乘法,GPU 友好到极致。
Scale 因子 √d_k 的必要性:当 d_k 很大时,点积的方差 ≈ d_k,softmax 输入过大导致输出趋向 one-hot,梯度 ≈ 0。除以 √d_k 是维持梯度健康的数学必然,不是经验技巧。
3.2 Multi-Head Attention:多角度并行理解
单头 Attention 的致命缺陷:一个 Query 只能表达一种"查找意图"。
考虑句子 "The cat sat on the mat because it was tired",这里的 "it" 需要同时关注 "cat"(语义指代)和 "sat"(句法结构)。单头 Attention 的权重分布只能是一个折中。
Multi-Head Attention 的解决方案:把 Q、K、V 投影到 h 个不同的低维子空间,每个子空间独立计算 Attention,最后拼接。
head_i = Attention( Q·W_Qi, K·W_Ki, V·W_Vi )
MultiHead = Concat(head_1, ..., head_h) · W_O
其中每个 W_Qi、W_Ki、W_Vi 将原始 d_model 维投影到 d_k = d_model / h 维。
类比:一个评审委员会,有人管语法正确性,有人管事实准确性,有人管逻辑连贯性,有人管风格一致性——每个人从不同角度打分,主席(W_O)综合所有人的意见给出最终评分。
实践中,h=8 或 h=32 是常见选择。LLaMA-7B 用 32 个头,每个头 d_k = 128(4096/32)。
3.3 Causal Attention(Masked Self-Attention)
Decoder 在自回归生成时,位置 t 的 token 不能看到 位置 t+1 及之后的 token(否则就是"作弊")。
实现方式:在 softmax 之前,将注意力矩阵的上三角部分置为 -∞,softmax 后这些位置的权重自然归零。
Masked_Attention(Q, K, V) = softmax( (Q·K^T / √d_k) + M ) · V
其中 M[i][j] = 0 (i ≥ j)
= -∞ (i < j)
3.4 Cross-Attention vs Self-Attention
| 类型 | Q 来源 | K、V 来源 | 用途 |
|---|---|---|---|
| Self-Attention | 自身序列 | 自身序列 | Encoder/Decoder 内部建模 |
| Cross-Attention | Decoder 当前状态 | Encoder 输出 | Decoder 从 Encoder 提取信息 |
| Causal Attention | 自身序列(Masked) | 自身序列(Masked) | Decoder 自回归生成 |
现代 GPT 类模型(Decoder-only)只保留了 Causal Self-Attention,不再使用 Cross-Attention 和独立的 Encoder。这也是为什么 GPT 被称为"Decoder-only"——它只有带 Causal Mask 的 Self-Attention。
四、架构流程
┌──────────────────────────────────────────────────┐
│ Multi-Head Attention 流程 │
├──────────────────────────────────────────────────┤
│ │
│ 输入 X (n × d_model) │
│ │ │
│ ├─→ Q = X·W_Q ─┐ │
│ ├─→ K = X·W_K ─┤ │
│ └─→ V = X·W_V ─┤ │
│ │ │
│ ┌───────────────▼───────────────┐ │
│ │ 拆分为 h 个头 (n × d_k) │ │
│ └───────────────┬───────────────┘ │
│ │ │
│ ┌───────────────┼───────────────┐ │
│ ▼ ▼ ▼ │
│ head_1 head_2 ... head_h │
│ scores scores scores │
│ Q_1·K_1^T/√d_k Q_2·K_2^T/√d_k │
│ │ │ │ │
│ ▼ ▼ ▼ │
│ softmax softmax softmax │
│ │ │ │ │
│ ▼ ▼ ▼ │
│ out_1 = out_2 = out_h = │
│ scores·V_1 scores·V_2 scores·V_h │
│ │ │ │ │
│ └───────────────┼───────────────┘ │
│ ▼ │
│ Concat + W_O 投影 │
│ │ │
│ ▼ │
│ 输出 (n × d_model) │
└──────────────────────────────────────────────────┘
五、与其他算法的关系
Attention 机制是大模型技术栈中优化最密集的模块:
| 优化方向 | 代表技术 | 解决的问题 |
|---|---|---|
| 计算效率 | Flash Attention | 减少 HBM 读写,用 SRAM 做分块计算 |
| 内存占用 | KV Cache | 缓存已计算的 K、V,避免重复计算 |
| 长序列 | Sparse Attention、Sliding Window | 将 O(n²) 降到 O(n·k) 或 O(n·log n) |
| 位置感知 | RoPE、ALiBi | 在 Attention 计算中注入相对位置信息 |
| 多查询 | MQA、GQA | 多个头共享 K、V 以减少 KV Cache |
其中 Flash Attention 和 KV Cache 是现代大模型推理加速的绝对核心。
六、关键论文
| # | 论文 | 贡献 |
|---|---|---|
| 1 | "Neural Machine Translation by Jointly Learning to Align and Translate" — Bahdanau et al., ICLR 2015 (arXiv:1409.0473) | 首次提出 Attention 机制用于 Seq2Seq |
| 2 | "Attention Is All You Need" — Vaswani et al., NeurIPS 2017 (arXiv:1706.03762) | 提出 Scaled Dot-Product Attention 和 Multi-Head Attention |
| 3 | "FlashAttention: Fast and Memory-Efficient Exact Attention" — Dao et al., NeurIPS 2022 (arXiv:2205.14135) | IO-Aware 的 Attention 实现,2-4x 加速 |
七、实践思考
三个常见理解误区:
| 误区 | 正解 |
|---|---|
| ❌ "Multi-Head 就是多算几次 Attention 取平均" | ✅ 每个头在不同的低维子空间独立计算,最后拼接(不是平均),再投影回原空间 |
| ❌ "Causal Mask 是训练时防止过拟合的正则化手段" | ✅ 它是自回归生成的定义性约束——生成第 t 个 token 时不能偷看第 t+1 个,是因果性的必然要求 |
| ❌ "去掉 Scale 因子 √d_k 只是数值不稳定而已" | ✅ 去掉后 softmax 会饱和为 one-hot,梯度趋近于零,模型根本学不动——这不是"不稳定",是"学不了" |