推理加速的重要性

大模型推理延迟直接影响用户体验。研究表明,响应时间每增加 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-100x

PagedAttention:高效的 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 Cache10-100x所有场景
FlashAttention2-4x长序列
INT8量化2-3xGPU推理
Speculative Decoding2-3x低延迟场景
算子融合1.2-1.5x所有场景
TensorRT-LLM3-5xNVIDIA GPU

实践建议

  1. 先测量再优化:使用 Profiler 找到真正的瓶颈
  2. 从简单开始:KV Cache 和 FlashAttention 是性价比最高的优化
  3. 量化优先:INT8 量化几乎无精度损失,加速效果显著
  4. 组合使用:多种技术组合使用效果更好
  5. 持续监控:优化后持续监控延迟和吞吐量

总结

推理加速是一个系统工程。建议从 KV Cache 和量化开始,逐步引入更高级的技术。核心原则是:减少计算(量化、投机解码)、减少显存访问(FlashAttention、算子融合)、提高并行度(连续批处理)。