RRLforge
计算器预计 45 分钟

显存账本: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——这一项在长回答场景下能超过所有模型权重的总和。

iPPO 还要再多一个

经典 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 的杠杆在这里

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,往往还要额外留一份中间量。

×「模型这么小怎么会 OOM」

答案八成在这里。序列长度翻倍、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
1610241.75 GiB
1640967.0 GiB
64409628 GiB ❌

并发和长度对显存是乘法关系。 想让模型「思考得更长」,代价是二次方式增长的显存压力 —— 这是 RL 训推理模型时最硬的一堵墙。

交互计算器

把上面六项串起来了。建议按顺序试这几组对比,比读文字有效得多:

  1. 看默认配置:Qwen3-0.6B 全参 + colocate。注意优化器状态那一条占了多大比例。
  2. 换 1.7B 全参:看它怎么爆的,以及是哪一条把它顶爆的。
  3. 1.7B 改 LoRA:看梯度和优化器两条同时塌下去。
  4. KV cache 的乘法:把 rollout 长度从 1024 拉到 8192,再把并发从 16 拉到 64,看警告栏说了什么。
  5. logits 陷阱:模型固定 0.6B,把训练序列长度拉到 8192、micro batch 拉到 8,看谁变成了最大的一块。
  6. 关掉参考模型:省了多少?(对应 KL 系数设 0 的场景)
计算器RL 训练显存账本
装得下
15.4 GiB
总共 32 GiB · 余量 16.6 GiB
训练参考推理context 与碎片预留
  • 策略模型权重1.12 GiB
    bf16,2 字节/参数
  • 梯度1.12 GiB
    全参:与权重同量级
  • 优化器状态6.71 GiB
    AdamW:fp32 主权重 + 一阶 + 二阶动量,12 字节/可训练参数
  • 激活0.13 GiB
    开了梯度检查点,只存层边界
  • logits 与 logprob1.16 GiB
    词表 151,936,长度 × 词表 × 2 份
  • 参考模型1.12 GiB
    bf16 冻结,算 KL 用
  • vLLM 权重1.12 GiB
    推理引擎自己持有一份权重
  • KV cache1.75 GiB
    112 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 了先动哪一刀

按「效果 ÷ 代价」排序,从上往下试:

顺序动作省多少代价
1micro batch 降到 1 + 梯度累积激活和 logits 成比例下降慢一点,数学上等价
2开梯度检查点激活降一个数量级前向多算约 30%
3降 rollout 并发KV cache 线性下降采样变慢
4降 max_completion_lengthKV cache 线性下降可能影响效果,慎用
5关参考模型(KL 系数设 0)一份完整权重失去 KL 约束,模型可能跑飞
6换 LoRA梯度 + 优化器几乎归零表达能力受限
7换更小的模型全线下降换了个问题,不是解决问题
5090 的实际可用显存不是 32 GiB

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 的配置实际会炸

延伸资料