先说一个我反复遇到的场景:你拿到了一张32GB显存的卡,兴冲冲跑LoRA微调一个7B模型,脚本刚启动,loss还没出来就收到OutOfMemoryError。更气人的是,别人同一份代码在24GB的卡上都能跑,你32GB反而爆了。这种问题我碰了不下十次,根子往往不是卡太小,而是大家默认“LoRA很省显存”,却没人告诉你LoRA到底省的是哪几笔钱、哪几笔钱一分没省。
这篇文章就把显存估算这件事彻底讲明白:给出能直接套用的估算公式和手算过程,给出32GB GPU上7B、13B这类主流模型的LoRA训练配置,再走一遍从OOM报错反推配置问题的完整排查链路。适合刚接触大模型微调、手头只有一张消费级或专业卡、想在开跑前把显存预算算清楚的读者,也适合已经跑起来但经常中途爆显存、不知道怎么调的人。
1. 先把显存账算明白:LoRA到底把哪些钱省了,哪些一分没省
很多人对LoRA省显存的理解是“模型变小了”,这是错的。LoRA没有改变base model的参数量和计算图,前向、反向照样要过完整模型,模型权重也照样占着显存。它真正省掉的东西,是“全参数微调里那笔最吓人的开销:优化器状态、梯度、以及中间主权重副本”。
1.1 全参微调的显存大头到底有多大
以7B模型为例,全参微调如果走混合精度训练,显存账本大致是这样的:
- 模型权重:bf16精度,7B参数量,每参数2字节,约14GB;
- fp32主权重副本:因为混合精度需要稳定的fp32参数,约28GB;
- AdamW优化器状态:每个参数需要保存fp32一阶动量、fp32二阶动量,约28GB+28GB;
- 梯度:fp32或者bf16都可能,按fp32算约28GB,按bf16算约14GB;
- 激活值:取决于batch size和序列长度,7B短文本下也很容易到5GB以上。
这几笔加起来轻松超过100GB。所以别说32GB卡,80GB卡跑全参微调7B都很勉强。这也是LoRA横空出世后被广泛使用的最核心原因——它把参加梯度计算和优化器更新的参数,从7B压缩到几千万乃至几百万。
1.2 LoRA是真的在“可训练参数”上省钱
LoRA冻结了base model,只训练插入的低秩矩阵。假设你把7B模型的q/k/v/o、gate/up/down投影都接上LoRA,rank=16时,实际可训练参数量大概在4000万左右,占全模型的0.5%到0.6%。
这笔账就完全不一样了:
- 主权重副本:只为LoRA参数保留fp32副本,4000万×4字节≈0.16GB;
- 优化器状态:AdamW每个训练参数约12字节,4000万×12≈0.48GB;
- 梯度:只对LoRA参数算梯度,4000万×2字节≈0.08GB。
也就是说,LoRA把全参微调里那几十GB的梯度+优化器状态,压缩到了1GB以内。这就是“LoRA很省显存”这句话的真正来源:它省的是可训练参数相关的开销,不是模型权重,也不是激活值。
1.3 为什么很多人仍然跑爆显存
理解了上面那笔账,你就会发现:在32GB卡上跑一个7B模型LoRA微调,bf16权重14GB是绕不开的固定成本。剩下的大头就是激活值。而激活值的大小完全由模型结构、batch size、序列长度和是否开启gradient checkpointing决定,和“LoRA”三个字没有半毛钱关系。
我见过不少新手Rank直接开64,target_modules一股脑全加进去,再配一个长序列和batch size 4,结果可训练参数从4000万涨到一亿多,激活值也涨到吓人,32GB卡当场去世。不是LoRA不行,是显存账没算清。
所以记住第一句话:LoRA优化的是优化器状态和梯度,不是模型推理时的激活值。后面所有估算和调参,都围绕权重、激活、可训练参数这三块展开。
2. Step by Step估一次显存:7B模型在32GB卡上的预算表
显存估算不用精确到MB,但你要能算出数量级,知道现在是“舒适”“边缘”还是“必爆”。我习惯把总显存拆成五笔账:
总显存 = 模型权重W + 梯度G + 优化器状态O + 激活值A + 框架开销C其中W几乎固定;G和O由可训练参数量决定;A最复杂,也最容易被忽略;C是CUDA context、PyTorch缓存池、临时中间变量等,大概500MB到1GB。
2.1 权重、梯度、优化器状态怎么算
权重最直接:参数量×精度字节数。bf16是2字节,fp16也是2字节,fp32是4字节,4bit NF4量化后大约0.5到0.6字节(因为还有量化缩放系数)。
| 模型规模 | bf16权重 | fp32权重 | 4bit NF4权重(约) |
|---|---|---|---|
| 3B | 6GB | 12GB | 1.8GB |
| 7B | 14GB | 28GB | 3.8GB |
| 13B/14B | 26GB~28GB | 52GB~56GB | 7.5GB左右 |
| 32B | 64GB | 128GB | 17GB左右 |
梯度只对可训练参数算,所以LoRA场景下:
G = LoRA参数量 × 梯度字节数(通常2字节) O = LoRA参数量 × 优化器状态字节数AdamW在混合精度下,每个训练参数要额外保存三类状态:fp32主权重副本、fp32一阶动量、fp32二阶动量,总共约12字节。如果你用bitsandbytes的8bit Adam,这个数字可以降到每参数2~3字节,但普通32GB卡跑7B LoRA一般不需要动优化器。
举个例子:7B模型,LoRA rank=16,7个投影层全接上,假设打印出来的trainable参数约4000万。那么:
- G=4000万×2B≈0.08GB;
- O=4000万×12B≈0.48GB。
就算你把rank提到64,可训练参数可能涨到1.2亿到1.6亿,G+O也就2GB出头。所以结论很明确:LoRA参数对显存的影响是百万到GB级别的扰动,远不是生死线。
2.2 激活值:最难算但也最该算的一笔
激活值的准确数字很难只靠理论公式手算,因为每个Transformer层中间有大量临时变量:QKV投影结果、注意力权重、MLP中间层、LayerNorm、dropout mask等等。它和batch size、序列长度、层数、隐藏维度成正比,也和是否使用memory-efficient attention、gradient checkpointing密切相关。
我常用的经验粗算是:先算一个“基数额”:
Base = batch × seq_len × hidden_size × num_layers × 2字节7B模型hidden_size约4096,num_layers约32,batch=1,seq=2048时:
Base = 1 × 2048 × 4096 × 32 × 2 ≈ 0.537GB然后乘经验系数。不开gradient checkpointing时,系数通常在15到25之间,意味着激活值大约8GB到12GB;开了gradient checkpointing,系数可以降到3到6,激活值大约2GB到3GB。配合flash attention或SDPA的memory-efficient kernel,系数还能再低。
这个系数看着粗糙,但足够让你做决策:同样的7B模型,开不开gradient checkpointing,前后可能差出6到9GB显存。而这恰恰是32GB卡能不能同时塞下14GB权重+其他开销的关键。
2.3 用代码确认,别用“感觉”确认
理论估算完,一定要用实际打印结果校正。PEFT提供了现成方法:
model.print_trainable_parameters() # 输出示例: trainable params: 40,000,000 || all params: 7,000,000,000 || trainable%: 0.5714训练中要盯峰值显存,可以在Trainer里挂一个回调:
from transformers import TrainerCallback import torch class MemCallback(TrainerCallback): def on_step_end(self, args, state, control, **kwargs): if state.global_step % args.logging_steps == 0: peak = torch.cuda.max_memory_allocated() / 1024**3 print(f"step {state.global_step} peak: {peak:.2f} GiB")把五笔账大概加起来,再对照这份实际峰值,你很快就能形成对“这个模型配这些参数需要多少显存”的直觉。我个人的判断线是:预算不超过总显存的85%,留出余量给碎片和临时变量。
3. 32GB卡上的训练配置模板:建议照抄再改
有了预算顺序,就可以落地配置了。下面的模板我默认你用的是transformers + PEFT,这是目前做LoRA微调最主流的组合。32GB显卡可以是A100的某个实例规格、V100 32GB,也可以是具备相近显存容量但支持bf16的卡。如果你的卡不支持bf16,把bf16=True改成fp16=True即可,其他逻辑不变。
3.1 环境版本先说清楚
- PyTorch建议2.1以上,2.4以上更好;
- transformers建议4.40以上;
- peft建议0.10以上;
- 如果要用4bit QLoRA,还需要bitsandbytes,0.43以上比较稳。
老卡(比如V100)不支持bf16,这类卡上跑长序列场景,速度可能不如新卡,但fp16配合gradient checkpointing依然能跑。建议在开跑前先用一个极小的batch和seq把链路跑通,再逐步加量。
3.2 一个能直接跑的脚本骨架
import torch from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training model_id = "Qwen/Qwen2.5-7B" # 如果做QLoRA,在from_pretrained里加 load_in_4bit=True model = AutoModelForCausalLM.from_pretrained( model_id, torch_dtype=torch.bfloat16, device_map="auto", use_cache=False, # 训练时关闭KV cache,省显存 low_cpu_mem_usage=True, ) # 激活值最大的敌人:不开gradient checkpointing,7B短文本都能吃几GB model.gradient_checkpointing_enable() model = prepare_model_for_kbit_training(model) # QLoRA场景必开,纯bf16可不开 lora_config = LoraConfig( r=16, lora_alpha=32, target_modules=[ "q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj" ], lora_dropout=0.05, bias="none", task_type="CAUSAL_LM", ) model = get_peft_model(model, lora_config) training_args = TrainingArguments( output_dir="./lora_out", per_device_train_batch_size=1, gradient_accumulation_steps=16, gradient_checkpointing=True, optim="adamw_8bit", # 也可用adamw_torch learning_rate=2e-4, max_steps=1000, logging_steps=10, save_steps=200, bf16=True, # 不支持bf16就改fp16=True dataloader_num_workers=0, remove_unused_columns=False, report_to=None, )几个关键点:
target_modules里的模块名必须和模型实际结构一致。Qwen和Llama系的命名一般是q_proj, k_proj, v_proj, o_proj, gate_proj, up_proj, down_proj;ChatGLM系则是query_key_value, dense, dense_h_to_4h, dense_4h_to_h。不确定就打印model.config或微信群里搜一下。device_map="auto"在单卡且显存足够时会把整个模型放到GPU,但在某些版本里可能因为config里的use_cache等问题做offload,反而多出额外开销。如果确认单卡能放下,直接model = model.to("cuda")更可控。- 先
model.gradient_checkpointing_enable(),再get_peft_model,兼容性最好。如果你已经在TrainingArguments里开了gradient_checkpointing,Trainer也会在训练时设置,脚本里二选一也行。 adamw_8bit和adamw_torch在后向时的显存峰值差异通常在1GB量级,初期没必要为了这点显存牺牲稳定性,32GB卡跑7B用adamw_torch通常就够。
3.3 不同量级模型的显存实测组合
我按“峰值显存不超过27GB,给系统留buffer”的原则,整理了一份常见组合。注意不同框架、不同CUDA版本、有没有装flash-attn都会带来2~3GB浮动,所以这个表是决策参考,不是精确值。
| 模型规模 | 精度 | LoRA rank | max_seq_len | batch | gradient checkpointing | 实测峰值(约) | 32GB卡结论 |
|---|---|---|---|---|---|---|---|
| 7B | bf16 | 16 | 2048 | 1 | 开 | 17~19GB | 舒适 |
| 7B | bf16 | 16 | 8192 | 1 | 开+flash attn | 24~27GB | 可跑但紧 |
| 13B/14B | bf16 | 16 | 1024 | 1 | 开 | 30~33GB | 边缘,很可能爆 |
| 13B/14B | NF4 QLoRA | 16 | 2048 | 1 | 开 | 13~16GB | 舒适 |
| 32B | NF4 QLoRA | 8 | 512 | 1 | 开 | 20~23GB | 可跑 |
这个表说明几件事:32GB卡跑7B,priority是随便浪;跑13B以上,老老实实上QLoRA,否则即使seq压到1024,bf16权重26GB+激活3GB+优化器1GB,基本贴在32GB边缘,稍有碎片就爆。
3.4 为什么我坚持bs=1 + 高梯度累积
有人会问:32GB卡跑7B,batch size开个4不香吗?
激活值几乎正比于batch size,而梯度累积是在显存不增大的前提下,把多个batch的梯度先攒再更新。用bs=1、gradient_accumulation=16,等效batch size是16,但单次forward和backward的峰值显存远小于bs=4、accumulation=4的组合。尤其在你seq很长的时候,bs=4的激活值可能是bs=1的三到四倍。
代价是训练总时长几乎不变,甚至略慢一点。因为梯度累积只是分次更新,GPU算力并没有浪费多少。我的习惯是:先把batch固定为1,用gradient_accumulation撑等效batch;如果显存还剩很多,再尝试bs=2并对比峰值和训练吞吐。
4. OOM排查链路:从报错出现的时机反推账本哪个科目超支
OOM最怕的是瞎试。看到OutOfMemoryError就不管三七二十一降batch,降完了还爆,又去降LoRA rank,一晚上没了。实际上,OOM出现的时机已经把答案写给你了。
4.1 加载模型阶段就OOM
这个阶段的报错通常在from_pretrained或脚本刚起没多久出现,日志里会显示加载权重时分配显存失败。原因只有三类:
- 卡上已经有人占了显存。先执行
nvidia-smi --query-compute-apps=pid,used_memory --format=csv看当前进程,别让僵尸进程偷偷占着显存。 - bf16权重算出来14GB,加上CUDA context和模型加载临时缓冲区,一下冲过32GB。这种情况先开
low_cpu_mem_usage=True,再试着用device_map="auto",最后再考虑4bit量化。 - 代码里重复加载了多个模型。这个很常见:在同一个python进程里先load一个7B,又load一个测试模型,两个只在函数里没释放,显存叠加。
排查方法很简单:加载模型后立刻打印torch.cuda.memory_summary(abbreviate=True),看allocated是多少、reserved是多少。
4.2 第一次前向就OOM
这是最典型的一类,报错“Tried to allocate XX MiB”时GPU已经用了28GB以上。问题几乎都在激活值,而不是权重。
优先检查三件事:
gradient_checkpointing到底开没开。很多人以为是开了,实际Trainer配置没生效,或者base model被prepare_model_for_kbit_training改动后没重新enable。max_seq_len是不是设置了虚高值。有些数据集里大量样本其实只有几百token,你设了2048,每个batch都按batch内最长样本padding,显存直接按最大seq算。batch_size是不是默认的8或更高。Trainer里per_device_train_batch_size没设置时会默认8,一个7B短文本可能直接14GB权重+十几GB激活,OOM是必然的。
先跑一个最小配置确认链路能通:bs=1、seq=128、gradient_checkpointing开、LoRA rank=8。如果这个都能OOM,说明环境有问题;如果这个能通,再逐步往上加seq和rank。
4.3 训练到一半才OOM
训练前100步没事,跑到500步突然爆,是最让人头疼的。它一般对应三种情况:
- 显存碎片化:PyTorch的显存分配器会在整个训练过程中不断分配和释放,长期运行后产生大量碎片。明明reserved里还有几GB闲置,但找不到连续块分给你一个大tensor,于是OOM。解法是设置环境变量
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,PyTorch 2.4及以上效果明显。也可以试max_split_size_mb=128,但通常不如expandable segments。 - 峰值波动:训练中不同step的激活值不同。某个最长样本、某次特殊padding、某个大batch的step会把峰值顶上去。这和碎片无关,是正常峰值在哪个step暴露的问题。做法就是用
MemCallback记录最高峰值,再回头调seq或batch。 - 内存泄漏:自定义callback或训练循环里不小心保留了每step的tensor,显存allocated持续上涨。可以观察
max_memory_allocated是否每个step单调增加。如果是,检查代码里有没有把loss.item()之外的大tensor存进list。
4.4 一个容易误判的“整卡没满但PyTorch OOM”案例
这是排查里最反直觉的一类:nvidia-smi看整卡显存显示只用了20GB,但PyTorch报OOM。原因通常是nvidia-smi里的“已用显存”包含了系统CUDA context、其他进程和你这个进程的reserved池,这个数字不等于PyTorch的allocated。
另外,PyTorch的显存池一旦reserved,不会自动释放给其他进程,哪怕当前所有allocated都快释放了。如果长期不重启进程,reserved池和别的东西竞争,就会出现“明明还有显存,却分配失败”。
所以排查时不要只看nvidia-smi,要把torch.cuda.max_memory_allocated()、torch.cuda.memory_reserved()和memory.summary()一起看。我的经验做法是:OOM后立刻在异常处理里打印这几项,记录到日志,再反过来判断该砍哪笔账。
5. 排错后我常用的优化顺序:从免费到伤身
当显存不足时,我有一套固定的优化顺序,基本不会浪费时间。原则是:先做零成本改动,再做有精度损失的改动。
5.1 优先级清单
| 优先级 | 操作 | 显存收益 | 速度影响 | 效果影响 |
|---|---|---|---|---|
| 1 | 开gradient checkpointing | 数GB到十几GB | 慢20%~30%,可配合flash attn回补 | 基本无 |
| 2 | 砍max_seq_len到任务实际需要 | 线性省激活 | 更快 | 取决于数据长度分布 |
| 3 | bs=1 + gradient_accumulation | 数GB | 几乎无 | 基本无 |
| 4 | 换adamw_8bit或paged_adamw_8bit | 1GB上下 | 略快或持平 | 可能轻微影响收敛 |
| 5 | 降LoRA rank(16→8) | 0.5~1GB | 略快 | 低rank可能影响效果 |
| 6 | base model做NF4量化 | 权重从14GB降到4GB | 更快/更慢取决于kernel | 量化噪声,需要学习率略调 |
第一步永远是gradient checkpointing,因为它是显存和计算的重分配,不牺牲精度。第二步是seq,很多人把seq设得很高只是为了“保险”,其实任务数据根本用不到那么多token,砍掉一半往往显存直接降几GB。
5.2 小心padding带来的隐性浪费
我调显存时发现一个规律:不少OOM不是模型真的跑不动,而是dataloader把padding到很长。比如你设max_seq_len=2048,训练集里80%样本不到512 token,它们都被pad到2048,激活值凭空多了好几倍。
解法有两种:一是数据预处理时按长度分桶,把接近的样本放一起;二是用seq2seq里常见的packing方法,把多个短样本拼到接近固定长度,减少padding。PEFT和部分开源训练库直接支持packing,但packing后的attention mask处理要注意,别让文本跨样本看到。
5.3 学会用峰值显存做决策,而不是感觉
最后分享一个工作习惯:任何一次训练,我都会在脚本里打一行峰值显存日志,并把启动参数、模型规模、seq、batch、rank都写进保存日志的文件名。
这么做的原因是,显存问题很容易重复踩。这次调好跑了,过一个月换一个数据集又爆了,如果没有当时的峰值记录,你又要从头猜。把这些数据记录成一张表,下次改参数之前先翻一翻,通常几分钟就能定位是该砍seq还是该上量化。
我自己在32GB卡上长期跑LoRA的体会是:显存问题十有八九不是卡不够大,而是账没算明白。开跑前花五分钟把五笔账列一下,OOM之后按时机反推,再按优先级调参,基本都能在20分钟内解决。最后再提醒一个小细节:gradient checkpointing打开后,第一次训练会明显比不开慢,因为前向要额外重算一遍;但只要把gradient_checkpointing_kwargs={"use_reentrant": False}配上,再装好flash attention,这个速度损失在7B量级是可以接受的——相比之下省下来的好几个GB显存,性价比高太多了。