在大模型推理的工程实践中,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 的数学原理和工程实现,已经是不可跳过的一课。
汤不热吧