显存账本:32GB 到底放得下什么
RL 训练要同时装下策略模型、优化器状态、参考模型和一个推理引擎。这一节把这笔账算清,用计算器验证。
学完这节你能做到
- 手算出「模型参数量 + 精度 + 优化器」需要多少显存,误差在 10% 以内
- 解释为什么 RL 比 SFT 多吃两份显存,以及 GRPO 省掉的是哪一份
- 给定一个模型,判断该走全参、LoRA 还是 QLoRA,并说出 KV cache 该留多少
RL 为什么比微调贵这么多
同样一个 0.6B 模型,做 SFT 你可能觉得显存绰绰有余;换成 RL,很容易就 OOM。原因是 RL 的显存里同时住着好几个模型:
SFT GRPO
┌──────────────────┐ ┌──────────────────────────────┐
│ 策略模型 │ │ 策略模型(要更新) │
│ 权重 │ │ 权重 + 梯度 + 优化器状态 │
│ 梯度 │ ├──────────────────────────────┤
│ 优化器状态 │ │ 参考模型(冻结,算 KL 用) │
│ 激活 │ ├──────────────────────────────┤
└──────────────────┘ │ 推理引擎(采样用) │
│ 又一份权重 + KV cache │
└──────────────────────────────┘
数一下:SFT 里权重出现 1 次,GRPO 里出现 3 次(训练一份、参考一份、推理引擎一份)。再加上 KV cache——这一项在长回答场景下能超过所有模型权重的总和。
经典 PPO 里还有第四个模型:critic(价值网络),通常和策略模型同尺寸,而且它自己也要训,所以带着梯度和优化器状态。
GRPO 砍掉的就是它。这一刀省下的显存,正是单卡能玩 RL 的直接原因。L1 会讲它是怎么砍掉的。
五个吃显存的部件
一、权重
权重字节数 = 参数量 × 每参数字节数
bf16 / fp16 → 2 字节
fp8 → 1 字节
int4 (QLoRA) → 约 0.55 字节(含量化元数据)
Qwen3-0.6B 按 bf16:0.6e9 × 2 = 1.2 GB ≈ 1.12 GiB。
二、梯度
只有可训练参数才有梯度。 全参训练时梯度和权重同量级;LoRA 时只有适配器那一小撮有梯度。
这就是 LoRA 省显存的第一刀。
三、优化器状态 —— 通常是最大的一块
AdamW 每个可训练参数要存:
fp32 主权重 4 字节
一阶动量 m 4 字节
二阶动量 v 4 字节
─────────────────
合计 12 字节 / 参数
0.6B 全参:0.6e9 × 12 = 7.2 GB ≈ 6.7 GiB。比权重本身大 6 倍。
这也解释了为什么 1.7B 全参在 32GB 上没戏:光这一项就 20 GiB。
LoRA rank=16 时可训练参数大约是全量的 0.4%。梯度和优化器状态同时掉到千分之四 —— 7 GiB 变成 30 MB。
代价是表达能力受限,以及基座权重仍然要完整驻留(LoRA 省的是「训练态」,不是「权重态」)。
四、激活
前向过程中为了反向传播而留下的中间结果。它和 micro_batch × 序列长度 成正比,和参数量关系不大。
梯度检查点(gradient checkpointing)只保存每层的输入,反向时重算,能把这一项压掉一个数量级,代价是多约 30% 的前向计算。RL 场景里基本是默认开。
五、logits —— 新手最容易漏算的一项
这一项经常被忽略,但在 RL 里特别致命:
logits 字节数 ≈ micro_batch × 序列长度 × 词表大小 × 2 字节 × 2 份
Qwen3 的词表是 151936。按 micro_batch=2、序列长 1024 算:
2 × 1024 × 151936 × 2 × 2 = 1.24 GB ≈ 1.16 GiB
比激活还大。而 RL 里你要算 logprob,往往还要额外留一份中间量。
答案八成在这里。序列长度翻倍、micro batch 翻倍,logits 就翻倍,而它和模型大小没关系 —— 小模型配大词表配长序列,是最典型的翻车组合。
所以显存紧张时的第一刀是降 micro batch(用梯度累积补回等效 batch),而不是换更小的模型。
番外:KV cache
推理引擎那一侧的账。每个 token 的 KV cache:
每 token 字节 = 2(K和V) × 层数 × KV头数 × head_dim × 2字节
Qwen3-0.6B:2 × 28 × 8 × 128 × 2 = 114,688 字节 ≈ 112 KiB/token
看着不多,但它要乘上并发数 × 序列长度:
| 并发 | 长度 | KV cache |
|---|---|---|
| 16 | 1024 | 1.75 GiB |
| 16 | 4096 | 7.0 GiB |
| 64 | 4096 | 28 GiB ❌ |
并发和长度对显存是乘法关系。 想让模型「思考得更长」,代价是二次方式增长的显存压力 —— 这是 RL 训推理模型时最硬的一堵墙。
交互计算器
把上面六项串起来了。建议按顺序试这几组对比,比读文字有效得多:
- 看默认配置:Qwen3-0.6B 全参 + colocate。注意优化器状态那一条占了多大比例。
- 换 1.7B 全参:看它怎么爆的,以及是哪一条把它顶爆的。
- 1.7B 改 LoRA:看梯度和优化器两条同时塌下去。
- KV cache 的乘法:把 rollout 长度从 1024 拉到 8192,再把并发从 16 拉到 64,看警告栏说了什么。
- logits 陷阱:模型固定 0.6B,把训练序列长度拉到 8192、micro batch 拉到 8,看谁变成了最大的一块。
- 关掉参考模型:省了多少?(对应 KL 系数设 0 的场景)
- 策略模型权重1.12 GiBbf16,2 字节/参数
- 梯度1.12 GiB全参:与权重同量级
- 优化器状态6.71 GiBAdamW:fp32 主权重 + 一阶 + 二阶动量,12 字节/可训练参数
- 激活0.13 GiB开了梯度检查点,只存层边界
- logits 与 logprob1.16 GiB词表 151,936,长度 × 词表 × 2 份
- 参考模型1.12 GiBbf16 冻结,算 KL 用
- vLLM 权重1.12 GiB推理引擎自己持有一份权重
- KV cache1.75 GiB112 KiB/token × 1024 × 16 并发
- →还剩 16.6 GiB。优先加 rollout 并发(采样是墙钟时间的大头),而不是加 micro batch。
真实占用还要算上:CUDA context(几百 MB)、cuBLAS/cuDNN 的 workspace、PyTorch allocator 的碎片、vLLM 预分配的块。
留 10% 到 15% 的余量。 算出来 31 GiB 的配置,实际跑起来大概率 OOM。
OOM 了先动哪一刀
按「效果 ÷ 代价」排序,从上往下试:
| 顺序 | 动作 | 省多少 | 代价 |
|---|---|---|---|
| 1 | micro batch 降到 1 + 梯度累积 | 激活和 logits 成比例下降 | 慢一点,数学上等价 |
| 2 | 开梯度检查点 | 激活降一个数量级 | 前向多算约 30% |
| 3 | 降 rollout 并发 | KV cache 线性下降 | 采样变慢 |
| 4 | 降 max_completion_length | KV cache 线性下降 | 可能影响效果,慎用 |
| 5 | 关参考模型(KL 系数设 0) | 一份完整权重 | 失去 KL 约束,模型可能跑飞 |
| 6 | 换 LoRA | 梯度 + 优化器几乎归零 | 表达能力受限 |
| 7 | 换更小的模型 | 全线下降 | 换了个问题,不是解决问题 |
nvidia-smi 显示 32607 MiB ≈ 31.8 GiB,减掉显示输出占用(如果这张卡还接着显示器)和 CUDA context,实际能给训练用的大约 30 GiB。
跑训练的机器建议用核显输出,把整张 5090 空出来。
Qwen3-0.6B 用 AdamW 全参训练,bf16 权重。权重 + 梯度 + 优化器状态大约多少?
0.6B 的模型训练时 OOM 了。下面哪些是合理的第一步排查方向?(多选)
把 rollout 的最大长度从 2048 加到 8192,KV cache 会怎么变?
这节课的落点
- RL 显存里同时住着策略、参考、推理引擎三份权重,PPO 还要加 critic
- AdamW 优化器状态是 12 字节/可训练参数,通常是最大的一块
- logits = 长度 × 词表 × 精度,与模型大小无关,是「小模型也 OOM」的头号嫌疑
- KV cache = 每token × 长度 × 并发,长度和并发是乘法关系
- LoRA 的杠杆在梯度和优化器状态,不在权重
- OOM 时第一刀是降 micro batch,最后一刀才是换模型
- 留 10%~15% 余量,算出来 31 GiB 的配置实际会炸