跳转至

推理优化

💾 KV Cache 原理与优化手段

🔹 原理:为什么需要 KV Cache?

自回归生成时,每生成一个新 token,都要和前面所有 token 做注意力计算。如果没有缓存,每一步都要把整条序列重新喂给模型,计算所有 token 的 Key 和 Value,计算量随序列长度平方级增长。

KV Cache 的思路很直接:把每一层的 Key 和 Value 存下来,新 token 只算自己的 K、V,然后拼到缓存后面。

数学上,注意力公式是:

image.png

有了缓存后:

  • 第 tt 步只需从输入 xtxt 算出 qt,kt,vtqt,kt,vt

  • 更新 K_cache = concat(K_cache, k_t), V_cache = concat(V_cache, v_t)

  • 用 QtQt 和完整的 K_cache, V_cache 计算注意力。

显存与计算的变化:

  • 无缓存:每步计算量 O(t2d)O(t2d),显存 O(td)O(td)(但每次重新计算全部)。

  • 有缓存:每步计算量 O(td)O(td),显存占用累计 O(2×层数×头数×头维度×最大长度)O(2×层数×头数×头维度×最大长度),成为主要瓶颈。

image.png

🔹 关键优化手段

① MHA → MQA / GQA:减少缓存体积

  • MHA(多头注意力):每个头都有独立的 K 和 V,缓存很大。

  • MQA(多查询注意力):所有头共享同一份 K 和 V,缓存大幅缩减,但模型容量略降。

  • GQA(分组查询注意力):将头分成若干组,组内共享 K、V,在速度和效果间取平衡。LLaMA 2 70B 就用了 GQA。

② PagedAttention(vLLM 的核心)

把连续的 KV Cache 切成固定大小的 block,像操作系统的内存分页一样管理。好处:

  • 动态分配:只分配实际需要的块,消除预分配浪费。

  • 前缀共享:多个请求若共用同一个 system prompt,只需存一份 KV 块,显存节省巨大。

  • 无碎片:请求完成后立即回收 block 给其他请求复用。

③ KV Cache 量化(KV8/ KV4) 在推理时对 K、V 做 INT8 甚至 INT4 量化,进一步压缩缓存占用。例如 HuggingFace Transformers 可搭配 bitsandbytes 对 KV Cache 做量化。

④ Flash Decoding / Split-K 技巧

在 Decode 阶段,利用分块并行进一步榨取 GPU 带宽,让长序列生成更快。

代码示例:手动实现 KV Cache(微型版)

import torch
import torch.nn.functional as F

def attention_with_cache(q, k, v, k_cache=None, v_cache=None):
    # q, k, v 形状: (batch, heads, 1, head_dim)
    if k_cache is not None:
        k = torch.cat([k_cache, k], dim=2)  # 沿序列维度拼接
        v = torch.cat([v_cache, v], dim=2)
    attn_scores = torch.matmul(q, k.transpose(-2, -1)) / (q.shape[-1] ** 0.5)
    attn_weights = F.softmax(attn_scores, dim=-1)
    output = torch.matmul(attn_weights, v)
    return output, k, v  # 返回更新后的缓存

# 自回归循环
k_caches = [None] * num_layers
v_caches = [None] * num_layers
for step in range(max_new_tokens):
    # ... 对每一层
    q, k, v = proj_q(x), proj_k(x), proj_v(x)
    out, k_caches[i], v_caches[i] = attention_with_cache(q, k, v, k_caches[i], v_caches[i])

一句话收束:

KV Cache 是自回归推理的基础,优化它的体积和访问效率,直接决定了我们能多快、多省地服务用户。


⚙️量化技术 INT8 / INT4 / GPTQ / AWQ

量化就是把高精度浮点权重(FP16/BF16)映射到低位宽整数,用更小的存储和更快的整数运算换推理速度与显存。

🔹 INT8 量化(典型的 PTQ 方案)

  • 原理:对权重张量按通道(或整个张量)计算缩放因子 ss 和零点 zz,使浮点值 xfloat≈s⋅(xint−z)xfloats⋅(xintz)。矩阵乘法时,先做 INT8 乘累加,再反量化回 FP32 做偏置和激活。

  • 特点:几乎无损,显存减半,计算提速(需硬件支持 INT8 Tensor Core)。

  • 动态量化:激活的量化参数在线计算;静态量化:提前校准。

示例:手工对称量化

import torch

def quantize_int8(tensor):
    s_max = tensor.abs().max() / 127.0
    q = torch.round(tensor / s_max).clamp(-128, 127).to(torch.int8)
    return q, s_max

def dequantize_int8(q, s):
    return q.float() * s

# 权重量化
w_fp16 = torch.randn(512, 512, dtype=torch.float16)
w_int8, scale = quantize_int8(w_fp16.float())
# 推理时反量化
w_deq = dequantize_int8(w_int8, scale)

🔹 INT4 量化(更激进)

  • 权重压缩到 4 位,显存仅 FP16 的 1/4。

  • 单纯四舍五入会导致较大精度损失,需要更聪明的量化方法。

🔹 GPTQ

  • 思想:基于最优脑量化的思想,逐列对权重进行量化,每量化一列,就用Hessian 信息更新其余未量化列的权重,补偿误差。

  • 特点:离线一次性处理,不需要反向传播,速度较快;支持分组量化(group size),在组内再校准,减少异常值影响。

  • 效果:INT4 GPTQ 通常困惑度上升 < 1,接近 FP16 表现。

🔹 AWQ (Activation-aware Weight Quantization)

  • 发现:并非所有权重对输出同等重要——激活值较大通道的权重,其微小误差会被放大。

  • 做法:对这些“显著通道”先乘以大于 1 的系数,等效于减小量化步长,从而保护重要权重;其他通道步长略增。

  • 优点:不需要标定数据上的反向传播,轻量快速,也能达到良好精度。

量化方法对比表:

查看内嵌表格

代码示例:用 AutoGPTQ 加载量化模型

from auto_gptq import AutoGPTQForCausalLM
model = AutoGPTQForCausalLM.from_quantized(
    "TheBloke/Llama-2-7B-GPTQ",
    use_triton=False,
    use_safetensors=True
)
# 直接推理
output = model.generate(**tokenizer("你好", return_tensors="pt"))

一个关键点: GPTQ 和 AWQ 都允许分组量化,平衡精度与压缩率;实际部署时,通常使用 4-bit 权重 + FP16 激活,以兼顾效率和精度。

收束:

量化是模型部署的“压缩术”,从粗粒度的 INT8 无损压缩到 GPTQ/AWQ 的极智压缩,每一步都是在“计算精度”和“运行效率”之间找到最优折衷。


🎲解码策略:Greedy / Beam Search / Top-K / Top-P

解码策略控制模型如何在每个时间步从词表中选择下一个 token,直接决定输出的确定性、多样性和质量。

🔹 Greedy Decoding(贪心)

  • 做法:每步直接选概率最高的 token。

  • 优点:最快,确定性强。

  • 缺点:容易陷入重复循环,缺乏多样性;一旦选错,无法回头。

def greedy_decode(logits):
    return torch.argmax(logits, dim=-1)
  • 做法:维护 k 条候选序列(beam width),每步扩展所有可能的下一个 token,保留总概率最高的 k 条。

  • 优点:比贪心更能找到全局高概率序列,适合翻译、摘要等确定性任务。

  • 缺点:计算量 k 倍于贪心;在开放式生成中可能产生重复和过于“安全”的句子。

def beam_search_step(beam_seqs, beam_scores, logits, k=4):
    # beam_seqs: 当前 beam 序列列表
    # 对每条 beam 取 top-k 候选,合并后保留总分数最高的 k 条
    candidates = []
    for i, seq in enumerate(beam_seqs):
        topk_scores, topk_ids = torch.topk(logits[i], k)
        for score, token in zip(topk_scores, topk_ids):
            candidates.append((seq + [token], beam_scores[i] + score))
    candidates.sort(key=lambda x: x[1], reverse=True)
    return candidates[:k]

🔹 Top-K Sampling

  • 做法:只从概率最高的 K 个 token 中随机采样,其余概率置零。

  • 优点:避免选到极低概率的无关词,增加可控多样性。

  • 缺点:K 是固定值,不适应分布变化——如果概率集中在 3 个词,Top-K=10 仍会引入噪音;如果分布平坦,K=10 可能截断合理选择。

🔹 Top-P (Nucleus Sampling)

  • 做法:对概率从高到低排序,选取累积概率达到 p 的最小 token 集合,从中采样。

  • 优点:动态适应分布——分布尖锐时集合小,分布平坦时集合大,灵活得多。

  • 缺点:p 值需要根据任务调优。

采样组合典型配置:

def sample_token(logits, temperature=0.7, top_p=0.9, top_k=0):
    # 温度缩放
    logits = logits / temperature
    # Top-K 过滤
    if top_k > 0:
        indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None]
        logits[indices_to_remove] = -float('Inf')
    # Top-P 过滤
    if 0 < top_p < 1:
        sorted_logits, sorted_indices = torch.sort(logits, descending=True)
        cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
        # 移除累积概率超过 p 的 token
        sorted_indices_to_remove = cumulative_probs > top_p
        sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
        sorted_indices_to_remove[..., 0] = 0
        indices_to_remove = sorted_indices[sorted_indices_to_remove]
        logits[:, indices_to_remove] = -float('Inf')
    # 采样
    probs = F.softmax(logits, dim=-1)
    return torch.multinomial(probs, 1)

场景与策略速查:

查看内嵌表格

收束:

选择解码策略,本质上是控制“探索与利用”的平衡。温度调“大胆程度”,Top-K/Top-P 划“安全边界”,而 Beam Search 在需要全局最优时是利器。用好它们,同样的模型能产出风格迥异、恰如其分的回答。

推理加速 FlashAttention/PagedAttention/Speculative Decoding

我们直接进入主题。这三种技术都是大模型推理优化的核心利器,分别从计算效率、显存管理、串行依赖三个角度动刀。下面我逐一拆解,并给出可现场画出的图示和关键代码骨架。


⚡ 1. Flash Attention:把 O(N²) 的显存读写打下来

问题在哪?

标准 self-attention 会生成一个 N×N 的注意力分数矩阵 S,然后对它做 softmax,再乘以 V。这个 N×N 矩阵需要反复在 GPU 的高带宽显存(HBM) 和片上缓存(SRAM) 之间搬运,而 HBM 的带宽远跟不上计算单元的速度。时间都花在 IO 上了,计算单元大量空转。

标准做法:
Q,K (HBM) → 算出 S=QK^T (N×N) 写出到 HBM
→ 从 HBM 读 S 做 softmax 得到 P 写出到 HBM
→ 从 HBM 读 P 乘 V 得到 O 写出到 HBM

核心思想

把整个计算切成小块,一次加载一个块到 SRAM,在 SRAM 内把该块需要的所有计算全部做完,中间那个 N×N 矩阵根本不存在 HBM 里。 利用 online softmax 技巧动态更新全局统计量,保证数学上完全等价。

Flash Attention:
循环加载 Q_block, K_block (HBM → SRAM)
  ├─ 在 SRAM 内算局部 S
  ├─ 用 online softmax 更新累积的 max 和分母
  ├─ 直接在 SRAM 内累加该块对 O 的贡献
  └─ 最终写出 O (SRAM → HBM)
全程不产生完整 S 矩阵

关键代码:online softmax 更新逻辑

def flash_attention_update(O_prev, m_prev, l_prev, K_block, V_block, Q_i):
    # 计算当前块注意力分数
    S_ij = Q_i @ K_block.T               # (1, block_size)
    m_curr = max(m_prev, S_ij.max())     # 更新最大值(数值稳定)

    # 修正之前的累计
    correction = torch.exp(m_prev - m_curr)
    l_curr = correction * l_prev + torch.exp(S_ij - m_curr).sum()  # 新分母

    # 计算当前块的输出贡献并累积
    P_ij = torch.exp(S_ij - m_curr)      # 当前块未归一化权重
    O_curr = correction * O_prev + P_ij @ V_block
    return O_curr, m_curr, l_curr

为什么快?

  • SRAM 带宽是 HBM 的 10 倍以上,把 N×N 矩阵的读写省掉,直接带来数倍加速。

  • 由于不存储完整注意力矩阵,显存从 O(N²) 降到 O(N),允许训练更长序列或更大 batch。


📄 2. PagedAttention:把 KV Cache 当成“虚拟内存”管理

问题在哪?

自回归生成时,需要为每个请求缓存完整的 KV Cache。传统做法为每个序列预分配一块最大长度的连续显存,造成:

  • 内部碎片:短序列浪费了分配但未使用的尾部。

  • 无法共享:多个请求的系统提示词即使完全相同,也要各存一份。

  • 分配不灵活:释放的显存无法高效复用给新请求,吞吐量极低。

核心思想

将 KV Cache 切割成固定大小的 block(如 16 个 token 一块),像操作系统的内存分页一样管理。 每个序列的 KV Cache 不再要求物理连续,而是通过 block table 记录逻辑位置到物理块的映射。

image.png

带来的好处

  • 零内部碎片:需要多少 token 就分配多少块,最后一块才可能有少许未用空间。

  • 前缀共享:所有请求共用同一个系统提示词的 KV blocks,显存节省 50% 以上。

  • 动态调度:请求完成或抢占后,释放的 block 立即进入空闲池,用于新请求,显存利用率可达 80% 以上。

关键代码:Block Table 管理

block_size = 16
free_blocks = list(range(total_blocks))  # 空闲物理块索引
block_tables = {}  # seq_id -> [phy_block_0, phy_block_1, ...]

def allocate(seq_id, num_tokens):
    needed = (num_tokens + block_size - 1) // block_size
    new_blocks = [free_blocks.pop() for _ in range(needed)]
    block_tables.setdefault(seq_id, []).extend(new_blocks)
    # 将 token 对应的 K,V 写入这些物理块

def free_sequence(seq_id):
    free_blocks.extend(block_tables.pop(seq_id))

注意力计算时,通过 block table 将逻辑位置翻译为物理地址,用定制 CUDA kernel 一次性完成,避免拷贝和碎片。这就是 vLLM 吞吐量能数倍于传统 HF 推理的根因。


🏎️ 3. Speculative Decoding:打破自回归的串行枷锁

问题在哪?

标准自回归生成一次只能出一个 token,因为每个 token 都依赖前文。这导致 Decode 阶段是严格的串行过程,GPU 算力大量闲置,整个生成的速度被内存带宽锁死。

核心思想

用小模型“猜”一串,用大模型“并行验”,一次通过多个 token。 因为验证一个序列的多个 token 是可并行的。

image.png

为什么加速?

  • Draft 模型很小,自回归生成 K 个 token 的代价远低于大模型。

  • 大模型只需要做一次前向就能验证 K 个 token,而正常情况需要 K 次前向。

  • 加速比 ≈ 大模型单次耗时 / (大模型验证耗时 + 小模型草稿耗时) × 平均接受长度 αK

  • 实际中,用小模型(或大模型早期层)做 draft,接受率可达 0.8 以上,典型加速 2~3 倍,且输出质量无损(最终分布由大模型控制)。

关键代码:Draft & Verify 循环

def speculative_decode(prompt, draft_model, target_model, K=5):
    generated = []
    while len(generated) < max_length:
        # 1) Draft 阶段
        draft_input = prompt + generated
        draft_tokens = draft_model.generate(draft_input, max_new_tokens=K)

        # 2) Verify 阶段:一次前向拿到 logits
        verify_input = draft_input + draft_tokens
        logits = target_model(verify_input)  # (seq_len, vocab)

        # 3) 逐个验证
        accepted = []
        for i, d_token in enumerate(draft_tokens):
            prob = softmax(logits[len(draft_input) + i - 1])
            if d_token == torch.argmax(prob):
                accepted.append(d_token)
            else:
                accepted.append(torch.multinomial(prob, 1))
                break
        generated.extend(accepted)
        if len(accepted) < K:
            # 被拒绝,停止本次草稿
            pass
    return generated

进一步省显存:可使用 Self-Speculative,用模型自身前几层作为 draft 模型,无需额外加载小模型。


🧩 三者如何协同?

这三个技术常常组合应用,构建高性能推理引擎:

  • Flash Attention 解决 Prefill 阶段的计算密集和显存瓶颈,也让长序列解码更快。

  • PagedAttention 解决多请求并发时的 KV Cache 显存浪费和调度问题,最大化吞吐。

  • Speculative Decoding 解决单请求 Decode 阶段的串行限制,大幅压低延迟。

高性能推理引擎 = Flash Attention (计算加速)
               + PagedAttention (显存管理)
               + Speculative Decoding (串行破局)

理解它们,你就真正握住了大模型从“跑起来”到“跑得快、跑得多”的钥匙。