A2 注意力机制全景¶
更新日期:2026-04-26
注意力机制是 Transformer 的核心,也是架构演进中变化最大的模块。本节从结构而非时间维度拆解。
三条独立路线(不是时间演进,是结构差异)¶
注意力变体演化 不是单线 MHA → MLA → KDA → Mamba 的"代际更替",是 3 条结构上正交的路线,各自优化不同的复杂度维度:
flowchart LR
base["注意力 = softmax(QK^T) V<br/>朴素 O(N²) 计算 / O(N) KV cache"]
base --> full["1. Full 路线<br/>保留 softmax<br/>压缩 K/V 表示"]
base --> sparse["2. Sparse 路线<br/>保留 softmax<br/>跳过部分 (q,k) 对"]
base --> linear["3. Linear / RNN 路线<br/>砍 softmax<br/>固定大小 state 替代 KV"]
full --> mla["MHA / MQA / GQA / MLA / HCA"]
sparse --> sw["SWA / BigBird / DSA / CSA"]
linear --> kda["GLA / KDA / Lightning / Mamba-2"]
classDef root fill:#f5f3eb,stroke:#bdb9ab;
classDef path fill:#fff,stroke:#cc785c;
class base root
class full,sparse,linear,mla,sw,kda path
每条线根本约束不同,下面给出"核心 trade-off"对照:
| 路线 | 核心 trade-off | 复杂度 | 代价 |
|---|---|---|---|
| Full | KV cache 大小 ↔ 每头表达力 | KV ∝ \(N\)(dim 可压) | dim 压太狠丢精度 |
| Sparse | 计算量 ↔ 信息可达性 | \(N \cdot k\)(k = 看几个) | 远距离信息可能漏 |
| Linear | KV cache 常数 ↔ 历史摘要损失 | KV = \(O(d^2)\) 常数 | 状态会"覆盖"旧记忆 |
关键 insight:3 条线不是替代关系,frontier 全部用 hybrid(详 §四)。理解每条线的失败模式才知道为什么混搭。
→ 各路线深读:
- A2-1 Full 路线 — MHA / MQA / GQA / MLA / HCA
- A2-2 Sparse 路线 — SWA / BigBird / DSA / CSA + FlashAttention IO 优化 + SageAttention 量化
- A2-3 Linear / RNN 路线 — GLA / KDA / Lightning / SSM (Mamba-2)
一、复杂度速查表¶
设 \(N\) = 序列长度、\(H\) = 头数、\(d_h\) = 头维度、\(G\) = GQA 组数、\(d_c\) = MLA latent 维度、\(W\) = sliding window、\(k\) = top-k、\(m\) = 序列压缩比。
| 变体 | 路线 | KV cache / token / layer | Decode / step | Prefill | 长 ctx 增长 |
|---|---|---|---|---|---|
| MHA | Full | \(2 H d_h\) | \(O(N H d_h)\) | \(O(N^2 H d_h)\) | KV ∝ \(N\) |
| MQA | Full | \(2 d_h\) | \(O(N H d_h)\) | \(O(N^2 H d_h)\) | KV ∝ \(N\) |
| GQA-G | Full | \(2 G d_h\) | \(O(N H d_h)\) | \(O(N^2 H d_h)\) | KV ∝ \(N\) |
| MLA | Full | \(d_c + d_R\)(V3=576) | \(O(N d_c)\) 吸收后 | \(O(N^2 d_c)\) | KV ∝ \(N\) |
| HCA-m' | Full + 序列压缩 | \((d_c + d_R)/m'\) | \(O((N/m') d_h)\)/head | \(O((N/m')^2 d_h)\) | \(1/m'^2\) quadratic |
| SWA-W | Sparse | \(\le 2 W H d_h\)(cap) | \(O(W H d_h)\) | \(O(N W H d_h)\) | KV ≤ \(W\) |
| Block Sparse | Sparse | \(2 H d_h\) | \(O(\sqrt{N} H d_h)\) | \(O(N \sqrt{N} H d_h)\) | sub-quadratic |
| DSA-k | Sparse + latent | \(d_c + d_R\) + indexer | \(O(N d_\text{idx} + k d_h)\) | \(O(N^2 d_\text{idx})\) | indexer 仍 ∝ \(N\) |
| CSA-(m,k) | Sparse + 序列压缩 | \((d_c + d_R)/m\) | \(O((N/m) d_\text{idx} + k d_h)\) | \(O((N/m)^2 d_\text{idx})\) | sub-linear |
| GLA / Linear | Linear | \(d_h^2\) (常数) | \(O(d_h^2)\) | \(O(N d_h^2)\) chunk-wise | 不增长 |
| SSM / Mamba-2 | Linear | \(O(N_\text{state})\) 常数 | \(O(d_h N_\text{state})\) | chunk matmul | 不增长 |
关键 take-away:
- 唯一真正常数 KV 的是 Linear 路线(GLA / SSM)—— 1M / 10M ctx 不爆显存
- Sparse 路线 prefill 仍 quadratic(indexer 要扫所有 K),只在 decode 时拿到线性
- HCA 是"先压缩再 dense",复杂度仍 quadratic,但常数 \(1/m'^2\) 小到可接受
- MQA/GQA decode FLOPs 跟 MHA 一样多,省的是 HBM 带宽(KV 字节少 → 加载更快),decode 阶段是带宽 bound 所以体感快
二、attention mask 原理图¶
每行是一个 query token,每列是 key token,色块表示该 (q, k) 对参与 softmax 计算。设 \(N=8\),causal。
1. MHA / MLA — Full causal(所有合法 (q,k) 都算)
2. SWA W=3 — Sparse 静态窗口(每个 q 只看前 W 个)
3. DSA top-3 — Sparse 动态选择(每个 q 由 lightning indexer 选 k=3 个最相关)
4. CSA m=2, k=2 — 序列 2× 压缩 + top-k 选块
5. GLA / Linear — 没有 mask,全压进 state
▲ S 是 \(d_h \times d_h\) 固定大小矩阵,不存 KV 序列。任何历史 token 通过 S 间接影响输出(有损)。Linear 路线没有 attention 矩阵,只有 state matrix —— 机理可解释性跟 Full / Sparse 完全不同。
5 张 SVG 由
scripts/gen_attention_masks.py生成(matplotlib),改 mask 逻辑就重跑一次。原始 ASCII 版本保留在 git 历史里。
三、Full 路线一图直观对比(来自论文)¶
DeepSeek-V2 paper Figure 3 把 MHA / GQA / MQA / MLA 四个变体并排画了出来——这是看懂 Full 路线最快的方式:

来源:DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model (Liu et al., 2024), Figure 3。版权归原作者,本站仅作学习引用。
重点:
- MHA:每个 Q head 配独立的 K/V head(H 套)
- GQA:每 N 个 Q head 共享一套 K/V(G 套,G < H)
- MQA:所有 Q head 共享同一套 K/V(1 套)
- MLA:K/V 压成 latent vector \(c_{kv}\)(dim \(d_c\),单一),推理时通过 \(W^{UK}, W^{UV}\) 上投影恢复,只缓存 \(c_{kv}\) + \(k^R\)(RoPE 部分)
四者按"压什么"分类清楚:MHA → MQA/GQA 压 head 数;MLA 改压 dim(latent)。
四、KV cache 横向对比(按公式推算)¶
按"每 token KV cache 大小"和"是否随 ctx 增长"两轴。下表数字按定义公式直接算出,不是测试数据:
| 方法 | 路线 | KV / token | 是否随 ctx 增长 | 1M ctx 单 layer 相对 MHA | 代表 |
|---|---|---|---|---|---|
| Full MHA | Full | \(2 H d_h\) | ✓ | 1× | GPT-3 |
| GQA-8 | Full | \(2 \cdot 8 \cdot d_h\) | ✓ | ~6% | LLaMA-3 |
| MLA | Full latent | \(d_c + d_R\) | ✓ | ~3% | DeepSeek V3 |
| DSA | Sparse + latent | MLA + indexer | ✓ (sparse) | ~30% of MLA | DeepSeek V3.2 |
| CSA + HCA | Full + Sparse + 压缩 | ~MLA / 50 | ✓ (压缩 50×) | ~2% | DeepSeek V4 |
| Sliding Window | Sparse | \(2 H d_h\) cap by W | \(W\) 内 | W-cap | Mistral 7B (v0.1) |
| GLA / KDA / Mamba | Linear | \(d^2\) (固定) | ✗ | 不变 | Kimi Linear / Nemotron |
| Lightning | Linear | \(d^2\) (固定) | ✗ | 不变 | MiniMax M1 |
几个 pattern:
- Full 路线 都"随 ctx 增长",区别在斜率(MHA 最陡 → CSA+HCA 最缓)
- Sparse 路线 把"随 ctx 增长"改成"随 W cap"或 sub-linear
- Linear 路线 唯一"不随 ctx 增长"
- frontier 选择:1M+ ctx 必须放弃"随 ctx 增长" → 上 hybrid(含 Linear 层)
五、Hybrid 设计 — frontier 现实¶
3 条路线单纯都不够好,frontier 全部用 hybrid。每家选不同混合 ratio,背后是对各路线失败模式的具体判断。
5.1 为什么必须 hybrid¶
| 路线 | 短 ctx | 长 ctx | 推理成本 | 失败模式 |
|---|---|---|---|---|
| Full | ★★★★★ | ★★ | ★ | KV cache 爆,长 ctx 不可行 |
| Sparse | ★★★ | ★★★ | ★★★ | mask pattern 选错 = 漏关键 token |
| Linear | ★★ | ★★★★ | ★★★★★ | state 容量限制,精确 recall 弱 |
hybrid = 在不同 layer 用不同路线:让 Full layer 补 Linear 的"精确 recall",让 Linear / Sparse layer 担起"低成本长 ctx"。
5.2 frontier 4 种 hybrid 模式¶
flowchart LR
Q{"hybrid 模式"}
Q --> A["KDA : MLA = 3 : 1<br/>(Kimi Linear)<br/>跨 layer"]
Q --> B["Lightning : MHA = 7 : 1<br/>(MiniMax M1)<br/>跨 layer"]
Q --> C["Mamba : Attn = 92 : 8<br/>(Nemotron-H)<br/>跨 layer,最激进"]
Q --> D["CSA + HCA<br/>(DeepSeek V4)<br/>同 layer 内 local + global"]
classDef p fill:#fff,stroke:#cc785c;
class Q,A,B,C,D p
关键观察:
- Kimi (3:1) / M1 (7:1) / Nemotron (92:8) 都是 跨 layer hybrid —— 不同 layer 用不同 attention 类型
- DeepSeek V4 (CSA+HCA) 是 同 layer 内 hybrid —— Full 路线下"局部细粒度(CSA)+ 全局粗粒度(HCA)"
- 跨 layer 的好处:每 layer 工程实现独立,kernel 可以专门优化
- 同 layer 的好处:每 query 同时拿到细+粗信息,没有"被 Linear 层压扁"的风险
5.3 选 ratio 的原则¶
- 更激进 linear 比例(M1 7:1, Nemotron 92:8):长 ctx + 长 output 任务(reasoning, agent multi-turn)
- 更保守 linear 比例(Kimi 3:1):精度敏感任务(math, code)
- 同 layer hybrid(V4):训练复杂度高但单 layer 信息密度最高,质量最稳
→ 最终是 engineering trade-off + 任务特性决定 ratio,没有"最优 hybrid"。
六、微观 ↔ 宏观全谱对比¶
每条路线在 kernel dtype / 数值稳定性 / 激活模式 / KV 布局 / 训练动力学 5 个维度有截然不同的代价。
6.1 计算精度 & 累加¶
每变体在不同精度下的稳定性 + 推荐配置:
| 变体 | 推荐 forward | accumulator | FP8 训练? | FP4 推理? | 关键陷阱 |
|---|---|---|---|---|---|
| MHA / GQA | BF16 | FP32 | ✅ E4M3 (FlashAttn-3) | ✅ NVFP4 PTQ | softmax 必须 FP32 partial sum |
| MLA | BF16 | FP32 | ✅(DeepSeek V3 验证) | ✅ FP4 重训 | latent c_kv 在 FP8 下需 per-tile scaling |
| DSA | BF16 | FP32 | ✅ | 待 V3.2 之外验证 | top-k indexer 用 FP16 即可,主 attention 跟 MLA 一致 |
| CSA + HCA | FP4 master + FP8 compute (V4 native) | FP32 | ✅ 设计原生 | ✅ 设计原生 | 压缩器 pooling 在 FP4 下要 dequantize 到 FP8 |
| Sliding Window | BF16 / FP16 | FP32 | ✅ | ✅ | 同 MHA |
| GLA | BF16 | FP32 必须(state 矩阵 matmul 累积长程) | 学术原型,未规模验证 | ✗ | \(S_t\) 长程数值漂移,必须 FP32 累加器 |
| Mamba-2 | BF16 | FP32 | ✅ (SSD framework 解决了 Mamba-1 的不稳定) | 实验中 | \(A\) 矩阵跟踪 \(\log A\) 防 underflow;离散化用 softplus |
| KDA | BF16 | FP32 | 待验证 | ✗ | DPLR 转移矩阵 \((I - \beta vk^\top)\) 必须 FP32(rank-1 误差累积) |
| Lightning | BF16 | FP32 | 实验中 | ✗ | \(K^\top V\) 累加 d² 个数,量级随 N 漂;feature_map ELU+1 防止除 0 |
重点 1:FlashAttention BF16 的精度坑¶
Saturn et al., 2024 "Is Flash Attention Stable?" 指出:BF16 下 FlashAttention 比 baseline attention 多 ~10× 数值偏差。原因:
- tile 越多 → rescaling 累积越多 → 误差增长
- softmax 出现
exp(0) = 1时 normalization 常数除以重复 max 会"抹平" → bias 累积
修复(FlashAttention-3 + Sliding Window):动态调整 normalization 常数,避免 exp 输出精确等于 1。
重点 2:Linear / SSM 路线 state 漂移¶
GLA / KDA / Mamba-2 / Lightning 都有同一个根本问题:state \(S_t\) 是 \(\sum\) 累加形式,长 ctx 下数值会漂。Mamba-2 的解法:
- \(A\) 学 \(\log A\) 而非 \(A\)(保 \(A < 0\) 稳定)
- 离散化 \(\bar A = \exp(\Delta A)\) 用 softplus 限 \(\Delta > 0\)
- SSD chunk-wise 形式让 FP32 累加器只在 chunk 内累,跨 chunk reset
→ 这就是 Mamba-2 比 Mamba-1 "训练能 scale" 的工程核心。
6.2 数值稳定性陷阱(per-variant 失效模式)¶
| 路线 | 典型 NaN / 不稳定 来源 | 缓解 |
|---|---|---|
| Full (MHA/GQA) | softmax overflow(attention scores 太大) | safe softmax: softmax(x - max(x)) |
| Full (MLA) | latent up-project 后 norm 过大 | RMSNorm 包裹 + W_uk init 缩放 |
| Full (CSA/HCA) | pooling 时所有 token 同号 → 信息丢失 | softmax-gated weights + position bias |
| Sparse (SWA) | causal mask 边界 token attention 全 -inf → softmax NaN | 至少保留 self-attention 一项 |
| Sparse (Top-k) | top-k 选择不可微 | Gumbel-softmax / straight-through |
| Linear (GLA) | \(S_t\) 长 ctx 后某些 dim 爆炸 | gate \(g_t\) sigmoid 强制 < 1 |
| Linear (KDA) | \((I - \beta vk^\top)\) 在 \(\beta\) 大时矩阵接近奇异 | \(\beta\) 用 sigmoid 限到 [0, 1) |
| Linear (Mamba-2) | \(A\) 接近 0 时 \(\bar A = \exp(\Delta A) \approx 1\),state 永不衰减 | \(A\) 初值远离 0(HiPPO init) |
| Linear (Lightning) | feature_map 输出 0 → normalization 除 0 | ELU + 1(保正)+ epsilon |
6.3 激活内存(forward 保存的 tensor 大小)¶
训练时反向需要保 forward 中间,对显存占用影响巨大:
| 变体 | 主要 activation | 大小(B = batch, S = seq, H = head, d_h, d_c, d = d_model) | 1B params 7B 模型 32K ctx 估计 |
|---|---|---|---|
| MHA (no FlashAttn) | attention 矩阵 + V | \(B H S^2 + B S H d_h\) | 巨大:32K² × 32 × bs8 = 100+ GB |
| MHA (FlashAttn) | softmax stat + Q/K/V | \(B H S \cdot 2 + 3 B S H d_h\) | ~5 GB(不存矩阵) |
| GQA | 同 MHA-FlashAttn 但 KV 压缩 | \(B H S \cdot 2 + B S G d_h + B S H d_h\)(Q) | ~3 GB |
| MLA | softmax stat + c_kv + k_R + Q | \(B H S \cdot 2 + B S (d_c + d_R) + B S H d_h\) | ~2 GB |
| Sliding Window | 同 MHA-FlashAttn 但 mask 切片 | \(B H S \cdot 2 + 3 B S H d_h\) | ~5 GB |
| GLA | state \(S\) + intermediate gate | \(B H d_h^2 + B S H d_h\) | ~1 GB(state 不随 S 增长) |
| Mamba-2 | \(h\) state + Δ B C 中间 | \(B H d_{state} d_h + B S \cdot \text{const}\) | ~1.5 GB |
| KDA | \(S_t\) 矩阵 + DPLR 中间 | \(B H d_h^2 + B S \cdot 2 d_h\)(α/β) | ~1.5 GB |
Linear 路线 激活内存唯一不随 seq 长度增长,是它训长 ctx 的核心优势。
6.4 GPU kernel & 训练并行性¶
| 变体 | 训练时主要 op | 并行模式 | Tensor Core 利用 | Production kernel |
|---|---|---|---|---|
| MHA / GQA | matmul + softmax | 全序列并行 | ★★★★★(FlashAttn-⅔) | FlashAttention |
| MLA | matmul + softmax + 吸收 trick | 全序列并行 | ★★★★(W_uk 吸收后接近 GQA) | FlashMLA (DeepSeek 自家) |
| DSA | matmul + top-k indexer | 全序列并行 + 不可微 selection | ★★★(top-k 是 thread-level 操作) | 自研 |
| CSA / HCA | pooling + matmul + softmax | 全序列并行 | ★★★★ | 自研 |
| Sliding Window | matmul + sliding mask | 全序列并行 | ★★★★(FlashAttn 支持) | FlashAttention 内置 |
| GLA | chunk-wise scan | chunk 内并行 + chunk 间序列 | ★★★(chunk size 影响利用率) | fla |
| Mamba-2 | SSD chunk matmul | chunk 内并行 + chunk 间序列 | ★★★★(SSD 比 selective_scan 快 2-8×) | mamba_ssm selective_scan + SSD |
| KDA | DPLR chunk matmul | chunk 内并行 + chunk 间序列 | ★★★ | Moonshot 自家 KDA kernel |
| Lightning | \(K^\top V\) + chunk | chunk 内并行 | ★★★ | MiniMax 自家 |
几个观察:
- Full / Sparse 路线整段并行(FlashAttention 风格)→ Tensor Core 利用率最高
- Linear / SSM 路线本质上是扫描算法,必须 chunk-wise 拆才能并行 → 利用率永远略输 Full
- Mamba-2 SSD 是目前 Linear 路线训练效率最高的 framework,因为它把 SSM 写成纯 matmul(跟 Linear Attention 等价)
- Production kernel 是工程门槛:MLA/MOE/DPLR 等变体性能能不能跑出理论值,决定于 kernel 质量;糟糕的 kernel 让 MLA 比 GQA 还慢
6.5 训练动力学¶
| 变体 | LR 敏感度 | 训练稳定性 | 长 ctx 学习 | RoPE 兼容 | hyperparameter |
|---|---|---|---|---|---|
| MHA / GQA | 中 | ★★★★★ | 直接走 YaRN / NTK 外推 | 原生 | 最少 |
| MLA | 中(需 c_kv dim 调) | ★★★★ | YaRN + decoupled RoPE | 需 decouple | \(d_c\), \(d_R\) |
| DSA | 高(top-k 不可微) | ★★★ | 与 MLA 一致 | 跟 MLA 同 | top-k, indexer dim |
| CSA / HCA | 高(双路径权重) | ★★★ | 1M+ ctx 设计原生 | 跟 MLA 同 | m, m', top-k, hybrid 比例 |
| Sliding Window | 低 | ★★★★ | 不能外推(W 固定) | 原生 | W |
| GLA | 中 | ★★★ | 100K+ recall 显著弱 | RoPE 不直接适用 | gate scale |
| Mamba-2 | 中(HiPPO init 救一波) | ★★★★ | 中等(hybrid 救) | 不需要(state 自带顺序) | \(d_{state}\), expand |
| KDA | 中 | ★★★ | 配合 MLA hybrid 较强 | 同 GLA | α/β init, DPLR rank |
| Lightning | 高(normalization 跟 init 强相关) | ★★★ | hybrid 后强 | 不直接适用 | feature map type |
重点:纯 Linear 路线在长 ctx recall 上永远输 Full(state 容量决定,不是 hyperparameter 调能解决)。这就是为什么 frontier 必然 hybrid:用 Linear 担成本、用 Full 担精度,分工明确。
七、选型决策¶
| 场景 | 选哪个 | 为什么 |
|---|---|---|
| 学术 baseline / 7B | GQA-8 | LLaMA-3 标配,工程最成熟 |
| 重视长 ctx 推理成本 | MLA | DeepSeek V3 验证 |
| 1M+ ctx 极致压缩 | CSA + HCA | DeepSeek V4,质量最稳 |
| 中等 ctx + agent / RL | KDA hybrid (3:1) | Kimi Linear,长 ctx + 精度兼顾 |
| 长 ctx + 长 output | Lightning hybrid (1:7) | MiniMax M1,最激进 linear |
| 想减 attention 依赖 | Mamba-Transformer (92:8) | Nemotron-H, 8% attn |
| 端侧 / 边缘 / 流式 | 纯 GLA / KDA | KV 常数,单步算力固定 |
八、追问与延伸¶
| 问题 | 答案方向 | 详见 | 为什么值得深入 |
|---|---|---|---|
| MLA 的 RoPE 为什么不能走压缩路径? | RoPE 的旋转矩阵 R(θ) 是位置相关的,W_uk @ R(θ) @ c_kv ≠ R(θ) @ W_uk @ c_kv,压缩后旋转会破坏低秩结构 | A4 | 这是 MLA 设计中最关键的约束:理解它才能理解为什么 KV Cache 是 d_c + d_rope 而不是只有 d_c |
| GLA 的 chunk-wise 并行怎么做? | 将序列分为固定大小 chunk(如 64 tokens),chunk 内用矩阵乘法并行计算(类似短序列 attention),chunk 间顺序传递状态矩阵 S | linear.md | 纯递推在 GPU 上效率极低(无法并行);chunk-wise 方法在训练时实现接近 FlashAttention 的 GPU 利用率 |
| 实际训练时怎么选注意力类型? | GQA 是最成熟选择(LLaMA/Mistral/Qwen 均已验证);MLA 质量更优但工程复杂度高(需要定制 kernel);GLA 适合流式/端侧等特定场景 | D3/D4 | 注意力类型一旦选定就无法更改(不像 LoRA 可后期添加),选型失误意味着整个预训练投入浪费 |
| 不同注意力的 Triton kernel 怎么写? | GQA 可直接复用 FlashAttention-⅔ 的开源 kernel;MLA 推荐用 TileLang 编写 FlashMLA(需处理吸收 trick);GLA 需要自定义 chunk-wise recurrence kernel | C4/C5 | 注意力变体的理论优势能否兑现完全取决于 kernel 实现质量;糟糕的 kernel 可能让 MLA 比 GQA 更慢 |
参考链接¶
- DeepSeek-V3 MLA Explained
- DeepSeek-V3 Technical Report
- GQA Paper
- GLA Paper
- Mamba-2 / SSD framework
- Kimi Linear (KDA)
- FlashMLA with TileLang
↑ 上级 · A. 基础理论