面试知识库
高 困难

注意力机制变体与位置编码#

一句话答案#

注意力变体解决的是两个瓶颈:KV Cache 显存(MHA→MQA→GQA→MLA,用更少/更压缩的 K、V 换显存和带宽)和计算 IO(FlashAttention 用 tiling + online softmax 不落地 N×N 矩阵,精确不近似);位置编码则从绝对走向 RoPE(用旋转把相对位置藏进 q·k 内积),长度外推靠改 RoPE 的旋转频率(PI/NTK/YaRN)。

核心要点

1. 瓶颈在哪:KV Cache 随 seq_len 线性爆炸#

自回归解码时每个新 token 都要读全部历史 K/V(见 LLM推理优化 的 KV Cache 基础)。按每个序列算:

KV Cache 字节数 = 2(K和V) × layers × kv_heads × head_dim × seq_len × bytes_per_elem
plaintext

量级感受(FP16,2 字节):Llama 2 70B 若用标准 MHA(80 层、64 头、head_dim 128)约 2.6 MB/token,4K 上下文就要 ~10 GB,128K 就要 ~340 GB——比权重本身(140 GB)还大;batch 再乘以并发数。解码阶段是访存受限(memory-bound):每步都要把整块 KV 从 HBM 搬进计算单元,KV 越大吞吐越低。所以所有变体的目标都是”少存 K/V”。

2. MHA → MQA → GQA → MLA 的演进#

方案做法KV 头数质量代表
MHA每个 Q 头配独立 K/V 头= n_heads基准早期 GPT/Llama 1
MQA所有 Q 头共享一组 K/V1略降,需从头训练或 uptrainPaLM、Falcon、StarCoder
GQAQ 头分 G 组,组内共享 K/VG(常 8)≈MHALlama 2 70B、Llama 3、Mistral、Qwen
MLAK/V 低秩压缩成潜在向量 c_KV,用时再升维等效 ~2 个 GQA 组≈/优于 MHADeepSeek-V2/V3
  • GQA 是 MHA 与 MQA 的插值:Llama 2 70B 把 64 个 Q 头分成 8 组,KV Cache 直接降 8 倍(上例从 10 GB→1.3 GB);可从 MHA checkpoint 把同组 K/V 头均值池化后少量续训得到。
  • MLA 的巧思:c_KV = W_DKV · h(维度 d_c=512,远小于 n_heads×head_dim=16384),K、V 由 W_UK·c_KV、W_UV·c_KV 升维。推理时 W_UK 可吸收进 W_Q、W_UV 吸收进 W_O,所以根本不必物化 K/V,只缓存 c_KV。难点是 RoPE 的旋转矩阵夹在 W_Q 和 W_UK 之间、和位置相关无法吸收,于是 解耦 RoPE:另开一个小维度(d_R=64)的 k_R 专门带位置信息(跨头共享,像 MQA),与压缩部分拼接。每 token 只缓存 512+64=576 个数,DeepSeek-V2 论文称比其 MHA 版本省 93% KV Cache,且效果不降。

3. FlashAttention:IO 感知的精确注意力#

朴素实现要把 S = QKᵀ(N×N)和 P = softmax(S) 写回 HBM 再读回来乘 V,显存 O(N²),且 GPU 的瓶颈不在算力而在 HBM 带宽(A100:HBM ~2 TB/s vs 片上 SRAM ~19 TB/s,但 SRAM 只有 ~20 MB)。FlashAttention 的三个点:

  1. Tiling:把 Q/K/V 切块装进 SRAM,逐块算,每块结果直接累加到输出 O;
  2. Online softmax:softmax 需要全行最大值 m 和归一化和 l,分块时维护运行中的 (m, l),新块到来时用 exp(m_old − m_new) 对已有累加量重缩放——数学上与一次性 softmax 等价;
  3. 反向重计算:不保存 P,反向传播时从 (Q, K, m, l) 重算,用算力换显存。

结果:显存从 O(N²) 降到 O(N),速度提升 2–4 倍,完全精确(和稀疏/线性注意力的”近似”本质不同);但 FLOPs 仍是 O(N²),它优化的是 IO 不是复杂度。FA2 改善了线程块的并行与 warp 划分,FA3 面向 Hopper 利用异步与 FP8。vLLM/SGLang 等框架默认集成。

4. 真正改复杂度的路线(只点名)#

  • 稀疏/滑动窗口注意力(SWA):Mistral 7B 每层只看前 4096 个 token,KV Cache 用环形缓冲固定大小;感受野随层数叠加(window × layers),但远距离信息是”间接”传递的。
  • 线性注意力 / SSM:Performer、RWKV、Mamba 把状态压成固定大小,O(n) 且无 KV Cache 增长,代价是精确召回(如 NIAH)偏弱;Jamba 等用 Transformer+Mamba 混合层折中。
  • 混合注意力进入主流开源模型:Qwen3.5 每 4 层里 3 层用 Gated DeltaNet 线性注意力、1 层用标准(门控)注意力,标准层仍用 GQA;DeepSeek-V4 用压缩稀疏注意力(CSA + HCA 混合)支撑 1M 上下文。思路都是”大部分层走便宜的近似/压缩注意力,少数层保留精确全注意力负责精确召回”。

5. 位置编码:从绝对到 RoPE#

注意力本身是置换不变的,必须注入位置。路线:正弦/可学习绝对编码(加在 embedding 上;可学习的超出训练长度直接没有向量)→ 相对位置(Shaw、T5 bias:在 score 上加与 i−j 相关的偏置)→ RoPE(Su 等,RoFormer)。

RoPE 机制:把 q、k 的第 (2i, 2i+1) 两维看作一个复数,位置 m 处乘以旋转 e^{i·m·θ_i},其中 θ_i = base^{−2i/d}(base 常取 10000,Llama 3 提到 500000)。于是:

⟨R_m q, R_n k⟩ = Re[ q · k̄ · e^{i(m−n)θ} ]   ← 只和相对距离 m−n 有关
plaintext

绝对形式、相对效果:不加参数、不改 attention 公式、只旋转 q 和 k(不动 v),缓存的 k 已带位置可直接复用,低频维度天然带远距离衰减——这就是 Llama/Qwen/DeepSeek/Mistral 等几乎全部主流开源模型采用它的原因。ALiBi 走另一条路:不加位置向量,直接在 score 上减 slope × |i−j|,外推性好但主流模型少用。

6. 长度外推:改旋转频率#

训练长 L、推理 >L 时,RoPE 会遇到训练中从未见过的相对角度,注意力分布失控,表现为 困惑度陡升、输出重复/乱码。三种修法本质都是”让推理时的角度落回训练见过的范围”(窗口多长够用、有效长度怎么评估,见 长上下文处理):

方法思路特点
位置插值 PI位置 m 缩放为 m/s(所有频率同乘 1/s)需少量微调;高频维度被压缩,损伤局部分辨率
NTK-aware改 base:base' = base × s^{d/(d−2)},高频基本不动、低频被插值无微调也能用一部分,动态版按当前长度调 s
YaRNNTK-by-parts:按每维”波长 vs 训练长度”决定不插值/全插值/线性过渡,再对 logits 做温度缩放补熵微调数据需求比 PI 小一个量级;DeepSeek-V2/V3、Qwen2.5 的 128K 配置采用,Qwen3 用它从原生 32K 扩到 131K、Qwen3.5 从原生 262K 扩到约 1M

面试回答(2分钟版)

注意力变体我从两个瓶颈讲。第一是 KV Cache 显存:解码每步要读全部历史 K/V,大小是 2×层数×KV头数×head_dim×序列长×字节数,70B 模型 MHA 下每 token 两三兆,长上下文时比权重还大。所以 MQA 让所有 Q 头共享一组 KV,GQA 折中分组——Llama 2 70B、Llama 3 都是 8 组,KV 降 8 倍而效果接近 MHA;DeepSeek 的 MLA 把 KV 低秩压缩成 512 维潜在向量,升维矩阵吸收进 Q、O 的投影,只缓存潜在向量,RoPE 不能被吸收就解耦成一个 64 维带位置的 key,省 93% KV。第二是计算 IO:FlashAttention 不改 O(n²),它针对 HBM 带宽瓶颈分块进 SRAM、用 online softmax 重缩放、反向重算,从不落地 N×N 矩阵,显存 O(N)、提速 2–4 倍且精确,和稀疏、线性注意力的近似本质不同。位置编码主流是 RoPE:把 q、k 相邻两维当复数乘以和位置成正比的旋转角,内积只剩相对距离 m−n,不加参数、兼容 KV Cache。超出训练长度困惑度会爆炸,修法是 PI 整体缩位置、NTK-aware 改 base 保高频、YaRN 按维度波长分段再加温度缩放。

追问与易错

追问方向:

  • GQA 为什么选 8 组而不是 1 组(MQA)? → MQA 质量有可见下降且训练不稳;8 组正好对应 8 卡张量并行时每卡一个 KV 头,通信友好;论文显示 8 组质量已≈MHA,再多收益小
  • MLA 和 GQA 谁更省?为什么 MLA 效果还能不降? → DeepSeek-V2 每 token 缓存 576 个数,等效约 2.25 组 GQA,比 8 组省;不降是因为 K/V 在升维后仍是”每头独立”的,表达力比共享头的 GQA 更强,只是表示被约束在低秩子空间
  • FlashAttention 是近似吗?能降低 O(n²) 吗? → 精确,数值上与标准 attention 等价;FLOPs 不变仍是 O(n²),它省的是 HBM 读写(显存 O(n))和时间。想降复杂度得用 SWA/线性注意力/SSM
  • RoPE 为什么只作用于 q、k 不作用于 v? → 位置信息的用途是决定”关注谁”,体现在 score=q·k;v 是被取走的内容,加旋转反而污染输出
  • NTK-aware 为什么比 PI 更能”免微调”外推? → PI 把高频维度(负责局部顺序)也压缩了,模型分不清相邻 token;NTK 改 base 使高频几乎不动、只让低频(远距离)维度插值,局部分辨率保住

易错点:

  • ❌ “FlashAttention 是稀疏注意力的一种” → 它是精确算法,优化的是 IO;稀疏/线性注意力才是近似
  • ❌ “GQA 是减少 Q 头数” → 减的是 K/V 头数,Q 头数不变,多个 Q 头共用一组 K/V
  • ❌ “RoPE 是相对位置编码所以外推天然好” → RoPE 外推并不好,超长直接崩,必须配 PI/NTK/YaRN 或续训
  • ❌ “KV Cache 大小 = 2×层数×hidden_dim×seq_len×字节” → 这是 MHA 特例,通用式要用 kv_heads×head_dim,GQA/MLA 下差几倍到几十倍