最近在尝试微调一个7B的模型,用的LoRA,rank设的8,batch size调到了2,看显存占用才10G左右(显卡是24G的),按理说应该稳的。但每次训练到大概四五百步的时候,突然就Out of Memory了,训练直接崩掉。我查了日志也没看到什么明显错误,就是显存突然飙升。
我怀疑是不是gradient checkpointing没开对,或者中间缓存没清?也试过把batch size降到1,但还是会在不同步数挂掉。有没有大佬遇到过类似情况?是不是LoRA本身在某个阶段会突然占更多显存?还是我用的peft库版本有bug?
先谢过,真有点被搞懵了。
用LoRA微调7B模型,显存够但训练到一半就OOM了,咋回事?
全部回复
共 27 条这问题我去年调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,牺牲点速度换稳定。
我最近刚踩过这个坑,7B模型用LoRA训练到一半OOM确实挺常见的,不一定是显存不够的问题。你提到每次都在四五百步左右崩,我怀疑是PyTorch的缓存分配器在搞鬼——有些中间变量虽然看起来释放了,但显存碎片化严重,累积到某个临界点就突然爆了。你可以试试在训练循环里手动清一下缓存,比如每隔一两百步调用一下torch.cuda.empty_cache(),或者把环境变量PYTORCH_CUDA_ALLOC_CONF设置成max_split_size_mb:128,能减少碎片问题。
另外gradient checkpointing不一定能解决瞬时峰值,LoRA本身在反向传播时可能对某些层产生临时的大张量,特别是如果你的模型用了长序列或者注意力计算没优化的话。我建议你开一下NVIDIA的nvidia-smi监控实时显存变化,看看到底是哪一步突然跳涨。peft库的版本也值得怀疑,我之前用0.6.0就遇到过类似问题,升到0.7.1之后明显稳定了,你可以试试更新。
最后一个小技巧:把optimizer换成AdamW 8bit或者bitsandbytes的优化器,能在相同batch size下省出好几G显存,而且对训练效果影响不大。
这种情况我也遇到过,挺搞心态的。其实不是LoRA本身显存波动,而是当训练到某个step时, optimizer state、中间激活值或者gradient accumulation的缓存会突然累积到峰值,尤其你用的7B模型就算rank低,梯度checkpoint没开对位置也容易爆。可以试试把gradient checkpointing显式加到transformers的model config里,或者把optimizer换成Adafactor省点显存,再不行就开torch.cuda.empty_cache()定期清缓存试试。
可能是gradient checkpointing没生效,试试显式设置model.gradient_checkpointing_enable()看看。
我最近也遇到过类似的情况,而且也是7B模型加LoRA,rank设的8,batch size甚至调到了1,结果还是会在某个点突然崩掉。后来排查发现,问题其实不在显存总量,而是训练过程中某些层的激活值会突然暴增,尤其是当序列长度不均匀或者某些样本特别长的时候,attention机制的缓存会猛吃显存。你可以试试打开gradient checkpointing,但记得要配合torch的显存优化器一起用,比如把max_split_size_mb设小一点,或者开一下torch.backends.cuda.enable_mem_efficient_sdp。另外检查一下你的dataloader里有没有shuffle导致的数据长度突变,如果数据集里混了特别长的样本,LoRA虽然只更新少量参数,但前向传播时完整模型的中间结果还是会撑爆显存。peft库的话,我建议升级到最新版,之前有个版本确实存在缓存未释放的bug,尤其是在多步训练后累积的梯度检查点会残留。如果还不行,可以试试在optimizer里加入gradient accumulation,配合更小的batch size,虽然慢点但能稳定很多。
这种情况我也碰到过,通常不是LoRA本身的问题,而是某个batch里数据长度不均匀,导致计算图临时变大,显存突然就爆了。可以试试把gradient checkpointing打开,同时检查一下数据加载器里有没有设置padding到固定长度,或者用dataloader_collate_fn把长序列截断。另外peft的新版本确实修过一些内存泄漏的bug,不妨升级到最新版试一下,有时候就是版本问题。
这个问题我也遇到过,大概率不是LoRA本身的问题,而是训练过程中某些中间变量(比如attention的key/value cache)没被正确释放,尤其是在长序列或者梯度累积时容易累积显存碎片。建议你检查一下dataloader的num_workers是不是设太高了,有时候多进程加载数据会悄悄吃显存;另外可以试试在优化器step之后手动调一下torch.cuda.empty_cache(),虽然治标不治本但能临时缓解。peft库的版本我也踩过坑,建议升级到最新版或者换个稳定版试试。
这种中途OOM我碰到过,大概率是PyTorch的缓存碎片化问题,不是显存真的不够。你可以在训练循环里手动调一下torch.cuda.empty_cache(),或者试试在peft的配置里把gradient_checkpointing显式设成True,有些版本默认不生效。另外检查下是不是有中间变量没及时释放,比如loss.item()这种操作偶尔也会卡住显存。
我遇到过,试试把gradient checkpointing打开,或者把优化器的状态存到CPU上。
这种情况我也遇到过,挺玄学的,感觉不完全是显存不够的问题。我上次是关了gradient checkpointing,手动加了个torch.cuda.empty_cache()在每个step结尾,反而稳住了。另外可以看看是不是dataloader的num_workers设太高,有时候数据加载线程会偷偷吃显存,把workers设成0试试。
我也遇到过类似的情况,最后发现是PyTorch的缓存分配器在搞鬼——训练过程中会不断累积一些临时张量,虽然单个不大,但到几百步之后碎片化严重,就会突然爆掉。你可以试试在训练循环里每隔一定步数手动调一下torch.cuda.empty_cache(),或者设置环境变量PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,把最大分割块调小一点,减少碎片堆积。另外,gradient checkpointing确实要确认是不是真的生效了,有时候加载模型后没正确配置,相当于没开。还有peft库的版本问题,我上次用0.6.0时也碰到过类似显存泄漏,升级到0.7.1之后就好了,你可以看看是不是版本太老。如果还不行,不妨试试把LoRA的rank降到4,或者用更小的投影层,虽然效果会弱一点,但至少能跑完。