推理加速的重要性
大模型推理延迟直接影响用户体验。研究表明,响应时间每增加 100ms,用户满意度下降 3-5%。对于实时对话应用,首 token 延迟(TTFT)在 200ms 以内是优秀,500ms 以内是可接受,超过 1s 用户会明显感到等待。
KV Cache:最基础的加速技术
KV Cache 是 Transformer 推理中最核心的优化。在没有 KV Cache 的情况下,每次生成新 token 都需要重新计算所有历史 token 的 Key 和 Value。KV Cache 将已计算的 K、V 缓存起来,每步只需计算新 token:
# KV Cache 工作原理
class KVCacheAttention:
def __init__(self):
self.k_cache = None
self.v_cache = None
def forward(self, q, k, v, use_cache=True):
if use_cache and self.k_cache is not None:
# 将新的 K、V 追加到缓存
k = torch.cat([self.k_cache, k], dim=-2)
v = torch.cat([self.v_cache, v], dim=-2)
# 更新缓存
self.k_cache = k
self.v_cache = v
# 计算注意力
attn = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k)
attn = F.softmax(attn, dim=-1)
return torch.matmul(attn, v)
# 效果:将 O(n²) 降低到 O(n),推理速度提升 10-100xPagedAttention:高效的 KV Cache 管理
vLLM 提出的 PagedAttention 通过分页管理 KV Cache,解决了显存碎片问题:
# PagedAttention 显存管理
class PagedKVCache:
def __init__(self, block_size=16, num_blocks=1024):
self.block_size = block_size
self.free_blocks = list(range(num_blocks))
self.block_table = {} # seq_id -> [block_ids]
def allocate(self, seq_id, num_tokens):
num_blocks_needed = (num_tokens + self.block_size - 1) // self.block_size
blocks = self.free_blocks[:num_blocks_needed]
self.free_blocks = self.free_blocks[num_blocks_needed:]
self.block_table[seq_id] = blocks
return blocks
# 效果:显存利用率从 20-40% 提升到接近 100%Speculative Decoding(投机解码)
使用小模型「猜测」多个 token,再用大模型一次性验证,实现 2-3x 的解码加速:
class SpeculativeDecoder:
def __init__(self, target_model, draft_model):
self.target_model = target_model # 大模型
self.draft_model = draft_model # 小模型
def generate(self, prompt, max_tokens=100):
tokens = tokenize(prompt)
while len(tokens) < max_tokens:
# 1. 小模型生成 K 个候选 token
draft_tokens = self.draft_model.generate(tokens, k=5)
# 2. 大模型一次性验证所有候选
logits = self.target_model.forward(tokens + draft_tokens)
# 3. 接受匹配的 token,拒绝不匹配的
accepted = self._verify_and_accept(logits, draft_tokens)
tokens.extend(accepted)
if len(accepted) < len(draft_tokens):
# 从不匹配的位置重新采样
correction = self.target_model.sample(logits[len(accepted)])
tokens.append(correction)
return detokenize(tokens)算子融合
将多个小算子融合为一个大算子,减少显存读写:
# 融合前:5次显存读写
# LayerNorm → Linear → Dropout → ReLU → Linear
# 融合后:1次显存读写
# FusedMLP (LayerNorm + Linear + GELU + Linear)
# 使用 torch.compile 自动融合
import torch
@torch.compile
def fused_mlp(x, w1, w2, b1, b2):
return torch.nn.functional.linear(
torch.nn.functional.gelu(
torch.nn.functional.linear(x, w1, b1)
),
w2, b2
)
# 效果:延迟降低 20-30%,显存带宽节省 40%技术选型对比
| 技术 | 加速比 | 实现难度 | 适用场景 |
|---|---|---|---|
| KV Cache | 10-100x | 低 | 所有场景 |
| FlashAttention | 2-4x | 低 | 长序列 |
| INT8量化 | 2-3x | 中 | GPU推理 |
| Speculative Decoding | 2-3x | 高 | 低延迟场景 |
| 算子融合 | 1.2-1.5x | 中 | 所有场景 |
| TensorRT-LLM | 3-5x | 高 | NVIDIA GPU |
实践建议
- 先测量再优化:使用 Profiler 找到真正的瓶颈
- 从简单开始:KV Cache 和 FlashAttention 是性价比最高的优化
- 量化优先:INT8 量化几乎无精度损失,加速效果显著
- 组合使用:多种技术组合使用效果更好
- 持续监控:优化后持续监控延迟和吞吐量
总结
推理加速是一个系统工程。建议从 KV Cache 和量化开始,逐步引入更高级的技术。核心原则是:减少计算(量化、投机解码)、减少显存访问(FlashAttention、算子融合)、提高并行度(连续批处理)。