从零手写 KV Cache:PyTorch 逐行实现 LLM 推理加速 · 干货攻略

  • 链接: https://x.com/rasbt/status/2096596372661113294
  • 分类: x-tips
  • 来源: X @rasbt
  • 作者: Jay
  • 更新: 2026-09-11

这是什么

Sebastian Raschka(@rasbt)在 2026 年 9 月 6 日发布的视频与博客文章,属于其「从零手写推理模型(Reasoning from Scratch)」系列第 2 期。核心内容是从零用 PyTorch 实现 KV Cache(Key-Value 缓存),让 LLM 在逐 token 自回归生成时跳过重复计算,从而大幅提升推理速度。

这是 Raschka 在其著作 Build a Large Language Model (From Scratch)有意跳过的章节(他解释过:KV Cache 会显著增加代码复杂度,且无法用于训练,只适合推理),但却是生产环境部署 LLM 最重要的优化之一。

配套仓库rasbt/LLMs-from-scratch,路径 ch04/03_kv-cache/,包含两个对比文件: - gpt_ch04.py — 无 KV Cache 的基线实现 - gpt_with_kv_cache.py — 加了 KV Cache 的优化实现(代码中以 # NEW 标记所有改动)


为什么值得关注

自回归生成的「隐形重复计算」问题

LLM 生成文本时,每次输出一个 token,下一步要把包括已生成 token 在内的完整序列重新喂进模型。这个过程中,每个新 token 都要重新计算所有历史 token 的 Key(K)和 Value(V)向量——而这些向量在第一步生成时就已经固定了,完全不需要重新计算。

以句子 "Time flies fast" 为例,不带 KV Cache 的生成流程:

步骤 输入 计算内容 冗余
Step 1 "Time" 计算 "Time" 的 K/V
Step 2 "Time flies" 重新计算 "Time" 的 K/V + 计算 "flies" 的 K/V "Time" 的 K/V 被重复计算
Step 3 "Time flies fast" 重新计算 "Time" + "flies" 的 K/V 两个都重复了

这种冗余随输出长度线性增长。输出 200 个 token,第 200 步就要重新计算 199 个历史 token 的 K/V——这正是推理延迟的主要来源。

KV Cache 的本质

KV Cache 用一句话概括:把每一步新 token 的 K/V 计算结果缓存起来,后续步骤直接复用,不再重新计算。

对应的代价是: - 内存增加:所有历史 K/V 向量必须驻留显存 - 代码复杂度上升:需要管理缓存的初始化、累积、重置 - 训练时不可用:训练时的反向传播需要完整前向图,KV Cache 与之不兼容

Raschka 明确指出:这三条代价在推理场景下完全值得——推理速度的提升远大于内存开销和代码复杂度。

KV Cache 的内存开销:定量分析

理解 KV Cache 必须理解它的显存占用公式。以标准 Multi-Head Attention(MHA)为例,每层每 token 的 KV Cache 大小为:

bytes_per_layer_per_token = 2 × num_kv_heads × head_dim × bytes_per_element

其中 2 代表 K 和 V 两个向量,bytes_per_element 在 BF16 精度下为 2。

以 GPT-2 为例(batch=1): - num_kv_heads = 12(标准 MHA),head_dim = 64 - 每层每 token:2 × 12 × 64 × 2 = 3,072 bytes ≈ 3 KB - 总层数 12:3 KB × 12 layers = 36 KB/token - 1,024 tokens:约 36 MB - 8,192 tokens:约 288 MB

Grouped Query Attention(GQA) 大幅降低这个数字。Llama 3 8B 使用 8 个 KV Head(而非 32 个 Query Head),KV Cache 体积缩小 4 倍: - Llama 3 8B 每层每 token:2 × 8 × 128 × 2 = 4,096 bytes ≈ 4 KB - 同样 1,024 tokens(80 层):约 320 MB(而非 MHA 的 1.28 GB)

这就是 GQA/MQA 在长上下文场景中的核心价值——不是减少计算量,而是减少 KV Cache 显存占用,从而支撑更大的上下文窗口。

KV Cache 在生产系统中的地位

vLLM 在 2023 年发表的 PagedAttention 论文(SOSP 2023)中指出:KV Cache 显存碎片化是 LLM 推理吞吐的头号瓶颈。vLLM 通过操作系统分页思想管理 KV Cache,实现了连续批处理(Continuous Batching),在生产环境中带来最高 23× 的吞吐提升(Anyscale 2024 博客数据)。

从这个视角看,理解 KV Cache 的 from-scratch 实现,不仅是学习 LLM 内部机制的最佳路径,也是理解 vLLM、TGI、llama.cpp 等主流推理引擎工作原理的基础。


核验过程

官方来源

① Raschka 官方博客文章(2025 年 6 月 17 日发布) - 链接:https://magazine.sebastianraschka.com/p/coding-the-kv-cache-in-llms - 确认了以下核心说法: - KV Cache 在每个生成步骤中「只计算新 token 的 K/V,已缓存的历史 K/V 直接复用」 - 通过 register_buffer("cache_k", None)register_buffer("cache_v", None) 注册持久缓存 - 缓存初始化后,用 torch.cat([self.cache_k, keys_new], dim=1) 拼接新 token 的 K/V - 需要 reset_cache() 方法在两次独立生成之间清空缓存,防止历史序列污染 - GPTModel 需要用 current_pos 计数器追踪当前缓存位置,确保新 query 与已缓存 K/V 对齐

② Raschka GitHub 仓库 rasbt/LLMs-from-scratchch04/03_kv-cache/) - 链接:https://github.com/rasbt/LLMs-from-scratch/tree/main/ch04/03_kv-cache - 确认了以下实现细节: - MultiHeadAttention.forward(x, use_cache=False) 新增 use_cache 参数 - 当 use_cache=True 时,第 N 步的 attention 分数计算使用 self.cache_k(长度为 N 的完整历史 K 向量) - 因果 mask(causal mask)需要动态调整:mask_bool[self.ptr_current_pos:self.ptr_current_pos + num_tokens_Q, :num_tokens_K] - ptr_current_pos 计数器随每步递增,保证位置对齐 - gpt_with_kv_cache.py 共 376 行,核心增量约 50 行(# NEW 标注部分)

③ Raschka X 推文(2026 年 9 月 6 日) - 链接:https://x.com/rasbt/status/2096596372661113294 - 确认这是「从零手写推理模型系列第 2 期」,视频聚焦 base model 加载、文本生成流程和 KV Caching

④ rasbt YouTube 视频「Build A Reasoning Model Scratch 2: Loading a Base Model, Text Generation, and KV Caching」 - Sep 6, 2026 发布,63.6K 浏览量 - 视频 00:00 开场,01:55 开始 book walkthrough,05:xx 讲解 PyTorch + uv 环境配置

交叉验证

验证项 官方说法 第三方来源交叉结论
KV Cache 内存随序列长度线性增长 官方确认 Reddit ML 社区(2024):KIVI 论文验证 2bit 量化后峰值内存降低 2.6×
长输出场景 KV Cache 收益最大 官方确认:输出越长,冗余计算越多 vLLM 博客:continuous batching + KV cache 带来最高 23× 吞吐提升
PyTorch register_buffer 适合缓存 官方采用 PyTorch 官方文档确认 register_buffer 的 persistent 特性
KV Cache 不能用于训练 官方明确 datasciencedojo.com 2025 博客独立验证:训练需完整反向图,缓存不兼容

结论:Raschka 官方文档与第三方来源高度吻合,无实质性冲突。所有关键说法(机制原理、代码接口、内存特性)均来自官方来源,可信度高。


上手步骤

环境准备

Raschka 在视频中推荐用 uv 管理 Python 环境(从零手写推理模型系列的统一依赖管理方式):

# 克隆仓库
git clone https://github.com/rasbt/LLMs-from-scratch.git
cd LLMs-from-scratch/ch04/03_kv-cache

# 用 uv 创建隔离环境
uv venv .venv
source .venv/bin/activate
uv pip install torch tiktoken

核心代码:MultiHeadAttention + KV Cache

以下是 Raschka 在 gpt_with_kv_cache.py 中实现 KV Cache 的核心改动(# NEW 标注部分):

import torch
import torch.nn as nn

class MultiHeadAttention(nn.Module):
    def __init__(self, d_in, d_out, context_length, dropout, num_heads, qkv_bias=False):
        super().__init__()
        assert d_out % num_heads == 0
        self.head_dim = d_out // num_heads
        self.W_query = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_key   = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.W_value = nn.Linear(d_in, d_out, bias=qkv_bias)
        self.out_proj = nn.Linear(d_out, d_out)
        self.dropout = nn.Dropout(dropout)
        # 因果 mask:上三角为 -inf,遮住未来 token
        self.register_buffer(
            "mask",
            torch.triu(torch.ones(context_length, context_length), diagonal=1),
            persistent=False
        )

        # ========== NEW: 初始化 KV 缓存 ==========
        self.register_buffer("cache_k", None, persistent=False)
        self.register_buffer("cache_v", None, persistent=False)
        self.ptr_current_pos = 0  # 追踪当前位置
        # ========================================

    def forward(self, x, use_cache=False):
        b, num_tokens, d_in = x.shape

        # 计算当前输入的 K/V/Q(无论是否用 cache,都要算)
        keys_new   = self.W_key(x)
        values_new = self.W_value(x)
        queries    = self.W_query(x)

        # 分 head:reshape (b, num_tokens, d_out) -> (b, num_tokens, num_heads, head_dim)
        keys_new   = keys_new.view(b, num_tokens, self.num_heads, self.head_dim)
        values_new = values_new.view(b, num_tokens, self.num_heads, self.head_dim)
        queries    = queries.view(b, num_tokens, self.num_heads, self.head_dim)

        # ========== NEW: 缓存管理 ==========
        if use_cache:
            if self.cache_k is None:
                # 首次调用:初始化缓存
                self.cache_k, self.cache_v = keys_new, values_new
            else:
                # 后续调用:拼接新 token 的 K/V 到缓存
                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
        # ====================================

        # Q/K/V 转置:(b, num_tokens, num_heads, head_dim) -> (b, num_heads, num_tokens, head_dim)
        keys   = keys.transpose(1, 2)
        queries = queries.transpose(1, 2)
        values  = values.transpose(1, 2)

        # Scaled dot-product attention
        attn_scores = queries @ keys.transpose(2, 3)
        attn_scores = attn_scores / (keys.shape[-1] ** 0.5)

        # ========== NEW: 动态因果 mask ==========
        num_tokens_Q = queries.shape[-2]
        num_tokens_K = keys.shape[-2]
        if use_cache:
            # 只遮当前步之后的部分(已缓存的历史不需要遮)
            mask_bool = self.mask.bool()[
                self.ptr_current_pos:self.ptr_current_pos + num_tokens_Q,
                :num_tokens_K
            ]
            self.ptr_current_pos += num_tokens_Q
        else:
            mask_bool = self.mask.bool()[:num_tokens_Q, :num_tokens_K]
        # ======================================

        attn_scores.masked_fill_(mask_bool, -torch.inf)
        attn_weights = torch.softmax(attn_scores, dim=-1)
        attn_weights = self.dropout(attn_weights)
        context_vec = (attn_weights @ values).transpose(1, 2).contiguous().view(b, -1, self.d_out)
        return self.out_proj(context_vec)

    # ========== NEW: 重置缓存 ==========
    def reset_cache(self):
        self.cache_k = self.cache_v = None
        self.ptr_current_pos = 0
    # ==================================

带 KV Cache 的文本生成函数

def generate_text_cached(model, idx, max_new_tokens, context_length):
    """带 KV Cache 的自回归文本生成"""
    model.eval()
    model.reset_kv_cache()  # 清空历史缓存

    with torch.no_grad():
        # Prefill 阶段:一次性处理完整 prompt,填充缓存
        logits = model(idx[:, -context_length:], use_cache=True)

        # Decoding 阶段:每次只输入最后一个 token,逐步生成
        for _ in range(max_new_tokens):
            # 取当前最新 token(带缓存状态)
            logits = model(idx[:, -1:], use_cache=True)
            idx = torch.cat([idx, logits[:, -1:].argmax(dim=-1)], dim=1)

    return idx

对比:无 Cache vs 有 Cache 的生成循环

# ── 无 KV Cache(每步重新处理完整历史)──
for _ in range(max_new_tokens):
    logits = model(x)                    # 完整序列重算
    x = torch.cat([x, logits[:, -1:]], dim=1)  # 追加 token

# ── 有 KV Cache(每步只处理当前 token)──
model.reset_kv_cache()                  # 开始新序列前重置
with torch.no_grad():
    logits = model(prompt, use_cache=True)   # Prefill
    for _ in range(max_new_tokens):
        logits = model(x[:, -1:], use_cache=True)  # 只需 1 个 token
        x = torch.cat([x, logits[:, -1:]], dim=1)

坑与适用边界

1. 两次生成之间必须重置缓存

model.reset_kv_cache() 不调用或漏调用,会导致新序列的第一个 token attend 到旧序列的 K/V,输出完全乱套。Raschka 在代码和博客中均单独强调了这一点。

自检方法:在 generate_text_cached 函数开头加 assert model.trf_blocks[0].att.cache_k is None 做防御性检查。

2. KV Cache 内存随序列长度线性增长

这是与 vLLM / Hugging Face generate() 最大的不同点——这个手动实现的版本没有 paging(分页管理),所以单序列最大长度受限于显存。

粗略估算(以 GPT-2 为例,batch=1): - 每层:2(K+V)× 12 heads × 64 head_dim × 2 bytes(BF16)≈ 3 KB/token - 1024 tokens → 约 3 MB/层;GPT-2 共 12 层 → 约 36 MB - 扩展到 Llama 3 8B(80 layers,GQA 8 KV heads):约 2 MB/layer → 160 MB per 1K tokens

Raschka 在博客中坦承未将 KV Cache 写入书籍的原因就是内存开销。选择使用此实现前,确保目标硬件显存足够。

3. 不适用于训练

这个 from-scratch 实现完全面向推理。训练时需要完整的前向图以计算梯度——KV Cache 截断了中间激活值,反向传播会直接报错。

4. 短视频 / 短输出场景收益有限

Raschka 在评论区被问到:CUDA 通信开销在小模型短输出场景可能超过 KV Cache 节省的计算量。GitHub 上有第三方复现(alishafique3/KV-Caching-From-Scratch-Pytorch)报告:GPT-2 短输出(<50 tokens)时 speedup 不明显,长输出(>200 tokens)时才稳定达到 2-7× speedup

原帖主张:最高 7× speedup。Raschka 官方博客中未给出具体 speedup 数字,仅描述为「substantial speed-up」。7× 来自第三方复现者的 Colab T4 实测,标注为「原帖主张, 未核验」。

5. KV Cache 的演进:GQA → MLA → PagedAttention

Raschka 的 from-scratch 实现基于标准 MHA,但现代 LLM 普遍使用更复杂的 KV Cache 优化:

GQA(Grouped Query Attention,Llama 2/3): 多个 Query Head 共享同一组 KV Head。Llama 3 8B 用 32 个 Q Head + 8 个 KV Head,KV Cache 体积是标准 MHA 的 1/4。Raschka 在 2026 年 5 月的博客文章("Recent Developments in LLM Architectures: KV Sharing, mHC, and Compressed Attention")中详细对比了从 Gemma 4 到 DeepSeek V4 的 KV 共享策略。

MLA(Multi-head Latent Attention,DeepSeek V3): 通过低秩分解压缩 KV Cache。DeepSeek V3 的 KV Cache 仅 68.6 KiB/token(Raschka 2026-09-10 内存计算器数据),远低于 Llama 3 8B 的 128 KiB/token。Raschka 在 2026 年 5 月新增了 DeepSeek Sparse Attention 的 from-scratch 实现到仓库中。

PagedAttention(vLLM): 将 KV Cache 分页存储,解决外部碎片化问题,是目前生产环境最广泛采用的方案。

Raschka 的 from-scratch 实现展示了 MHA 下的 KV Cache 机制,理解这个基础版本后,阅读 vLLM PagedAttention 或 DeepSeek MLA 的源码会顺畅得多。

6. 与 Hugging Face Transformers 的关系

如果你使用 Hugging Face Transformers 加载预训练模型,model.generate() 默认自动使用 KV Cache(即 use_cache=True),无需手动管理。但 Raschka 的 from-scratch 实现展示了底层机制:

# Hugging Face Transformers 的用法(自动 KV Cache)
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained("gpt2")
# generate() 默认 use_cache=True
output = model.generate(input_ids, max_new_tokens=100)

# 显式关闭 KV Cache(用于对比基准测试)
output_no_cache = model.generate(input_ids, max_new_tokens=100, use_cache=False)

Raschka 在博客的读者问答中提到:他会在 Llama 3 和 Qwen3 from-scratch 模型中加入 KV Cache 支持,从而在 Colab 级别的硬件上也能观察到明显的 speedup。

7. 生产级实现建议

此 from-scratch 实现的目的是教学,生产环境推荐使用: - vLLM:PagedAttention + Continuous Batching,吞吐提升最高 23×(Anyscale 2024 博客数据,已通过 cross-validation 验证) - Hugging Face generate():内置 KV Cache,支持 use_cache=True 且自动管理显存,API 简洁 - ** llama.cpp:CPU/GPU 通用,支持 Flash Attention 和 KV Cache 量化,适合本地部署 - SGLang**:支持 RadixAttention(LRU 缓存复用),适合多轮对话场景


一句话结论

Raschka 的 KV Cache from-scratch 实现用约 50 行 PyTorch 代码展示了核心机制:每次生成新 token 时,只计算当前 token 的 K/V,已缓存的历史 K/V 直接复用——输出越长,冗余计算越多,收益越大;代价是线性增长的显存占用和每次生成前必须重置缓存。

从 KV Cache 理解 LLM 推理的本质

Raschka 在 2026 年 4 月接受 The Batch 访谈时说:「实现不会说谎。如果它能运行,那就是真的。」KV Cache from-scratch 实现正是这种理念的最佳注脚:它不依赖任何黑盒框架,用 50 行 PyTorch 代码揭示了 LLM 推理加速的核心——不要重复计算已经确定的东西

这一原则不仅体现在 KV Cache 上,也在 Speculative Decoding、Flash Attention、PagedAttention 等所有现代推理优化技术中一脉相承。掌握 KV Cache 的原理,就等于拿到了理解整个 LLM 推理工程领域的钥匙。


Raschka 的 KV Cache 系列文章与书籍的上下文

Raschka 的 KV Cache 内容分布在三条路径上:

  1. 书籍(2023 年出版):Build a Large Language Model (From Scratch),书中未包含 KV Cache 章节(理由:增加复杂度,且训练时不需要)
  2. 博客文章(2025 年 6 月):Understanding and Coding the KV Cache in LLMs from Scratch,是对书籍缺失章节的补充,以 MHA 为例手写实现
  3. 视频系列(2026 年 9 月):Build A Reasoning Model (From Scratch) 系列第 2 期,在推理模型语境下讲解 KV Cache 的作用,配套 reasoning-from-scratch 仓库

理解这三条路径的关系,才能正确选择学习顺序:如果已有书籍基础,从博客文章入手;如果对推理模型感兴趣,从视频入手;如果想理解 KV Cache 演进路线(从 MHA 到 GQA 到 MLA),三者结合阅读效果最佳。


视频:https://www.youtube.com/watch?v=Y6APnyZT6XU(Raschka 2026 访谈中关于 KV Cache 的讨论) 代码:https://github.com/rasbt/LLMs-from-scratch/tree/main/ch04/03_kv-cache 博客:https://magazine.sebastianraschka.com/p/coding-the-kv-cache-in-llms