LLM 推理优化实战:把响应时间从 2s 降到 200ms 的全过程
用 Llama 3 8B 做的聊天机器人项目,通过量化、Flash Attention、Speculative Decoding 等优化,最终延迟降低 91%。
# LLM 推理优化实战:把响应时间从 2s 降到 200ms 的全过程
去年我们做的项目里有个聊天机器人功能,后端用 Llama 3 8B 模型,部署在自家 GPU 服务器上。一开始没太在意性能,想着"反正用户也不会等太久"。结果上线测试的时候,用户反馈明显:等待时间太长,体验很差。
我们最后把平均响应时间从 2.1 秒降到了 180 毫秒。这个过程里踩了不少坑,现在整理出来。
第一步:先搞清楚慢在哪里
别急着优化,先 profiling。我们用 OpenTelemetry 给推理流程打了点,发现:
主要问题在 Token 生成阶段。我们的模型是 Llama 3 8B,在 A10G 上跑,batch size 为 1。
第二步:换了个量化方案
一开始用的是 FP16 精度,模型权重占 16GB 显存。考虑过 INT8 量化,但担心质量损失太大。后来试了 GPTQ-INT4,把模型权重降到 4GB,显存压力小了很多,精度损失在可接受范围内。
代码改动很简单:
from auto_gptq import GPTQForCausalLM
from transformers import AutoTokenizer
model = GPTQForCausalLM.from_quantized(
"llama-3-8b-gptq-4bit",
use_triton=True,
device="cuda:0"
)
tokenizer = AutoTokenizer.from_pretrained("llama-3-8b-gptq-4bit")
这一步就把生成速度提升了约 30%。
第三步:KV Cache 优化
这是最关键的优化。默认情况下,每个新 token 生成时,模型都要重新计算之前所有 token 的 KV cache。对于长对话,这很浪费。
我们用了 Flash Attention 2,它能高效计算 attention 而不需要显式存储完整的 KV cache。
model = AutoModelForCausalLM.from_pretrained(
model_path,
torch_dtype=torch.float16,
device_map="auto",
attn_implementation="flash_attention_2"
)
这一步提升约 40%,但要注意:Flash Attention 2 需要较新的 CUDA 版本和 GPU 架构支持(Ampere 及以上)。
第四步:Speculative Decoding
这个技术有点前沿,但效果显著。基本原理是:用一个小的草稿模型(draft model)先预测几个 token,然后用大模型验证。如果验证通过,这些 token 就一次性接受;否则回退到正常生成。
我们用了 Llama 3 8B 作为主模型,Llama 3 1B 作为草稿模型:
from vllm import LLM, SamplingParams
main_model = LLM(model="llama-3-8b", speculative_model="llama-3-1b", num_speculative_tokens=5)
sampling_params = SamplingParams(temperature=0.7, max_tokens=512)
这一步提升约 50%,但部署复杂度增加不少。
第五步:PagedAttention
vLLM 的 PagedAttention 机制把 KV cache 分成多个 block,按需分配和回收。这对 batch 推理特别有效,因为不同请求的序列长度不同,传统的静态分配会浪费很多显存。
我们切换到 vLLM 后,同样的显存可以支持更大的 batch size:
from vllm import LLM, SamplingParams
llm = LLM(
model="llama-3-8b",
tensor_parallel_size=1,
gpu_memory_utilization=0.9
)
这一步把吞吐提升了约 3 倍,延迟进一步降低。
第六步:Prompt 优化
这不是算法优化,但效果最直接。我们发现有些用户的 prompt 特别长,包含了大量不相关的背景信息。我们在 API 层加了一个预处理,把 prompt 截断到合理长度,并提取关键信息。
def optimize_prompt(raw_prompt: str) -> str:
# 简单的规则截断
if len(raw_prompt) > 2000:
raw_prompt = raw_prompt[:2000] + "..."
# 移除重复内容
lines = set(raw_prompt.split('\n'))
return '\n'.join(lines)
这一步提升了约 10%。
最终效果
| 优化项 | 延迟降低 |
|--------|----------|
| 量化 | -30% |
| Flash Attention 2 | -40% |
| Speculative Decoding | -50% |
| PagedAttention | -60% |
| Prompt 优化 | -10% |
| **总计** | **-91%** |
从 2100ms 降到 180ms,用户体验明显改善。
踩过的坑
2. **Speculative Decoding 的兼容性**:不是所有模型都能用 spec decoding,需要草稿模型和主模型是同一家族。Llama 系列支持得最好,其他模型可能要自己适配。
3. **vLLM 的显存管理**:gpu_memory_utilization 设置太高会导致 OOM,太低又浪费显存。我们最后是 0.9,但需要根据具体 GPU 型号调整。
结语
推理优化是个系统工程,没有银弹。我们用的方法组合起来效果最好,单独用任何一个提升都很有限。如果你是做产品,建议从 Prompt 优化和量化开始,成本最低,效果最直接。