跳转至

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(详 §四)。理解每条线的失败模式才知道为什么混搭。

→ 各路线深读:


一、复杂度速查表

\(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) 都算)

MHA / MLA full causal mask

2. SWA W=3 — Sparse 静态窗口(每个 q 只看前 W 个)

SWA W=3 sliding window mask

3. DSA top-3 — Sparse 动态选择(每个 q 由 lightning indexer 选 k=3 个最相关)

DSA top-3 dynamic mask

4. CSA m=2, k=2 — 序列 2× 压缩 + top-k 选块

CSA compressed sparse mask

5. GLA / Linear — 没有 mask,全压进 state

Linear attention state matrix update

▲ 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 路线最快的方式

MHA / GQA / MQA / MLA 对比(DeepSeek-V2 paper Fig 3)

来源: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\) 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× 数值偏差。原因:

  1. tile 越多 → rescaling 累积越多 → 误差增长
  2. 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 自家

几个观察

  1. Full / Sparse 路线整段并行(FlashAttention 风格)→ Tensor Core 利用率最高
  2. Linear / SSM 路线本质上是扫描算法,必须 chunk-wise 拆才能并行 → 利用率永远略输 Full
  3. Mamba-2 SSD 是目前 Linear 路线训练效率最高的 framework,因为它把 SSM 写成纯 matmul(跟 Linear Attention 等价)
  4. 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 更慢

参考链接


上级 · A. 基础理论