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

作者: 光仔玩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 的核心改进。
分两步:
- Lightning Indexer:用一组轻量 indexer 头给每个候选 token 打一个相关分
- 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) VO(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 模型我没见过谁真用。
容易踩坑的几个点
写几条我自己反复想过的反直觉的事情,免得你照搬时翻车:
- KV Cache 不能用于训练。训练时所有 token 一次性算,缓存反而是累赘。这玩意儿是纯推理优化。
- 小模型上 KV Cache 收益会被 device overhead 吞掉。 124M 模型在 CUDA 上几乎没加速,M4 上反而 5×。真正的甜区是大模型 + 长序列。
- MoE 不是 KV 优化。MoE 省的是 FFN 激活显存,KV 缓存大小根本没变。两者正交,生产里经常一起上但别搞混。
- 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」,仅供学习交流使用。
觉得内容不错?我要