LLM推理优化

阶段3 | 第25-28周

📅 4周详细学习计划

第25周:Attention基础

• Self-Attention原理
• 复杂度分析
• KV Cache动机
• 内存占用计算

第26周:FlashAttention

• IO感知原理
• 分块计算(Tiling)
• Online Softmax
• 实现源码分析

第27周:PagedAttention

• vLLM核心原理
• 内存分页管理
• Copy-on-Write
• 连续批处理

第28周:高级优化

• Continuous Batching
• 投机解码
• 多卡推理
• 性能对比测试

1. KV Cache详解

1.1 问题背景

1.2 内存占用计算

// KV Cache内存计算公式
// 对于单个序列:
KV_Cache_Size = 2 × num_layers × num_heads × seq_len × head_dim × dtype_size

// 示例:Llama-2-70B
// num_layers=80, num_heads=64, head_dim=128, seq_len=4096, FP16
KV_Cache = 2 × 80 × 64 × 4096 × 128 × 2 bytes
         = 2 × 80 × 64 × 4096 × 128 × 2
         = 10.7 GB (单序列!)

// 这就是为什么大模型推理需要巨大显存

1.3 KV Cache优化策略

2. FlashAttention详解

2.1 核心思想

2.2 内存对比

// Standard Attention
// 1. 计算 S = Q × K^T (N×N矩阵,O(N²)内存)
// 2. 计算 P = softmax(S) (N×N矩阵)
// 3. 计算 O = P × V (输出)
// 总内存: O(N²) + O(N×d) = O(N²)

// FlashAttention
// 1. 分块加载Q,K,V到SRAM
// 2. 在SRAM中计算部分注意力
// 3. Online Softmax更新统计量
// 4. 逐步累加输出
// 总内存: O(N) - 只需要存储输出和统计量

2.3 使用示例

# FlashAttention Python接口
from flash_attn import flash_attn_func

# q, k, v: [batch, seqlen, nheads, headdim]
# 不同head_dim需要不同版本的FlashAttention
output = flash_attn_func(q, k, v, causal=True)

# FlashAttention-2 (更快)
from flash_attn import flash_attn_func
output = flash_attn_func(q, k, v, causal=True, window_size=(-1, -1))

3. PagedAttention (vLLM)

3.1 问题

3.2 解决方案

# PagedAttention内存管理
# 传统方法:每个序列预分配max_seq_len内存
# PagedAttention:按需分配Block

# 示例:
# 序列1: 100 tokens → 7个Block (每个Block 16 tokens)
# 序列2: 50 tokens → 4个Block
# 序列3: 200 tokens → 13个Block
# 总共: 24个Block (而不是350个Block预分配)

3.3 vLLM使用

# vLLM安装和使用
pip install vllm

# 基本使用
from vllm import LLM, SamplingParams

# 初始化模型
llm = LLM(
    model="meta-llama/Llama-2-7b-hf",
    tensor_parallel_size=2,  # 2卡并行
    max_model_len=4096,
    gpu_memory_utilization=0.9
)

# 批量生成
prompts = ["Hello, how are you?", "What is AI?"]
sampling_params = SamplingParams(temperature=0.7, max_tokens=100)
outputs = llm.generate(prompts, sampling_params)

4. Continuous Batching

4.1 Static vs Continuous

4.2 实现原理

5. 投机解码 (Speculative Decoding)

5.1 原理

5.2 实现方案

💡 关键洞察:LLM推理的主要瓶颈是内存带宽(读取KV Cache),而不是计算。优化重点是减少内存访问和数据传输。

🛠️ 实践项目

📚 学习资源