Zer0e's Blog

去TM的AI其三:什么是多头潜在注意力Multi-Head Latent Attention (MLA)

字数统计: 3.4k阅读时长: 15 min
2026/07/09 Share

Multi-Head Latent Attention(MLA,多头潜在注意力)是由 DeepSeek 团队在 DeepSeek-V2 模型中提出的一种创新注意力机制。它在保持多头注意力强大表达能力的同时,显著降低了推理时的 KV Cache 内存占用,是大型语言模型效率优化的重要突破。

为什么要提出 MLA?

在传统的 Multi-Head Attention(MHA)中,每个注意力头都需要独立存储 Key 和 Value 的 KV Cache。对于大语言模型来说,这带来了严重的内存瓶颈:

传统 MHA 的 KV Cache 问题

假设我们有一个模型:

  • 模型维度 $$(d_{\text{model}} = 5120)$$
  • 注意力头数 $$(h = 128)$$
  • 序列长度 $$(L = 8192)$$
  • 每个头的维度 $$(d_k = d_v = 64)$$

KV Cache 的大小计算:
$$
[\text{KV Cache Size} = 2 \times L \times h \times d_k \times \text{bytes}]
$$

代入数值:
$$
[2 \times 8192 \times 128 \times 64 \times 2 \text{ bytes} \approx 26.8 \text{ GB}]
$$

这意味着仅仅存储 KV Cache 就需要 26.8 GB 的显存!这对于长序列生成任务来说是巨大的负担。

flowchart TD
    A[传统 MHA] --> B[每个头独立存储 KV]
    B --> C[128 个头 × 64 维]
    C --> D[KV Cache 巨大]
    D --> E1[显存占用高]
    D --> E2[批处理大小受限]
    D --> E3[推理速度慢]
    
    E1 --> F[部署成本高昂]
    E2 --> F
    E3 --> F
    
    style A fill:#ffcdd2
    style D fill:#ff8a65
    style F fill:#d32f2f,color:#fff

MLA 的核心洞察

MLA 的核心思想是:大部分 Key 和 Value 的信息可以共享,不需要每个头都独立存储。通过引入低秩压缩技术,MLA 将 KV Cache 大幅压缩,同时保持模型性能几乎不损失。

MLA 的核心思想

MLA 通过以下三个关键技术实现高效的注意力计算:

1. 低秩 KV 压缩(Low-Rank KV Compression)

MLA 将 Key 和 Value 投影到一个低维的潜在空间(Latent Space),然后在这个压缩空间存储 KV Cache:

flowchart TD
    subgraph Traditional["传统 MHA"]
        A1[输入 X] --> B1[Q = X·Wᵠ]
        A1 --> C1[K = X·Wᴷ]
        A1 --> D1[V = X·Wⱽ]
        C1 --> E1[存储完整 KV
h × d_k 维] D1 --> E1 end subgraph MLA["MLA"] A2[输入 X] --> B2[Q = X·Wᵠ] A2 --> C2[K_latent = X·Wᴷ_latent] A2 --> D2[V_latent = X·Wⱽ_latent] C2 --> E2[存储压缩 KV
低维潜在空间] D2 --> E2 E2 --> F2[解码时恢复
K = K_latent·Wᴷ_up] F2 --> G2[计算注意力] end style Traditional fill:#ffebee style MLA fill:#e8f5e9 style E1 fill:#ffcdd2 style E2 fill:#c8e6c9

2. 解耦的 Query 表示(Decoupled Query Representation)

MLA 将 Query 分为两部分:

  • 共享部分:所有头共享的潜在表示
  • 独立部分:每个头特有的表示

这种解耦设计使得模型既能捕捉共享的语义信息,又能保持每个头的独特性。

3. 高效的注意力计算

在推理时,MLA 通过以下步骤计算注意力:

  1. 从压缩的潜在空间恢复 Key 和 Value
  2. 使用恢复的 KV 计算注意力分数
  3. 得到加权求和的输出
flowchart LR
    subgraph Decode["解码阶段"]
        A1[潜在向量 K_c] --> B1[上投影 W_k_up]
        B1 --> C1[恢复 K]
        A2[潜在向量 V_c] --> B2[上投影 W_v_up]
        B2 --> C2[恢复 V]
    end
    
    C1 --> D["计算注意力 softmax(QK^T/√d)·V"]
    C2 --> D
    D --> E[输出]
    
    style Decode fill:#e3f2fd
    style D fill:#fff3e0
    style E fill:#e8f5e9

MLA 的数学原理

1. 潜在空间压缩

MLA 首先将输入 (X) 投影到低维潜在空间:

$$
[K_c = X W_c^K, \quad V_c = X W_c^V]
$$

其中:

  • $$(K_c, V_c \in \mathbb{R}^{d_c}) 是压缩后的潜在向量$$
  • $$(d_c \ll d_k) 是潜在空间维度(通常 (d_c) 约为 (d_k) 的 1/10)$$
  • $$(W_c^K, W_c^V) 是压缩投影矩阵$$

2. Query 的解耦设计

MLA 将 Query 分为共享部分和独立部分:

$$
[Q_i = X W_{shared}^Q + X W_i^Q]
$$

其中:

  • $$(X W_{shared}^Q) 是所有头共享的 Query 部分$$
  • $$(X W_i^Q) 是第 (i) 个头特有的 Query 部分$$

3. Key 和 Value 的恢复

在计算注意力时,MLA 从潜在空间恢复完整的 Key 和 Value:

$$
[K_i = K_c W_{i}^{K,\text{up}}, \quad V_i = V_c W_{i}^{V,\text{up}}]
$$

其中 $$(W_{i}^{K,\text{up}}) 和 (W_{i}^{V,\text{up}}) $$是上投影矩阵,用于从低维空间恢复到高维空间。

4. 注意力计算

恢复后的 Key 和 Value 用于标准的注意力计算:

$$
[\text{Attention}(Q_i, K_i, V_i) = \text{softmax}\left(\frac{Q_i K_i^T}{\sqrt{d_k}}\right) V_i]
$$

完整的 MLA 公式

将所有头合并,MLA 的完整计算过程为:

$$
[\text{MLA}(Q, K_c, V_c) = \text{Concat}(\text{head}_1, …, \text{head}_h) W^O]
$$

其中第 (i) 个头:

$$
[\text{head}_i = \text{Attention}(Q_i, K_c W_i^{K,\text{up}}, V_c W_i^{V,\text{up}})]
$$

MLA 的工作流程

让我们通过详细的流程图来理解 MLA 的完整计算过程:

flowchart TD
    subgraph Training["训练阶段"]
        A1["输入 X"] --> B1["压缩投影"]
        B1 --> C1["K_c = X·W_c^K\nV_c = X·W_c^V"]
        A1 --> D1["Query 解耦投影"]
        D1 --> E1["Q_i = X·W_shared^Q + X·W_i^Q"]
        C1 --> F1["上投影恢复"]
        F1 --> G1["K_i = K_c·W_i^{K,up}\nV_i = V_c·W_i^{V,up}"]
        E1 --> H1["计算注意力"]
        G1 --> H1
        H1 --> I1["softmax(Q_i·K_i^T/√d)·V_i"]
    end
    
    subgraph Inference["推理阶段-KV Cache 存储"]
        A2["只存储压缩的 KV"] --> B2["K_c: d_c 维\nV_c: d_c 维"]
        B2 --> C2["显存大幅降低"]
    end
    
    subgraph Decode["推理阶段-解码"]
        A3["从 K_c,V_c 恢复"] --> B3["K_i = K_c·W_i^{K,up}"]
        A3 --> C3["V_i = V_c·W_i^{V,up}"]
        B3 --> D3["计算注意力"]
        C3 --> D3
    end
    
    style Training fill:#e3f2fd
    style Inference fill:#e8f5e9
    style Decode fill:#fff3e0
    style C2 fill:#66bb6a,color:#fff

详细步骤解析

步骤 1:低秩压缩(Low-Rank Compression)

flowchart LR
    A[输入 X
d_model 维] --> B[低秩投影 W_cᴷ] A --> C[低秩投影 W_cⱽ] B --> D[K_c
d_c 维, d_c << d_k] C --> E[V_c
d_c 维, d_c << d_k] style A fill:#bbdefb style D fill:#c8e6c9 style E fill:#c8e6c9

这是 MLA 的核心创新:将高维的 Key 和 Value 压缩到低维潜在空间,大幅减少存储需求。

步骤 2:Query 解耦投影(Decoupled Q Projection)

flowchart TD
    A[输入 X] --> B[共享 Query 投影 W_sharedᵠ]
    A --> C[独立 Query 投影 W_iᵠ]
    B --> D[Q_shared]
    C --> E[Q_i_independent]
    D --> F[相加合并]
    E --> F
    F --> G[Q_i = Q_shared + Q_i_independent]
    
    style A fill:#ffe0b2
    style F fill:#ce93d8
    style G fill:#81c784

这种设计使得所有头共享一部分语义理解,同时保持各自的独特性。

步骤 3:训练时的上投影恢复(Up-Projection Recovery)

flowchart LR
    subgraph Latent["潜在空间 (存储)"]
        A1[K_c: d_c 维]
        A2[V_c: d_c 维]
    end
    
    subgraph Recovery["恢复"]
        A1 --> B1[上投影 W_iᴷ_up]
        A2 --> B2[上投影 W_iⱽ_up]
        B1 --> C1[K_i: d_k 维]
        B2 --> C2[V_i: d_k 维]
    end
    
    style Latent fill:#e8f5e9
    style Recovery fill:#fff3e0
    style C1 fill:#bbdefb
    style C2 fill:#bbdefb

训练时,模型学习如何通过上投影矩阵从低维空间恢复高维表示。

步骤 4:注意力计算(Attention Computation)

flowchart LR
    A[Q_i] --> C[计算注意力分数]
    B[K_i] --> C
    C --> D["softmax(Q_i·K_i^T/√d_k)"]
    D --> E[加权求和]
    F[V_i] --> E
    E --> G[Output_i]
    
    style A fill:#ffcc80
    style B fill:#ffcc80
    style F fill:#ffcc80
    style G fill:#a5d6a7

使用恢复的 Key 和 Value 进行标准的缩放点积注意力计算。

KV Cache 对比:MHA vs MLA

让我们通过具体数值对比两种方法的 KV Cache 大小:

参数设置

  • 序列长度 $$(L = 8192)$$
  • 注意力头数 $$(h = 128)$$
  • 每个头维度 $$(d_k = 64)$$
  • 潜在空间维度 $$(d_c = 8)(MLA 使用)$$
  • 数据类型:FP16(2 bytes)

传统 MHA 的 KV Cache

$$
[\text{KV Cache}_{\text{MHA}} = 2 \times L \times h \times d_k \times 2]
$$
$$
[= 2 \times 8192 \times 128 \times 64 \times 2 = 26.8 \text{ GB}]
$$

MLA 的 KV Cache

$$
[\text{KV Cache}_{\text{MLA}} = 2 \times L \times d_c \times 2]
$$
$$
[= 2 \times 8192 \times 8 \times 2 = 0.25 \text{ GB}]
$$

压缩效果

flowchart TD
    A[KV Cache 对比] --> B1["MHA: 26.8 GB"]
    A --> B2["MLA: 0.25 GB"]
    B1 --> C1[❌ 显存占用巨大]
    B2 --> C2[✅ 显存占用极小]
    C2 --> D[压缩比: 107 倍]
    
    style B1 fill:#ffcdd2
    style B2 fill:#c8e6c9
    style D fill:#66bb6a,color:#fff

MLA 将 KV Cache 压缩了约 107 倍! 这是一个巨大的进步。

MLA 的优势

1. 显存效率大幅提升

flowchart LR
    A[MLA 低秩压缩] --> B[KV Cache 减少 100+ 倍]
    B --> C1[更大批处理]
    B --> C2[更长序列]
    B --> C3[更低部署成本]
    
    C1 --> D[推理吞吐量提升]
    C2 --> D
    C3 --> D
    
    style A fill:#e3f2fd
    style B fill:#4caf50,color:#fff
    style D fill:#ff9800

2. 支持更长的序列

由于 KV Cache 大幅减少,MLA 使得模型能够处理更长的序列而不会遇到显存瓶颈。

3. 提高推理吞吐量

更小的 KV Cache 意味着:

  • 可以使用更大的批处理大小(batch size)
  • 减少显存带宽压力
  • 提高 GPU 利用率

4. 保持模型性能

尽管 KV Cache 被大幅压缩,MLA 通过精心设计的低秩投影和解耦 Query 机制,保持了与标准 MHA 相当的性能。

flowchart TD
    A[MLA 设计] --> B1[低秩压缩
减少存储] A --> B2[解耦 Query
保持表达力] A --> B3[上投影恢复
完整注意力] B1 --> C[高效 + 高性能] B2 --> C B3 --> C C --> D1[✅ 显存效率] C --> D2[✅ 推理速度] C --> D3[✅ 模型质量] style A fill:#ff8a65 style C fill:#4db6ac,color:#fff style D1 fill:#81c784 style D2 fill:#64b5f6 style D3 fill:#ba68c8

MLA 与传统注意力的架构对比

flowchart TD
    subgraph MHA["Multi-Head Attention"]
        A1[输入 X] --> B1[Q = X·Wᵠ]
        A1 --> C1[K = X·Wᴷ]
        A1 --> D1[V = X·Wⱽ]
        C1 --> E1[存储 K: h×d_k]
        D1 --> E1
        B1 --> F1[计算注意力]
        E1 --> F1
    end
    
    subgraph MLA["Multi-Head Latent Attention"]
        A2[输入 X] --> B2[Q_i = X·W_sharedᵠ + X·W_iᵠ]
        A2 --> C2[K_c = X·W_cᴷ
V_c = X·W_cⱽ] C2 --> D2[存储 K_c, V_c: d_c] D2 --> E2[恢复 K_i = K_c·W_iᴷ_up] D2 --> F2[恢复 V_i = V_c·W_iⱽ_up] B2 --> G2[计算注意力] E2 --> G2 F2 --> G2 end style MHA fill:#ffebee style MLA fill:#e8f5e9 style E1 fill:#ffcdd2 style D2 fill:#c8e6c9

代码实现(PyTorch 简化版)

下面是一个简化版的 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
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
import torch
import torch.nn as nn
import math

class MultiHeadLatentAttention(nn.Module):
def __init__(self, d_model=512, num_heads=8, d_latent=8):
super(MultiHeadLatentAttention, self).__init__()
self.num_heads = num_heads
self.d_model = d_model
self.d_k = d_model // num_heads
self.d_latent = d_latent # 潜在空间维度

# 1. 低秩压缩投影
self.W_c_K = nn.Linear(d_model, d_latent)
self.W_c_V = nn.Linear(d_model, d_latent)

# 2. Query 解耦投影
self.W_shared_Q = nn.Linear(d_model, d_model)
self.W_head_Q = nn.Linear(d_model, d_model)

# 3. 上投影矩阵(每个头独立)
self.W_K_up = nn.Parameter(torch.randn(num_heads, d_latent, self.d_k))
self.W_V_up = nn.Parameter(torch.randn(num_heads, d_latent, self.d_k))

# 4. 输出投影
self.W_O = nn.Linear(d_model, d_model)

def forward(self, x, mask=None):
batch_size, seq_len, _ = x.size()

# 步骤 1: 低秩压缩 Key 和 Value
K_c = self.W_c_K(x) # (batch, seq_len, d_latent)
V_c = self.W_c_V(x) # (batch, seq_len, d_latent)

# 步骤 2: 解耦 Query 投影
Q_shared = self.W_shared_Q(x)
Q_head = self.W_head_Q(x)
Q = Q_shared + Q_head # 解耦合并

# 步骤 3: 重塑为多头
Q = Q.view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2)

# 步骤 4: 从潜在空间恢复 Key 和 Value
# K_i = K_c @ W_K_up[i]
K = torch.einsum('bld,hdk->bhlk', K_c, self.W_K_up)
V = torch.einsum('bld,hdk->bhlk', V_c, self.W_V_up)

# 步骤 5: 计算注意力
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)

if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)

attention_weights = torch.softmax(scores, dim=-1)
attention_output = torch.matmul(attention_weights, V)

# 步骤 6: 拼接并输出变换
attention_output = attention_output.transpose(1, 2).contiguous().view(
batch_size, seq_len, self.d_model
)

output = self.W_O(attention_output)

# 推理时只存储 K_c 和 V_c(低维)
return output, attention_weights, K_c, V_c

MLA 在 DeepSeek 系列中的应用

MLA 首次在 DeepSeek-V2 中提出,并在后续版本中持续优化:

flowchart LR
    A[DeepSeek-V1] --> B[DeepSeek-V2
引入 MLA] B --> C[DeepSeek-V2-Lite
优化 MLA] C --> D[DeepSeek-V3
增强 MLA] style A fill:#e0e0e0 style B fill:#81c784 style C fill:#64b5f6 style D fill:#ffb74d

DeepSeek-V2 的 MLA 配置

  • 模型维度:$$(d_{\text{model}} = 5120)$$
  • 注意力头数:$$(h = 128)$$
  • 潜在空间维度:$$(d_c = 512)(相比原始 (d_k = 40),压缩比为 12.8 倍)$$
  • KV Cache 减少:约 100 倍

性能表现

根据 DeepSeek-V2 论文报告:

  • 推理吞吐量:提升 2-3 倍
  • 显存占用:降低约 90%
  • 模型质量:与标准 MHA 相当(在某些任务上甚至更好)

实际应用场景

1. 长文本生成

flowchart TD
    A[长文本生成任务] --> B[需要存储长序列 KV]
    B --> C1["MHA: 显存爆炸"]
    B --> C2["MLA: 显存可控"]
    C2 --> D[支持 128K+ 序列长度]
    
    style C1 fill:#ffcdd2
    style C2 fill:#c8e6c9
    style D fill:#66bb6a,color:#fff

2. 大规模并发推理

MLA 使得在单个 GPU 上可以同时服务更多用户请求:

  • 更大的 batch size
  • 更高的吞吐量
  • 更低的延迟

3. 边缘设备部署

KV Cache 的大幅减少使得大模型可以在资源受限的边缘设备上部署。

MLA 的局限性

尽管 MLA 有显著优势,但也存在一些局限性:

1. 训练复杂度增加

MLA 需要学习额外的投影矩阵,训练时的计算图更复杂。

2. 低秩假设的约束

MLA 假设 Key 和 Value 可以被低秩表示捕获,这在某些复杂任务中可能不够充分。

3. 超参数调优

潜在空间维度 (d_c) 的选择需要权衡:

  • (d_c) 太小:信息损失严重
  • (d_c) 太大:压缩效果不明显
flowchart TD
    A[选择 d_c] --> B1[d_c 太小]
    A --> B2[d_c 适中]
    A --> B3[d_c 太大]
    B1 --> C1[❌ 模型性能下降]
    B2 --> C2[✅ 效率与性能平衡]
    B3 --> C3[⚠️ 压缩效果有限]
    
    style C1 fill:#ffcdd2
    style C2 fill:#c8e6c9
    style C3 fill:#fff9c4

总结

Multi-Head Latent Attention (MLA) 是 DeepSeek 团队在注意力机制效率优化方面的重要创新。通过以下三个核心技术:

  1. 低秩 KV 压缩:将 Key 和 Value 投影到低维潜在空间,大幅减少 KV Cache
  2. 解耦 Query 设计:共享与独立部分结合,保持模型表达力
  3. 上投影恢复:在计算注意力时从潜在空间恢复完整表示
flowchart TD
    A[MLA 创新] --> B1["KV Cache 减少 100+ 倍"]
    A --> B2["推理吞吐量提升 2-3 倍"]
    A --> B3["模型质量保持不变"]
    
    B1 --> C[高效大模型推理]
    B2 --> C
    B3 --> C
    
    C --> D1[长序列支持]
    C --> D2[大规模并发]
    C --> D3[边缘部署]
    
    style A fill:#ff7043
    style C fill:#42a5f5,color:#fff
    style D1 fill:#66bb6a
    style D2 fill:#ab47bc
    style D3 fill:#26c6da

MLA 为大规模语言模型的高效推理提供了新的思路,是继 Multi-Query Attention (MQA) 和 Grouped-Query Attention (GQA) 之后的又一重要进展。随着大模型向更长序列、更高并发方向发展,MLA 及其衍生技术将发挥越来越重要的作用。

参考资料

  1. DeepSeek-V2 Team. “DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model”
  2. Vaswani, A., et al. (2017). “Attention Is All You Need”
  3. Shazeer, N. (2019). “Fast Transformer Decoding: One Write-Head is All You Need” (MQA)
  4. Ainslie, J., et al. (2023). “GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints”
  5. DeepSeek Technical Reports
CATALOG
  1. 1. 为什么要提出 MLA?
    1. 1.1. 传统 MHA 的 KV Cache 问题
    2. 1.2. MLA 的核心洞察
  2. 2. MLA 的核心思想
    1. 2.1. 1. 低秩 KV 压缩(Low-Rank KV Compression)
    2. 2.2. 2. 解耦的 Query 表示(Decoupled Query Representation)
    3. 2.3. 3. 高效的注意力计算
  3. 3. MLA 的数学原理
    1. 3.1. 1. 潜在空间压缩
    2. 3.2. 2. Query 的解耦设计
    3. 3.3. 3. Key 和 Value 的恢复
    4. 3.4. 4. 注意力计算
    5. 3.5. 完整的 MLA 公式
  4. 4. MLA 的工作流程
    1. 4.1. 详细步骤解析
      1. 4.1.1. 步骤 1:低秩压缩(Low-Rank Compression)
      2. 4.1.2. 步骤 2:Query 解耦投影(Decoupled Q Projection)
      3. 4.1.3. 步骤 3:训练时的上投影恢复(Up-Projection Recovery)
      4. 4.1.4. 步骤 4:注意力计算(Attention Computation)
  5. 5. KV Cache 对比:MHA vs MLA
    1. 5.1. 参数设置
    2. 5.2. 传统 MHA 的 KV Cache
    3. 5.3. MLA 的 KV Cache
    4. 5.4. 压缩效果
  6. 6. MLA 的优势
    1. 6.1. 1. 显存效率大幅提升
    2. 6.2. 2. 支持更长的序列
    3. 6.3. 3. 提高推理吞吐量
    4. 6.4. 4. 保持模型性能
  7. 7. MLA 与传统注意力的架构对比
  8. 8. 代码实现(PyTorch 简化版)
  9. 9. MLA 在 DeepSeek 系列中的应用
    1. 9.1. DeepSeek-V2 的 MLA 配置
    2. 9.2. 性能表现
  10. 10. 实际应用场景
    1. 10.1. 1. 长文本生成
    2. 10.2. 2. 大规模并发推理
    3. 10.3. 3. 边缘设备部署
  11. 11. MLA 的局限性
    1. 11.1. 1. 训练复杂度增加
    2. 11.2. 2. 低秩假设的约束
    3. 11.3. 3. 超参数调优
  12. 12. 总结
  13. 13. 参考资料