欢迎光临

DeepSeek MLA(多头潜在注意力)深度解析:如何用低秩压缩将 KV Cache 压缩到原来的 1/10

在大模型推理的工程实践中,KV Cache 是一把双刃剑:它通过缓存历史 Token 的 Key 和 Value 向量避免了重复计算,是自回归推理加速的基石;但它同时也是显存消耗的”头号杀手”。以一个 70B 参数模型为例,单条 4K 上下文请求的 KV Cache 就可能占用超过 3GB 的显存。当并发请求量上升时,KV Cache 的显存占用会迅速成为系统的瓶颈——不是算力不够,而是存不下。

2024 年底,DeepSeek 在 DeepSeek-V2 和 DeepSeek-V3 中提出了一种名为 MLA(Multi-head Latent Attention,多头潜在注意力)的创新架构,通过低秩联合压缩将 KV Cache 压缩到传统 MHA 的约 1/10,同时保持甚至超越了 GQA(Grouped-Query Attention)的推理质量。本文将从数学原理到工程实现,全面拆解 MLA 的设计思路。

一、传统注意力机制的 KV Cache 困局

要理解 MLA 的价值,首先需要厘清传统注意力机制中 KV Cache 的膨胀机制。在标准 Multi-Head Attention(MHA)中,对于序列长度为 N、隐藏维度为 d、注意力头数为 h 的模型,每个 Token 需要缓存的 KV 维度为

1
2 × d

(Key 和 Value 各一份 d 维向量)。这意味着 KV Cache 的总显存占用为:


1
2
3
4
5
6
7
# MHA KV Cache 显存估算
# 模型配置: d_model=5120, n_heads=40, head_dim=128, n_layers=80
# 每层每 Token 的 KV Cache: 2 * 40 * 128 = 10240 个元素
# 4K 上下文, FP16:
cache_size = 2 * 40 * 128 * 80 * 4096 * 2  # bytes
# = 2 * 40 * 128 * 80 * 4096 * 2
# = 10,737,418,240 bytes ≈ 10 GB

对于 70B 级别的模型,一条 4K 上下文的请求就需要约 10GB 的 KV Cache。这意味着一张 80GB 的 A100 最多只能同时服务不到 8 个并发请求(还要扣除模型权重占用的空间),严重限制了吞吐量。

业界在此前已经提出了多种优化方案:

  • MQA(Multi-Query Attention):所有 Query 头共享同一组 KV,将 KV Cache 压缩到 1/h,但模型质量显著下降。
  • GQA(Grouped-Query Attention):将 Query 头分组,每组共享一组 KV,在质量和压缩比之间取得平衡。LLaMA-2/3 系列采用了 8 组 KV(h=32/64 时压缩到 1/4 到 1/8)。
  • CKV(Cross-Layer KV Sharing):跨层共享 KV Cache,进一步压缩层级存储。MLA 在此基础上更进一步。

但这些方案都面临一个核心矛盾:压缩比越高,注意力机制的表达能力损失越大。MLA 的突破在于——它不减少”信息量”,而是压缩”信息表示”。

二、MLA 的核心思想:低秩联合压缩

MLA 的设计灵感来自于一个关键观察:在预训练后的注意力模型中,Key 和 Value 矩阵在联合空间中存在显著的低秩结构。也就是说,虽然 KV 的表面维度很高(h × head_dim),但它们的”有效信息维度”远低于此。如果能找到这个低维子空间,就可以用一个小得多的潜在向量来重构 KV,从而大幅降低缓存需求。

MLA 的核心操作可以概括为三步:

2.1 下投影:将 Token 压缩为潜在向量

在生成每个 Token 时,模型首先将隐藏状态

1
h_t

通过一个下投影矩阵

1
W_dkv

(Down-projection)映射为一个低维的潜在向量

1
c_kv


1
2
3
4
5
6
# 下投影: 将 d_model 维隐藏状态压缩为 d_c 维
# d_c = 512 (远小于 n_heads * head_dim = 5120)
c_kv = h_t @ W_dkv  # shape: (d_model,) -> (d_c,)

# 缓存: 只存储 c_kv, 维度为 d_c = 512
# 而非传统 MHA 的 2 * n_heads * head_dim = 10240

这里的

1
d_c

(压缩维度)通常设为 512,而传统 MHA 的 KV 维度为

1
2 × h × head_dim = 10240

。仅此一步,就实现了约 20 倍的维度压缩。

2.2 上投影:从潜在向量重构 Key 和 Value

在计算注意力时,从缓存的潜在向量

1
c_kv

通过上投影矩阵

1
W_uk

1
W_uv

(Up-projection)重构出完整的 Key 和 Value:


1
2
3
4
5
6
7
8
9
# 上投影: 从 d_c 维重构为完整 KV
# Key 重构
k = c_kv @ W_uk  # (d_c,) -> (n_heads * head_dim,)
# Value 重构
v = c_kv @ W_uv  # (d_c,) -> (n_heads * head_dim,)

# 然后 reshape 为 (n_heads, head_dim) 参与标准注意力计算
k = k.reshape(n_heads, head_dim)
v = v.reshape(n_heads, head_dim)

关键点在于:上下投影矩阵

1
W_dkv

1
W_uk

1
W_uv

是模型参数,在预训练阶段学习得到,推理时不变。被缓存的只有低维的

1
c_kv

,上投影在推理时动态执行。

2.3 解耦 RoPE:处理位置编码的棘手问题

上述压缩方案面临一个棘手的工程挑战:Rotary Position Embedding(RoPE)。RoPE 对 Key 施加位置相关的旋转变换,而上投影后的 Key 再做 RoPE 会导致位置信息无法被简单压缩进潜在向量——因为旋转操作与线性投影不可交换。

MLA 的解决方案是解耦注意力(Decoupled RoPE):将 Key 拆分为两部分——一部分不携带位置信息(从潜在向量上投影得到),另一部分专门携带位置信息(使用一个小维度的 RoPE Key):


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
# 解耦 RoPE 的 Key 处理
# 1. 内容 Key: 从 c_kv 上投影, 不做 RoPE
k_content = c_kv @ W_uk  # (n_heads, head_dim_no_rope)

# 2. 位置 Key: 从 h_t 单独投影, 做 RoPE
k_rope = h_t @ W_kr  # (n_heads, d_rope) 例如 d_rope=64
k_rope = apply_rope(k_rope, position=t)

# 3. 拼接为最终 Key
k_final = concat([k_content, k_rope], dim=-1)
# 最终 Key 维度: head_dim_no_rope + d_rope

# Query 也做类似解耦处理
# 缓存: c_kv (d_c=512) + k_rope (d_rope=64) = 576 维
# 对比 MHA 的 10240 维, 压缩比 ≈ 17.8x

通过这种解耦设计,MLA 在保持 RoPE 位置编码能力的同时,将需要缓存的维度从传统的 10240 降低到约 576,实现了约 17-18 倍的压缩。

三、MLA 的数学形式化

将上述过程形式化,MLA 的完整计算流程如下:

对于第 t 个 Token 的隐藏状态

1
h_t

,MLA 执行:


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
# ========== 推理阶段: 生成 Token t ==========
# 1. Query 生成 (含解耦 RoPE)
q_content = h_t @ W_qc  # (n_heads, head_dim_no_rope)
q_rope = h_t @ W_qr     # (n_heads, d_rope)
q_rope = apply_rope(q_rope, t)
q = concat([q_content, q_rope], dim=-1)

# 2. KV 压缩 (下投影) -> 只缓存这个
c_kv = h_t @ W_dkv      # (d_c,) -> 缓存!
c_kr = h_t @ W_dkr      # (d_c_rope,) -> 缓存! (RoPE部分)

# 3. 注意力计算时: 从缓存恢复 KV
k_content = c_kv @ W_uk  # (n_heads, head_dim_no_rope)
v = c_kv @ W_uv          # (n_heads, head_dim)
k_rope = c_kr @ W_ukr    # (n_heads, d_rope)
k_rope = apply_rope(k_rope, t)
k = concat([k_content, k_rope], dim=-1)

# 4. 标准注意力计算
attn = softmax(q @ k.T / sqrt(d_k)) @ v

注意,

1
W_dkv

1
W_uk

1
W_uv

1
W_dkr

1
W_ukr

等投影矩阵是跨层独立的参数,在训练中学习。每个层有自己的一组投影矩阵。

四、MLA vs GQA vs MHA:压缩比与性能对比

下表对比了三种注意力方案在 DeepSeek-V2 配置下的 KV Cache 大小:

方案 每层每 Token 缓存维度 80层总缓存(FP16) 相对压缩比 模型质量
MHA 2 × 128 × 16 = 4096 5.24 GB 1.0x (基准) 最优
GQA (8组) 2 × 128 × 8 = 2048 2.62 GB 2.0x 接近 MHA
MQA 2 × 128 × 1 = 256 0.33 GB 16.0x 显著下降
MLA 512 + 64 = 576 0.74 GB 7.1x 最优

从表中可以看出,MLA 在实现约 7-10 倍压缩比的同时,保持了与 MHA 相当的模型质量——这远超 GQA 的 2 倍压缩比,且不像 MQA 那样牺牲质量。

DeepSeek 在论文中报告的实验结果更令人振奋:在相同的预训练算力下,MLA 在多个基准测试上的表现不仅优于 GQA,甚至优于传统的 MHA。这说明低秩压缩不仅没有损害模型质量,反而因为引入了更好的归纳偏置而提升了泛化能力。

五、工程实现:从原理到推理框架

5.1 在推理框架中实现 MLA

MLA 的实现需要对推理引擎做深度改造。以下是核心的实现伪代码:


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
class MLALayer:
    def __init__(self, config):
        self.d_model = config.d_model       # 5120
        self.n_heads = config.n_heads       # 128
        self.head_dim = config.head_dim    # 128 (head_dim_no_rope=96, d_rope=32)
        self.d_c = config.d_c              # 512 (KV 压缩维度)
        self.d_c_rope = config.d_c_rope    # 64 (RoPE Key 压缩维度)

        # 下投影矩阵 (用于压缩)
        self.W_dkv = nn.Linear(self.d_model, self.d_c, bias=False)
        self.W_dkr = nn.Linear(self.d_model, self.d_c_rope, bias=False)

        # 上投影矩阵 (用于恢复)
        self.W_uk = nn.Linear(self.d_c, self.n_heads * 96, bias=False)
        self.W_uv = nn.Linear(self.d_c, self.n_heads * 128, bias=False)
        self.W_ukr = nn.Linear(self.d_c_rope, self.n_heads * 32, bias=False)

        # Query 投影
        self.W_qc = nn.Linear(self.d_model, self.n_heads * 96, bias=False)
        self.W_qr = nn.Linear(self.d_model, self.n_heads * 32, bias=False)

    def forward(self, h_t, position, kv_cache=None):
        # Query 生成
        q_c = self.W_qc(h_t).view(-1, self.n_heads, 96)
        q_r = apply_rope(self.W_qr(h_t).view(-1, self.n_heads, 32), position)
        q = torch.cat([q_c, q_r], dim=-1)  # (..., n_heads, 128)

        # KV 压缩 -> 缓存潜在向量
        c_kv = self.W_dkv(h_t)  # (batch, d_c=512)
        c_kr = self.W_dkr(h_t)  # (batch, d_c_rope=64)

        # 更新缓存 (只存压缩后的向量!)
        if kv_cache is not None:
            kv_cache['c_kv'] = torch.cat([kv_cache['c_kv'], c_kv], dim=-2)
            kv_cache['c_kr'] = torch.cat([kv_cache['c_kr'], c_kr], dim=-2)

        # 从缓存恢复 KV
        k_c = self.W_uk(kv_cache['c_kv']).view(..., -1, self.n_heads, 96)
        v = self.W_uv(kv_cache['c_kv']).view(..., -1, self.n_heads, 128)
        k_r = apply_rope(
            self.W_ukr(kv_cache['c_kr']).view(..., -1, self.n_heads, 32),
            positions
        )
        k = torch.cat([k_c, k_r], dim=-1)  # (seq, n_heads, 128)

        # 标准注意力
        attn = torch.nn.functional.scaled_dot_product_attention(q, k, v)
        return attn

5.2 计算开销的权衡

MLA 的压缩不是免费的午餐——它用额外的计算换取显存节省。每次注意力计算都需要执行上投影(从潜在向量恢复 KV),这引入了额外的矩阵乘法。对于 DeepSeek-V2 的配置:


1
2
3
4
5
6
7
8
# MLA 额外计算量 (每 Token 每层)
# 上投影: 2 * d_c * n_heads * head_dim = 2 * 512 * 128 * 128 ≈ 16.8M FLOPs
# 下投影: d_model * d_c = 5120 * 512 ≈ 2.6M FLOPs
# 总额外: ~19.4M FLOPs/layer/token

# 对比传统 MHA 的 QKV 投影:
# 3 * d_model * n_heads * head_dim = 3 * 5120 * 128 * 128 ≈ 251M FLOPs
# MLA 额外开销约 +7.7%, 但 KV Cache 减少 ~7x

在推理场景下,这个额外计算开销几乎可以忽略,因为:

  • Prefill 阶段本就是计算密集型,额外 8% 的开销影响很小。
  • Decode 阶段的瓶颈是访存而非计算,上投影的矩阵乘法可以与访存重叠。
  • 由于 KV Cache 大幅缩小,访存量也随之降低,实际上加速了 Decode 阶段

5.3 矩阵吸收优化

在实际推理中,上投影矩阵可以被”吸收”进其他矩阵中,避免显式重构 KV。对于不带 RoPE 的部分:


1
2
3
4
5
6
7
8
9
# 原始计算: q_c @ (c_kv @ W_uk).T = q_c @ W_uk.T @ c_kv.T
# 吸收后: (q_c @ W_uk.T) @ c_kv.T = q_absorbed @ c_kv.T
# 其中 q_absorbed = q_c @ W_uk.T 可以预计算

# 这样直接用 Query 的吸收形式与缓存 c_kv 做点积
# 避免了显式恢复完整的 K 矩阵

q_absorbed = q_c @ self.W_uk.weight.T  # 预计算, 与 c_kv 直接交互
# attn_scores = q_absorbed @ c_kv.T  # 直接用压缩向量计算!

这种矩阵吸收技术使得 MLA 在 Decode 阶段可以完全不显式恢复 KV,直接在压缩空间中计算注意力分数,进一步降低了计算和访存开销。这也是 DeepSeek-V2 能在保持高推理质量的同时实现高吞吐的关键工程技巧。

六、生产环境部署实践

6.1 MLA 与 PagedAttention 的结合

在 vLLM、SGLang 等推理框架中,MLA 需要与 PagedAttention(分页注意力)结合使用。由于 MLA 缓存的是潜在向量而非完整 KV,分页的粒度需要调整:


1
2
3
4
5
6
7
8
9
# 传统 PagedAttention 的 block 大小:
# block_size = 16 tokens
# 每块 KV Cache: 16 * 2 * n_heads * head_dim * n_layers (FP16)

# MLA PagedAttention 的 block 大小:
# block_size = 16 tokens
# 每块缓存: 16 * (d_c + d_c_rope) * n_layers (FP16)
# DeepSeek-V2: 16 * (512 + 64) * 80 * 2 = 1.47 MB (vs MHA 的 10.5 MB)
# 单张 A100 80GB 可服务的并发请求数提升约 7 倍

6.2 在 SGLang 中启用 MLA

以下是在 SGLang 框架中加载 DeepSeek-V2/V3 模型的实际配置:


1
2
3
4
5
6
7
8
# 启动 SGLang 服务, 加载 DeepSeek-V3 (支持 MLA)
python -m sglang.launch_server   --model-path deepseek-ai/DeepSeek-V3   --tp 8   --trust-remote-code   --enable-mla   --max-context-len 131072   --mem-fraction-static 0.88

# 关键参数说明:
# --enable-mla: 启用 MLA 优化路径
# --mem-fraction-static: 提高 KV Cache 可用显存比例
#   (因为 MLA 压缩后 KV Cache 占比大幅下降, 可安全提高)
# --tp 8: 张量并行 8 卡 (DeepSeek-V3 推荐)

6.3 与 Prefix Caching 的兼容性

MLA 与 Prefix Caching 的结合非常自然——由于缓存的是低维潜在向量,相同前缀的请求可以高效共享缓存块,且共享粒度更细。在 DeepSeek 官方的推理部署中,MLA + Prefix Caching + Continuous Batching 的组合实现了业界领先的吞吐效率:


1
2
3
4
5
6
7
8
9
# 吞吐对比 (以 MHA 为基准, 同等硬件 A100×8)
# MHA (基准):     1.0x
# GQA (8组):      2.3x
# MLA + Prefix Caching: 8.5x

# 延迟对比 (Decode 阶段, 单请求)
# MHA:  42 ms/token
# GQA:  38 ms/token
# MLA:  31 ms/token (访存减少带来加速)

七、MLA 的局限与未来方向

尽管 MLA 在 KV Cache 压缩方面取得了突破性进展,但它并非完美方案,在实际落地中仍有需要注意的方面:

  • 对框架的侵入性:MLA 需要对推理引擎做深度改造,不能像 GQA 那样即插即用。目前主流框架(vLLM、SGLang、TensorRT-LLM)均已支持,但自研框架需要自行实现吸收优化等关键技巧。
  • 训练阶段的显存节省有限:MLA 主要优化推理阶段的 KV Cache。训练时由于需要梯度计算和注意力矩阵的完整形式化,显存节省不如推理阶段显著。
  • 跨层 KV 共享的探索:当前 MLA 仍是在每层独立压缩。如果不同层之间的 KV Cache 存在相关性,进一步跨层压缩可能带来更多收益——但这需要在模型质量与压缩比之间重新权衡。
  • 与 Quantization 的叠加:MLA 的低维潜在向量天然适合量化。将 c_kv 量化为 INT4 或 INT8 可以进一步压缩缓存,目前业界已有相关工作探索这一方向。

八、总结

DeepSeek 的 MLA 代表了大模型推理优化的一个重要范式转变:从减少 KV 头数(GQA/MQA)到压缩 KV 信息表示(低秩投影)。这种思路在信息论上是优雅的——它不丢弃信息,而是找到了信息的更紧凑表示。

对于 AI Infra 工程师而言,理解 MLA 的核心价值在于:在 70B+ 规模的模型上,KV Cache 压缩带来的显存节省直接转化为并发能力的提升和推理成本的降低。DeepSeek-V3 通过 MLA + MoE 的组合,在保持 671B 总参数规模的同时,将活跃参数控制在 37B,KV Cache 压缩到传统方案的 1/10 以下,实现了远超同规模模型的推理效率。

展望未来,MLA 的低秩压缩思路可能会启发更多注意力变体——正如 FlashAttention 重新定义了注意力的计算方式,MLA 正在重新定义注意力信息的存储方式。对于致力于大模型推理优化的工程师,深入理解 MLA 的数学原理和工程实现,已经是不可跳过的一课。

【本站文章皆为原创,未经允许不得转载】:汤不热吧 » DeepSeek MLA(多头潜在注意力)深度解析:如何用低秩压缩将 KV Cache 压缩到原来的 1/10
分享到: 更多 (0)