MLA 是什么

MLA封面

多头潜注意力(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 都存下来?只存一个"浓缩精华"的潜向量 cc 就够了,要用时再"稀释"还原出 K/V

类比:与其在桌上摊 100 本厚厚的资料,不如把它们拍照扫描、压缩成一个 50 MB 的 zip 包存硬盘,要用时再解压拿出来。桌面(显存)瞬间清爽。

具体怎么做?分两步。

1. 下采样(压缩):把 K/V 联合压成一个潜向量

cKV,t=WDKVhtc_{KV,t} = W^{DKV} \cdot h_t

符号逐项解读:

  • tt:token 在序列里的位置编号(第几个词),从 1 数到序列长度;本文所有带 tt 下标的量,都特指「第 tt 个 token」的那一份。
  • hth_t:第 tt 个 token 的输入向量(模型里流动的隐藏状态),维度很高,比如 dmodel=5120d_{model} = 5120
  • WDKVW^{DKV}下采样矩阵("DKV" 表示 Down-projection for KV),形状是 dc×dmodeld_c \times d_{model},负责把高维压到低维。
  • cKV,tc_{KV,t}:压缩后的 KV 潜向量,维度 dcd_c 远小于 dmodeld_{model}

关键:K 和 V 共用同一个潜向量 cKV,tc_{KV,t},这就是"联合压缩"。

这个 cKV,tc_{KV,t} 就是推理时要缓存进显存的全部内容,替代了原本的 K 和 V 两个大矩阵。

2. 上采样(还原):从潜向量重建 K 和 V

用到的时候,再用两个上采样矩阵把 cKV,tc_{KV,t} 分别还原成 K 和 V:

kt=WUKcKV,t,vt=WUVcKV,tk_t = W^{UK} \cdot c_{KV,t}, \quad v_t = W^{UV} \cdot c_{KV,t}
  • WUKW^{UK}:把潜向量还原成 K 的上采样矩阵("UK" = Up-projection for K)。
  • WUVW^{UV}:把潜向量还原成 V 的上采样矩阵("UV" = Up-projection for V)。
  • cKV,tc_{KV,t}:上一步压好的 KV 潜向量,在这里作为还原的「原料」。
  • dnd_n:每个注意力头的维度(单个头里 K/V 向量的长度)。
  • ktk_tvtv_t:重建出来的第 tt 个 token 的 K 和 V,喂进标准注意力公式算权重。

通俗理解:下采样是"打包压缩成 zip",上采样是"解压还原"。推理时只把 zip(潜向量)存进显存,要用时才解压。

3. 一个聪明的优化:上采样矩阵可以"吸收"

按上面写法,每次推理似乎要先解压出 K,再算 QKQ \cdot K^\top。但 DeepSeek 团队观察到一个数学事实:矩阵乘法可以合并

kt=WUKcKV,tk_t = W^{UK} c_{KV,t} 代入 QktQ \cdot k_t^\top,得到 Q(WUK)cKV,tQ \cdot (W^{UK})^\top \cdot c_{KV,t}^\top(右上角的 ^\top 是「转置」记号,表示把矩阵的行与列互换,这里出现是因为点积 QktQ \cdot k_t^\top 要求维度对齐)。这意味着可以把 WUKW^{UK} 提前合并进 Q 的投影矩阵里(叫做"吸收",absorb),推理时根本不用真的解压出 K,直接用潜向量 cKV,tc_{KV,t} 算注意力。这把上采样的算力开销也省掉了。

一组真实数字:压缩到底有多狠?

DeepSeek-V2 的配置(论文 3.1.2 节):

参数含义
nhn_h128注意力头数
dhd_h128每头维度
dcd_c512KV 潜向量维度
dhRd_h^R64给 RoPE 位置编码额外留的维度

对比标准多头注意力(MHA):单个 token、单层的 KV 缓存元素数 = 2nhdh=2×128×128=32,7682 \cdot n_h \cdot d_h = 2 \times 128 \times 128 = 32{,}768

MLA:单个 token、单层只需要缓存 dc+dhR=512+64=576d_c + d_h^R = 512 + 64 = 576 个元素(潜向量 + RoPE 部分)。

比例576/32,7681.8%576 / 32{,}768 \approx 1.8\%,相当于只剩零头。论文官方给出的等效水平:和 GQA 只有约 2.25 个组时一样省,但模型质量反而比标准 MHA 更强。

MLA 与 MHA 单 token KV 缓存元素数比例条

图 1:单个 token 单层的 KV 缓存元素数对比——MHA 要 32,768 个,MLA 只需 576 个(≈ 1.8%),MLA 缓存大约只有 MHA 的 1/57。

和 GQA 的核心区别

维度GQAMLA
省显存的招多个 Q 头共享同一组 K/V 头把 K/V 压缩成低维潜向量,缓存潜向量
压缩上限受限于头数(最多压到 1 个组 = MQA)受限于潜维度(可压到极小)
长上下文表现头共享会损失信息,长上下文受限潜变量保留信息更充分,长上下文优势显著
用在哪些模型Llama 2/3、Qwen 等DeepSeek-V2、DeepSeek-V3

一句话:GQA 靠"少几个头",MLA 靠"压一维"

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,进一步推广这个机制。

一个简化的计算示例

为了看清压缩怎么发生,用一组极小的数字演示。

假设 dmodel=4d_{model} = 4,潜维度 dc=2d_c = 2,每头维度 dn=3d_n = 3,1 个头。输入 h=[1,0,2,1]h = [1, 0, 2, -1]

下采样矩阵2×42 \times 4,挑出第 0、2 维):

WDKV=[10000010]W^{DKV} = \begin{bmatrix} 1 & 0 & 0 & 0 \\ 0 & 0 & 1 & 0 \end{bmatrix}

压缩:

cKV=WDKVh=[1,2]c_{KV} = W^{DKV} \cdot h = [1, 2]

原本 4 维,现在压成 2 维。

上采样矩阵3×23 \times 2):

WUK=[100111],WUV=[110110]W^{UK} = \begin{bmatrix} 1 & 0 \\ 0 & 1 \\ 1 & 1 \end{bmatrix}, \quad W^{UV} = \begin{bmatrix} 1 & 1 \\ 0 & 1 \\ 1 & 0 \end{bmatrix}

重建:

k=WUKcKV=[1,2,3],v=WUVcKV=[3,2,1]k = W^{UK} \cdot c_{KV} = [1, 2, 3], \quad v = W^{UV} \cdot c_{KV} = [3, 2, 1]

自检:原本要缓存 K、V 各 3 个元素,共 6 个浮点数;现在只缓存 cKVc_{KV} 的 2 个。压缩比 6/2=36 / 2 = 3 倍。真实场景下,DeepSeek-V2 的压缩比是这个示例的几十倍。

MLA 压缩还原计算示例向量流

图 3:压缩示例——4 维输入 h 经下采样压成 2 维潜向量 c_KV(只缓存它),再分别上采样还原成 3 维的 k 和 v;原本要缓存 6 个数,现在只缓存 2 个。

PyTorch 代码示例

先看核心步骤对应的代码片段。下采样 + 上采样:

python
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 重建后照常算):

python
# 注意力权重 = 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 解耦部分,聚焦"压缩—还原"主流程):

python
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 的关键技术之一。

参考资料

  1. DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model - DeepSeek-AI https://arxiv.org/abs/2405.04434
  2. Multi-Head Latent Attention (MLA) - Sebastian Raschka https://sebastianraschka.com/llms-from-scratch/ch04/05_mla/
  3. A Gentle Introduction to Multi-Head Latent Attention (MLA) - Machine Learning Mastery https://machinelearningmastery.com/a-gentle-introduction-to-multi-head-latent-attention-mla/