LLM推理优化实战:从KV Cache到Speculative Decoding,把延迟打下来

发布时间:2026/7/25 0:39:10
LLM推理优化实战:从KV Cache到Speculative Decoding,把延迟打下来 做过LLM应用的都知道模型效果再好推理速度跟不上也是白搭。尤其是在对话场景里用户等3秒以上就开始不耐烦了——这还不是我瞎说的Google的研究数据表明页面加载时间从1秒增加到3秒跳出率增加32%。推理优化这个话题很大从模型量化、KV Cache管理、到注意力机制优化、再到调度策略每个方向都有不少值得深挖的东西。这篇文章我想从实际工程的角度把目前主流的推理加速技术串起来讲一遍。为什么LLM推理这么慢要理解怎么加速先得搞清楚为什么慢。LLM推理慢的核心原因就两个内存带宽瓶颈和自回归解码。LLM推理瓶颈内存带宽瓶颈自回归解码每次推理需要加载全部模型权重到显存权重 IO 时间远超计算时间Memory-Bound而非 Compute-Bound每次只生成1个token生成100个token需要100次前向传播无法并行化token间有依赖内存带宽瓶颈以LLaMA-2 70B为例FP16精度下模型权重约140GB。即便是H1003TB/s带宽光加载权重就需要约47ms。而实际计算只需要约10ms。换句话说80%的时间都在等数据GPU的计算单元大部分时间都在闲着。自回归解码LLM生成文本是逐token的每个token的生成都依赖前面所有token。这意味着生成长度为N的回复需要N次串行的前向传播。每次前向传播都要重新计算整个序列的注意力——这就是著名的二次复杂度问题。KV Cache推理加速的第一板斧KV Cache是LLM推理中最基础也最重要的优化。原理不复杂但很巧妙在自回归解码过程中每生成一个新token之前所有token的Key和Value矩阵其实不需要重新计算。importtorchimporttorch.nnasnnimporttorch.nn.functionalasFclassKVCacheAttention(nn.Module):带 KV Cache 的自注意力实现def__init__(self,d_model:int,n_heads:int):super().__init__()self.d_modeld_model self.n_headsn_heads self.d_headd_model//n_heads self.q_projnn.Linear(d_model,d_model,biasFalse)self.k_projnn.Linear(d_model,d_model,biasFalse)self.v_projnn.Linear(d_model,d_model,biasFalse)self.o_projnn.Linear(d_model,d_model,biasFalse)# KV Cacheself.k_cache:torch.Tensor|NoneNoneself.v_cache:torch.Tensor|NoneNonedefforward(self,x:torch.Tensor,use_cache:boolTrue):batch,seq_len,_x.shape qself.q_proj(x).view(batch,seq_len,self.n_heads,self.d_head)kself.k_proj(x).view(batch,seq_len,self.n_heads,self.d_head)vself.v_proj(x).view(batch,seq_len,self.n_heads,self.d_head)ifuse_cacheandself.k_cacheisnotNone:# 拼接历史 KV 和新 KV避免重复计算ktorch.cat([self.k_cache,k],dim1)vtorch.cat([self.v_cache,v],dim1)# 更新 Cacheifuse_cache:self.k_cachek self.v_cachev# 注意力计算简化版qq.transpose(1,2)# [batch, heads, seq, d_head]kk.transpose(1,2)vv.transpose(1,2)scaleself.d_head**-0.5attntorch.matmul(q,k.transpose(-2,-1))*scale attnF.softmax(attn,dim-1)outtorch.matmul(attn,v)outout.transpose(1,2).contiguous().view(batch,seq_len,self.d_model)returnself.o_proj(out)defreset_cache(self):开始新对话时清空 Cacheself.k_cacheNoneself.v_cacheNoneKV Cache带来的加速效果是显著的。对于长度为N的序列不用Cache时注意力计算量是O(N²)用了Cache后每次只计算新token与历史token的注意力复杂度降为O(N)。在实际场景中KV Cache可以把推理速度提升2-5倍。但KV Cache也有代价——显存占用。对于LLaMA-2 70B80层、64个注意力头、128维1K token的KV Cache占用约2.5GB显存。如果上下文长度是32K光KV Cache就要80GB。这就是为什么很多推理框架都在做KV Cache的量化压缩。PagedAttentionvLLM的核心创新KV Cache的内存管理是推理引擎的关键问题。传统做法是预分配一块连续显存但这样浪费严重——不同的请求序列长度不同预分配多了浪费少了不够用。vLLM团队从操作系统的虚拟内存管理中获得了灵感提出了PagedAttention。核心思想是把KV Cache分成固定大小的页Page每个页可以独立分配和释放就像操作系统的内存分页一样。PagedAttention请求1: 2K token分配 8个 Page每Page 256 token请求2: 512 token分配 2个 Page请求3: 1K token分配 4个 Page统一 Page 池无碎片传统方式请求1: 2K token预分配 2K 显存块请求2: 512 token预分配 512 显存块碎片化严重PagedAttention带来的好处是立竿见影的显存利用率从20-40%提升到接近100%支持更大的batch size吞吐量提升2-4倍不同请求可以共享相同的Page比如相同的system promptSpeculative Decoding用草稿模型加速前面说到自回归解码是串行的这是推理速度的根本瓶颈。但有没有办法打破这个限制Speculative Decoding提供了一个巧妙的思路。核心想法是用小模型快速生成多个候选token然后用大模型并行验证这些token是否正确。importtorchfromtransformersimportAutoModelForCausalLM,AutoTokenizerclassSpeculativeDecoder:投机解码实现def__init__(self,target_model:str,draft_model:str):self.targetAutoModelForCausalLM.from_pretrained(target_model,torch_dtypetorch.float16,device_mapauto)self.draftAutoModelForCausalLM.from_pretrained(draft_model,torch_dtypetorch.float16,device_mapauto)self.tokenizerAutoTokenizer.from_pretrained(target_model)torch.no_grad()defgenerate(self,prompt:str,max_new_tokens:int256,gamma:int5)-str: gamma: 每次投机解码生成的候选 token 数量 越大则并行度越高但接受率可能下降 input_idsself.tokenizer(prompt,return_tensorspt).input_ids input_idsinput_ids.to(self.target.device)generated[]whilelen(generated)max_new_tokens:# 步骤1: 用小模型快速生成 gamma 个候选 tokendraft_outputself.draft.generate(torch.cat([input_ids,torch.tensor([generated])],dim-1)ifgeneratedelseinput_ids,max_new_tokensgamma,do_sampleFalse,pad_token_idself.tokenizer.eos_token_id)draft_tokensdraft_output[0,-gamma:]# 步骤2: 用大模型并行验证所有候选 tokentarget_outputself.target(torch.cat([input_ids,draft_tokens],dim-1))target_logitstarget_output.logits[0,-gamma-1:-1]# 步骤3: 接受匹配的 token拒绝不匹配的accepted0foriinrange(gamma):target_tokentarget_logits[i].argmax().item()iftarget_tokendraft_tokens[i].item():accepted1else:# 拒绝当前位置但从大模型采样作为替代generated.append(target_token)accepted1breakgenerated.append(target_token)ifaccepted0:# 全被拒绝回退到大模型单步生成next_tokentarget_logits[0].argmax().item()generated.append(next_token)returnself.tokenizer.decode(generated,skip_special_tokensTrue)Speculative Decoding在实践中通常能带来1.5-2.5倍的加速而且不损失任何精度——因为最终验证还是由大模型完成的。Google的Gemini、OpenAI的GPT-4 Turbo都用了类似的技术。Flash Attention从算法层面优化注意力注意力机制的计算量和显存占用是O(N²)的这在大上下文场景中是不可接受的。Flash Attention通过重排计算顺序把注意力计算从HBM搬到SRAM中完成避免了中间结果的显存读写。通俗地说Flash Attention的核心技巧是分块计算Tiling——把Q、K、V矩阵切成小块每次只加载一小块到SRAM中计算算完立即写回HBM不保存中间结果。这样虽然计算量没变但显存读写量从O(N²)降到了O(N)。渲染错误:Mermaid 渲染失败: Parse error on line 6: ...计算 ×V] E -- F[O(N²) 显存读写] ----------------------^ Expecting SQE, DOUBLECIRCLEEND, PE, -), STADIUMEND, SUBROUTINEEND, PIPE, CYLINDEREND, DIAMOND_STOP, TAGEND, TRAPEND, INVTRAPEND, UNICODE_TEXT, TEXT, TAGSTART, got PSFlash Attention 2.0进一步优化了并行策略把序列长度维度也并行化在A100上达到了理论峰值算力的73%。Flash Attention 3则针对H100的新架构Tensor Memory Accelerator做了适配在FP8精度下进一步提速。实际部署的选型建议说了这么多技术最后聊聊实际场景怎么选场景推荐方案预期加速单用户对话KV Cache Flash Attention2-3x高并发API服务vLLM (PagedAttention)3-5x 吞吐量长文本生成Speculative Decoding1.5-2.5x极致延迟优化TensorRT-LLM INT4量化4-8x综合方案vLLM Flash Attention AWQ量化5-10x说实话对于大多数开发团队来说不需要从零实现这些技术。直接上vLLM或者TensorRT-LLM开箱即用比自己折腾效率高得多。但理解背后的原理还是有用的——至少出了问题你知道从哪排查。写在最后LLM推理优化的核心思路可以用一句话概括把串行变并行把显存IO降到最低。KV Cache减少了重复计算PagedAttention优化了显存管理Speculative Decoding打破了自回归的串行限制Flash Attention从算法层面减少了IO。这些技术加在一起让秒级响应从不可能变成了可能。用过的都懂从5秒降到1秒用户体验的差别不是线性的——是质变。标签LLM推理优化、KV Cache、vLLM、Speculative Decoding、Flash Attention