尧图精选

32GB显存不够?LoRA/QLoRA微调大模型OOM自救指南

🕒 发布时间:2026/9/13 8:24:37 📁 来源:尧图网络
如果你手里刚好有一张32GB显存的卡V100 32GB、A100 40GB这种级别想微调一个7B或者13B的大模型第一次跑训练脚本就撞上CUDA out of memory这几乎是每个人都要经历的一关。我不止一次在群里看到有人把batch size调到1之后还在OOM然后绝望地问还有救吗有救。这篇东西就围绕LoRA/QLoRA微调这个话题讲清楚显存到底花在哪、怎么用32GB把微调跑起来、以及在OOM现场怎么一步一步自救。不管你是给学生做演示、自己研究模型能力还是给公司内部业务做指令微调只要你的GPU显存不是特别宽裕这篇文章里的东西都能直接用。1. 32GB显存为什么也会OOM先算一笔显存账1.1 全参数微调的显存去向很多人拿到卡第一反应是32GB这么大塞个7B模型权重才14GBFP16怎么跑个训练就爆了这是因为训练和推理完全不是一回事。推理只需要模型权重加激活值训练还要额外存梯度、优化器状态和中间激活开销直接翻好几倍。具体拆开看训练时显存主要被四类东西吃掉模型权重本身、反向传播需要的梯度、优化器状态Adam里的一阶动量m和二阶动量v、前向过程中产生的激活值activations。这里面最容易被忽略的就是优化器状态。用AdamW做全参微调时每个参数要额外存两个FP32的动量值也就是8字节。7B模型光这一步就是56GB加上14GB权重、14GB梯度还没算激活值就已经84GB了。这时候再回头看32GB显存OOM几乎是必然的。1.2 一份可复用的显存估算公式我习惯用一个粗略公式来预估全参微调的显存下限显存 ≈ 权重(FP162字节) 梯度(FP162字节) Adam状态(FP32×28字节) 激活值拿7B模型算每参数12字节约84GB。13B模型就是156GB。33B更不用想。所以32GB显存跑全参数微调基本是死刑除非用极其激进的多卡并行加Offload方案但那已经超出轻量微调的范畴了。1.3 所以LoRA/QLoRA解决的是哪一层LoRA的思路很简单不训练全部参数只在一个低秩空间里学一个增量。相当于原来要重新写整本教材现在只在原教材旁边加几页注释。这样需要计算梯度和优化器状态的参数量从几十亿降到几十万上百万省掉的是梯度加优化器这两块大头。而QLoRA更进一步把冻结的原始权重压到4bit存储省掉的是权重这块大头。两头一省32GB就跑得动了。理解了这个逻辑后面所有的配置选择就都围绕一件事哪块显存能省、哪块不能省。2. LoRA与QLoRA的显存节省逻辑2.1 LoRA冻结主干、只练增量LoRA全称Low-Rank Adaptation核心假设是模型微调时权重的更新量可以是低秩的。给定原始权重W维度d×k训练时不再更新W而是引入两个小矩阵Ad×r和Br×k其中r远小于d和k前向时计算h Wx BAx。训练时W冻结只更新A、B。这样做的显存变化非常直观模型权重还在显存里占14GB但梯度和优化器状态只需要对应A、B这两个小矩阵。拿Llama 7B举例如果只把attention的q、v做LoRAr8时可训练参数量大约400万只占全部参数的0.5%左右优化器状态从56GB降到几十MB级别。2.2 QLoRA把权重压到4bitQLoRAQ代表Quantization的核心改动是把冻结的原始模型权重用4bit量化存储计算时反量化回高精度做前向和反向。7B模型原来FP16要14GBNF4量化后只需要3.5GB左右省下来的空间非常可观。这里有个关键概念NF4NormalFloat4。它不是简单的四舍五入而是基于数据分布做分位数量化让量化误差尽量小。简单理解把权重的分布切成16个区间每个区间用一个4bit的码表示取值更密集的地方划分更细。相比FP4这种均匀量化NF4在LLM上的效果更稳。QLoRA论文里还有一个细节值得注意输入给模型的数据并不会全程保持4bit而是先把4bit权重反量化成BF16再用BF16做计算。所以QLoRA省的是存量计算过程仍然需要较高精度的临时变量这也是为什么显存虽然比LoRA少但还是需要一定余量。注意如果GPU不支持BF16比如V100要显式把compute_dtype设成fp16否则会遇到类型相关的问题。2.3 Double Quantization与Paged OptimizerQLoRA还做了两个容易被忽略的优化。第一个是Double Quantization双重量化量化时需要一个缩放常数scale每个block一个这个常数本身也用8bit再量化一次省下的显存看起来不大但以block为粒度累计下来一个7B模型能省几百MB。第二个是Paged Optimizer分页优化器。它利用CUDA的显存分页机制当优化器状态瞬间超过显存容量时自动换页到CPU内存避免直接OOM。这个机制不需要你主动干预只要优化器选paged_adamw_8bit或paged_adamw_32bit就能生效。它在训练初期损失值突然跳高、显存曲线出现尖峰时特别有用。2.4 三种方案显存对比方案7B模型13B模型33B模型全参数FP16约84GB约156GB约400GBLoRAFP16基座约16~20GB约28~34GB约70GBQLoRANF4基座约6~10GB约10~15GB约22~30GB注意表格里的QLoRA是训练中的实际占用加上了激活值和优化器状态且假设开了gradient checkpointing、batch size比较小。这个量级意味着32GB显卡7B和13B随便跑33B卡着边也能跑70B就没戏了那需要更大的卡或者换其他方案。3. 32GB GPU显存预算与训练配置3.1 不同规模模型能训到哪个量级基于上面的估算32GB显存的实操底线大概是这样7B模型QLoRA随便训甚至可以塞下batch_size4、seq_len2048。13B模型QLoRA是甜点位batch_size1加梯度累积稳。33B模型QLoRA能跑但要把seq_len压到1024以内、开gradient checkpointing、打开paged optimizer整体局促但没问题。70B32GB不够除非用CPU offload到极限或者多卡但那已经不是轻量方案的范畴。这里我要强调一点如果你用的是V100 32GB而不是A100注意V100不支持BF16只能用FP16。FP16在数值稳定性上更娇气建议调低学习率并且打开loss scaling。如果用BF16友好型显卡如A100、RTX 40系列训练过程会省心很多。3.2 LoRA关键参数怎么选先说r秩。r越大可训练参数越多、表达能力越强但显存和计算压力也越大。QLoRA论文里实验用了r64但我个人经验是业务场景先r8开始实验跑通再往上试。r8对应7B模型大概0.1%的可训练参数量对大部分指令微调任务已经够用。再说alpha缩放参数。常见做法是alpha 2 * r也就是r8时alpha16。这个比例不是死的但2:1是社区里用得最普遍、最不容易出问题的起点。最后说target_modules作用目标。Llama结构里常见的候选是q_proj、k_proj、v_proj、o_proj、gate_proj、up_proj、down_proj。QLoRA原论文只改了attention的q和v实测效果已经不错但如果你想在中文数据上做更强的能力迁移把全部linear都纳入会更稳代价是可训练参数从0.1%涨到0.5%左右显存多几百MB可以接受。3.3 训练超参与显存的trade-off微调时最常调的几个参数每一个都和显存强相关per_device_train_batch_size这是显存消耗的第一杀手。32GB跑7B QLoRA时batch_size从1涨到4显存涨得肉眼可见。gradient_accumulation_steps不改变单步显存但等效扩大batch。batch_size1、accumulation16就是等效batch 16。seq_lenmax_seq_length激活值显存和序列长度的平方成正比attention部分从2048降到1024省下的显存非常可观。gradient_checkpointing用30%左右的训练速度换取大幅降低激活值显存近两年微调基本必开。优化器paged_adamw_8bit比普通adamw省一半以上优化器状态还带分页兜底。我常用的一个稳妥起步配置是batch_size1、gradient_accumulation16、seq_len2048、gradient_checkpointingTrue、optimpaged_adamw_8bit、学习率2e-4。这个配置在7B和13B上都能跑起来然后根据显存余量再逐步放开batch_size或seq_len。4. QLoRA微调完整实操从环境到代码4.1 环境与依赖安装QLoRA微调最常用的组合是transformers peft bitsandbytes accelerate trl。安装本身不复杂但版本坑很多我给一个我实测可用的组合transformers4.35 peft0.7 bitsandbytes0.41 accelerate0.25 trl0.7transformers版本太老会不支持NF4量化bitsandbytes在Windows上早期版本有兼容性问题建议用0.41以上。安装命令就是常规的pip install有CUDA环境的话bitsandbytes会自动识别驱动版本。注意bitsandbytes依赖CUDA运行时如果你在服务器上用conda环境先确认nvidia-smi里的驱动版本足够新再确认PyTorch自带的CUDA版本和bitsandbytes的预期一致。最省心的做法是直接用官方PyTorch镜像里的CUDA版本环境。4.2 加载4bit模型加载的核心是BitsAndBytesConfig直接把代码放出来import torch from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig bnb_config BitsAndBytesConfig( load_in_4bitTrue, bnb_4bit_quant_typenf4, bnb_4bit_use_double_quantTrue, bnb_4bit_compute_dtypetorch.bfloat16, # V100等老卡改成torch.float16 ) model AutoModelForCausalLM.from_pretrained( meta-llama/Llama-2-7b-hf, quantization_configbnb_config, device_mapauto, trust_remote_codeTrue, ) tokenizer AutoTokenizer.from_pretrained(meta-llama/Llama-2-7b-hf) tokenizer.pad_token tokenizer.eos_token这里几个参数挨个说load_in_4bitTrue打开4bit加载。bnb_4bit_quant_type选nf4而不是fp4效果更好。bnb_4bit_use_double_quant省几百MB显存建议打开。bnb_4bit_compute_dtype计算精度BF16友好型显卡用bfloat16老卡用float16。device_mapauto让库自动分配层到GPU/CPUQLoRA场景一般全在GPU里。4.3 训练代码与参数加载完模型后需要先调用prepare_model_for_kbit_training它会帮我们处理4bit模型训练需要的冻结、混合精度等设置。然后挂上LoraConfigfrom peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training model prepare_model_for_kbit_training(model, use_gradient_checkpointingTrue) lora_config LoraConfig( r8, lora_alpha16, lora_dropout0.05, biasnone, task_typeCAUSAL_LM, target_modules[q_proj, v_proj, k_proj, o_proj, gate_proj, up_proj, down_proj], ) model get_peft_model(model, lora_config) model.print_trainable_parameters()print_trainable_parameters会打印可训练参数量和占比建议第一次跑的时候留意一下确认LoRA真的挂上了。接着是TrainingArgumentsfrom transformers import TrainingArguments training_args TrainingArguments( output_dir./lora-7b-zh, per_device_train_batch_size1, gradient_accumulation_steps16, gradient_checkpointingTrue, learning_rate2e-4, bf16True, max_grad_norm0.3, optimpaged_adamw_8bit, num_train_epochs3, logging_steps10, save_strategysteps, save_steps500, warmup_ratio0.03, lr_scheduler_typecosine, report_tonone, )这几个参数里learning_rate2e-4是QLoRA团队在论文里验证过的4bit微调常用值比普通LoRA的1e-4要大一档。max_grad_norm0.3也是从论文里抄的防止梯度爆炸。bf16True是为了在支持BF16的卡上更稳老卡改成fp16True。数据部分建议直接用trl的SFTTrainer它对序列长度截断、指令对话格式处理得更省心显存控制也比较友好代码可以省不少from trl import SFTTrainer trainer SFTTrainer( modelmodel, argstraining_args, train_datasetdataset, tokenizertokenizer, max_seq_length2048, dataset_text_fieldtext, packingFalse, ) trainer.train()4.4 显存监控技巧训练中除了用nvidia-smi我还会在训练里主动打印PyTorch侧的显存统计。原因是nvidia-smi显示的是进程占用的总显存包括缓存而PyTorch的allocated才是实际在用的张量两者差值越大说明缓存碎片越严重。import torch def print_gpu_memory(tag): allocated torch.cuda.memory_allocated() / 1024**3 reserved torch.cuda.memory_reserved() / 1024**3 max_allocated torch.cuda.max_memory_allocated() / 1024**3 print(f[{tag}] allocated{allocated:.2f}GB, reserved{reserved:.2f}GB, max{max_allocated:.2f}GB)训练到一定步数时调用一次能清楚看到显存占用峰值在哪个阶段出现。实测QLoRA训7B时显存峰值通常不出现在stepping开始而出现在前向计算长序列的时候也就是activation peaks。这也是为什么压seq_len比压batch_size更有效的原因。5. OOM现场排查与优化手段5.1 OOM报错常见形态第一个是大家最熟悉的RuntimeError: CUDA out of memory. Tried to allocate 200.00 MiB括号里会告诉你GPU总容量、已分配、剩余、保留了多少。关键看already allocated和reserved如果reserved很大而allocated不大说明PyTorch的缓存机制占了很多显存但没有及时还给驱动这时可以用torch.cuda.empty_cache()或者调小缓存分配策略。第二个是CUDA error: device-side assert triggered这个不是显存不够通常是数据问题或者某个张量出现了不合法值却经常被误判成OOM。看到这个报错先去检查数据标签范围、形状对不对别急着减显存。第三个是进程直接被杀掉比如Killed或者Out of memory出现在dmesg里。这种通常是系统级别的OOM不只是显存而是CPU内存或进程内存超限要用CPU内存Offload时特别容易出现。这种情况的排查思路是看系统内存大小别让主机内存也爆了。5.2 显存优化手段按优先级排序我按实战性价比从高到低排一下开gradient_checkpointing。一行代码激活值显存立减代价是慢30%左右性价比最高。压seq_len。attention的激活值和序列长度近似平方关系2048降到1024省下的显存比改batch_size多得多。前提是训练数据的有效信息长度能承载。换paged_adamw_8bit优化器。省优化器状态还有分页保底。减小batch_size用gradient_accumulation补。这招最直接但batch_size1是极限不能无限减。检查padding。数据加载时如果没做好padding mask会在无效token上白白消耗显存长尾序列尤其明显。打开flash_attention_2。新版transformers直接传attn_implementationflash_attention_2就行能减少attention部分的激活显存还更快前提是GPU支持A100、4090等Ampere以上架构都支持V100不行。5.3 实战一次OOM到收敛的调参记录说个我最近的例子。手头一台机器是V100 32GB要在这台机上微调一个13B模型做中文指令数据。第一次全参微调自然秒OOM这不算意外。我改成QLoRA后配置是batch_size2、seq_len2048、gradient_checkpointing开、NF4。启动后大约第3个step就OOM。当时的处理顺序第一步把per_device_train_batch_size从2降到1显存确实降了不少但峰值还是离32GB很近尤其到长样本时仍然会爆。第二步把max_seq_length从2048降到1536准确说是按整个训练集统计了样本长度分布后发现95%的样本在1200 token以内压到1536完全不损失信息。改完训练稳定了。第三步再把optim换成paged_adamw_8bit并把数据里的padding统一截断到最短可用的batch长度不pad到固定长度最后峰值稳定在26GB左右留了6GB余量做推理缓存备用。还有个细节V100不支持bf16所以我把compute_dtype和训练精度都设为fp16同时把learning_rate从2e-4微调到1.5e-4因为fp16的数值范围更窄学习率太高容易出NaN。这个组合最终把loss跑到了正常收敛区间。6. 常见问题速查与避坑6.1 环境与版本问题问加载模型报bitsandbytes requires... CUDA版本不匹配怎么办答先看bitsandbytes报错里要求的CUDA版本再对比你环境里PyTorch自带的CUDA版本。大多数情况是conda环境的CUDA和bitsandbytes预期不一致。建议用pip重新安装匹配版本的bitsandbytes或者干脆重建一个环境用官方推荐的torch镜像能省很多时间。问Windows上bitsandbytes装不上或者训练时卡死答Windows下bitsandbytes在0.40之前有较大兼容性问题尽量用0.41以上。另外很多Windows上的问题来自显卡驱动和WSL混用建议要么纯Windows环境要么纯WSL2环境不要两个混着来。6.2 训练效果问题问LoRA微调后loss降了但生成质量没提升怎么回事答常见原因有三个一是数据量太小或格式混乱lora记住的是格式而不是能力二是r太小导致容量不足尤其是领域术语密度高的数据三是只微调了q、v而没有动MLP对知识型任务来说信息注入有限。可以尝试target_modules全量覆盖、适当提高r。问4bit量化微调后效果比全参差很多答正常。QLoRA在大多数公开基准上和全参微调的差距已经很小但在小样本、高领域密度任务上仍然会吃亏。如果试过r64仍然不满意可以考虑用8bitload_in_8bit或fp16基座LoRA显存占用高一截但效果更接近全参。6.3 微调后的保存与部署LoRA训练完保存的是增量权重不是完整模型。保存和后续使用要记得合并或加载回基座# 保存 model.save_pretrained(./lora-7b-zh) tokenizer.save_pretrained(./lora-7b-zh) # 推理时加载 from peft import PeftModel base_model AutoModelForCausalLM.from_pretrained(基座模型路径) model PeftModel.from_pretrained(base_model, ./lora-7b-zh) # 如果想导出为完整模型 merged_model model.merge_and_unload() merged_model.save_pretrained(./merged-7b-zh)merge_and_unload会把LoRA增量合并回原始权重导出后就是一个普通模型可以直接用transformers加载或者转其他格式部署。但注意合并后会丢失4bit量化带来的省显存优势如果你的部署环境显存也紧张建议保留LoRA增量文件推理时在原位加载。关于推理显存测算很多人问我这张卡能跑多大模型。推理相对训练简单很多模型权重按精度来加KV cache加输入输出激活。一个粗略经验是7B FP16推理大约需要14~16GB4bit量化后大约6~8GB13B FP16大约26GB左右4bit后大约10~13GB。这个可以作为部署选型参考。7. 一些更进阶的优化方向到这里32GB基本能跑通7B/13B的QLoRA微调了。如果你想在同样的显存下挑战更大的模型或者想把速度再提一提还有几个方向可以试PEFT DeepSpeed ZeRO-3把优化器状态和梯度分片到多卡单卡显存压力更小但配置复杂度高不少。使用FlashAttention前提是Ampere架构以上能把attention的激活占用从平方级降下来。用Unsloth等优化库它们对LoRA训练做了算子级优化显存占用能进一步下降训练速度也有提升接口基本兼容transformers值得零成本试一下。数据侧动手脚清洗数据、合理截断、packing不padding这些在长序列场景下比任何显存工具都管用。说实话微调这个事方案比蛮力重要。我自己最早也在32GB卡上试图全参微调7B撞了无数OOM后才彻底理解显存是一笔账每一层都要算清楚这句话。后来用QLoRA从7B到13B再到33B反而越用越顺就是因为账算明白了配置手段也能一条一条对号入座。最后分享一个实用小技巧如果某一步突然OOM不要急着把整个配置推倒重来先把batch_size减半、seq_len减半再逐步加回来找到临界点。这个二分法比凭感觉乱调高效得多而且每次改动都记录一下峰值显存几次之后你对自己模型的显存习性就门儿清了。
上一篇/下一篇内容由系统自动关联 返回资讯列表 →