推理阶段报OOM,是做大模型部署时绕不开的坎。训练阶段的显存问题往往有成熟的混合精度方案可以套用,而推理阶段的OOM更隐蔽:模型加载成功不代表推理能跑通,上下文一长、并发一高,显存占用就会悄无声息地翻倍增长。要真正解决问题,第一步不是急着换更大的显卡,而是搞清楚显存到底花在了哪里。

一、推理显存到底由哪几部分构成
很多人对推理显存的直觉认知是“模型多大就占多少显存”,这其实只对了一小部分。一个完整的推理服务,显存占用大致分为四块:模型权重、KV Cache、激活值和框架运行时开销。模型权重是最容易估算的部分,一个7B参数的模型在FP16精度下,权重占用约为 7 × 2 = 14GB。如果是INT8量化则减半,INT4量化则约为四分之一。
KV Cache是最容易被低估的部分。自回归解码时,每生成一个token,都需要把每一层注意力的Key和Value缓存下来,避免重复计算。它的规模随着序列长度和并发数线性增长。以Llama-7B为例,每token的KV Cache大小可以用下面的公式估算:
# KV Cache 每token显存估算(字节) # 2 表示 K 和 V 两个张量 kv_cache_per_token = 2 * num_layers * num_kv_heads * head_dim * dtype_size # 以 Llama-2-7B 为例 # num_layers=32, num_kv_heads=32, head_dim=128, fp16 每元素2字节 kv = 2 * 32 * 32 * 128 * 2 # = 524288 字节,约 0.5MB/token # 4096 长度的上下文,单条请求的 KV Cache 就是约 2GB
也就是说,一个7B模型跑8K上下文、并发10路请求,光是KV Cache就可能吃掉30GB以上显存,远超权重本身。这就是为什么模型加载得好好的,一旦并发上来就OOM。此外还有激活值(前向传播中间张量)和CUDA context、内存池碎片等运行时开销,通常预留1到2GB比较稳妥。
二、显存排查:先定位再动手
优化之前必须先确认瓶颈在哪。盲目上量化可能发现瓶颈根本不在权重。排查的第一步是用nvidia-smi观察基础占用,但更精确的方式是用PyTorch的内存分析工具。torch.cuda.max_memory_allocated能拿到峰值分配值,结合torch.cuda.memory_summary可以输出详细的分配分布,帮你区分是权重占了大头还是运行时的临时张量在作怪。
import torch
# 推理前重置统计
torch.cuda.reset_peak_memory_stats()
# ... 执行一次完整的推理 ...
print(f"当前分配: {torch.cuda.memory_allocated() / 1024**3:.2f} GB")
print(f"峰值分配: {torch.cuda.max_memory_allocated() / 1024**3:.2f} GB")
print(f"预留总量: {torch.cuda.memory_reserved() / 1024**3:.2f} GB")
有一个常见误区值得提醒:显存碎片。PyTorch的缓存分配器会保留已释放的块,导致reserved远大于allocated,新的大块分配请求(比如一个很长的上下文)申请不到连续显存而报OOM,哪怕nvidia-smi显示还有“空闲”。这种情况下调用torch.cuda.empty_cache()或设置环境变量PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True能明显缓解。另一个排查点是长序列:如果OOM只出现在超长输入时,基本可以锁定是KV Cache或激活值的问题,而不是权重。
三、优化手段一:量化压缩模型权重
如果排查确认权重占了显存大头,量化是最直接的手段。推理量化不像训练量化那样复杂,社区方案已经非常成熟。GPTQ和AWQ是两种主流的训后量化方法,都能把模型压到4bit而精度损失可控。以HuggingFace的transformers配合bitsandbytes为例,一行配置就能加载INT4模型:
from transformers import AutoModelForCausalLM, AutoTokenizer
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-2-7b-hf",
device_map="auto",
load_in_4bit=True, # 4bit 量化加载
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_quant_type="nf4", # NF4 量化类型,精度更稳
)
量化后7B模型的权重占用从14GB降到约4GB,一张24GB的卡就能轻松跑7B模型并留出充足的KV Cache空间。需要注意的权衡是:bitsandbytes的动态量化推理速度略慢于原始精度,因为它需要在计算时反量化;而GPTQ、AWQ这类离线量化方案预先完成了权重变换,推理速度反而可能更快。如果追求极致吞吐,建议优先尝试AWQ;如果只是想省显存快速验证,bitsandbytes最省事。
四、优化手段二:KV Cache与注意力机制优化
当瓶颈在KV Cache时,思路有两个方向:减小Cache本身,或者改进注意力计算方式。第一个方向里最有效的是GQA(分组查询注意力)和MQA(多查询注意力)。Llama-2-70B、Qwen等新模型原生支持GQA,把KV头的数量从32降到8甚至更少,KV Cache直接缩小到原来的四分之一。选型时优先选择支持GQA的模型,是零成本的显存优化。
另一个方向是注意力实现层面的优化。标准注意力在长序列下的显存占用是序列长度的平方级,而FlashAttention通过分块计算避免了生成完整的注意力矩阵,显存从平方级降到线性级,同时速度还更快。大多数推理框架只需一个参数就能开启。如果用的是vLLM这类专业推理引擎,它默认就集成了FlashAttention,并且通过PagedAttention机制把KV Cache切分成固定大小的页来管理,彻底解决了显存碎片和浪费问题,官方数据是在相同显存下把吞吐提升了数倍。
# 使用 vLLM 启动服务,自动启用 PagedAttention 和 FlashAttention
python -m vllm.entrypoints.openai.api_server \
--model /path/to/llama-2-7b \
--gpu-memory-utilization 0.9 \
--max-model-len 8192 \
--enable-prefix-caching
上面命令里有两个关键参数值得说明。--gpu-memory-utilization 0.9告诉vLLM最多使用90%的显存,它会据此预分配KV Cache的池子,避免和系统其他进程冲突。--enable-prefix-caching开启前缀缓存,对于多轮对话这种系统提示词重复出现的场景,相同前缀的KV Cache可以直接复用,进一步节省计算和显存。
五、优化手段三:批处理与工程层面的策略
显存优化不只是压模型,调度策略同样重要。传统的静态批处理会按最长序列填充,短请求浪费大量显存和算力;而连续批处理(continuous batching)让请求动态进出批次,显存利用率大幅提升。vLLM、TGI这些引擎都内置了这一能力,自建服务的话可以考虑直接采用而不是重复造轮子。
此外还有几个工程层面的实用招数。第一,限制最大上下文长度,不要无脑设置成模型支持的最大值,按业务实际需要来,KV Cache的占用和它直接成正比。第二,多卡场景下优先考虑张量并行而不是数据并行,推理阶段数据并行意味着每张卡都装一份完整模型。第三,如果允许,把非热点路径的请求分流到CPU推理或更小的模型上,做分级服务。最后别忘了监控,生产环境建议用Prometheus采集GPU显存指标,设置阈值告警,OOM问题最好在压测阶段就暴露出来,而不是等线上流量高峰时才崩。
总结一下排查路径:先用内存统计工具确认显存构成,权重超标就上量化,KV Cache超标就选GQA模型加PagedAttention,碎片问题就调分配器参数,吞吐不够就优化批处理策略。大部分推理OOM问题,都能在这几步里找到答案。