7B模型微调到底要多少显存?我拿一张3090算了笔账,发现多数人第一步就搞反了

🔑 关键词:7B模型微调显存, LoRA, QLoRA, CUDA out of memory, 梯度检查点

📖 摘要:拆解混合精度训练下7B模型微调的真实显存开销,对比全参/LoRA/QLoRA三种方案在24G单卡上的可行边界,附一套从便宜到贵、按顺序执行的OOM排查清单。

先把账算清楚,再决定要不要买卡

图片

网上只要有人问“微调7B要什么卡”,评论区大概率会飘来一句“没8卡A100别玩”。我去年也信了这一套,直到某天真的把显存账单一项一项列出来,才发现很多人从第一步就把问题搞反了。

混合精度训练(AMP,bf16前向 + fp32优化器状态)下,每个参数占的字节数是固定的:bf16权重2字节,fp32 master weight 4字节,AdamW一阶动量m 4字节,二阶动量v 4字节,加起来14字节。这个系数跟模型多大没关系,是PyTorch + AdamW这套组合的“底价”。

7B就是7×10^9 × 14 ≈ 98GB。也就是说,全参微调7B,光优化器状态就吃掉84GB,是模型权重的6倍。这才是“必须多卡”的真正原因——不是模型装不下,是Adam装不下。所以看到有人用ZeRO-1在8卡上跑7B全参,不用惊讶,ZeRO-1切的就是优化器状态,8张卡正好把84GB摊平。

插一句跑题的:这也是为什么SGD + momentum在某些场景又被人捡回来了,它每个参数只要8字节(权重fp32 4 + 动量4),能省掉差不多43%。收敛会慢,但如果你只是做一个领域适配,未必不能接受。

图片

我那张3090上跑出来的实际数字

环境交代一下:单张RTX 3090,nvidia-smi显示总显存24576 MiB,但Ubuntu 22.04桌面加显示输出要吃掉400到600 MiB,实际能用的大概23800 MiB,这个细节很多人会忽略,然后对着“为什么我只加载了23G就OOM”发呆。软件栈是torch 2.2 + transformers 4.40 + peft 0.10 + bitsandbytes 0.43。

方案A,全参微调:直接OOM,没有讨论空间。换8bit Adam也救不回来——它把m和v各从4字节压到1字节,每参数从14字节降到8字节,7B也就是56GB左右,仍然超。

图片

方案B,bf16 + LoRA。常规做法是在q_proj / k_proj / v_proj / o_proj上加LoRA,r=8。这几个投影层的可训练参数算下来是4 × 8 × 8192 × 32层 ≈ 8.4M,占7B的0.12%。权重占14GB(bf16),剩下大约10GB留给激活和临时张量,batch=1、max_seq_length=1024、开启gradient checkpointing的情况下,峰值显存落在17到18GB之间。

方案C,QLoRA。load_in_4bit + nf4 + double quant,权重压到约3.5GB;把r开到64,并且把FFN的gate_proj / up_proj / down_proj也加进target_modules,可训练参数大约160M,优化器状态也就1.9GB上下,峰值显存11到13GB。

这里说一个我个人的、可能不太合群的看法:在24GB这个档位上,我不建议无脑上4bit。4bit权重每做一次forward都要现场反量化,bitsandbytes那套kernel在40系卡上开销不小,我实测QLoRA比bf16 LoRA慢25%到40%(同样的step数和等效batch)。质量上QLoRA论文说nf4能追平16bit,但那个结论是在65B规模上得的,7B上我体感能追回来大部分,可没到“完全无损”。所以我的一般判断是:显存够就bf16 LoRA,真的只有16GB及以下才考虑4bit。

梯度检查点不是免费的,而且很多人用错了时机

图片

gradient_checkpointing=True的原理是不保存中间激活,反向传播时重算一遍前向。它确实省显存,代价是20%到35%的额外时间,模型越深越亏。

问题在于:你的瓶颈到底是不是激活?如果batch=1、seq=512,激活可能只占1到2GB,而权重占14GB,那你开检查点就是在用30%的速度换1GB空间,非常不划算。反过来,seq=4096的时候标准attention的s²项会爆炸,这时候开检查点再加flash attention 2才是正解。

我自己的判定标准很土:先跑一个step,用torch.cuda.max_memory_allocated() / 1024**3打出峰值,再把这个数和权重占用比一比。激活占比超过40%才值得动检查点,不到20%就别折腾了,去调别的地方。

图片

顺便提一个容易忽略的点:开了gradient checkpointing之后,use_cache必须关掉,Peft里如果不设model.config.use_cache = False,会直接报一个和显存八竿子打不着的错误,我第一次踩的时候查了两小时。

遇到OOM的排查顺序,从便宜的开始

按这个顺序试,基本能覆盖九成情况。顺序反了会很痛苦,我见过太多人一上来就QLoRA 4bit,然后抱怨训练慢得像蜗牛——其实他的卡根本不需要量化。

  1. 先降batch,但配合梯度累积保持等效batch不变。per_device_train_batch_size=1, gradient_accumulation_steps=16和bs=16在数学上等价,Transformer里没有BatchNorm这类依赖真实batch统计的层,所以基本可以放心换。
  2. 再砍max_seq_length,从2048降到1024,激活内存大致线性下降。注意这是有损的,长样本会被截断,先看看你的数据长度分布再动手。
  3. 开gradient checkpointing,记得同时关use_cache。
  4. 把优化器换成adamw_bnb_8bit或paged_adamw_8bit,这一步常常能白捡2到4GB。
  5. 上flash attention 2(attn_implementation="flash_attention_2",需要Ampere及以上),这个是真提速,不是取舍,能上就上。
  6. 最后才考虑量化权重。

图片

还有一个特别隐蔽的坑:DataLoader的num_workers设太大,每个worker会持有自己的一份数据,虽然占的是CPU内存,但当pin_memory=True且worker数超过8时,page-locked内存的挤占会影响CUDA分配,间接引发看起来莫名其妙的OOM。我一般num_workers=4配persistent_workers=True就够,不是越多越快。

最后一个习惯:实验跑完一定要del model; torch.cuda.empty_cache()。不然在同一个notebook里跑第二次实验,第一次留下的显存碎片还在,报出来的错会让你怀疑是不是自己代码写错了。

说到底,显存不是买出来的,是算出来的。90%的OOM,问题不在batch size,在于你从来没认真算过优化器状态那笔账。