Full 路线:保留 softmax,压缩 K/V 表示¶
更新日期:2026-04-27
核心原理:保留 \(\text{softmax}(QK^\top/\sqrt{d}) V\) 的标准 attention 语义,只在 K/V 表示上做文章。所有 query 仍能 attend 到所有 historical token,没信息损失,但 KV cache 大小可以从 \(O(N \cdot H \cdot d_h)\) 压到 \(O(N \cdot d_c)\) 甚至更小。
一、四种变体一图对比(MHA / GQA / MQA / MLA)¶
DeepSeek-V2 paper Figure 3 给出了 Full 路线四个变体的可视化对比,这张图是看懂全章最快的入口:

来源:DeepSeek-V2 paper (Liu et al., 2024), Figure 3。版权归原作者,本站作学习引用。
GQA paper 自己的对比图(更早的版本,3 栏 MHA / GQA / MQA):

来源:GQA: Training Generalized Multi-Query Transformer Models (Ainslie et al., 2023), Figure 1。
一图看懂区别:
- MHA:\(H\) 个独立的 KV head(图中 GQA paper Fig 1 左),KV cache ∝ \(H\)
- GQA-G:只有 \(G < H\) 套 K/V(中间),每组被 \(H/G\) 个 Q head 共享,推理时 broadcast 复用
- MQA:所有 Q head 共享同一套 K/V(右),\(G = 1\),是 GQA 的极端情形
- MLA:K 和 V 都从一个 \(d_c\) latent vector 通过 \(W^{UK}\), \(W^{UV}\) 上投影恢复(DeepSeek-V2 Fig 3 最右);推理时利用矩阵结合律吸收 \(W^{UK}\) 进 Q,KV cache 只剩 \(c_{kv}\) + \(k^R\)
二、演化逻辑¶
Full 路线始终保留 softmax over all KV,只压缩 KV 表示本身。所谓"演化"其实就是同一个方向(KV cache 减少)的不同 trick:
flowchart LR
MHA["MHA<br/>2 H d_h"]
MQA["MQA<br/>2 d_h"]
GQA["GQA-G<br/>2 G d_h"]
MLA["MLA<br/>d_c + d_R"]
HCA["HCA m'=128<br/>(d_c+d_R)/m'"]
MHA -->|"共享 K/V 头"| MQA
MQA -->|"分组共享"| GQA
GQA -->|"低秩 latent"| MLA
MLA -->|"序列维压缩"| HCA
classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
class MHA,MQA,GQA,MLA,HCA stage
注意:DSA / CSA 在 V3.2 / V4 里跟 MLA 配合用,但它们引入了 top-k 选择(不是所有 KV 都 attend),按定义属于 Sparse 路线 —— 见 A2-2 Sparse 路线。Full 路线压缩"维度"(dim),Sparse 路线压缩"对数"(pairs),两件事正交。
每一步都是前一步的瓶颈被打破:
| 变体 | 压缩什么 | KV / token / layer | 解决了什么 |
|---|---|---|---|
| MHA | — (基准) | \(2 H d_h\) | 表达力上限 |
| MQA | KV head=1 | \(2 d_h\) | head 数太多导致 KV 大 |
| GQA | KV head=G | \(2 G d_h\) | MQA 表达力损失太大 |
| MLA | low-rank latent | \(d_c + d_R\) | head 数 × dim 仍线性 |
| HCA | 序列维 m'=128 压缩 | \((d_c + d_R)/m'\) | 长 ctx 仍线性 |
三、MHA / GQA — HuggingFace 参考实现¶
这两者在 HF transformers 中是同一份代码:num_key_value_heads 等于 num_attention_heads 时是 MHA,小于时是 GQA。MQA 是 GQA 的极端(num_key_value_heads=1)。
HF 路径:transformers/src/transformers/models/llama/modeling_llama.py 中的 LlamaAttention 类。
关键片段(HF 风格,省略部分细节):
# https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py
class LlamaAttention(nn.Module):
def __init__(self, config: LlamaConfig, layer_idx: int):
super().__init__()
self.config = config
self.layer_idx = layer_idx
self.head_dim = config.head_dim
self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads
self.scaling = self.head_dim ** -0.5
self.attention_dropout = config.attention_dropout
self.q_proj = nn.Linear(config.hidden_size, config.num_attention_heads * self.head_dim, bias=config.attention_bias)
self.k_proj = nn.Linear(config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias)
self.v_proj = nn.Linear(config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias)
self.o_proj = nn.Linear(config.num_attention_heads * self.head_dim, config.hidden_size, bias=config.attention_bias)
def forward(self, hidden_states, position_embeddings, attention_mask, past_key_value=None, **kwargs):
input_shape = hidden_states.shape[:-1]
hidden_shape = (*input_shape, -1, self.head_dim)
# [B, S, H_q*d_h] -> [B, H_q, S, d_h] (q)
# [B, S, H_kv*d_h] -> [B, H_kv, S, d_h] (k, v)
query_states = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2)
key_states = self.k_proj(hidden_states).view(hidden_shape).transpose(1, 2)
value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
cos, sin = position_embeddings
query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)
if past_key_value is not None:
cache_kwargs = {"sin": sin, "cos": cos, "cache_position": kwargs.get("cache_position")}
key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)
# GQA broadcast: 把 H_kv 沿 head 维 repeat 到 H_q
# eager_attention_forward 内部会处理(或调 SDPA / FlashAttn backend)
attn_output, attn_weights = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation](
self, query_states, key_states, value_states, attention_mask,
dropout=self.attention_dropout, scaling=self.scaling, **kwargs,
)
attn_output = attn_output.reshape(*input_shape, -1).contiguous()
return self.o_proj(attn_output), attn_weights
def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
"""GQA broadcast: [B, H_kv, S, d_h] -> [B, H_kv * n_rep, S, d_h]"""
batch, num_kv, slen, head_dim = hidden_states.shape
if n_rep == 1:
return hidden_states
hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_kv, n_rep, slen, head_dim)
return hidden_states.reshape(batch, num_kv * n_rep, slen, head_dim)
关键点:
num_key_value_heads == num_attention_heads:MHA(LLaMA-1 7B/13B、LLaMA-2 7B/13B)num_key_value_heads < num_attention_heads:GQA(LLaMA-2 70B、LLaMA-3 全尺寸、Qwen、Mistral)num_key_value_heads == 1:MQA(PaLM、Falcon-1B/7B/40B)- 实际 attention 计算交给
ALL_ATTENTION_FUNCTIONS[backend],可选eager/sdpa/flash_attention_2/flash_attention_3,都是 GQA-aware 的(调用前repeat_kv做 broadcast)
KV Cache 大小:2 × num_kv_heads × d_h × seq_len per layer。LLaMA-3 70B(H=64, n_kv=8, d_h=128, L=80)在 128K context:2 × 8 × 128 × 128K × 80 × 2 bytes ≈ 32 GB。
四、MLA (Multi-head Latent Attention) — DeepSeek 核心创新¶
参考 DeepSeek-V2 Technical Report。
4.1 动机¶
GQA 通过减少 KV 头数来压缩,但信息损失与头数减少成正比。MLA 的想法不同:不减少头数,而是压缩 KV 的表示本身到一个 latent vector。
4.2 HF 参考实现¶
HF 路径:transformers/src/transformers/models/deepseek_v3/modeling_deepseek_v3.py 中的 DeepseekV3Attention。
关键片段:
# https://github.com/huggingface/transformers/blob/main/src/transformers/models/deepseek_v3/modeling_deepseek_v3.py
class DeepseekV3Attention(nn.Module):
def __init__(self, config: DeepseekV3Config, layer_idx: int):
super().__init__()
self.config = config
self.layer_idx = layer_idx
self.num_heads = config.num_attention_heads # H = 128
self.qk_nope_head_dim = config.qk_nope_head_dim # d_h_nope = 128
self.qk_rope_head_dim = config.qk_rope_head_dim # d_R = 64
self.q_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim # 192
self.v_head_dim = config.v_head_dim # d_v = 128
self.q_lora_rank = config.q_lora_rank # 1536 (Q 也压缩)
self.kv_lora_rank = config.kv_lora_rank # d_c = 512
self.scaling = self.q_head_dim ** -0.5
# === Q 路径:先压成 q_lora_rank,再上投影 ===
self.q_a_proj = nn.Linear(config.hidden_size, self.q_lora_rank, bias=config.attention_bias)
self.q_a_layernorm = DeepseekV3RMSNorm(self.q_lora_rank)
self.q_b_proj = nn.Linear(self.q_lora_rank, self.num_heads * self.q_head_dim, bias=False)
# === KV 路径:压成 d_c (kv_lora_rank) + 单独 RoPE 部分 d_R ===
self.kv_a_proj_with_mqa = nn.Linear(
config.hidden_size,
self.kv_lora_rank + self.qk_rope_head_dim, # d_c + d_R = 576,只 cache 这个!
bias=config.attention_bias,
)
self.kv_a_layernorm = DeepseekV3RMSNorm(self.kv_lora_rank)
self.kv_b_proj = nn.Linear(
self.kv_lora_rank,
self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), # 上投影回 K_nope 和 V
bias=False,
)
self.o_proj = nn.Linear(self.num_heads * self.v_head_dim, config.hidden_size, bias=config.attention_bias)
def forward(self, hidden_states, position_embeddings, attention_mask, past_key_value=None, **kwargs):
bsz, q_len, _ = hidden_states.size()
# ===== Q:从 hidden 压到 q_lora_rank 再上投影到 [H, d_h_nope + d_R] =====
q = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states)))
q = q.view(bsz, q_len, self.num_heads, self.q_head_dim).transpose(1, 2)
# 拆出 nope (与 c_kv 内积) 和 rope (与 k_R 内积) 两部分
q_nope, q_rope = torch.split(q, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)
# ===== KV:低秩压缩 =====
compressed_kv = self.kv_a_proj_with_mqa(hidden_states)
# 拆出 c_kv(512) 和 k_R(64)
compressed_kv, k_pe = torch.split(
compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1
)
compressed_kv = self.kv_a_layernorm(compressed_kv) # [B, S, 512]
# k_R 走 RoPE 路径(不压缩)
k_pe = k_pe.unsqueeze(1).expand(-1, self.num_heads, -1, -1) # [B, H, S, d_R]
# ===== 上投影 c_kv 恢复 K_nope 和 V =====
kv = self.kv_b_proj(compressed_kv).view(bsz, q_len, self.num_heads, -1).transpose(1, 2)
k_nope, value_states = torch.split(kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1)
# ===== 应用 RoPE 到 q_rope 和 k_pe =====
cos, sin = position_embeddings
q_rope, k_pe = apply_rotary_pos_emb(q_rope, k_pe, cos, sin, unsqueeze_dim=2)
# 拼回完整 K = [k_nope ; k_pe]
key_states = torch.cat([k_nope, k_pe], dim=-1) # [B, H, S, d_h_nope + d_R]
query_states = torch.cat([q_nope, q_rope], dim=-1)
# ===== KV cache 只缓存 (compressed_kv, k_pe) =====
if past_key_value is not None:
# 注意:标准 HF 缓存 (key_states, value_states),DeepSeek 自家 kernel 才走压缩缓存路径
key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, ...)
# ===== Attention =====
attn_output, attn_weights = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation](
self, query_states, key_states, value_states, attention_mask,
dropout=0.0, scaling=self.scaling, **kwargs,
)
attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()
return self.o_proj(attn_output), attn_weights
4.3 推理时的吸收 trick¶
HF 的实现仍然显式地从 c_kv 重建 K 和 V(标准 PyTorch 路径),KV cache 也存的是完整 K/V。真正的"压缩缓存"路径只在 DeepSeek 自家 FlashMLA kernel 才启用。
数学上,吸收 trick 是利用矩阵乘法结合律:
把 W_uk 吸收进 Q 的投影矩阵,推理时只需缓存 c_kv(512 dim)和 k_pe(64 dim),每 token 每层 KV cache = (512 + 64) × 2 = 1152 bytes。详见 TileLang FlashMLA 实现。
4.4 MLA 的实际 KV Cache 对比¶
以 DeepSeek-V3(128 头, \(d_h=128\), \(d_c=512\), \(d_R=64\))和 LLaMA-3 70B(GQA, 64 Q + 8 KV, \(d_h=128\))为例,单 token 单层 KV cache 字节数(BF16, 2 bytes/elem):
| 方案 | KV cache (bytes/token/layer) | 1M ctx 单层 GB | 相对 MHA |
|---|---|---|---|
| MHA (假想 V3 用) | \(2 \cdot 128 \cdot 128 \cdot 2 = 65536\) | 65 GB | 1.0× |
| GQA-8 (LLaMA-3 70B) | \(2 \cdot 8 \cdot 128 \cdot 2 = 4096\) | 4 GB | 16× |
| MLA (V3) | \((512 + 64) \cdot 2 = 1152\) | 1.15 GB | 57× |
V3 总 61 层 → 1M ctx 全网络 KV cache 仅约 70 GB,单卡 H100 (80 GB) 即可装下。同等 ctx 的 MHA 等价模型需要 ~4 TB,根本不可能 serve。
4.5 吸收 trick 的实际收益(RTX 4070 Ti SUPER 实测)¶
PyTorch reference 实现:mla_naive(重建 K/V)vs mla_absorbed(在压缩空间算)。完整代码见 tests/test_attention_kernels.py,跑法:
DeepSeek-V3 shape (H=128, \(d_c\)=512, \(d_h^\text{nope}\)=\(d_v\)=128, \(d_R\)=64, BF16, S=1 decode, B=1):
| past KV (T) | naive 重建 K/V | absorbed | 加速 | 内存节省 |
|---|---|---|---|---|
| 4096 | 3.55 ms / 335 MB | 0.14 ms / 4.7 MB | 25.6× | 71× |
| 16384 | 14.61 ms / 1342 MB | 0.34 ms / 19 MB | 42.8× | 71× |
为什么 speedup 比 71× 内存比小:
- 内存节省是结构性的(\(H \cdot d_h^\text{nope}\) 不再每个 token 各存一份)
- 速度收益主要是带宽 bound:decode 时读 KV cache 是瓶颈,naive 要扫 1.3 GB 才能算 1 个 token,absorbed 只扫 19 MB
- T 越长,naive 越受 KV 带宽 bound → speedup 越大(4K → 26×,16K → 43×)
这就是为什么 DeepSeek 部署时硬要用吸收 kernel —— 不只是 KV cache 小,decode 速度也跟着翻几十倍。
数值正确性验证(test_mla_numerical):absorbed vs naive 在 FP32 下 atol/rtol = 1e-4 完全一致,说明吸收 trick 是等价变换而非近似。
五、HCA (Heavily Compressed Attention) — DeepSeek V4 全局粗粒度¶
MLA 压 KV 表示但长度仍 ∝ N,1M ctx 仍很大。HCA 沿序列维做 128× 激进压缩,压完之后做 dense softmax(仍 Full 路线,没引入 selection)。参考 DeepSeek-V4 tech report。
# 简化伪代码(HF 尚未上 V4,参考 V4 tech report Section 3)
class HCA(nn.Module):
"""Heavily Compressed Attention: m'=128 压缩,然后 dense MLA attention。"""
def __init__(self, d_model, d_c, m_prime=128):
super().__init__()
self.m_prime = m_prime
# softmax-gated pooling + 学习 positional bias
self.gate = nn.Linear(d_c, m_prime)
self.pos_bias = nn.Parameter(torch.zeros(m_prime, d_c))
self.mla = MLA(d_model, n_heads=128, d_h=128, d_c=d_c)
def compress(self, c_kv):
# 每 m'=128 个 c_kv 用 softmax-gated pool + positional bias 合成 1 个块
B, T, d = c_kv.shape
T_blocks = T // self.m_prime
c_kv_blocks = c_kv[:, : T_blocks * self.m_prime].reshape(B, T_blocks, self.m_prime, d)
gate = self.gate(c_kv_blocks).softmax(dim=-2) # [B, Tb, m', m']
weighted = c_kv_blocks + self.pos_bias # 加 positional bias
compressed = (gate.unsqueeze(-1) * weighted.unsqueeze(-2)).sum(-2).mean(-2)
return compressed # [B, T/m', d]
def forward(self, x, c_kv_cache):
compressed = self.compress(c_kv_cache) # 128× 压缩
return self.mla(x, compressed) # dense attention
为什么是 Full 而不是 Sparse:压缩之后每个 query 仍 attend 到全部 (N/m') 个压缩块,softmax 覆盖完整历史(只是粗粒度),没有 top-k selection。
V4 实际架构 = HCA(全局粗)+ CSA(局部细)双路径并行,CSA 是 Sparse 路线(见 sparse.md §四)。
六、CUDA 性能验证 / kernel 实现¶
每个变体的"理论 KV cache 节省"能不能兑现,完全取决于 kernel 质量。糟糕的 GQA kernel 比 MHA 还慢;糟糕的 MLA kernel 因为 latent 上投影开销反而吃亏。
GQA / MLA 的 Triton + PyTorch reference 实现 + 在 RTX 4070 上的实测对比,连同 Linear 路线(GLA / Lightning / KDA / Mamba-2)的 kernel 一起放在横跨三路线的独立文档:
涵盖:
- Triton FlashAttention-2 (GQA-aware) 完整实现
- TileLang FlashMLA(DeepSeek V3)核心 kernel
- 数值正确性验证(vs PyTorch eager)+ 性能基准(RTX 4070 Ti SUPER BF16 实测)
参考链接¶
- DeepSeek-V2 Technical Report — MLA 起源
- DeepSeek-V3 Technical Report
- DeepSeek-V4 tech report (PDF) — HCA + CSA
- GQA Paper
- HuggingFace 实现:
modeling_llama.py— MHA / MQA / GQAmodeling_deepseek_v3.py— MLA
- 高性能 kernel:
- FlashAttention — Tri Dao 团队
- FlashMLA — DeepSeek Hopper 优化
- TileLang FlashMLA — DSL 实现参考
↑ 上级 · A2 注意力机制全景