KV Cache 全攻略

本文摘要KV Cache 全攻略本文来源: 光仔玩AI(公众号:光仔玩AI) 原文链接: https://mp.weixin.qq.com/s/v6YYxCf_N4Q6c7qfD_lZDw 发布时间: 2026-08-18 11:29作者: 光仔玩AI  发布时间: 2026-08-18 11:29KV Cache 让推理从 O(n²) 变成 O(n),代价是显存从 0 变成线性增长。当上下文从 4k 涨...

KV Cache 全攻略

本文来源: 光仔玩AI(公众号:光仔玩AI)
原文链接: https://mp.weixin.qq.com/s/v6YYxCf_N4Q6c7qfD_lZDw
发布时间: 2026-08-18 11:29

KV Cache 全攻略

作者: 光仔玩AI  发布时间: 2026-08-18 11:29



KV Cache 让推理从 O(n²) 变成 O(n),代价是显存从 0 变成线性增长。当上下文从 4k 涨到 128k、模型从 7B 涨到 70B,这块缓存就不是有点大,而是能直接 OOM。

从 2023 年到现在,工业界其实只做了一件事:怎么把这块缓存压小。

这篇文章把 GQA、MLA、SWA、跨层共享、DSA、Gated DeltaNet 六种主流方案一次性摊开——每节一张手绘图、一段最小核心可运行的 PyTorch 代码、一组仓库实测的显存/速度数据。

一、先看看到底有多浪费

我第一次跑通自己的 GPT 模型想生成点东西的时候,看 log 里 token 一个蹦出来等了大概半秒——心想这也太慢了。翻开代码一看,注意力那块居然把前面所有 token 的 K/V 都重算了一遍。

具体点说。生成"Time flies fast"这三个 token,第 1 步喂"Time flies"算 K₁V₁/K₂V₂;第 2 步喂"Time flies fast",K₁V₁ 和 K₂V₂ 完全一样,又被算了一遍;第 3 步喂"Time flies fast when",前两个 token 的 K/V 又被算第三遍。把这件事推一下:生成 n 个 token,累计做 1+2+3+…+n ≈ O(n²) 次 K/V 投影。n 一上 1000,光这些矩阵乘就够把 GPU 拖死。

我当时第一反应是:这么蠢的事情,肯定有解吧......

二、KV Cache:算一次存一辈子

思路其实贼朴素:每层 attention 配两个 buffer cache_k、cache_v,每一步只算新 token 的 K/V,然后 cat 进 buffer;算注意力时从 buffer 里读全量。三步走完上面的例子:

  • step 1:算 K₁V₁/K₂V₂ → 缓存里有 2 个
  • step 2:只算 K₃V₃ → 缓存里有 3 个(前面 2 个读现成的)
  • step 3:只算 K₄V₄ → 缓存里有 4 个

每步计算量从「随长度增长」变成「恒定 1 个 token」,O(n²) → O(n)。在 M4 Mac Mini 上拿 124M 小模型 + 200 token 实测: 27 tok/s,加 KV Cache 后 144 tok/s,约5 倍。

实现改的地方不多,主要集中在 MultiHeadAttention:

class MultiHeadAttention(nn.Module):
    def __init__(self, ...):
        super().__init__()
        # ... W_key/W_value/W_query 不变 ...
        self.register_buffer("cache_k", None)   # 缓存 K
        self.register_buffer("cache_v", None)   # 缓存 V
        self.ptr_current_pos = 0                # 写到哪了
    def forward(self, x, use_cache=False):
        keys_new, values_new = self.W_key(x), self.W_value(x)
        queries = self.W_query(x)
        if use_cache:
            if self.cache_k is None:
                self.cache_k, self.cache_v = keys_new, values_new
            else:
                self.cache_k = torch.cat([self.cache_k, keys_new],   dim=1)
                self.cache_v = torch.cat([self.cache_v, values_new], dim=1)
            keys, values = self.cache_k, self.cache_v
        else:
            keys, values = keys_new, values_new
        # 因果 mask 跟着偏移
        if use_cache:
            mask_bool = self.mask.bool()[
                self.ptr_current_pos:self.ptr_current_pos + queries.shape[1],
                :keys.shape[1]
            ]
            self.ptr_current_pos += queries.shape[1]
        else:
            mask_bool = self.mask.bool()[:queries.shape[1], :keys.shape[1]]
        # ... 后面 attn 公式不变 ...

TransformerBlock 和 GPTModel 还要把 use_cache 一路传下去,循环里维护 current_pos 给位置编码用,再加个 reset_kv_cache() 让不同请求之间干净重置。

这里有个我踩过的坑:第 1 步的预填充(prompt 一次性塞进去)不能用 use_cache 跳过,否则位置编码算错,会输出乱码。正确做法是预填充开 cache、新 token 接着走 cache。

三、快是快了,代价呢?

我跑完 KV Cache 一看显存——懵了。原本 8GB 的 batch 一下子飙到 14GB,把 batch size 砍回 1 才勉强塞下。打开 profile 才发现,KV 缓存占的显存比模型权重还大。

显存公式就一行:

bytes ≈ batch × seqlen × n_layers × n_kv_heads × head_dim × 2(K,V) × bytes_per_elem

只要 seqlen 涨一项,缓存就跟着涨。在 128k 上下文 + 70B 模型上,这玩意儿能吃掉几十 GB。所以后面所有技术都在回答同一个问题:怎么把这块缓存压小。

四、砍头数:GQA / MQA

GQA 是 MHA 和 MQA 的连续插值——把 K/V 头分成几组,每组服务多个 Q 头。g=1 就是 MQA,g=Q 头数就是 MHA,中间任何 g 都是 GQA。

Llama 2/3/4、Qwen3、Gemma 3 全是 GQA。你想想为啥这堆顶流模型全押同一张牌——因为它在「质量 ≈ MHA、显存 ≈ MQA」之间找到了甜蜜点。

数据上,在标准配置下(batch=1, seqlen=8192, 64 层, head_dim=64),MHA 17.18 GB → GQA 4.29 GB,省 75%。MQA 更狠,能到 2.86 GB。

改动集中在初始化和 forward 头里:

class GroupedQueryAttention(nn.Module):
    def __init__(self, d_in, d_out, n_heads, n_kv_groups, ...):
        super().__init__()
        self.n_heads     = n_heads
        self.n_kv_groups = n_kv_groups
        self.group_size  = n_heads // n_kv_groups   # 每个 KV 头服务几个 Q 头
        self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
        # KV 投影按 n_kv_groups 分配,不是 n_heads
        self.W_key   = nn.Linear(d_in, d_out // self.group_size, bias=qkv_bias)
        self.W_value = nn.Linear(d_in, d_out // self.group_size, bias=qkv_bias)
    def forward(self, x, use_cache=False):
        # ... 算 Q, K, V ...
        # 把 KV 头 expand 到和 Q 头一样多
        keys   = keys.repeat_interleave(self.group_size, dim=1)
        values = values.repeat_interleave(self.group_size, dim=1)
        # ... 标准多头注意力 ...

注意 group_size = n_heads // n_kv_groups 这个数一定要整除,不然会报错。

五、压维度:MLA

GQA 解决的是少存几套 K/V,MLA 解决的是每套都更小。DeepSeek V2 论文里那个 ablation 我看完是有被震撼到——MLA 在建模能力上比 MHA 还稍好,而 GQA 是比 MHA 略差。这就是为啥 DeepSeek 没选 GQA。

MLA 的精髓就一句话:把 K/V 投影到一个低维 latent 空间再存。

class MultiHeadLatentAttention(nn.Module):
    def __init__(self, d_in, d_out, n_heads, latent_dim, ...):
        super().__init__()
        self.W_down = nn.Linear(d_in, latent_dim, bias=False)   # 唯一的 down 投影
        self.W_uk   = nn.Linear(latent_dim, d_out, bias=False)  # 解 K
        self.W_uv   = nn.Linear(latent_dim, d_out, bias=False)  # 解 V
        self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)    # Q 直出,不走 latent
    def forward(self, x, use_cache=False):
        c_t = self.W_down(x)                          # (b, tokens, latent_dim)
        if use_cache:
            self.cache_c = torch.cat([self.cache_c, c_t], dim=1) if self.cache_c is not None else c_t
            c_full = self.cache_c
            keys   = self.W_uk(c_full)                # 一次性解压所有
            values = self.W_uv(c_full)
        else:
            keys   = self.W_uk(c_t)
            values = self.W_uv(c_t)
        # ... 标准多头注意力 ...

跑出来

3.25 GB → 0.81 GB,省 75%

——和 GQA 数字几乎一样,但MLA 比 GQA 强在效果不是显存。

代价是多一次矩阵乘,推理期每次都要 W_uk(c) 和 W_uv(c) 现场解压。但显存省得太值了,没人在乎那点算力。

六、砍长度:SWA

如果说 GQA/MLA 是改存多少,SWA 是改存哪儿——只看最近的 W 个 token,远的直接扔。我第一次看这图的时候反应是:这模型不就瞎了吗?前文全忘了?

Gemma 2 的论文,人家是用5:1 混合(5 个 SWA 层 + 1 个全局注意力层)兜底。每隔几个 SWA 层插一个完整注意力层,让全局层去抓那些 W 窗口外的关键信息。所以你不会真瞎。

# 构造 SWA mask:每行只保留最近 W 个 + 因果
seq_len = x.shape[1]
mask = torch.full((seq_len, seq_len), float("-inf"))
for i in range(seq_len):
    start = max(0, i - window_size + 1)   # 滑动窗起点
    mask[i, start:i+1] = 0                # 滑动窗内可见

显存公式里 seqlen 被替换成 W,W 远小于 seqlen。Gemma 2 用 W=4096,Gemma 3 收紧到 W=1024。

数据上:GQA+SWA(5:1)只要 0.78 GB,相比纯 MHA 省了 22 倍。这是我目前看到的最猛性价比组合之一。

七、砍层数:跨层共享

这一节我第一次读完是困惑的。让好几层用同一份 K/V——这不是偷懒吗?模型能力不是要崩?

仔细看才发现它跟 GQA 是同一思路的不同维度——GQA 是「同一层内头共享」,KV Sharing 是「跨层共享」。Gemma 4 E2B 35 层里只有 8 层产 KV,每个消费层复用本组内最近那份 KV。代价是某些层失去自己的 K/V 投影,模型容量会有损失,但叠上 GQA + SWA 之后这部分损失被摊薄了。

class CrossLayerKVSharingAttention(nn.Module):
    def __init__(self, d_in, d_out, n_heads, kv_producer_id, layer_idx, ...):
        super().__init__()
        self.is_producer = (kv_producer_id == layer_idx)   # 只有 producer 算 K/V
        self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
        if self.is_producer:
            self.W_key   = nn.Linear(d_in, d_out, bias=qkv_bias)
            self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)
    def forward(self, x, shared_cache):
        queries = self.W_query(x)
        if self.is_producer:
            keys = self.W_key(x); values = self.W_value(x)
            shared_cache['k'] = torch.cat([shared_cache['k'], keys], dim=1) if shared_cache['k'] is not None else keys
            shared_cache['v'] = torch.cat([shared_cache['v'], values], dim=1) if shared_cache['v'] is not None else values
        else:
            keys, values = shared_cache['k'], shared_cache['v']   # 复用别人的
        # ... 标准注意力 ...

Gemma 4 E2B:35 层只 8 层产 KV · 显存 ∝ n_kv_producing_layers(而非 n_layers) · KV 缓存省到原来的 ~23%。

如果你跑的是单层 demo,看不出任何收益——只有真正多层的模型才感受到。所以这一招一般用户感知不到,但生产里是真省。

八、先打分再挑:DSA

DSA(DeepSeek Sparse Attention)的方向又不一样:索性不让模型看所有 token,只让打分最高的 K 个进来。这是 DeepSeek-V3.2 的核心改进。

分两步:

  1. Lightning Indexer:用一组轻量 indexer 头给每个候选 token 打一个相关分
  2. Token Selector:softmax 之前把所有非 top-K 的位置 mask 成 -∞
class DeepSeekSparseAttention(nn.Module):
    def __init__(self, d_in, d_out, n_heads, n_index_heads, top_k, ...):
        super().__init__()
        self.top_k = top_k
        self.W_q_idx = nn.Linear(d_in, n_index_heads * d_in // n_heads)
        self.W_k_idx = nn.Linear(d_in, n_index_heads, bias=False)
        self.W_gate  = nn.Linear(d_in, n_index_heads)            # per-head gate
    def _indexer_scores(self, x):
        q_idx = self.W_q_idx(x)           # (b, tokens, H_I, d_I)
        k_idx = self.W_k_idx(x)           # (b, tokens, H_I)
        scores = torch.relu(torch.einsum('bthd,bsh->bts', q_idx, k_idx)) / (d_in ** 0.5)
        gate   = self.W_gate(x).sigmoid()
        scores = (gate.unsqueeze(2) * scores).sum(dim=-1) / (self.n_index_heads ** 0.5)
        return scores
    def forward(self, x, use_cache=False):
        # ... 标准 K/V 投影(走 MLA 的 latent 也行) ...
        scores = self._indexer_scores(x)
        topk = scores.topk(self.top_k, dim=-1).indices
        mask = torch.full(scores, float("-inf"), device=x.device)
        mask.scatter_(2, topk, 0.0)
        # 后续走标准 softmax(QK^T + mask) V

O(L²) → O(L·K),K 一般取 2048 左右,L=128k 时省 60×+。

九、不缓存 K/V,缓存一个固定状态:DeltaNet

如果前面几招都是让 KV 缓存少存点,DeltaNet 的方向是干脆不要 K/V 缓存—— 改成维护一个固定大小的循环状态 S。

我看 Qwen3-Next 技术报告的时候反复琢磨这个状态 S 是怎么工作的,后来意识到一个关键差异:没有 softmax。传统注意力里 softmax(QK^T)V 是个概率加权平均,DeltaNet 里就是 S · q,线性加权。

class GatedDeltaNet(nn.Module):
    def __init__(self, d_in, d_out, n_heads, ...):
        super().__init__()
        self.W_q = nn.Linear(d_in, d_out, bias=False)
        self.W_k = nn.Linear(d_in, d_out, bias=False)
        self.W_v = nn.Linear(d_in, d_out, bias=False)
        self.alpha = nn.Parameter(torch.zeros(n_heads))     # 每头一个标量 gate
        self.beta  = nn.Parameter(torch.ones(n_heads))
    def forward(self, x):
        Q, K, V = self.W_q(x), self.W_k(x), self.W_v(x)
        b, t, h, d = Q.shape
        S = torch.zeros(b, h, d, d, device=x.device)        # 固定状态
        outputs = []
        for i in range(t):
            q_i, k_i, v_i = Q[:, i], K[:, i], V[:, i]
            S = self.alpha[None,:,None,None] * S \
                + self.beta[None,:,None,None] * torch.einsum('bhd,bhe->bhde', k_i, v_i)
            o_i = torch.einsum('bhde,bhd->bhe', S, q_i)
            outputs.append(o_i)
        return torch.stack(outputs, dim=1)

S 永远就那么大,不随长度增长。代价是检索精度比完整注意力弱——没有 softmax 那一步的概率归一,纯靠加性更新记住过去。

所以生产里都是3:1 混合(3 个 DeltaNet + 1 个完整注意力),让 DeltaNet 扛长上下文、注意力兜底精确检索。Qwen3-Next、Kimi Linear 都是这套。纯 DeltaNet 模型我没见过谁真用。

容易踩坑的几个点

写几条我自己反复想过的反直觉的事情,免得你照搬时翻车:

  1. KV Cache 不能用于训练。训练时所有 token 一次性算,缓存反而是累赘。这玩意儿是纯推理优化。
  2. 小模型上 KV Cache 收益会被 device overhead 吞掉。 124M 模型在 CUDA 上几乎没加速,M4 上反而 5×。真正的甜区是大模型 + 长序列。
  3. MoE 不是 KV 优化。MoE 省的是 FFN 激活显存,KV 缓存大小根本没变。两者正交,生产里经常一起上但别搞混。
  4. PagedAttention 和 KV 量化是另一条线。PagedAttention 改显存分配(vLLM 核心),KV 量化把 FP16 砍到 INT8/FP8——这俩和本文讲的所有技术正交,下次单独聊。

参考资源
书籍《Build a Large Language Model (From Scratch)》
GQA 论文:https://arxiv.org/abs/2305.13245
DeepSeek-V2 MLA 论文:https://arxiv.org/abs/2405.04434
Gated DeltaNet:https://arxiv.org/abs/2412.06464
Kimi Linear:https://arxiv.org/abs/2510.26692

以上,既然看到这里了,
如果觉得不错,随手点个赞、在看、转发三连吧,
如果想第一时间收到推送,也可以给我个星标⭐~
谢谢你看我的文章

你的关注是我持续更新的动力~


本文转载自微信公众号「光仔玩AI」,仅供学习交流使用。

觉得内容不错?我要

打赏杯咖啡或蜜雪冰城吧
微信扫一扫
微信赞赏码
支付宝扫一扫
支付宝赞赏码
评论 暂无评论
请登录后参与评论