最近在尝试微调一个7B的模型,用的LoRA,rank设的8,batch size调到了2,看显存占用才10G左右(显卡是24G的),按理说应该稳的。但每次训练到大概四五百步的时候,突然就Out of Memory了,训练直接崩掉。我查了日志也没看到什么明显错误,就是显存突然飙升。
我怀疑是不是gradient checkpointing没开对,或者中间缓存没清?也试过把batch size降到1,但还是会在不同步数挂掉。有没有大佬遇到过类似情况?是不是LoRA本身在某个阶段会突然占更多显存?还是我用的peft库版本有bug?
先谢过,真有点被搞懵了。
用LoRA微调7B模型,显存够但训练到一半就OOM了,咋回事?
全部回复
共 165 条这情况多半不是LoRA本身,是数据加载或者某个batch特别长导致的,试试固定随机种子看还崩不崩。
我之前跑13B LoRA也踩过一模一样的坑,显存看监控一直稳定,结果一到某个step直接爆掉。你提到gradient checkpointing,我怀疑问题不在它本身,而是它和某些操作组合后触发了动态图内存的峰值,比如当序列长度或者attention的计算图在某个batch里恰好遇到长尾数据时,临时缓存会指数级涨一波,这种情况在7B上很常见。另外peft库的版本确实有坑,旧版在梯度累积时会漏掉释放一些中间变量,建议先升级到最新版试试,同时把optimizer换成AdamW8bit或者Sophia,能省下不少峰值显存。还有个偏方,你可以试着把modeling文件里的past_key_value缓存改成显式清理,或者干脆在每N步后手动调一下torch.cuda.empty_cache(),虽然治标不治本但有时候能扛过去。最省心的方案其实是开accelerate的offload策略,让optimizer状态走CPU,但代价是速度会慢一半,不过总比崩了强。你试着把lr稍微调低点,有时候loss剧烈震荡也会触发异常显存分配,我上次就是降了0.0001就稳了。
我之前跑13B也遇到过一模一样的,后来发现是eval时把eval_dataloader也塞进显存了,加上loss计算时的激活值峰值正好撞一起。你查下是不是每个step都跑eval,或者看看数据里有超长样本,LoRA虽然参数少但激活值跟batch size和序列长度直接挂钩,偶尔来一条长文本就爆了。另外试试最新版peft,之前有个版本确实有缓存没释放的bug,换掉就好了。
八成是序列长度突然变长导致的激活值爆炸,试试max_length锁死加梯度裁剪。
这情况我也踩过,把optimizer换成adafactor或者开offload能稳不少。
我之前也踩过类似的坑,不是LoRA本身的问题,而是PyTorch的缓存分配器在搞鬼。你看到的显存占用10G是“已分配”的,但CUDA context和碎片化预留的显存可能早就超过24G了,特别是训练到中途某个特别大的activation被临时创建时,峰值直接顶爆。
gradient checkpointing确实能缓解,但你得确认它真的作用在了所有子模块上,有些自定义层会绕过这个开关。另外试试torch.cuda.empty_cache()在每N步手动清一下,虽然治标不治本,但能帮你定位是不是碎片化在累积。
还有个更隐蔽的原因——如果用了多进程数据加载,worker进程会复制CUDA上下文,每个worker都占一点显存,几百步之后累积起来就爆了。把num_workers设成0跑一次看看还崩不崩。
peft库版本的话,0.5.0之前有个已知bug,会在反向传播时保存不必要的梯度张量,你升级到最新版或者换到0.6.0试试。我上次就是升了个版本就稳了,rank和batch都没动。
最后建议你开一下CUDA的out-of-memory回溯,设置CUDA_LAUNCH_BLOCKING=1,让报错时能打出具体的分配点,不然光看日志真的很难猜。
我之前跑13B也遇到过一模一样的,不是LoRA本身的问题,大概率是数据加载或loss计算里某个tensor在特定step突然撑爆了,比如变长序列padding没处理好。你可以试试在训练循环里每50步打印一次当前allocated和reserved显存,盯着看到底是哪个操作涨的。另外peft的gradient checkpointing有时候和transformers版本不兼容,建议先单独开/关对比一下,如果还不行就换个旧版peft试试,这坑我踩过。
写得挺好,建议补充一些性能数据。
八成是数据加载或者loss缩放那块的显存峰值,查查是不是eval时把缓存带进去了。
遇到过,多半是loss spikes导致中间激活暴涨,试试gradient checkpointing加max_grad_norm限一下。另外看看是不是eval时开了full batch。
我之前跑7B也遇到过一模一样的情况,后来发现是数据加载那块的锅,某个batch的序列长度特别长,导致激活值爆了,跟LoRA本身关系不大。你可以看看是不是有特别长的样本,或者试试在tokenizer那边加个max length硬截断。另外gradient checkpointing开着的话,理论上显存曲线应该更平缓才对,突然飙升更像是某个操作没被包进去。还有个小技巧,把optimizer换成adamw8bit或者sgd能省不少临时显存,说不定就能撑过去了。
我之前跑7B也遇到过一模一样的症状,不是LoRA本身吃显存,大概率是激活值在某些特殊token序列上突然炸了。你可以试试开torch.utils.checkpoint,再把gradient_accumulation_steps调大点,batch size保持1,这样能稳住峰值。另外检查下是不是dataloader里有padding到固定长度,偶尔有个超长样本会直接把缓存顶爆。peft版本倒不太可能背锅,你先排除数据问题再说。
我之前调13B模型也撞过一模一样的墙,后来发现不是LoRA本身的问题,而是某些batch的序列长度特别长,导致激活值峰值远超平均值。你检查一下数据里有没有特别长的样本,或者试试max length是不是设了固定值,如果没截断干净,到长样本那一两步显存就会瞬间炸掉。
另外gradient checkpointing如果没生效,你可以直接打印一下模型是否真的在重算激活,有时候peft包版本和transformers版本不匹配,会导致checkpointing被静默跳过。我上次就是升级了transformers后,LoRA的forward钩子把checkpointing的缓存逻辑搞乱了。
还有个小技巧,把optimizer改成Adafactor,或者把优化器状态分片打开,能省下不少显存。如果还是不行,你可以在训练循环里手动清一下torch.cuda.empty_cache(),虽然治标不治本,但至少能确认是不是碎片化问题。最后怀疑peft库有bug的话,直接换到最新版,或者用transformers自带的peft集成试试,我遇到过旧版peft在特定步数后缓存不释放的case。
八成是某个batch的序列特别长,activation峰值炸了,查查数据里有没有超长样本。
我之前跑13B也遇到过一模一样的鬼情况,稳定训练十几个小时突然OOM,后来发现是数据加载那边有个缓存没释放,batch size降了只是缓解不是解决。你试试在训练循环里手动清一下torch的cache,或者用max_memory参数限制一下,很多框架默认不限制峰值容易爆。另外peft版本确实坑多,我之前从0.5.0升到0.9.0才解决一个诡异的显存泄漏,建议你直接看下transformers和peft的版本匹配,不兼容的话换套组合可能就好了。
我之前也踩过这个坑,症状一模一样,后来发现是数据加载那边出了问题,某个batch的序列特别长,导致激活值突然爆掉,跟LoRA本身关系不大。你可以看看是不是有异常长的样本,或者试试在collator里做动态padding,把max length卡死。另外gradient checkpointing确实得开,但光开它不够,建议把optimizer state的offload也打开,能省不少。还有个小技巧,就是定期清一下CUDA缓存,torch.cuda.empty_cache()在step循环里隔几百步调一次,有时候能救急。
我之前也遇到过一模一样的坑,也是7B+LoRA,跑着跑着显存曲线突然起飞。后来发现不是LoRA的问题,是DataLoader的num_workers在某个step触发了内存碎片化,加上PyTorch的缓存分配器没及时回收,试试在训练循环里手动调一下torch.cuda.empty_cache(),或者把环境变量PYTORCH_CUDA_ALLOC_CONF设成max_split_size_mb:128,大概率能撑过去。另外你确认下是不是用了最新的peft,之前有个版本在梯度累积时会泄漏激活缓存,回退到0.9.0试试。
我之前跑13B也遇到过这种半路暴毙的OOM,跟显存峰值有关,LoRA虽然省显存但激活值在forward阶段是动态的,可能某个batch的序列长度或者padding触发了峰值。你可以试试开torch.cuda.empty_cache()定期清缓存,或者检查一下是不是tokenizer没设pad token导致不同batch长度差异大。另外peft版本确实偶尔有内存泄漏的issue,建议换个版本或者用transformers自带的多卡策略试试。我当时是改成gradient accumulation加固定max length解决的,你可以参考下。
多半是某个batch的序列特别长,激活值瞬间爆炸,试试max length截断或者按长度排序batch。
我之前也遇到过一模一样的情况,你试试把optimizer的state清理一下或者换用Adafactor,有时候Adam的动量缓存会在某个step突然爆掉。另外检查下dataloader是不是有随机长度的样本,最后几步如果sequence特别长,就算LoRA也会瞬间吃满显存。
另一种风格
这问题八成不是LoRA的锅,而是你某个batch恰好碰上了超长序列,加上activation checkpointing没生效的话,峰值显存直接翻倍。建议你用torch.cuda.max_memory_allocated()打一下峰值,大概率能发现是尾部几个batch在作妖。
这情况我遇到过,多半不是LoRA本身的问题,而是数据加载或者某一batch的序列特别长导致的峰值显存暴涨。你可以查一下是不是有个别样本特别长,或者试试在DataLoader里设个max_length硬截断,顺便开一下gradient checkpointing,显存占用会平缓很多。
另外peft版本确实有过类似bug,建议直接升到最新版再看看。如果还崩,就用torch.cuda.reset_peak_memory_stats()配合max_memory_allocated打点,看是不是某个step突然跳上去的。
我上次是卡在embedding层的梯度累积上,换个方式初始化就稳了,你也可以往这个方向排查下。