kv-cache 的优化学习

背景 多头注意力(Multi-Head Attention)的核心原因是:单头注意力一次只能形成一种“关注模式”,而真实序列中往往同时存在多种不同的关系需要捕捉。

1. 单头注意力的局限

在单头注意力中,对于同一个 Query,模型只会产生一个注意力权重分布:

 

 

image

 这意味着它只能表达一种“平均”的注意力模式。

例如句子:

“苹果很好吃,但苹果公司最近股价下跌。”

对于“苹果”这个词,它有时需要关注“吃”,有时需要关注“公司”。单头注意力可能把这两种关系混在一起,形成一个模糊的平均权重,导致语义、语法、指代等信息无法被清晰区分。

2. 多头注意力的动机

多头注意力通过多个不同的投影矩阵,把 Q、K、V 分别映射到不同的子空间,然后并行计算多个注意力:

image

 最后拼接并线性变换:

image

 这样做的目的是让不同的头学习不同的关系。

 

3. 多头注意力的主要好处

① 捕捉多种不同的关系

不同头可以关注不同位置、不同特征子空间。

例如:

  • 一个头关注语法依存关系,比如形容词修饰名词;

  • 一个头关注指代关系,比如“it”指向“animal”;

  • 一个头关注语义相似性,比如“苹果”和“水果”;

  • 一个头关注远距离依赖,比如句首主语和句尾谓语。

单头很难同时做好这些。

② 类似 CNN 中的多个卷积核

在卷积神经网络中,一个卷积层通常有多个卷积核,每个卷积核提取不同特征,例如边缘、纹理、形状等。

多头注意力类似:每个头相当于一个不同的特征提取器,它们从不同角度观察输入序列,从而提高模型的表达能力。

③ 提升模型的容量和灵活性

如果只有一个头,模型只能在一个表示空间中计算注意力。多头相当于把模型宽度拆成多个并行的注意力通道,每个通道维度较小,但组合起来可以表达更复杂的函数

而且由于每个头的维度更小,总计算量与单头全维度注意力大致相当,但效果通常更好。

④ 增强稳定性和鲁棒性

多个头可以看作多次注意力的集成。即使某个头学到了不太好的注意力模式,其他头仍然可以提供有效信息,降低单头注意力可能出现的偏差风险。

 

4. 直观总结

可以这样理解:

  • 单头注意力:一个人只能从一个角度看问题,容易片面。

  • 多头注意力:多个人从不同角度看同一个问题,最后把意见综合起来,判断更全面。

因此,Transformer 使用多头注意力,是为了让模型在多个子空间中并行学习不同类型的依赖关系,从而获得更强的表达能力和更好的性能。这也是《Attention Is All You Need》中提出多头注意力的重要原因。

 
=================================== 穿插
多头注意力计算中为什么把 (B, L, H, D) 变成 (B*H, L, D) 呢?
本质上是为了把“多个头”当成“多个独立的样本”来做批量矩阵乘法从而一次性、高效地完成所有头的注意力计算
 

1. 多头注意力的计算需求

假设:

  • B:batch size,句子数量

  • L:序列长度

  • H:注意力头的数量

  • D:每个头的维度

对于每一个头,注意力计算都是

image

 

其中:

  • 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)

所以我们希望 QK 的形状是:

  • 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

 

 

 

 

 

 
 
 
 
 

 

posted on 2026-09-08 17:16  zhangkele  阅读(5)  评论(0)    收藏  举报

导航