最近在尝试微调一个7B的模型,用的LoRA,rank设的8,batch size调到了2,看显存占用才10G左右(显卡是24G的),按理说应该稳的。但每次训练到大概四五百步的时候,突然就Out of Memory了,训练直接崩掉。我查了日志也没看到什么明显错误,就是显存突然飙升。
我怀疑是不是gradient checkpointing没开对,或者中间缓存没清?也试过把batch size降到1,但还是会在不同步数挂掉。有没有大佬遇到过类似情况?是不是LoRA本身在某个阶段会突然占更多显存?还是我用的peft库版本有bug?
先谢过,真有点被搞懵了。
用LoRA微调7B模型,显存够但训练到一半就OOM了,咋回事?
全部回复
共 165 条这问题我去年调Qwen-7B的时候也踩过坑,说下我的排查思路供参考。
首先,LoRA本身确实不会在训练中途突然涨显存,但有两个常见元凶:一是transformers库的gradient checkpointing实现有bug,特别是配合peft的旧版本时,某些算子(比如flash attention的兼容层)会在反向传播时产生临时缓存膨胀。你检查下peft和transformers的版本,如果是0.6.x之前的peft,建议升到0.9.0以上,同时确保gradient checkpointing是在model.enable_input_require_grads()之后才启用的。
另一个更隐蔽的点:你的数据集是不是在某个batch里突然出现了特别长的序列?LoRA虽然只训练两个低秩矩阵,但前向推理时完整权重还是要加载的。如果数据中有padding长度不均匀的情况,即使设置了batch size=1,单条序列超过2048 token也可能导致显存尖峰。建议在dataloader里加个max_length截断,或者用dynamic padding统一填充长度。
另外检查下是不是用了torch.compile或者混合精度训练?AMP的grad_scaler在某些场景下会累积显存碎片,跑几百步后碎片化严重导致申请连续内存失败。可以尝试在optimizer.step()后手动清缓存:torch.cuda.empty_cache(),或者改用bitsandbytes的8-bit Adam优化器,它对显存碎片更友好。
如果还崩,建议开NVIDIA的nsight system跑一下,看显存到底是哪个算子突然飙起来的。我上次查出来是peft的lora_dropout在前向时分配了一个临时tensor没释放,换用dropout=0.0或者用peft的r=16但alpha=32来缓解。
我之前也碰到过类似的情况,跑着跑着显存突然炸掉,后来发现是PyTorch的缓存分配器在搞鬼,尤其是用了gradient checkpointing之后,某些中间变量没及时释放。你试试在训练循环里手动调一下torch.cuda.empty_cache(),或者把环境变量PYTORCH_CUDA_ALLOC_CONF设成max_split_size_mb:128,有时候能缓解。另外peft的版本确实有坑,我之前0.5.0遇到过内存泄漏,升级到0.7.0就好了,你也可以看看是不是这个问题。
我最近也碰到过类似的情况,后来发现是梯度累积和优化器状态在长时间训练后会累计缓存碎片,导致显存突然炸掉。你可以试试在训练脚本里加上torch.cuda.empty_cache()定期清理一下,或者把优化器的betas参数调一调,有时候能缓解。另外peft的0.6版本有个已知的显存泄漏问题,建议升级到最新版或者换个分支试试。
试试把gradient accumulation steps设小点,或者换个peft版本,我之前也是类似问题。
我遇到过一模一样的情况,也是7B加LoRA,显存够但跑到几百步就炸。后来发现是PyTorch的缓存没及时释放,试试在训练循环里加个torch.cuda.empty_cache(),或者检查一下dataloader的num_workers是不是设太高了。另外peft库的版本确实有坑,0.6.0之前有个显存泄漏的bug,升级到最新版可能就解决了。你用的什么优化器?AdamW有时候会额外吃显存,换Adam8bit或者Adafactor能省不少。
这情况我也碰到过,大概率不是LoRA本身的问题,而是某些中间变量或者优化器状态在累积。可以试试开启gradient checkpointing,同时把optimizer换成Adafactor,它对显存占用更友好。另外检查下是不是数据加载时缓存没释放,用torch.cuda.empty_cache()手动清一下试试。
这个情况我也遇到过,挺玄学的。LoRA本身虽然省显存,但如果你用的peft版本比较老,某些实现里在优化器状态或中间激活值上会有内存泄漏问题,建议先升级到最新版试试。另外gradient checkpointing确实得确认一下是不是正确开启了,有时候你开了但模型某些层没被包进去,到特定步数显存就会突然炸掉。还有一个可能性是数据加载那边,比如某个batch的序列长度特别长,导致计算图临时膨胀,你可以试试给tokenizer加个max_length硬截断。你提到batch size降到1还是会崩,那大概率是框架层面的缓存没清,像transformers的cache或者xformers的memory efficient attention有时候会累积碎片,可以加个torch.cuda.empty_cache()在每个step后手动清一下。最后实在不行,试试把LoRA的rank降到4或者用更小的投影维度,虽然效果会打折扣,但至少能跑完。
这情况我也遇到过,大概率不是LoRA本身的问题,而是训练过程中某些中间变量(比如attention的key/value cache)在长期运行后没释放干净。建议试试在训练循环里显式调用torch.cuda.empty_cache(),或者检查一下dataloader是不是有内存泄漏,batch size降到1还崩说明不是单纯显存不够。另外peft的版本确实坑多,可以试试回退到0.6.0或者升级到最新版,有时候是gradient checkpointing和某些层的兼容性导致显存突然爆炸。
这种OOM中途炸掉的情况我也遇到过,大概率不是LoRA本身的问题,而是某个batch的输入序列特别长导致显存峰值暴涨。你可以检查下数据里有没有异常长的文本,或者试试把max_length设小一点,同时开启gradient_checkpointing并配合torch的显存优化器。另外peft库0.6之后的版本有修复一些缓存泄漏的问题,可以升级试试。
这情况我也遇到过,挺玄学的。24G显存跑7B+LoRA rank8按理说确实应该够,但问题可能出在训练过程中的中间变量上。比如某些激活值或者优化器状态在反向传播时会突然膨胀,尤其是如果你的序列长度不固定或者数据里有特别长的样本,显存峰值会比平均占用高出一大截。建议你检查一下是不是用了dynamic padding或者某个batch里混进了长文本,这会导致计算图在特定步数突然变大。gradient checkpointing确实能缓解,但要确认是不是只对前向有效,反向时有些库的缓存策略还是会把中间结果保留下来。另外,peft库的版本和transformers的兼容性也容易出幺蛾子,我之前升级到最新版后反而更稳定了。还有个偏方:试试在训练脚本里手动清一下torch.cuda.empty_cache(),虽然不是根本解决,但有时能撑过峰值。如果还不行,可以考虑把LoRA的rank降到4或者用梯度累积来等效减小batch size,牺牲点速度换稳定。
这种情况我也遇到过,后来发现是PyTorch的缓存分配器在搞鬼,训练过程中显存碎片越积越多,到某个临界点就炸了。你可以试试在训练循环里手动调一下torch.cuda.empty_cache(),或者设置环境变量PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128。另外PEFT的gradient checkpointing确实有些版本有内存泄漏问题,建议升级到最新版或者切到transformers原生的实现看看。
我也是用LoRA微调7B遇到过类似情况,后来发现是PyTorch的缓存分配器在训练过程中没及时释放中间变量,导致某个checkpoint后显存突然爆掉。可以试试在训练循环里显式调用torch.cuda.empty_cache(),或者调低gradient_accumulation_steps让梯度累积更均匀。另外peft库的版本确实有坑,我之前升到0.10.0之后问题就少了很多,你也可以检查下是不是某个自定义dataset加载时没做好缓存清理。
我最近也遇到过类似情况,后来发现是PyTorch的缓存分配器在作祟——显存占用虽然显示只有10G,但实际预留的缓存块会随着训练步数累积,到某个临界点突然爆掉。试着手动调一下torch.cuda.empty_cache()的频率,或者在训练循环里加个gc.collect()看看能不能稳住。另外PEFT的LoRA实现确实在梯度更新时会临时分配额外显存,建议把use_reentrant参数显式设为True试试。
这个问题我之前也踩过坑,大概率不是LoRA本身的问题,而是PyTorch的缓存分配器在搞鬼——训练到一定步数后,某些中间变量的显存没被及时释放,累积到某个阈值就爆了。建议你在训练循环里手动加个torch.cuda.empty_cache(),或者试试把optimizer的betas调低一点,有时候AdamW的动量缓存也会突然膨胀。另外peft 0.7.0之前有个版本确实有梯度累积时的显存泄露bug,更新一下库版本可能直接解决。
我之前也遇到过类似的情况,排查下来发现是gradient checkpointing和LoRA的某些实现细节冲突了,导致中间变量没有被正确释放。你可以试试把gradient_checkpointing设成False跑一小段看看是否还崩,如果稳定了就是这个问题。另外peft库0.6.0左右有个版本在反向传播时会有显存泄漏,建议升级到最新版试试。
这种突然OOM我碰到过,大概率不是batch size的问题,而是PyTorch的显存碎片化或者某些中间变量没释放干净。你可以试试在训练循环里手动加个torch.cuda.empty_cache(),或者用gradient checkpointing的时候,检查一下model.enable_input_require_grads()是不是正确调用了。另外,peft库0.6版本以后对LoRA的显存管理有优化,看看你版本是不是太旧了。
试试把gradient accumulation拆成两步,或者换下optimizer的缓存策略,有些库的默认配置会搞事情。
我也遇到过类似情况,后来发现是PyTorch的缓存没及时释放,特别是LoRA的adapters权重在反向传播时会临时分配额外显存。你可以在每个step后手动调一下torch.cuda.empty_cache()试试,或者试试把gradient_checkpointing设为True的同时再开个optimizer的offload。另外peft的0.6.0版本有个显存泄漏的bug,更新到最新版也可能管用。
我之前也遇到过类似的情况,不一定只是显存问题,有时候是PyTorch的缓存没及时释放,可以试试在训练循环里手动调一下torch.cuda.empty_cache()。另外PEFT库的几个版本确实有内存泄漏的bug,建议更新到最新版或者换个稳定版本看看。还有你gradient checkpointing开了没?如果没开的话加上能省不少显存,尤其LoRA在某些层计算时会突然爆一波。
你这情况我遇到过好几次,基本可以锁定是gradient checkpointing和中间变量缓存的问题。LoRA本身不会突然爆显存,但PyTorch的autograd在某些操作(比如attention计算)上会保留中间激活值,哪怕你开了checkpointing,如果代码里有用到一些自定义的forward钩子或者peft的某些版本对特定层没正确处理,步数一多缓存就炸了。建议你先用torch.cuda.empty_cache()在每步之后手动清一下,或者试试把gradient_accumulation_steps设成2,这样等效batch size不变但实际显存波动会更平滑。另外peft的0.6版本有个已知bug,LoRA的dropout在训练中期会莫名其妙分配额外显存,更新到0.7或者直接换diffusers的LoRA实现试试。如果还崩,可以在训练脚本里加个显存监控,看是哪一步突然跳的,大概率是某个层的输出尺寸没对齐导致中间张量膨胀。我之前调7B时也卡过,最后发现是数据加载时padding不一致造成的,你可以检查下是不是某个batch的序列长度突然变长了。