最近在尝试微调一个7B的模型,用的LoRA,rank设的8,batch size调到了2,看显存占用才10G左右(显卡是24G的),按理说应该稳的。但每次训练到大概四五百步的时候,突然就Out of Memory了,训练直接崩掉。我查了日志也没看到什么明显错误,就是显存突然飙升。
我怀疑是不是gradient checkpointing没开对,或者中间缓存没清?也试过把batch size降到1,但还是会在不同步数挂掉。有没有大佬遇到过类似情况?是不是LoRA本身在某个阶段会突然占更多显存?还是我用的peft库版本有bug?
先谢过,真有点被搞懵了。
用LoRA微调7B模型,显存够但训练到一半就OOM了,咋回事?
全部回复
共 165 条这情况我也遇到过,折腾了好几天才找到原因。你降到batch size 1还崩,基本可以排除batch size本身的问题。我怀疑是gradient checkpointing配合LoRA时,某些中间激活值没有被正确释放,特别是如果你的模型用了比较长的序列长度,某个时刻累积的缓存会突然爆掉。可以试试显式开启model.gradient_checkpointing_enable(),并且把torch.cuda.empty_cache()放进训练循环里每几十步调用一次,手动清一下缓存。另外检查一下是不是dataloader的num_workers设得太大,有时候多进程加载数据会偷偷占显存,四五百步刚好是数据轮转到新批次的时间点。peft库版本确实有坑,之前0.6.0有过显存泄漏的bug,建议升到0.9.0以上试试。我最后是换成了bitsandbytes的4bit量化+LoRA,显存直接降到6G而且再也没崩过,你可以试试这条路。
这个现象我也遇到过,挺坑的。LoRA本身不会突然涨显存,但问题可能出在梯度检查点上——如果没配合activation checkpointing正确使用,显存会在反向传播时堆积,尤其到几百步后某些中间变量开始累积。你试试把gradient_checkpointing_enabled=True加到模型加载那里,同时确认一下是不是在PeftModel包装前就开启了,顺序搞反了效果会打折扣。另外可以看看是不是数据加载器的问题,比如某些长序列样本在后期才出现,导致单个batch突然变大,虽然你设了batch size但token长度不一致也会爆。还有个小技巧:用torch.cuda.empty_cache()在每步后手动清一下缓存,虽然治标不治本,但能帮你定位是不是缓存泄漏。如果还崩,建议换个peft版本试试,0.6到0.8之间有些版本确实有显存管理bug。
这个问题我也踩过坑,大概率不是LoRA本身的问题,而是训练过程中某些中间变量或者优化器状态在积累。你显存占用10G看起来稳,但可能只是初始状态,随着步数增加,比如attention的key-value cache、梯度累积的中间结果,或者数据加载时的pin memory,都会慢慢把剩余空间填满。我建议你开一下torch.cuda.empty_cache()在每个step后手动清一下,或者用torch.cuda.memory_summary()看看到底哪一步突然暴涨。另外你用的peft版本是0.12.0吗?之前有个版本在合并LoRA权重时确实有内存泄漏的bug,升到最新版试试。还有一个骚操作是把optimizer换成AdamW的8-bit版本,能省出不少显存余量。对了,检查下是不是混合精度训练没开,fp16能大幅降低中途显存峰值。要是还崩,试试把gradient checkpointing的层级调得更细,别只开全局开关。
跟你情况很像,我上次用7B+LoRA也这样,后来发现是peft库在保存中间checkpoint的时候显存会突然暴涨一波。可以试试把save_steps调大,或者手动清一下torch.cuda.empty_cache()。
另外建议看看你的DataLoader是不是num_workers设太高了,有时候卡在数据加载那一步也会偷偷吃显存。
这情况我也碰到过,大概率不是LoRA本身的问题,而是PyTorch的显存分配策略在某个step触发了碎片化或者缓存没释放干净。可以试试在dataloader里加个pin_memory=False,或者手动调一下torch.cuda.empty_cache()的调用频率,有时候gradient checkpointing没配合好forward时的临时张量也会突然爆显存。另外peft最近几个版本确实有显存泄漏的issue,可以回退到0.6.0试试看。
遇到过类似情况,后来发现是PyTorch的CUDA缓存没及时释放,尤其是某些中间变量在反向传播时突然膨胀。建议试试在训练循环里手动加torch.cuda.empty_cache(),或者把gradient_checkpointing打开后配合更大的梯度累积步数。另外peft的某些旧版本确实有内存泄漏问题,更新到最新版说不定能解决。
检查下是不是梯度累积步数设太大了,显存峰值会随着步数积累翻倍。另外试试把optimizer换成Adafactor,能省不少显存。
检查下是不是数据加载到显存没释放,我遇到过类似情况,加个torch.cuda.empty_cache()在每步后试试。
这种情况我也遇到过,跟你配置差不多,24G显存跑7B LoRA,一样会突然爆掉。我后来发现是peft的某些版本在梯度累积时,中间变量不会被及时释放,建议你检查一下transformers和peft的版本,升级到最新试试。另外可以试试把gradient checkpointing显式打开,再用torch.cuda.empty_cache()在每步后清理缓存,至少我这样改完就没再中途崩过了。
这个情况我也遇到过,简直一模一样,24G显存跑7B LoRA按理说确实稳得很,但就是会在某个拐点突然爆掉。我后来排查下来,发现最可能的原因其实是PyTorch的显存分配策略——它会在训练过程中不断缓存一些中间变量,比如优化器状态或者梯度累积的临时buffer,LoRA虽然参数少,但前向传播的中间激活值并没有减少,尤其当序列长度变化或者数据集中有特别长的样本时,那个缓存的显存会突然撑爆。你可以试试在训练脚本里加上torch.cuda.empty_cache()放到每个step之后,或者用环境变量PYTORCH_CUDA_ALLOC_CONF设置max_split_size_mb把小碎片合并一下,我调成128之后就没再崩了。另外peft库确实有几个版本有memory leak的问题,特别是0.6到0.8之间的某些commit,建议你升级到0.11以上看看,他们修了那个gradient checkpointing和LoRA交互时的显存泄漏。你batch size降到1还崩的话,基本排除数据并行问题,大概率是库的bug或者缓存没及时释放。
我遇到过,可能是显存碎片化或者某个中间变量没释放,试试加个torch.cuda.empty_cache()。
哈哈这个问题我太熟了,之前用7B跑LoRA也遇到过一模一样的“定时炸弹”式OOM。显存占用看着平稳,但一到某个步数突然暴涨,八成是中间某个缓存没清理干净,比如activation checkpointing虽然开了,但可能没配合好gradient accumulation的步数,导致中间变量在某个batch积压。你试过把gradient checkpointing的配置换成更激进的版本吗?比如在peft里显式设置enable_gradient_checkpointing(True)后再加个model.gradient_checkpointing_enable(),有时候库的默认实现会漏掉某些模块。另外,也有可能是数据加载时某些样本序列长度不均匀,到特定步数遇到超长文本导致中间张量暴增,可以检查下数据集的token长度分布,或者加个动态padding截断。peft库最近几个版本确实有显存泄漏的issue,尤其是和transformers 4.38+搭配时,建议降到4.36试试,我之前降版本后就没再随机OOM了。还有个小技巧,试试在训练循环里手动调一下torch.cuda.empty_cache(),虽然治标不治本,但能帮你定位是不是缓存累积的问题。实在不行,干脆把LoRA的rank降到4或者8再加一层量化,7B模型用4bit量化后显存压力会小很多,基本不会半路崩了。
这情况我也遇到过,大概率不是LoRA本身的问题,而是PyTorch的缓存没及时释放。你可以试试在训练循环里手动调一下torch.cuda.empty_cache(),或者把gradient checkpointing打开再配合gradient accumulation,这样能平滑显存波动。另外peft库最近几个版本确实有显存泄漏的issue,建议升到最新版或者换个稳定版本试试看。
可能是数据集里某些长文本导致中间变量突然爆了,试试调低max_length或者加个动态padding。
遇到过类似的情况,我怀疑是显存碎片化或者中间激活值缓存没释放干净。建议你试试在训练循环里手动清一下torch.cuda.empty_cache(),或者把gradient checkpointing的配置再检查下,有时候默认设置不一定生效。另外可以看看peft版本,我之前用0.6.0遇到过类似bug,升到0.7.2就稳了。
这种情况我也遇到过,后来发现是PyTorch的缓存分配器在搞鬼,训练到一定步数后累积的中间变量没及时释放,显存会突然暴涨。可以试试在训练循环里手动调一下torch.cuda.empty_cache(),或者把环境变量PYTORCH_CUDA_ALLOC_CONF设成max_split_size_mb:128。另外peft的gradient checkpointing确实容易踩坑,建议检查一下是不是只在模型上开了但没对LoRA层生效。
可能是中间变量没清干净,试试每步结束后手动调下torch.cuda.empty_cache()。
梯度累积设太大了?也可能是中间变量没释放,试试torch.cuda.empty_cache()每步清一下。
试试把batch size设为1、开启gradient checkpointing,然后调低learning rate,可能是梯度累积导致显存峰值。
这种突然爆显存的情况我也遇到过,大概率不是LoRA本身的问题,而是某些中间变量(比如attention的key/value缓存)没有及时释放。你可以试试在trainer里加上remove_unused_columns=False,或者手动调一下gradient_accumulation_steps,有时候步数累积到一定程度缓存会叠上去。另外检查下是不是用了torch.compile,某些版本会有显存泄漏。