微调一跑就OOM,报错弹出来那一刻,我敢说绝大多数人的第一反应是“模型太大了”。但等我真正把显存账算清楚之后发现,这个直觉基本是错的——尤其是用LoRA这类高效微调方案时,模型参数本身根本不是什么大头。真正让显存原地爆炸的,是训练态特有的那几块开销:优化器状态、梯度、激活值。这篇文章我就把这笔账一条一条拆开算,再把低显存微调能用得上的工具和参数配置全盘托出。
先说一个典型的误导场景:一张RTX 3060 12G,跑Qwen 7B推理,BF16加载也就14G,量化到4bit大概6G,很轻松。然后你想微调,心想“模型也就14G,24G卡总够了吧”,结果一跑训练,瞬间OOM。问题出在哪?因为你只算了模型权重的账,没算训练时多出来的那三笔开销。尤其全参微调,AdamW优化器状态下,每多一个参数就要多吞8字节的FP32状态,加上梯度、激活值,7B模型全参微调轻轻松松吞掉80到120G显存——这不是你在普通显卡上能想象的事。
这篇文章适合手里只有6G、8G、12G这类中低显存卡,又想自己试试大模型微调的人。我先把显存消耗的物理构成讲清楚,然后给出一套从框架选型到参数配置的完整实操路径。保证你看完能动手,也知道每一步为什么要这么做。
1. 显存到底被谁吃掉了?别再把锅甩给模型参数
1.1 训练态和推理态的最大区别:那些“隐藏”的内存占用
推理的时候,你的显存里只有模型权重外加KV Cache,所以你看7B模型BF16权重14GB,显存占用也就十五六个G,顶多加一点KV Cache。
但一进入训练态,事情就完全变了。以全参微调为例,显存里至少要有四样东西:
- 模型权重:就是你加载进来的参数,BF16下2字节/参数,7B约14GB。
- 梯度:反向传播算出来的梯度,和数据一样形状,通常也要一份。BF16或者FP32看实现,但至少要再占一份权重同等大小的空间。
- 优化器状态:AdamW这类优化器要为每个参数维护两份状态,动量项和二阶矩项,都是FP32,也就是每参数8字节。此外混合精度下还会再留一份FP32的“主权重”副本,约4字节/参数。
- 激活值:前向传播过程中每层算出来的中间张量,这些在反向传播时要用来算梯度,所以默认都会存下来。
推理时你只需要第一项,训练时四项都要。而且第三项和第四项通常比你想象的大得多。
1.2 用数字化一下:7B模型从加载到训练的显存账本
咱们直接列个表,按7B模型、BF16加载来算:
| 项目 | 每参数开销 | 总开销(7B) | 说明 |
|---|---|---|---|
| 模型权重(BF16) | 2字节 | 14GB | 推理态就有的部分 |
| 梯度 | 2~4字节 | 14~28GB | 取决于是否混合精度 |
| AdamW优化器状态 | 12字节 | 84GB | FP32主权重4B+动量4B+二阶矩4B |
| 激活值 | 取决于序列长度和batch | 数GB~数十GB | 和batch、seq_len强相关 |
明白了吧,全参微调7B,光是模型权重加优化器状态就要接近70~100GB。所以大家天天说的“全参微调至少要80G显存”不是危言耸听,你这是跟物理规律在打仗,省不掉的。
而LoRA这类PEFT方法关键就在于:冻结原权重,只训练极少量新增的低秩矩阵。7B模型如果rank=16,可训练参数大概只有几百万到一千多万,是原来的1/700。优化器状态和梯度开销瞬间从几十GB降到几十MB,剩下的显存大头又变回了模型权重本身。这才是LoRA能低显存跑的根本原因。
1.3 激活值为什么是最阴险的“隐形成本”
如果说优化器状态是显存爆炸的第一元凶,那激活值就是第二元凶,而且它最阴险——因为很多人都没意识到它存在。
激活值的公式大致可以写成:激活值占用约等于 batch × seq_len × hidden_size × 层数 × 一个常数。
举个例子,7B模型一般是32层,hidden_size约4096,MLP中间维度约11008。如果你用seq_len=2048、batch=4硬跑,激活值占用很容易来到20GB以上。有人以为batch和seq_len只影响计算量,不影响显存,这是错的。它们对显存的影响往往比“换个更大的模型”还猛。
那同样7B模型,为什么有人用6G卡也能跑微调?他们做的事就是:LoRA把优化器状态削掉,4bit量化把权重从14GB压到4GB左右,再把seq_len掐到512、batch=1,激活值压到1GB以内。显存就被这么硬生生挤出来了。
2. 微调显存优化的三个真正的杠杆:精度、秩、序列长度
2.1 别上来就标配BF16,NF4量化可能是你的改命手段
很多人习惯了推理时用BF16、FP16,一说到微调也直接按这个精度来。对于24G以上的卡这没大问题,但对于12G以下的卡,你要是还坚持BF16全量加载7B,14GB的权重本身就塞不进去,后面全是白搭。
解决办法就是把权重量化之后再做微调,这就是QLoRA的思路:用4bit的NF4格式存放原始权重,同时只对LoRA新增的低秩矩阵保持BF16精度做训练。原始权重虽然被量化了,但LoRA分支在学习更新,最终把LoRA合并回去,效果上损失通常很小。
举个例子,7B模型用NF4量化,权重降到4GB左右,加上LoRA的梯度/优化器开销,12G卡甚至可以开seq_len=2048跑。6G卡把seq_len压到512也不是不能跑。这就是为什么QLoRA几乎成了低显存微调的默认选项。
这里顺便提一下,现在量化格式也在快速演进,像NVFP4这类针对新架构的4bit浮点格式也开始进入工具链。如果你用的是比较新的卡,可以留意一下bitsandbytes或者对应框架对新格式的支持情况,训练时占用会更低,精度也比NF4更稳一点。
2.2 LoRA rank不是越大越好,很多人的显存死在rank=64甚至128上
LoRA的核心是“用低秩矩阵近似权重增量”,rank就是低秩矩阵的维度。很多人有个本能冲动:rank设大一点,是不是学得更充分?从效果上看确实存在“rank上限”的说法,但对显存来说,rank每翻一倍,可训练参数的量就翻一倍,优化器状态、梯度、甚至某些实现里的激活开销都会涨。
我的经验是低显存场景下优先从rank=8开始试。7B模型rank=8的可训练参数量大约在500万到800万,rank=16大约是1000万到1600万,rank=64则能到4000万以上。对大多数指令微调任务来说,rank=16已经能匹配绝大部分任务需求。你如果先跑通流程,再往上加rank也不迟。
同样的问题也出现在target_module的选择上。有人默认把all linear全挂了LoRA,这没问题,但如果你显存已经见底,可以考虑只挂q_proj和v_proj,或者至少把embedding层排除掉。embedding层参数量巨大且很少是微调的重点,在部分框架里默认不参与LoRA训练就能省不少。
2.3 序列长度是比batch_size更凶的显存杀手
我再强调一遍这个结论:训练显存峰值对序列长度的敏感程度远高于对batch_size的敏感程度。原因很简单,Transformer的Self-Attention激活显存随序列长度线性涨(FlashAttention之后),而MLP的激活值和批量大小线性相关。所以在低显存卡上,优先砍seq_len,其次才砍batch。
具体怎么选?我给出一个调参顺序,照着做就行:
- 先把seq_len设成你下游任务能接受的最短长度,比如128/256/512。
- batch_size从1开始,这是底线。
- 如果batch=1还OOM,就把seq_len再砍半。
- 如果batch=1能跑但速度太慢,把gradient_accumulation_steps抬高到4/8/16,用多步累计来模拟更大的batch效果——注意这一步不会增加单步显存峰值,只是等效增大了batch,省显存的同时不牺牲更新质量。
还有一个被很多人忽略的开关叫gradient_checkpointing,也叫激活重计算。它会把前向传播的中间激活值扔掉一部分,反向传播时再重算一遍,代价是多花约30%计算量,但激活值显存通常能降到原来的1/3左右。对于长序列任务,这个开关基本是必开的。
3. 低显存微调的主流框架怎么选?从LLaMA-Factory到Unsloth
3.1 先明确一件事:你的显存究竟是多少
不同显存档位,能选的路径差别很大。我先画一个简单的分档:
| 显存档位 | 能跑什么 | 推荐路径 |
|---|---|---|
| 6G | 7B以内模型QLoRA,seq_len需限制在512左右 | Unsloth / bitsandbytes + PEFT |
| 8G | 7B模型QLoRA较稳,可上seq_len=1024 | LLaMA-Factory / Unsloth,开gradient_checkpointing |
| 12G | 7B模型QLoRA舒服,13B模型限量跑 | LLaMA-Factory / PEFT |
| 16G | 13B模型QLoRA、7B全参微调勉强 | Unsloth / DeepSpeed Stage2 |
| 24G | 7B全参微调、13B全参微调紧张 | DeepSpeed ZeRO-2/3 |
如果是6G卡,就别碰BF16全参了,那等于直接建一座自己跨不过去的墙。老老实实用QLoRA + Unsloth。
3.2 6-8G显存:QLoRA + bitsandbytes/PEFT 是基本功
6-8G这个档位,如果只是纯PEFT库自己拼配置,也完全能跑,但体验会和顺滑差很远。建议优先试Unsloth,它的底层kernel针对FlashAttention和低秩训练做了很多工程优化,训练速度比传统PEFT快2-3倍,显存占用低不少。
Unsloth的API和HuggingFace Transformers高度兼容,基本就是把from_pretrained换成UnslothMistralForCausalLM.from_pretrained之类,再用get_peft_model包一层,微调流程几乎不变。它自己集成的LoRA实现,比直接调PEFT要更省显存。
如果不想折腾新东西,那就用经典组合:transformers+peft+bitsandbytes。加载时load_in_4bit=True,然后get_peft_model包一层LoRA。这个组合兼容性最好,网上教程最多,踩坑容易找到答案。
3.3 12-16G显存:LLaMA-Factory 是目前最顺手的全家桶
到了12G以上,能做的事就多了,这时候我强烈推荐LLaMA-Factory。它的优势在于把数据集处理、训练配置、推理测试打包成了一个整体,支持CLI和WebUI两种方式,默认配置就为低显存优化过了。你只要给它一个JSON格式的数据集,写一段YAML配置,它就能自动处理LoRA/QLoRA、gradient checkpointing、序列长度、学习率等所有参数。
它的WebUI对新手尤其友好,选模型、选微调方法、填rank和learning rate,点开始就训练。我见过不少完全没有代码基础的人,靠它在12G卡上完成了自己第一个领域微调模型。当然,如果你习惯命令行,LLaMA-Factory的CLI同样好用,也方便写脚本批量实验。
如果目标是跑更大模型,13B或更大的量级,可以考虑DeepSpeed ZeRO Stage 2甚至Stage 3,把优化器状态和梯度分摊到多卡或offload到CPU内存。单卡12G也是可以启动DeepSpeed的,重点是把zero_optimization.stage设成2,并把offload_optimizer.device设为cpu,这样优化器状态不占显存,但会牺牲一些训练速度。
3.4 24G及以上:也别急着全参微调
很多拿到3090/4090的人觉得自己可以全参微调了。24G显存跑7B全参微调,用BF16 + AdamW + gradient checkpointing,确实是挤得进去的,但非常紧张,而且训练时会非常慢——因为你要让梯度和优化器状态在显存里来回倒腾。
我的建议是:除非你的任务是真正需要全量微调(比如改变了模型结构或需要生成全新的能力边界),否则24G卡也优先用LoRA。24G可以支持的LoRA配置更宽裕:rank=32甚至64、seq_len=2048、batch=4以上、全量all_linear,这些配置下LoRA的效果已经逼近全参。你省下来的显存可以用于加大batch、加长上下文,往往比全参微调收益更明显。
4. 实操路径:从OOM到跑通一次低显存微调
4.1 压显存的标准动作顺序
我通常的排查顺序是这样,也建议你照着这个顺序来:
- 第一步先关掉所有“锦上添花”的功能,确保训练能起步。比如先关掉gradient checkpointing以外的所有高级特性。
- 加载4bit量化权重,确认模型能不能load进来。
- 设batch=1、seq_len=最短可接受长度,开跑一步。
- 如果OOM,把seq_len再砍半,或者把模型换成更小的量级。
- 如果跑通了,再把gradient checkpointing打开,观察显存降没降。
- 显存有余量,再逐步加seq_len、batch、rank,直到找到你显卡的甜点区间。
这个过程看着机械,但特别有效。我见过很多人一上来就把rank=32、seq_len=2048、batch=8一股脑开满,结果就是OOM,然后开始怀疑框架有问题。其实框架很无辜,只是你没有给显存做预算。
4.2 一份可以直接抄的QLoRA微调配置
这里给出一个基于HuggingFace PEFT库的QLoRA配置,适合12G显存跑7B模型:
from transformers import ( AutoModelForCausalLM, AutoTokenizer, TrainingArguments, BitsAndBytesConfig ) from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training from trl import SFTTrainer bnb_config = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_use_double_quant=True, bnb_4bit_compute_dtype="bf16" ) model = AutoModelForCausalLM.from_pretrained( "Qwen/Qwen2-7B-Instruct", quantization_config=bnb_config, device_map="auto", trust_remote_code=True ) model = prepare_model_for_kbit_training(model) 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="./qwen-lora", per_device_train_batch_size=1, gradient_accumulation_steps=8, gradient_checkpointing=True, learning_rate=2e-4, num_train_epochs=3, logging_steps=10, save_steps=500, bf16=True, optim="paged_adamw_8bit", max_steps=-1, warmup_ratio=0.03, lr_scheduler_type="cosine", ) trainer = SFTTrainer( model=model, args=training_args, train_dataset=train_dataset, max_seq_length=2048, dataset_text_field="text", packing=False, ) trainer.train()几个关键点解释一下:
bnb_4bit_quant_type="nf4":NF4是4bit量化里效果更稳的一种格式,优先用它。bnb_4bit_compute_dtype="bf16":量化权重做计算时临时转回BF16,能兼得速度和精度。optim="paged_adamw_8bit":用8bit版本的AdamW,优化器状态直接减半,省下来的显存很可观。gradient_checkpointing=True:压低激活值,几乎必开。max_seq_length=2048:如果你手里是8G卡,把它降到1024或512,就能腾出余量。
4.3 训练过程中如何实时监控显存分布
很多人看显存只知道用nvidia-smi看一眼整体占用,然后就是一脸懵。我的建议是训练过程中配合PyTorch自带的内存分析工具,可以更清楚地看到显存到底消耗在哪个环节:
# 在训练循环的某个关键步骤后执行 print(torch.cuda.memory_summary(device=None, abbreviated=True))会看到每个分配点大概占了多少,以及缓存池的状态。如果发现某个激活值张量特别占地方,十有八九就是seq_len或batch导致的,回头把这两个参数降一降。
我自己在排查问题时的标准动作是:第一步开一个交互式会话,加载模型跑一个step,然后立刻看memory_summary,确认哪个环节最吃显存,再去动对应的配置项。这种“对症下药”比盲目调参数高效得多。
5. MoE模型、CLIP微调以及其他显存陷阱
5.1 MoE架构是不是只要激活一部分专家就能省显存?
很多人觉得MoE(混合专家)模型推理时只激活一小部分专家,那训练时是不是也能只把激活专家放进显存?这个想法听着很美,但实操上完全行不通——至少在绝大多数框架里是行不通的。
原因很简单:MoE模型的所有专家权重在物理上都在同一个模型结构中,虽然Single Token只激活top-k个专家,但整个模型权重依然要被加载到显存里。训练的时候,还要为所有参数维护梯度(至少是那些参与计算的),不能“只加载一部分权重”。所以MoE模型的微调,显存压力一点也不小,甚至因为专家数量多、优化器状态更庞大而更头疼。
如果你只有12G卡还想跑MoE架构模型,我能给出的可行路径只有两条:首选用QLoRA,并且在target_modules里手动指定需要训练的层,跳过所有或大部分专家层的LoRA挂载;其次就是选择参数量更小的MoE子版本。千万别指望“只载入部分专家”这种妙招存在。
5.2 CLIP和多模态微调的显存优化逻辑
多模态模型的微调显存问题和大语言模型是同一个底层逻辑,Load进来的权重、梯度、优化器状态、激活值四大块,一样不少。比如CLIP模型微调,视觉编码器和文本编码器都要参与前向,激活值会比你想象的更占地方,尤其是高分辨率图像输入时。
解决办法依然是LoRA/QLoRA + 限制输入尺寸 + 限制batch。对图像模型,输入分辨率对显存的影响和LLM的序列长度本质是一回事,降低分辨率就是从根上减少激活值。
我自己的经验是:视觉任务的低显存微调,第一步一定是把输入分辨率降到数据允许的底线(比如224x224,甚至112x112先跑通),然后逐步往上加。这比折腾任何花哨优化都来得实在。
5.3 生成模型、视频模型与“轻量化部署”的交叉思路
现在的热点已经不仅限于LLM了,像视频生成、图片生成、人物替换这类任务,显存压力只会更大。前段时间很多人讨论“mocha-gguf”这类视频项目的8G显存轻量化部署,本质上走的还是量化 + 低分辨率 + 减少推理颗粒度的路子。GGUF原本是LLM推理的量化格式,现在被越来越多地引入到其他生成任务中,原因就是它把模型压缩之后,能直接塞进小显存卡里。
这类项目里经常能见到8G显存跑视频人物替换的演示,但说实话,那是“能跑”,不是“跑得好”。你在参考这类方案时,要把“轻量部署”和“高质量生成”分开看:前者是一个演示可行性,后者还需要你真正控制显存预算并做大量效果调优。
6. 那些年我们踩过的显存优化坑:经验清单
6.1 “我以为但实际不是”的经典误解
这里我把这几年见到最多的误解列出来,每一条背后都有真实的翻车经历:
- 以为模型参数是显存大头:全参微调下优化器状态往往才是大头,而LoRA下激活值和权重并重。
- 以为batch越大效果越好:盲目增大batch只会OOM,用gradient_accumulation做等效增大,效果一样但显存完全可控。
- 以为gradient checkpointing省显存必须牺牲很多速度:实测一般只多花20%-30%的计算时间,换来的显存下降非常值。
- 以为FP16比BF16省显存:对显存占用来说两者几乎一致(都是2字节),但BF16的数值稳定范围更大,训练Loss不容易飞掉,尤其是低精度量化场景下。
- 以为量化后微调效果会崩:QLoRA在绝大多数指令微调和领域适配任务上,效果损失能控制在很小的范围内,远没有想象中那么大。
6.2 给不同显存用户的最终建议
如果你只有6G显存,别贪心。你的甜点区是7B模型的QLoRA,rank=8,seq_len=512,公式化处理,稳扎稳打。别想着全参,也别想着13B。老老实实把数据集做干净,效果不会差的。
如果你是12G显存,这是当前性价比最高的微调档位。7B模型QLoRA + seq_len=2048 + rank=16 + gradient checkpointing + 8bit优化器,这套配置能覆盖大部分实际任务。如果你想挑战13B,就把seq_len砍到1024,其他不动。
如果你有24G以上,恭喜你,不要浪费这个优势。但还是建议优先LoRA,除非你明确知道自己在做什么。显存是用来换数据吞吐、换上下文长度、换更高秩的,不是用来证明“我能全参”的。
最后说一个我自己很深的体会:微调项目里,最浪费的不是显存,而是时间——反复用一套错误配置反复OOM的时间。真正高效的做法是先用“最小可行配置”把整个链路跑通,再一点点把参数加到甜点区。这个思路,不管你是用Unsloth还是LLaMA-Factory,不管是微调Qwen还是CLIP,都一样适用。先把流程跑通,再谈优化效果,永远是最靠谱的。