多头潜注意力(MLA)是什么
一句话定义:多头潜注意力(Multi-head Latent Attention,MLA)是 DeepSeek 在 DeepSeek-V2 中提出的一种注意力机制,它把推理时要缓存的 K、V 两个大矩阵联合压缩成一个低维的"潜向量",需要用时再"上采样"还原成 K、V,从而把 KV 缓存压到极小。
为什么需要它:KV 缓存的烦恼
先说一个背景知识。LLM 在做"下一个词预测"时,对前面每一个 token 都要保留两份信息:一把钥匙(Key,K)和一个数值(Value,V)。新 token 进来时,拿自己的"问题"(Query,Q)去和所有历史 token 的 K 比对,算出注意力权重,再从 V 里取信息。
这些历史 K/V 就是所谓的 KV 缓存(KV cache)。它的麻烦在于:上下文越长,要存的 K/V 越多,且成线性增长。
举一组直观数字(DeepSeek-V2 配置):128 个注意力头,每头维度 128,单个 token 的 K/V 就要存 2 × 128 × 128 = 32,768 个元素。如果上下文有 12.8 万个 token,光是 KV 缓存就要存数十亿个浮点数,显存压力极大,长文本场景下尤其吃紧。
类比:你在写一篇长论文,每翻一份资料就摊在桌上不收。资料越多,桌面越爆满,最后根本铺不开。KV 缓存就是这张越来越满的"桌子"。
前一个解法:GQA 靠"共享头"
在 MLA 之前,业界主流的省显存方案是 分组查询注意力(Grouped-Query Attention,GQA),Llama 系列就采用它。
GQA 的思路:多个 Query 头共享同一组 K/V 头。比如 32 个 Q 头共享 8 个 K/V 头,KV 缓存直接降到原来的四分之一。
这思路有效但有限:它只能靠"少几个头"来省,压缩比有天花板。类比一下,相当于让几个同事共用一本字典——字典本数少了,但每本还是原封不动的厚度。
MLA 的核心思路:低秩潜变量压缩
MLA 不走"共享头"那条路,而是换了个思路:K/V 本身就能被压缩。
直觉是这样的:K 和 V 都是从同一个输入向量算出来的,它们之间有大量冗余信息。既然如此,何必非要把完整 K、V 都存下来?只存一个"浓缩精华"的潜向量
类比:与其在桌上摊 100 本厚厚的资料,不如把它们拍照扫描、压缩成一个 50 MB 的 zip 包存硬盘,要用时再解压拿出来。桌面(显存)瞬间清爽。
具体怎么做?分两步。
1. 下采样(压缩):把 K/V 联合压成一个潜向量
符号逐项解读:
:token 在序列里的位置编号(第几个词),从 1 数到序列长度;本文所有带 下标的量,都特指「第 个 token」的那一份。 :第 个 token 的输入向量(模型里流动的隐藏状态),维度很高,比如 。 :下采样矩阵("DKV" 表示 Down-projection for KV),形状是 ,负责把高维压到低维。 :压缩后的 KV 潜向量,维度 远小于 。
关键:K 和 V 共用同一个潜向量
这个
2. 上采样(还原):从潜向量重建 K 和 V
用到的时候,再用两个上采样矩阵把
:把潜向量还原成 K 的上采样矩阵("UK" = Up-projection for K)。 :把潜向量还原成 V 的上采样矩阵("UV" = Up-projection for V)。 :上一步压好的 KV 潜向量,在这里作为还原的「原料」。 :每个注意力头的维度(单个头里 K/V 向量的长度)。 、 :重建出来的第 个 token 的 K 和 V,喂进标准注意力公式算权重。
通俗理解:下采样是"打包压缩成 zip",上采样是"解压还原"。推理时只把 zip(潜向量)存进显存,要用时才解压。
3. 一个聪明的优化:上采样矩阵可以"吸收"
按上面写法,每次推理似乎要先解压出 K,再算
把
一组真实数字:压缩到底有多狠?
DeepSeek-V2 的配置(论文 3.1.2 节):
| 参数 | 值 | 含义 |
|---|---|---|
| 128 | 注意力头数 | |
| 128 | 每头维度 | |
| 512 | KV 潜向量维度 | |
| 64 | 给 RoPE 位置编码额外留的维度 |
对比标准多头注意力(MHA):单个 token、单层的 KV 缓存元素数 =
MLA:单个 token、单层只需要缓存
比例:
图 1:单个 token 单层的 KV 缓存元素数对比——MHA 要 32,768 个,MLA 只需 576 个(≈ 1.8%),MLA 缓存大约只有 MHA 的 1/57。
和 GQA 的核心区别
| 维度 | GQA | MLA |
|---|---|---|
| 省显存的招 | 多个 Q 头共享同一组 K/V 头 | 把 K/V 压缩成低维潜向量,缓存潜向量 |
| 压缩上限 | 受限于头数(最多压到 1 个组 = MQA) | 受限于潜维度(可压到极小) |
| 长上下文表现 | 头共享会损失信息,长上下文受限 | 潜变量保留信息更充分,长上下文优势显著 |
| 用在哪些模型 | Llama 2/3、Qwen 等 | DeepSeek-V2、DeepSeek-V3 |
一句话:GQA 靠"少几个头",MLA 靠"压一维"。
图 2:两种省 KV 缓存的思路对比——GQA 靠「少几个头」(多个 Q 头共享同一组 K/V 头,每本字典还是原厚度),MLA 靠「压一维」(把 K/V 压成低维潜向量,缓存极小、用时再还原)。
在 AI 体系中的位置
MLA 属于大模型注意力机制这一支的工程优化,是对标准 MHA 和 GQA 的进一步演进。它属于"推理时显存优化"领域,和 Paged Attention、Flash Attention 等技术互补(一个省存储,一个省算力)。
提出方 DeepSeek 在 V2、V3 两代模型上都用 MLA,验证了它在超长上下文场景(百万级 token 上下文)下的显存优势。后续研究(如 MHA2MLA)也在探索如何把已有的 MHA 模型微调成 MLA,进一步推广这个机制。
一个简化的计算示例
为了看清压缩怎么发生,用一组极小的数字演示。
假设
下采样矩阵(
压缩:
原本 4 维,现在压成 2 维。
上采样矩阵(
重建:
自检:原本要缓存 K、V 各 3 个元素,共 6 个浮点数;现在只缓存
图 3:压缩示例——4 维输入 h 经下采样压成 2 维潜向量 c_KV(只缓存它),再分别上采样还原成 3 维的 k 和 v;原本要缓存 6 个数,现在只缓存 2 个。
PyTorch 代码示例
先看核心步骤对应的代码片段。下采样 + 上采样:
import torch
import torch.nn as nn
# 下采样:把 h 压成潜向量 c_KV(对应公式 c_KV = W_DKV · h)
self.W_DKV = nn.Linear(d_model, d_c, bias=False)
c_KV = self.W_DKV(h) # shape: (batch, seq, d_c) —— 推理时只缓存这个
# 上采样:从潜向量还原 K 和 V(对应 k = W_UK · c_KV, v = W_UV · c_KV)
self.W_UK = nn.Linear(d_c, n_heads * d_n, bias=False)
self.W_UV = nn.Linear(d_c, n_heads * d_n, bias=False)
k = self.W_UK(c_KV).view(batch, seq, n_heads, d_n)
v = self.W_UV(c_KV).view(batch, seq, n_heads, d_n)标准注意力计算(K/V 重建后照常算):
# 注意力权重 = Q · K^T / sqrt(d_n),再 softmax,再加权 V
scores = torch.einsum('bqhd,bkhd->bhqk', q, k) / (d_n ** 0.5)
attn = torch.softmax(scores, dim=-1)
out = torch.einsum('bhqk,bkhd->bqhd', attn, v)完整代码
下面给一份可复制可跑的极简版(省略 RoPE 解耦部分,聚焦"压缩—还原"主流程):
import torch
import torch.nn as nn
class MLASimpleAttention(nn.Module):
"""MLA 的极简实现(省略 RoPE 解耦部分,只展示核心压缩-还原流程)"""
def __init__(self, d_model=64, d_c=16, n_heads=4, d_n=16):
super().__init__()
self.n_heads = n_heads
self.d_n = d_n
# Q 路径(也可压缩,这里简化为直接投影)
self.W_Q = nn.Linear(d_model, n_heads * d_n, bias=False)
# KV 联合压缩:d_model -> d_c
self.W_DKV = nn.Linear(d_model, d_c, bias=False)
# 上采样:d_c -> K 和 V 各 n_heads * d_n
self.W_UK = nn.Linear(d_c, n_heads * d_n, bias=False)
self.W_UV = nn.Linear(d_c, n_heads * d_n, bias=False)
# 输出投影
self.W_O = nn.Linear(n_heads * d_n, d_model, bias=False)
def forward(self, x):
batch, seq, _ = x.shape
# 1) 算 Q
q = self.W_Q(x).view(batch, seq, self.n_heads, self.d_n)
# 2) 下采样:把 x 压成潜向量 c_KV —— 推理时只缓存它
c_KV = self.W_DKV(x) # (batch, seq, d_c) ← KV cache 只存这个
# 3) 上采样:从 c_KV 还原 K、V
k = self.W_UK(c_KV).view(batch, seq, self.n_heads, self.d_n)
v = self.W_UV(c_KV).view(batch, seq, self.n_heads, self.d_n)
# 4) 标准注意力(为清晰起见,这里没展示"把 W_UK 吸收进 W_Q"的优化)
scores = torch.einsum('bqhd,bkhd->bhqk', q, k) / (self.d_n ** 0.5)
attn = torch.softmax(scores, dim=-1)
out = torch.einsum('bhqk,bkhd->bqhd', attn, v).reshape(batch, seq, -1)
return self.W_O(out)
# 跑一下
torch.manual_seed(0) # 固定随机种子,结果可复现
model = MLASimpleAttention()
x = torch.randn(2, 10, 64) # batch=2, seq=10, d_model=64
out = model(x)
print(out.shape) # torch.Size([2, 10, 64]) —— 输入输出同形状
# 训练一步(演示反向传播能跑通)
loss = out.mean()
loss.backward() # 反向传播,计算所有参数的梯度
print("反向传播成功,梯度已计算")小结
MLA 的核心就一句话:别再傻乎乎地存完整的 K 和 V,把它们压成一个低维潜向量存起来,要用时再还原。它把"省 KV 缓存"这件事从 GQA 的"共享头"推进到了"低秩压缩"的新维度,让超长上下文成为可能,是 DeepSeek-V2/V3 的关键技术之一。
参考资料
- DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model - DeepSeek-AI https://arxiv.org/abs/2405.04434
- Multi-Head Latent Attention (MLA) - Sebastian Raschka https://sebastianraschka.com/llms-from-scratch/ch04/05_mla/
- A Gentle Introduction to Multi-Head Latent Attention (MLA) - Machine Learning Mastery https://machinelearningmastery.com/a-gentle-introduction-to-multi-head-latent-attention-mla/