kv-cache 的优化学习
背景 多头注意力(Multi-Head Attention)的核心原因是:单头注意力一次只能形成一种“关注模式”,而真实序列中往往同时存在多种不同的关系需要捕捉。
1. 单头注意力的局限
在单头注意力中,对于同一个 Query,模型只会产生一个注意力权重分布:

这意味着它只能表达一种“平均”的注意力模式。
例如句子:
“苹果很好吃,但苹果公司最近股价下跌。”
对于“苹果”这个词,它有时需要关注“吃”,有时需要关注“公司”。单头注意力可能把这两种关系混在一起,形成一个模糊的平均权重,导致语义、语法、指代等信息无法被清晰区分。
2. 多头注意力的动机
多头注意力通过多个不同的投影矩阵,把 Q、K、V 分别映射到不同的子空间,然后并行计算多个注意力:

最后拼接并线性变换:

这样做的目的是让不同的头学习不同的关系。
3. 多头注意力的主要好处
① 捕捉多种不同的关系
不同头可以关注不同位置、不同特征子空间。
例如:
-
一个头关注语法依存关系,比如形容词修饰名词;
-
一个头关注指代关系,比如“it”指向“animal”;
-
一个头关注语义相似性,比如“苹果”和“水果”;
-
一个头关注远距离依赖,比如句首主语和句尾谓语。
单头很难同时做好这些。
② 类似 CNN 中的多个卷积核
在卷积神经网络中,一个卷积层通常有多个卷积核,每个卷积核提取不同特征,例如边缘、纹理、形状等。
多头注意力类似:每个头相当于一个不同的特征提取器,它们从不同角度观察输入序列,从而提高模型的表达能力。
③ 提升模型的容量和灵活性
如果只有一个头,模型只能在一个表示空间中计算注意力。多头相当于把模型宽度拆成多个并行的注意力通道,每个通道维度较小,但组合起来可以表达更复杂的函数。
而且由于每个头的维度更小,总计算量与单头全维度注意力大致相当,但效果通常更好。
④ 增强稳定性和鲁棒性
多个头可以看作多次注意力的集成。即使某个头学到了不太好的注意力模式,其他头仍然可以提供有效信息,降低单头注意力可能出现的偏差风险。
(B, L, H, D) 变成 (B*H, L, D) 呢?1. 多头注意力的计算需求
假设:
-
B:batch size,句子数量 -
L:序列长度 -
H:注意力头的数量 -
D:每个头的维度
对于每一个头,注意力计算都是:

其中:
-
Q_h形状:(L, D) -
K_h形状:(L, D) -
V_h形状:(L, D)
每个头需要计算:
-
Q_h @ K_h^T:(L, D) @ (D, L) -> (L, L) -
再乘
V_h:(L, L) @ (L, D) -> (L, D)
如果有 B*H 个这样的头,当然可以写一个 for 循环逐个计算,但效率很低。
如果直接使用 (B, L, H, D) 的形状,要计算每个头的 Q_h @ K_h^T,通常需要写循环: for h in range(H): scores_h = Q[:, :, h, :] @ K[:, :, h, :].transpose(1, 2) ... 这样效率较低,且无法充分利用 GPU 的并行计算能力。
更好的方式是:把 B*H 个 (L, D) 矩阵看作一个批量,执行批量矩阵乘法 torch.bmm。
2. 为什么需要变成 (B*H, L, D)
PyTorch / NumPy 等库提供的批量矩阵乘法 bmm,要求输入是三维张量:
-
第一个输入:
(batch, m, n) -
第二个输入:
(batch, n, p) -
输出:
(batch, m, p)
其中 batch 维会被并行处理,每个 (m, n) 和 (n, p) 独立相乘。
对于我们的注意力计算:
-
m = L -
n = D -
p = L(因为K^T的形状是(D, L))
所以我们希望 Q 和 K 的形状是:
-
Q:(batch, L, D) -
K:(batch, L, D)
这样 Q @ K^T 就是:
-
(batch, L, D) @ (batch, D, L) -> (batch, L, L)
这里的 batch 如果等于 B*H,就一次性计算了所有 batch 和所有 head 的注意力分数。
因此,需要把原来的四维张量 (B, L, H, D) 合并成三维 (B*H, L, D)。
============================================ 前面的学习都是为了 下面这个几个 kvcache 计算优化的方法做铺垫的
1. MHA(Multi-Head Attention,多头注意力)
-
KV cache 大小:每个头都有独立的 K 和 V,因此 cache 元素数为
batch × 序列长度 × 头数 × 头维度 × 2。 -
结构上没有对 cache 做任何优化,是标准的基准方案。
-
所有其他变体都是基于它来减少 KV cache 的内存消耗。
2. MQA(Multi-Query Attention,多查询注意力)
-
核心优化:所有查询头共享同一份 K 和 V(即 K、V 为单头)。
-
KV cache 大小:变为 MHA 的
1/H(H 为头数),大幅降低。 -
代价:损失了 K、V 的多头表达能力,但 Q 仍多头,因此整体性能下降较小。
3. GQA(Grouped-Query Attention,分组查询注意力)
-
核心优化:将查询头分成若干组,每组共享一个 K、V 头。
-
KV cache 大小:变为 MHA 的
G/H,其中 G 是分组数(G=1 时退化为 MQA,G=H 时就是 MHA)。 -
优点:在 MHA 和 MQA 之间取得平衡,既减少 cache,又保留一定的多头表达能力。
-
目前被 LLaMA 2/3 等主流大模型采用。
4. MLA(Multi-Latent Attention,多潜在注意力)
-
核心优化:将 K 和 V 压缩成低维的潜在向量(latent)进行缓存,计算时再实时解压还原。
-
KV cache 大小:取决于潜在向量的维度,通常远小于完整 K、V。
-
优点:可以在几乎不损失性能的情况下大幅降低 cache 内存,并且无需像 MQA/GQA 那样牺牲 K、V 的多头性。
-
被 DeepSeek-V2 等模型采用。
总结
| 注意力机制 | 是否针对 KV cache 优化 | 优化思路 |
|---|---|---|
| MHA | 否(baseline) | 每个头独立 K、V |
| MQA | 是 | 所有头共享同一 K、V |
| GQA | 是 | 分组共享 K、V |
| MLA | 是 | 压缩 K、V 为低维潜在向量 |
所以,MQA、GQA、MLA 属于 KV cache 优化,而 MHA 是被优化的基础方案。
这个demo 只是一个简单的例子 学习的
import torch import torch.nn.functional as F import math # ==================== 公共参数 ==================== B, L, H, D = 1, 6, 8, 4 # batch, seq_len, heads, head_dim D_model = H * D # 32 torch.manual_seed(40) # 固定随机种子,方便复现 x = torch.randn(B, L, D_model) # 输入 (1, 6, 32) # 随机生成 Q/K/V 投影矩阵(无 bias) W_q = torch.randn(D_model, D_model) W_k = torch.randn(D_model, D_model) W_v = torch.randn(D_model, D_model) print(f"输入形状: {x.shape} (B={B}, L={L}, D_model={D_model})\n") # ================================================== # 1. MHA – 每个头独立 KV # ================================================== def demo_mha(x, W_q, W_k, W_v): Q = x @ W_q K = x @ W_k V = x @ W_v Q = Q.view(B, L, H, D) # (1,6,8,4) K = K.view(B, L, H, D) V = V.view(B, L, H, D) #print(f"输入: (Q={Q}\n") kv_elements = K.numel() + V.numel() print(f"【MHA】KV Cache 元素数: {kv_elements} (K: {K.shape}, V: {V.shape})") # 计算 Attention Q_flat = Q.reshape(B * H, L, D) K_flat = K.reshape(B * H, L, D) V_flat = V.reshape(B * H, L, D) #print(f"输入: (Q_flat={Q_flat}\n") scores = torch.bmm(Q_flat, K_flat.transpose(1, 2)) / math.sqrt(D) attn = F.softmax(scores, dim=-1) out = torch.bmm(attn, V_flat) out = out.reshape(B, L, H, D).transpose(1, 2).contiguous().reshape(B, L, D_model) return out print("=" * 50) _ = demo_mha(x, W_q, W_k, W_v) # ================================================== # 2. MQA – 所有 Q 头共享 1 份 KV # ================================================== def demo_mqa(x, W_q, W_k, W_v): Q = x @ W_q K = x @ W_k V = x @ W_v # 先全部拆成 H 个头 Q = Q.view(B, L, H, D) K_full = K.view(B, L, H, D) V_full = V.view(B, L, H, D) # MQA:只取第 1 个头(索引0)作为共享 KV K = K_full[:, :, :1, :] # (1,6,1,4) V = V_full[:, :, :1, :] kv_elements = K.numel() + V.numel() print(f"【MQA】KV Cache 元素数: {kv_elements} (K: {K.shape}, V: {V.shape}) —— 仅为 MHA 的 1/{H}") # 计算时把 1 份 KV 广播复制成 H 份 Q_flat = Q.reshape(B * H, L, D) K_expanded = K.expand(-1, -1, H, -1) # (1,6,8,4) V_expanded = V.expand(-1, -1, H, -1) K_flat = K_expanded.reshape(B * H, L, D) V_flat = V_expanded.reshape(B * H, L, D) scores = torch.bmm(Q_flat, K_flat.transpose(1, 2)) / math.sqrt(D) attn = F.softmax(scores, dim=-1) out = torch.bmm(attn, V_flat) out = out.reshape(B, L, H, D).transpose(1, 2).contiguous().reshape(B, L, D_model) return out print("-" * 50) _ = demo_mqa(x, W_q, W_k, W_v) # ================================================== # 3. GQA – 分组共享 KV(这里分为 2 组) # ================================================== def demo_gqa(x, W_q, W_k, W_v, G=2): assert H % G == 0 heads_per_group = H // G Q = x @ W_q K = x @ W_k V = x @ W_v Q = Q.view(B, L, H, D) K_full = K.view(B, L, H, D) V_full = V.view(B, L, H, D) # 只保留前 G 个头作为 KV 组 K = K_full[:, :, :G, :] # (1,6,2,4) V = V_full[:, :, :G, :] kv_elements = K.numel() + V.numel() print(f"【GQA】KV Cache 元素数: {kv_elements} (K: {K.shape}, V: {V.shape}) —— 约为 MHA 的 {G}/{H}") # 每组内复制 heads_per_group 次,扩展回 H 个头 填充回去为 H 个头 K_expanded = K.repeat_interleave(heads_per_group, dim=2) # (1,6,8,4) V_expanded = V.repeat_interleave(heads_per_group, dim=2) Q_flat = Q.reshape(B * H, L, D) K_flat = K_expanded.reshape(B * H, L, D) V_flat = V_expanded.reshape(B * H, L, D) scores = torch.bmm(Q_flat, K_flat.transpose(1, 2)) / math.sqrt(D) attn = F.softmax(scores, dim=-1) out = torch.bmm(attn, V_flat) out = out.reshape(B, L, H, D).transpose(1, 2).contiguous().reshape(B, L, D_model) return out print("-" * 50) _ = demo_gqa(x, W_q, W_k, W_v, G=2) # ================================================== # 4. MLA – 压缩成潜在向量 (Latent) # ================================================== def demo_mla(x, W_q, W_k, W_v, D_c=8): # D_c 是压缩后的维度(这里设为 8,远小于 D_model=32) W_down_k = torch.randn(D_model, D_c) W_up_k = torch.randn(D_c, D_model) W_down_v = torch.randn(D_model, D_c) W_up_v = torch.randn(D_c, D_model) Q = x @ W_q # (1,6,32) # 压缩成潜在向量(这就是实际缓存的) latent_k = x @ W_down_k # (1,6,8) latent_v = x @ W_down_v # (1,6,8) kv_elements = latent_k.numel() + latent_v.numel() print(f"【MLA】实际缓存的元素数: {kv_elements} (Latent_K: {latent_k.shape}, Latent_V: {latent_v.shape})") # 计算时实时解压回完整维度 K_reconstructed = latent_k @ W_up_k # (1,6,32) V_reconstructed = latent_v @ W_up_v # 拆成多头计算 Q_heads = Q.view(B, L, H, D) K_heads = K_reconstructed.view(B, L, H, D) V_heads = V_reconstructed.view(B, L, H, D) print(f" (计算时临时解压的 K 形状: {K_heads.shape}, 不长期缓存)") Q_flat = Q_heads.reshape(B * H, L, D) K_flat = K_heads.reshape(B * H, L, D) V_flat = V_heads.reshape(B * H, L, D) scores = torch.bmm(Q_flat, K_flat.transpose(1, 2)) / math.sqrt(D) attn = F.softmax(scores, dim=-1) out = torch.bmm(attn, V_flat) out = out.reshape(B, L, H, D).transpose(1, 2).contiguous().reshape(B, L, D_model) return out print("-" * 50) _ = demo_mla(x, W_q, W_k, W_v, D_c=8) print("=" * 50)
root@/self_mini_vllm# python MHA.py
输入形状: torch.Size([1, 6, 32]) (B=1, L=6, D_model=32)
==================================================
【MHA】KV Cache 元素数: 384 (K: torch.Size([1, 6, 8, 4]), V: torch.Size([1, 6, 8, 4]))
--------------------------------------------------
【MQA】KV Cache 元素数: 48 (K: torch.Size([1, 6, 1, 4]), V: torch.Size([1, 6, 1, 4])) —— 仅为 MHA 的 1/8
--------------------------------------------------
【GQA】KV Cache 元素数: 96 (K: torch.Size([1, 6, 2, 4]), V: torch.Size([1, 6, 2, 4])) —— 约为 MHA 的 2/8
--------------------------------------------------
【MLA】实际缓存的元素数: 96 (Latent_K: torch.Size([1, 6, 8]), Latent_V: torch.Size([1, 6, 8]))
(计算时临时解压的 K 形状: torch.Size([1, 6, 8, 4]), 不长期缓存)
==================================================
前面废话太多:给个表格
| 场景 | 推荐方法 | 理由 |
|---|---|---|
| 通用部署 | GQA + INT8 量化 | 广泛支持,效果好 |
| 极致效率 | MLA (DeepSeek 风格) | 需要从头训练 |
还有一个就是现在的优化就是,存储在kvcache 中的是 MLA形式 ,在prefill 阶段展开成MHA,在 decode 阶段用MLA
浙公网安备 33010602011771号