提到KV Cache,很多人只会联想到“参数缓存”,但是它具体是怎么工作的却不知,本篇博文会从0到1详解KV Cache的工作原理。

原理

当前主流大模型,均采用Decoder-only架构,文本生成过程依托自回归机制运行。具体来说,模型不会一次性输出所有文本,而是每次仅生成1个新token,将这个新token拼接至原有输入序列后,迭代循环生成下一个token,直至文本生成完成。
Transformer解码器的自注意力机制是模型推理的核心,依托三个矩阵Q、K、V完成计算,其核心计算公式为:
在这里插入图片描述
在未引入KV Cache的原生推理模式下,大模型自回归生成存在极其严重的重复计算问题。结合实际场景,当输入一段固定的Prompt文本后,模型需要逐字生成后续内容,此时每生成一个全新的token,模型都需要加载完整的历史序列,重新计算之前所有token对应的Q、K、V特征。
这一机制的弊端十分明显,对话过程中已经生成的历史序列是完全固定、不会发生任何变化的,对应的K、V 特征数值也不会改变,但原生推理逻辑会持续对这些固定特征进行重复计算。随着生成步数增加,文本序列长度不断拉长,整体计算量会呈现平方级暴涨,直接导致长文本计算量与生成速度大幅变慢,无法满足实际落地中的长对话、长文本推理需求,无KV Cache的原生推理时间复杂度为 O(N*N)。
而以“空间换取时间”的思想,用一个额外的空间存储之前已计算的K、V,解决原生推理的重复计算问题;在每一轮迭代中,只有最新生成的单个token是全新变量,仅需对该token完成特征计算。KV Cache的核心逻辑可以完整概括为:模型首次加载Prompt进行推理时,通过Prefill阶段,一次性并行的计算所有初始token的K、V特征并完成缓存存储,后续每一轮自回归迭代生成过程中,不再重复计算历史序列特征,仅计算新生成token的K、V,将新计算的特征拼接到历史缓存数据中,使用拼接后的完整K、V特征完成注意力计算,解决推理过程中重复计算问题,大幅提升推理速度。具体如图所示:在生成第一个token的Q、K、V后会将K、V存储在KV Cache中,供下一轮推理使用。
在这里插入图片描述
在这里插入图片描述
在这里插入图片描述

在这里插入图片描述
由于Transformer模型由多层解码层堆叠而成,KV Cache采用分层存储的架构设计,模型的每一层都会独立缓存当前层对应的K、V特征数据,不会出现层间数据混淆的情况。缓存的标准维度格式为「模型层数, 批次大小, 注意力头数, 序列长度, 单头维度」,整体分为 k_cache 和 v_cache 两部分,分别用于存储每一层所有历史 token 的K、V。
KV Cache 的完整工作流程可以分为两个核心阶段,全程逻辑清晰且高效。第一阶段为 Prompt 预填充阶段Prefill,也就是模型首次推理的过程,此时输入完整的初始提示词序列,模型遍历计算序列内所有 token 的 K、V 特征,将所有特征存入缓存,同时完成第一轮推理,生成第一个新 token。第二阶段为自回归迭代生成阶段Decode,也是模型持续输出文本的核心阶段。在这一阶段中,模型每一轮仅输入上一步生成的单个新 token,无需加载完整历史序列,仅针对该新 token 完成 K、V 特征计算,随后将全新特征拼接在历史缓存数据尾部,通过拼接后的完整 K、V 矩阵计算注意力权重,输出下一个 token。整个迭代过程会持续循环,直至模型生成终止符,或者序列长度达到预设最大值,最终完成完整文本的生成。

代码

下面是一个KV Cache的简单代码应用,

python
import torch
import torch.nn as nn
import torch.nn.functional as F

BATCH_SIZE = 1        # 推理批次
NUM_LAYERS = 2       # Transformer层数
NUM_HEADS = 2        # 注意力头数
HEAD_DIM = 16        # 单个注意力头维度
HIDDEN_DIM = NUM_HEADS * HEAD_DIM  # 模型隐藏层总维度
SEQ_LEN = 4          # 初始输入序列长度

class KVCache:
    def __init__(self, device: torch.device):
        self.device = device
        self.k_cache: list[torch.Tensor] = []
        self.v_cache: list[torch.Tensor] = []

    def update(self, k: torch.Tensor, v: torch.Tensor, layer_idx: int) -> tuple[torch.Tensor, torch.Tensor]:
        # 首次初始化缓存Prefill阶段
        if layer_idx >= len(self.k_cache):
            self.k_cache.append(k)
            self.v_cache.append(v)
        else:
            self.k_cache[layer_idx] = torch.cat([self.k_cache[layer_idx], k], dim=-2)
            self.v_cache[layer_idx] = torch.cat([self.v_cache[layer_idx], v], dim=-2)

        return self.k_cache[layer_idx], self.v_cache[layer_idx]

    def clear(self):
        self.k_cache.clear()
        self.v_cache.clear()

class SimpleAttentionLayer(nn.Module):
    def __init__(self):
        super().__init__()
        self.q_proj = nn.Linear(HIDDEN_DIM, HIDDEN_DIM)
        self.k_proj = nn.Linear(HIDDEN_DIM, HIDDEN_DIM)
        self.v_proj = nn.Linear(HIDDEN_DIM, HIDDEN_DIM)

    def forward(self, x: torch.Tensor, kv_cache: KVCache, layer_idx: int) -> torch.Tensor:
        B, T, C = x.shape
        q = self.q_proj(x).view(B, T, NUM_HEADS, HEAD_DIM).transpose(1, 2)
        k = self.k_proj(x).view(B, T, NUM_HEADS, HEAD_DIM).transpose(1, 2)
        v = self.v_proj(x).view(B, T, NUM_HEADS, HEAD_DIM).transpose(1, 2)
        k, v = kv_cache.update(k, v, layer_idx)
        attn_scores = torch.matmul(q, k.transpose(-2, -1)) / (HEAD_DIM ** 0.5)
        attn_probs = F.softmax(attn_scores, dim=-1)
        output = torch.matmul(attn_probs, v)
        return output.transpose(1, 2).contiguous().view(B, T, C)


def simulate_autoregressive_generation():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    kv_cache = KVCache(device)
    layers = nn.ModuleList([SimpleAttentionLayer() for _ in range(NUM_LAYERS)]).to(device)
    print("=" * 60)
    print("KV Cache 自回归推理模拟启动")
    print("=" * 60)
    x = torch.randn(BATCH_SIZE, SEQ_LEN, HIDDEN_DIM).to(device)
    print(f"初始输入序列形状: {x.shape}")
    for step in range(3):
        print(f"\n【第 {step+1} 轮生成】")
        x_new = x[:, -1:, :]
        print(f"当前推理输入形状: {x_new.shape}")
        for idx, layer in enumerate(layers):
            out = layer(x_new, kv_cache, layer_idx=idx)
        cache_len = kv_cache.k_cache[0].shape[-2]
        print(f"当前KV缓存序列长度: {cache_len}")
        x = torch.cat([x, x_new], dim=1)
    print("\n✅ KV Cache 推理模拟完成!")
    print(f"最终序列长度: {x.shape[1]}")

if __name__ == "__main__":
    simulate_autoregressive_generation()

运行结果

============================================================
KV Cache 自回归推理模拟启动
============================================================
初始输入序列形状: torch.Size([1, 4, 32])

【第 1 轮生成】
当前推理输入形状: torch.Size([1, 1, 32])
当前KV缓存序列长度: 5

【第 2 轮生成】
当前推理输入形状: torch.Size([1, 1, 32])
当前KV缓存序列长度: 6

【第 3 轮生成】
当前推理输入形状: torch.Size([1, 1, 32])
当前KV缓存序列长度: 7

✅ KV Cache 推理模拟完成!
最终序列长度: 7

缺点

在实际工业落地场景中,KV Cache会占用大量显存资源,显存开销也是大模型推理落地需要重点权衡的问题。KV Cache的显存占用可以通过固定公式精准计算,整体显存大小由模型层数、推理批次、注意力头数、序列长度、单头维度共同决定,同时因为需要同时存储K、V 两份特征数据,最终总占用量需要在基础计算结果上乘以2,对应公式为 “显存大小”=层数×批次×头数×序列长度×单头维度×2。
结合实际场景举例来看,常规7B参数大模型、2k上下文长度的推理场景中,KV Cache单独占用的显存就可达数百MB甚至1GB以上,当上下文长度拉长至8k、32k时,显存占用会成倍增长,为了解决 KV Cache带来的显存开销与性能问题,工业界衍生出了多种成熟的优化方案。混合精度量化通过 FP16、BF16 精度存储KV特征,能够直接将缓存显存占用减半,在几乎无精度损失的前提下大幅降低显存压力。PagedAttention通过分页式内存管理机制,解决传统KV Cache的内存碎片问题,大幅提升显存利用率和推理吞吐量。FlashAttention从算子层面做深度优化,将KV缓存数据放入高速缓存中读取,降低内存访问开销,进一步提升推理速度。除此之外,动态缓存淘汰机制会针对超长文本场景,智能淘汰无效的老旧缓存数据,在保证推理效果的同时,控制显存占用上限。
从适用场景维度区分,KV Cache 是大模型推理阶段的标配技术,所有自回归文本生成场景都必须启用该优化;但在模型训练阶段不会使用 KV Cache,原因是训练过程中序列长度固定,不存在重复计算的问题,无需通过缓存优化性能。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐