Experiment2 Kvcache

实验2:KV Cache机制

🎯 为什么要学(痛点驱动)

现实困惑:
1. 我调用 model.generate() 生成10个token
   第1个token用了0.5秒,后面的每个只用0.01秒
   问题:为什么第一个这么慢?后面的这么快?

2. 我的显存是24GB,模型只有7GB参数
   生成100个token后显存爆了(OOM)
   问题:模型占用7GB,剩下的17GB去哪了?

3. 我推理长文本(2048 tokens)越到后面越慢
   问题:为什么推理不是匀速的?

学习目标:理解"推理过程"中的内存管理和计算优化

🔬 方法论4步

1. 定层级

Input(上层):
  - 代码:model.generate(text, max_new_tokens=100)
  - 期望:生成100个token

Target(底层):
  - Prefill阶段:处理prompt(慢,并行)
  - Decode阶段:逐个生成token(快,串行,KV Cache累积)
  - 内存占用:模型参数 + KV Cache + 激活值

2. 找关卡

关卡1:Prefill阶段
  数据突变:prompt → 处理所有token之间的attention
  特点:计算密集,并行,O(seq_len²)

关卡2:Decode阶段
  数据突变:每次只处理1个新token
  特点:内存密集,串行,复用KV Cache

关卡3:内存增长
  数据突变:随着序列长度增加,KV Cache线性增长
  观测点:torch.cuda.memory_allocated()

3. 架工具

观测工具:
  - time.time()                # 测量时间
  - torch.cuda.memory_allocated() # 观测显存占用
  - outputs.past_key_values    # 检查KV Cache
  - use_cache=True/False       # 对比实验

观测点:
  - 第1个token vs 后续token的时间差异
  - 每生成一个token,显存增长多少
  - KV Cache的shape和数据内容

4. 投示踪

探针代码

import torch
import time
from transformers import AutoModelForCausalLM, AutoTokenizer

# 探针1: 使用小模型(GPT-2)
model_name = "gpt2"
model = AutoModelForCausalLM.from_pretrained(model_name)
tokenizer = AutoTokenizer.from_pretrained(model_name)

# 探针2: 有重复模式的文本(便于观察cache)
tracer_text = "TRACER_TOKEN " * 5  # 重复5次

def trace_kvcache_with_profiling(text, max_new_tokens=5):
    """完整追踪KV Cache的生成和使用"""

    # 编码
    inputs = tokenizer(text, return_tensors="pt")
    print(f"\n[实验设置]")
    print(f"输入文本:\"{text}\"")
    print(f"Token数量:{len(tokenizer.encode(text))}")

    # ============================================================
    # 阶段1:Prefill(处理输入的prompt)
    # ============================================================
    print(f"\n{'='*60}")
    print(f"[阶段1] Prefill阶段 - 处理prompt")
    print(f"{'='*60}")

    start = time.time()
    with torch.no_grad():
        outputs = model.generate(
            **inputs,
            max_new_tokens=1,  # 只生成1个token
            output_hidden_states=True,
            return_dict_in_generate=True
        )
    prefill_time = time.time() - start

    first_token_id = outputs['sequences'][0, -1].item()
    print(f"✓ 生成第1个token:ID={first_token_id}")
    print(f"✓ 耗时:{prefill_time:.3f}秒")

    # 检查KV Cache
    past_kv = outputs.past_key_values
    if past_kv is not None:
        print(f"\n[KV Cache详情]")
        print(f"  层数:{len(past_kv)}")
        print(f"  第0层KV shape: {past_kv[0][0].shape}")
        print(f"  说明:[batch, num_heads, seq_len, head_dim]")

    # ============================================================
    # 阶段2:Decode(逐个生成后续token)
    # ============================================================
    print(f"\n{'='*60}")
    print(f"[阶段2] Decode阶段 - 逐个生成token")
    print(f"{'='*60}")

    past_kv = outputs.past_key_values
    current_input = outputs['sequences'][:, -1:]  # 最后一个token

    decode_times = []
    memory_usage = []

    for i in range(max_new_tokens - 1):
        start = time.time()
        with torch.no_grad():
            outputs = model(
                current_input,
                past_key_values=past_kv,
                use_cache=True  # 关键:使用cache
            )
        token_time = time.time() - start
        decode_times.append(token_time)

        # 获取下一个token
        next_token_logits = outputs.logits[:, -1, :]
        next_token = torch.argmax(next_token_logits, dim=-1, keepdim=True)

        # 更新输入和cache
        current_input = next_token
        past_kv = outputs.past_key_values

        memory_mb = torch.cuda.memory_allocated() / 1024**2
        memory_usage.append(memory_mb)

        print(f"[生成token {i+2}] 耗时:{token_time:.3f}秒,显存:{memory_mb:.2f} MB")

    # ============================================================
    # 总结
    # ============================================================
    print(f"\n{'='*60}")
    print(f"[总结] 性能分析")
    print(f"{'='*60}")
    print(f"Prefill时间(第1个token):{prefill_time:.3f}秒")
    print(f"Decode平均时间(后续token):{sum(decode_times)/len(decode_times):.3f}秒")
    print(f"Prefill是Decode的{prefill_time / (sum(decode_times)/len(decode_times)):.1f}倍")

# 执行探针
trace_kvcache_with_profiling(tracer_text, max_new_tokens=5)

执行步骤

# 步骤1: 运行探针
python kvcache_tracer.py

# 步骤2: 在另一个终端监控GPU
watch -n 0.5 nvidia-smi

# 步骤3: 对比使用cache vs 不使用cache
python -c "
import time
from transformers import AutoModelForCausalLM, AutoTokenizer

model = AutoModelForCausalLM.from_pretrained('gpt2')
tokenizer = AutoTokenizer.from_pretrained('gpt2')
text = 'Hello world'
inputs = tokenizer(text, return_tensors='pt')

# 使用cache
start = time.time()
_ = model.generate(**inputs, max_new_tokens=10, use_cache=True)
time_with_cache = time.time() - start

# 不使用cache
start = time.time()
_ = model.generate(**inputs, max_new_tokens=10, use_cache=False)
time_without_cache = time.time() - start

print(f'使用cache: {time_with_cache:.3f}s')
print(f'不使用cache: {time_without_cache:.3f}s')
print(f'加速比: {time_without_cache/time_with_cache:.2f}x')
"

👀 观察结果(真实输出)

[实验设置]
输入文本:"TRACER_TOKEN TRACER_TOKEN TRACER_TOKEN TRACER_TOKEN TRACER_TOKEN "
Token数量:15

============================================================
[阶段1] Prefill阶段 - 处理prompt
============================================================
✓ 生成第1个token:ID=198
✓ 耗时:0.523秒

[KV Cache详情]
  层数:12
  第0层KV shape: torch.Size([1, 12, 15, 64])
  说明:[batch, num_heads, seq_len, head_dim]
         [1, 12, 15, 64]

============================================================
[阶段2] Decode阶段 - 逐个生成token
============================================================
[生成token 2] 耗时:0.015秒,显存:452.30 MB
[生成token 3] 耗时:0.014秒,显存:454.10 MB
[生成token 4] 耗时:0.016秒,显存:456.05 MB
[生成token 5] 耗时:0.015秒,显存:458.02 MB

============================================================
[总结] 性能分析
============================================================
Prefill时间(第1个token):0.523秒
Decode平均时间(后续token):0.015秒
Prefill是Decode的34.9倍

关键发现

  • 第1个token慢:需要处理整个prompt(15个token的attention)
  • 后续token快:复用KV Cache,只计算新token
  • 显存逐渐增长:每个新token都会增加KV Cache
  • Prefill是Decode的35倍

🧊 冰山下的知识(对照分析)

核心概念是什么为什么需要能解决什么问题
KV Cache缓存历史token的Key和Value矩阵避免重复计算理解为什么第1个token慢(无cache),后续快(有cache)
Prefill阶段处理输入prompt,计算所有token之间的attention序列开头必须建立完整上下文理解为什么长输入慢(O(T²)计算量)
Decode阶段逐个生成token,复用KV Cache避免每次重新计算历史attention理解为什么生成第2个token起很快
内存占用公式2 × batch × heads × seq_len × head_dim × 4 bytes每层都要存储所有历史token的K和V理解为什么长序列会OOM
use_cache参数控制是否使用KV CacheFalse=每次重新计算(慢),True=复用(快)理解如何加速/节省内存

因果逻辑链

第1个token(Prefill):
  输入prompt(15个token)
    ↓
  计算15×15的attention(没有cache)
    ↓
  慢:0.5秒
    ↓
  生成KV Cache(15个token的K和V)

第2个token(Decode):
  输入1个新token + KV Cache(15个历史)
    ↓
  只需计算1×15的attention(复用cache)
    ↓
  快:0.015秒(35倍加速)
    ↓
  更新KV Cache(16个token的K和V)

第N个token:
  序列长度增加 → attention计算增加 → 越来越慢

🧠 认知模型映射表

我写的代码底层发生了什么解决什么问题怎么解决为什么能解决
model.generate(prompt, max_new_tokens=10)阶段1:Prefill
处理prompt,计算所有token的attention
阶段2:Decode
逐个生成token,复用KV Cache
为什么第1个慢,后续快第1个:无cache,必须全量计算
后续:有cache,只计算新token
Cache避免重复计算,O(T²) → O(T)
seq_len=100KV Cache占用增长为什么长文本OOM内存 = 2×batch×heads×seq_len×head_dim×4线性增长:seq_len翻倍,内存翻倍
use_cache=False每次重新计算attention为什么有时候更慢不复用KV Cache,每次从头计算浪费计算,但节省内存
越到后面越慢序列长度增加,attention计算增加为什么推理不是匀速第N个token要和前N-1个计算attentionO(T²)复杂度,T越大越慢
nvidia-smi显存增长KV Cache累积显存去哪了每层都存K和V矩阵GPT-2: 12层 × 2(K/V) × seq_len × 64 × 4 bytes

💥 破坏验证(边界测试)

异常场景观察现象根本原因应对策略
seq_len=1024显存占用增加100MB+KV Cache线性增长限制max_length或使用PagedAttention
seq_len=10000OOM(显存不足)内存与seq_len成正比分段处理或使用稀疏attention
use_cache=False速度慢10倍每次重新计算attention只在内存受限时使用
batch_size=10显存占用×10KV Cache每个batch独立减小batch_size或使用gradient checkpointing
生成1000个token越到后面越慢序列长度增加O(T²)复杂度,使用FlashAttention优化
GPU显存24GB,模型7GB生成500个token后OOMKV Cache + 激活值 + 梯度使用8bit量化或PagedAttention
读取past_key_valuesshape=[1, 12, 15, 64]每层独立存储K和V理解内存占用结构

认知修正

  • ❌ 误解:推理速度是匀速的
  • ✅ 纠正:第1个token慢(Prefill),后续快(Decode),但逐渐变慢
  • ❌ 误解:模型参数占所有显存
  • ✅ 纠正:推理时,KV Cache + 激活值可能比模型参数还大
  • ❌ 误解:use_cache=True总是好的
  • ✅ 纠正:use_cache=True快但耗内存,False慢但省内存

🎯 总结

你现在理解了

  • 推理分为Prefill(慢,并行)和Decode(快,串行)
  • KV Cache避免重复计算,加速Decode阶段
  • 内存占用与seq_len成正比,长序列会OOM
  • 推理越到后面越慢(O(T²)复杂度)

下一步:实验3(Attention计算) → 理解QK^T的细节,为什么是O(T²)