跳转至

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 路线四个变体的可视化对比,这张图是看懂全章最快的入口

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

来源:DeepSeek-V2 paper (Liu et al., 2024), Figure 3。版权归原作者,本站作学习引用。

GQA paper 自己的对比图(更早的版本,3 栏 MHA / GQA / MQA):

MHA / GQA / MQA 架构(GQA paper Fig 1)

来源: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 是利用矩阵乘法结合律:

\[ Q_\text{nope} \cdot K_\text{nope}^\top = (W_{uq}^\top x_q)^\top \cdot W_{uk} c_{kv}^\top = x_q^\top \underbrace{(W_{uq} W_{uk})}_{\text{预合并}} c_{kv}^\top \]

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,跑法:

python -m pytest tests/test_attention_kernels.py::test_mla_bench -xvs --capture=no

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 一起放在横跨三路线的独立文档:

A2-4 attention kernel 实现与验证

涵盖:

  • Triton FlashAttention-2 (GQA-aware) 完整实现
  • TileLang FlashMLA(DeepSeek V3)核心 kernel
  • 数值正确性验证(vs PyTorch eager)+ 性能基准(RTX 4070 Ti SUPER BF16 实测)

参考链接


上级 · A2 注意力机制全景