课程0基础Agent开发课 / LLM基础 / LLM推理优化-KV-Cache量化与Flash-Attention
— 18 min read

LLM推理优化-KV-Cache量化与Flash-Attention

> **[进阶选读]** 本文适合想深入理解 LLM 内部机制的读者。KV Cache、量化和 Flash Attention 是 LLM 推理优化的三大技术,理解它们能帮你理解为什么不同模型的推理速度和成本差异如此之大。路径 A 的读者可以跳过。

LLM 推理优化:KV Cache、量化与 Flash Attention

[进阶选读] 本文适合想深入理解 LLM 内部机制的读者。KV Cache、量化和 Flash Attention 是 LLM 推理优化的三大技术,理解它们能帮你理解为什么不同模型的推理速度和成本差异如此之大。路径 A 的读者可以跳过。


[进阶选读] 本篇讲解推理服务优化技术,适合需要自建推理服务或深入理解模型性能瓶颈的读者。如果你的目标是应用开发(路径 A),可以跳过本篇,不影响后续学习。

语言模型推理慢而贵,根本原因在于自回归生成(Auto-regressive Generation):每生成一个 Token,都需要对整个已有序列做一次完整的 Transformer 前向计算(Forward Pass)。生成 500 个 Token 就需要 500 次前向传播。

这篇介绍四类核心推理加速技术:KV Cache、量化(Quantization)、Flash Attention 和投机采样(Speculative Decoding)。


1.1 逐 Token 生成的计算本质

在 Transformer 的自注意力机制中,每个 Token 都需要与序列中所有其他 Token 进行注意力计算:

KV Cache 工作原理
图 6.14:KV Cache 工作原理——无缓存每步重算全量 vs 有缓存只计算新 Token

code
Attention(Q, K, V) = softmax(QK^T / √d_k) × V

生成第 $n$ 个 Token 时,需要计算它与前面所有 $n-1$ 个 Token 的注意力得分。如果序列长度翻倍,计算量近似翻倍。对于 1000 Token 的上下文,生成最后一个 Token 要处理前面 999 个 Token 的 Key 和 Value 矩阵。

问题的核心:每一步生成都在重复计算前面 Token 的 K、V 矩阵,而这些矩阵从未改变。


1.2 KV Cache:缓存,不重复计算

1.2.1 原理

KV Cache 的思想极其直接:把已经计算过的 Key 和 Value 矩阵缓存起来,下一步直接复用

有KV Cache - 增量计算

Token 1-5 已生成
Cache: K1-5, V1-5
生成 Token 6

仅计算 K6,V6,Q6
从缓存读取 K1-5, V1-5

Attention Q6, K1-6, V1-6

输出 Token 6
更新 Cache: K1-6, V1-6

无KV Cache - 每步重算

Token 1-5 已生成
生成 Token 6

计算 K1,V1
计算 K2,V2
计算 K3,V3
计算 K4,V4
计算 K5,V5
计算 Q6

Attention Q6, K1-5, V1-5

输出 Token 6

python
# 展示有无 KV Cache 时的计算量差异(概念性演示)
import numpy as np

def compute_attention_flops(seq_len, d_model, n_heads, use_kv_cache=False):
    """
    计算生成第 seq_len 个 token 时的注意力 FLOPs。
    d_model: 模型维度
    n_heads: 注意力头数
    """
    d_k = d_model // n_heads

    if use_kv_cache:
        # 有 KV Cache:只需计算新 token 的 Q、K、V,然后与缓存的 K、V 做注意力
        # Q 投影: d_model × d_model(仅对 1 个新 token)
        # 注意力计算: 1 × seq_len(新 Q 与所有缓存 K)
        flops_projection = d_model * d_model * 3   # Q、K、V 投影(1 个 token)
        flops_attention = seq_len * d_k * n_heads  # 注意力得分计算
        return flops_projection + flops_attention
    else:
        # 无 KV Cache:需要重新计算所有 token 的 Q、K、V
        # 每个 token 的投影:d_model × d_model
        # 注意力矩阵:seq_len × seq_len
        flops_projection = d_model * d_model * 3 * seq_len
        flops_attention = seq_len * seq_len * d_k * n_heads
        return flops_projection + flops_attention

# 模拟生成 200 个 token 的总计算量
d_model, n_heads = 4096, 32
total_no_cache = sum(compute_attention_flops(i, d_model, n_heads, use_kv_cache=False)
                     for i in range(1, 201))
total_with_cache = sum(compute_attention_flops(i, d_model, n_heads, use_kv_cache=True)
                       for i in range(1, 201))

print(f"无 KV Cache 总 FLOPs: {total_no_cache:.2e}")
print(f"有 KV Cache 总 FLOPs: {total_with_cache:.2e}")
print(f"节省比例: {1 - total_with_cache/total_no_cache:.1%}")
# 在长序列下节省比例接近 99%

1.2.2 KV Cache 的内存代价

KV Cache 并非免费,它消耗大量显存:

code
KV Cache 内存 = 2 × batch_size × seq_len × n_layers × d_head × n_heads × bytes_per_element

以 Llama 3 8B(FP16)为例:

  • 32 层,32 头,d_head = 128
  • batch_size=1,seq_len=8192
  • 2 × 1 × 8192 × 32 × 128 × 32 × 2 bytes ≈ 4 GB

这意味着在 8B 模型本身占用约 16 GB 显存的基础上,8K 上下文的 KV Cache 再额外需要 4 GB,长上下文场景会很快耗尽显存。

1.2.3 Prompt Caching vs KV Cache 的区别

维度 KV Cache(架构层) Prompt Caching(应用层)
发生位置 推理引擎内部,每次请求内 API 服务端,跨请求复用
作用范围 单次生成过程中的增量计算 多次 API 调用之间复用相同前缀
用户感知 透明,自动发生 需要保持相同的 System Prompt 前缀
成本影响 降低单次推理的计算量 Anthropic/OpenAI 对缓存命中的 Token 折扣计费
典型场景 所有 LLM 推理 长 System Prompt 的多次调用

1.3 量化(Quantization):降低数值精度

1.3.1 精度与内存的权衡

神经网络的参数默认以 FP32(32 位浮点数,精度高但占用内存大)存储,量化(Quantization,将高精度浮点数压缩为低精度整数,以牺牲少量精度换取大幅节省内存)将其压缩为更低精度。GPU(图形处理器)是深度学习推理的主要硬件,其显存(VRAM)决定了能加载多大的模型:

格式 位数 内存占用(7B 模型) 相对精度损失 适用场景
FP32 32 位 ~28 GB 基准 训练
BF16 / FP16 16 位 ~14 GB 极小 推理标准
INT8 8 位 ~7 GB 生产服务器
INT4 4 位 ~3.5 GB 中等 消费级 GPU/CPU
INT2 / 2bit 2 位 ~1.75 GB 较大 极端资源受限

量化的基本原理是将浮点数映射到整数区间:

python
import numpy as np

def quantize_int8(tensor):
    """
    INT8 对称量化示例。
    将浮点张量映射到 [-127, 127] 范围。
    """
    # 计算缩放因子:用最大绝对值映射到 INT8 范围
    scale = np.max(np.abs(tensor)) / 127.0

    # 量化:浮点 → INT8(有损压缩)
    quantized = np.round(tensor / scale).astype(np.int8)

    # 反量化:INT8 → 浮点(用于实际计算)
    dequantized = quantized.astype(np.float32) * scale

    quantization_error = np.mean(np.abs(tensor - dequantized))
    return quantized, scale, quantization_error

# 示例:随机权重矩阵
weights = np.random.randn(100, 100).astype(np.float32)
q_weights, scale, error = quantize_int8(weights)
print(f"原始内存: {weights.nbytes / 1024:.1f} KB")
print(f"量化内存: {q_weights.nbytes / 1024:.1f} KB")
print(f"压缩比: {weights.nbytes / q_weights.nbytes:.1f}x")
print(f"平均量化误差: {error:.6f}")

1.3.2 GGUF 格式:为什么 llama.cpp 能在消费级硬件跑大模型

GGUF(GPT-Generated Unified Format) 是 llama.cpp 使用的模型格式,核心特点:

  1. 混合精度量化:不同层使用不同量化精度(注意力层保持 FP16,FFN 层量化到 INT4),在精度和速度之间精细权衡。
  2. CPU 友好:针对 x86/ARM 的 SIMD 指令集优化,无需 GPU 也能运行。
  3. 内存映射(mmap):支持将模型映射到内存而非完整加载,支持超出 RAM 的大模型。
bash
# 用 llama.cpp 下载并运行量化模型的典型流程
# 安装
git clone https://github.com/ggerganov/llama.cpp
cd llama.cpp && make -j4

# 下载 Q4_K_M 量化版本(4 bit,约 4.1 GB,适合 8GB 内存设备)
# Llama 3 8B 原始 FP16 版本约 16 GB
huggingface-cli download \
    bartowski/Meta-Llama-3-8B-Instruct-GGUF \
    Meta-Llama-3-8B-Instruct-Q4_K_M.gguf

# 运行推理
./llama-cli -m Meta-Llama-3-8B-Instruct-Q4_K_M.gguf \
    --prompt "解释量子纠缠" \
    -n 200

1.3.3 AWQ vs GPTQ vs GGUF 的适用场景

格式 量化方式 适用硬件 优势 劣势
GGUF (Q4_K_M) 混合 4bit CPU + GPU 消费级设备,极易部署 GPU 推理速度非最优
GPTQ 4/8bit 逐层 NVIDIA GPU 成熟工具链,支持广泛 量化速度慢(需数小时)
AWQ 4bit 激活感知 NVIDIA GPU 精度损失最小 仅限 NVIDIA,工具链较新
BitsAndBytes 动态 INT8/4bit NVIDIA GPU 集成 HuggingFace,一行代码 推理速度较慢

1.4 Flash Attention:解决内存瓶颈

1.4.1 标准 Attention 的问题

标准自注意力计算的内存复杂度为 $O(n^2)$($n$ 为序列长度),因为需要显式存储 $n \times n$ 的注意力矩阵:

python
# 标准 Attention 的内存问题演示
import torch

def standard_attention(Q, K, V):
    """
    标准注意力实现。
    关键问题:注意力矩阵 A 的大小是 [seq_len, seq_len],
    seq_len=8192 时 A 占用 8192² × 2 bytes ≈ 128 MB(每层)
    32 层模型的注意力矩阵总计 ~4 GB,严重挤占显存。
    """
    d_k = Q.size(-1)
    # 这里产生 [batch, heads, seq_len, seq_len] 的大矩阵
    scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5)
    attn_weights = torch.softmax(scores, dim=-1)  # 必须存储完整矩阵
    return torch.matmul(attn_weights, V)

1.4.2 分块计算(Tiling):Flash Attention 的核心

Flash Attention(Tri Dao 等,2022)的关键思想:不需要存储完整的注意力矩阵,通过分块计算(Tiling,将大矩阵切成小块分批处理)和在线 softmax 技术,在一次 GPU 内存扫描中完成注意力计算

python
# Flash Attention 的使用(PyTorch 2.0+ 内置)
import torch
import torch.nn.functional as F

# 方法 1:直接使用 PyTorch 的 scaled_dot_product_attention
# PyTorch 2.0+ 会自动选择 Flash Attention(如果 CUDA 版本支持)
def flash_attention_pytorch(Q, K, V):
    """
    PyTorch 2.0+ 内置了 Flash Attention 的调用。
    is_causal=True 表示因果掩码(自回归生成必须使用)。
    """
    return F.scaled_dot_product_attention(
        Q, K, V,
        attn_mask=None,
        dropout_p=0.0,
        is_causal=True  # 自动使用 causal mask,无需显式创建大矩阵
    )

# 方法 2:使用 flash-attn 库(更完整的实现)
# pip install flash-attn
from flash_attn import flash_attn_func

output = flash_attn_func(
    Q, K, V,
    dropout_p=0.0,
    causal=True  # 因果注意力
)

Flash Attention 的核心创新是将注意力计算分解为小块,每块在 GPU 的 SRAM(Static Random-Access Memory,GPU 芯片上的超高速缓存)中完成,避免了频繁读写 HBM(High Bandwidth Memory,即显存,容量大但速度较慢)。

性能提升数据:

序列长度 标准 Attention Flash Attention 2 速度提升
1K tokens 基准 ~2x 约 2 倍
4K tokens 基准 ~4x 约 4 倍
16K tokens 基准 ~6-8x 约 6-8 倍
64K tokens OOM(显存不足) 可运行 关键突破

Flash Attention 2 进一步优化了并行化策略,是当前所有主流推理框架(vLLM、TGI——Text Generation Inference、TensorRT-LLM 等开源LLM推理加速框架)的标配。


1.5 投机采样(Speculative Decoding)

1.5.1 思路

自回归生成的本质限制是串行——必须等上一个 Token 生成后才能生成下一个。投机采样打破了这一限制:

  1. 起草(Draft):用一个小型模型(Draft Model,如 68M 参数)快速生成 $k$ 个候选 Token
  2. 验证(Verify):用大型目标模型(Target Model,如 7B 参数)并行验证这 $k$ 个 Token
  3. 接受/拒绝:接受与目标模型分布吻合的 Token,从第一个不接受的位置重新生成
python
# 投机采样的概念性实现
import torch

def speculative_decoding_step(draft_model, target_model, input_ids, k=4):
    """
    draft_model: 小型快速模型(如同系列的小版本)
    target_model: 大型精确模型
    k: 每步起草的 token 数量
    """
    # 步骤 1:小模型快速起草 k 个 token
    draft_tokens = []
    draft_probs = []
    current_ids = input_ids.clone()

    for _ in range(k):
        with torch.no_grad():
            draft_logits = draft_model(current_ids).logits[:, -1, :]
            draft_prob = torch.softmax(draft_logits, dim=-1)
            next_token = torch.multinomial(draft_prob, 1)
            draft_tokens.append(next_token)
            draft_probs.append(draft_prob)
            current_ids = torch.cat([current_ids, next_token], dim=1)

    # 步骤 2:大模型并行验证(一次 forward pass 处理 k 个 token)
    all_tokens = torch.cat([input_ids] + draft_tokens, dim=1)
    with torch.no_grad():
        target_logits = target_model(all_tokens).logits
        # 获取大模型对 k 个位置的概率分布
        target_probs = torch.softmax(
            target_logits[:, len(input_ids[0])-1:-1, :], dim=-1
        )

    # 步骤 3:接受/拒绝(简化版,实际需要处理分布差异)
    accepted_tokens = []
    for i, (draft_tok, draft_p, target_p) in enumerate(
        zip(draft_tokens, draft_probs, target_probs)
    ):
        # 接受率:min(1, target_prob / draft_prob)
        accept_ratio = (target_p.gather(-1, draft_tok) /
                        draft_p.gather(-1, draft_tok)).clamp(max=1.0)
        if torch.rand(1) < accept_ratio:
            accepted_tokens.append(draft_tok)
        else:
            # 拒绝,从修正后的分布重新采样
            correction = (target_p - draft_p).clamp(min=0)
            correction = correction / correction.sum()
            accepted_tokens.append(torch.multinomial(correction, 1))
            break

    return torch.cat(accepted_tokens, dim=1)

加速效果: 对于长文本生成任务(输出 500+ Token),投机采样可带来 2-3 倍的速度提升,且输出质量与大模型完全一致(数学上等价)。


1.6 MoE(混合专家)架构:DeepSeek 为什么便宜

Mixture of Experts(MoE) 是另一类推理效率优化,通过稀疏激活(Sparse Activation)实现:

code
总参数量大,但每次前向传播只激活一小部分参数

DeepSeek-V3 有 671B 参数,但每次前向传播只激活约 37B 参数(通过路由网络选择激活哪些"专家"),推理成本接近 37B 密集模型,但知识容量接近 671B 模型。

架构 代表模型 总参数 激活参数 推理成本
密集(Dense) Llama 3 70B 70B 70B
MoE DeepSeek-V3 671B 37B 中等
MoE Mixtral 8x7B 46.7B 12.9B
MoE GPT-4(推测) ~1.8T ~220B

1.7 小结:各优化技术的定位

推理优化技术

减少计算量
KV Cache
Flash Attention

减少内存占用
量化 INT4/INT8
GGUF/AWQ/GPTQ

提高并行度
投机采样
Continuous Batching

稀疏激活
MoE 架构

速度提升 2-10x
(序列长度相关)

内存降低 2-8x
(精度轻微下降)

吞吐量提升 2-3x
(长输出任务)

同算力支持更大知识容量

这些技术并不互斥,现代 LLM 推理服务(如 vLLM——一个高性能的开源LLM推理引擎)同时使用 KV Cache + Flash Attention + Continuous Batching(连续批处理,让多个请求共享同一批次动态调度,提升GPU利用率)+ 量化,叠加效果可达单纯自回归生成速度的 20 倍以上。知道这几类优化的存在,选择部署方案时就不会一头雾水。

本页目录