推理优化
💾 KV Cache 原理与优化手段¶
🔹 原理:为什么需要 KV Cache?¶
自回归生成时,每生成一个新 token,都要和前面所有 token 做注意力计算。如果没有缓存,每一步都要把整条序列重新喂给模型,计算所有 token 的 Key 和 Value,计算量随序列长度平方级增长。
KV Cache 的思路很直接:把每一层的 Key 和 Value 存下来,新 token 只算自己的 K、V,然后拼到缓存后面。
数学上,注意力公式是:

有了缓存后:
-
第 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×层数×头数×头维度×最大长度),成为主要瓶颈。

🔹 关键优化手段¶
① 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)xfloat≈s⋅(xint−z)。矩阵乘法时,先做 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。
-
优点:最快,确定性强。
-
缺点:容易陷入重复循环,缺乏多样性;一旦选错,无法回头。
🔹 Beam Search(束搜索)¶
-
做法:维护 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 记录逻辑位置到物理块的映射。

带来的好处¶
-
零内部碎片:需要多少 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 是可并行的。

为什么加速?¶
-
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 阶段的串行限制,大幅压低延迟。
理解它们,你就真正握住了大模型从“跑起来”到“跑得快、跑得多”的钥匙。