最近在尝试微调一个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突然爆了,你试下把max_length调短点,或者用torch.cuda.empty_cache()在每步后清一下。另外peft版本确实有过显存泄漏的bug,建议升到0.9以上或者换transformers自带的那种PEFT集成试试。还有个小技巧,把optimizer换成adafactor能省不少显存,说不定就苟过去了。
十有八九是序列长度引起的激活值波动,试试固定max length并开gradient checkpointing,能压不少显存。
我之前跑13B的LoRA也撞过一模一样的墙,后来发现不是显存总量的问题,而是某个特定算子或者中间激活值在长序列上突然爆了。你试试把max sequence length缩到原来的一半,如果提前崩或者不崩了,那基本就是序列长度在某些batch里特别长,导致注意力矩阵那一下冲上天。gradient checkpointing确实能压峰值,但如果你用的是peft比较新的版本,它和某些transformers版本会有点兼容问题,建议锁一下peft的版本,或者换成transformers自带的PEFT集成试试。另一个隐蔽点是,虽然batch size是2,但如果你在DataLoader里没设drop_last,最后那个不完整的batch可能会因为padding差异导致计算图异常膨胀。你可以开一下显存日志,看崩之前是不是正好在换到下一个epoch或者某个特定step,有时候是优化器状态在累积到一定步数后突然翻倍,比如AdamW的二阶动量。我之前还遇到过CPU内存碎片导致swap到显存的情况,那就更玄学了,直接把dataloader的num_workers调成0试试。总之别只盯显存占用曲线,把torch.cuda.max_memory_allocated打出来看峰值,比实时占用准得多。
跑长点试试开max_grad_norm截断,或者盯一下loss是不是突然炸了触发重计算。
可能是某个batch的序列特别长,把激活值顶爆了,试试固定max_length。
我之前跑13B也碰到过这种半路暴毙的OOM,后来发现是数据里有个别超长序列在某个step把activation撑爆了,跟LoRA本身关系不大。你可以试试按序列长度做个bucket或者直接过滤掉过长的样本,或者开一下torch的显存缓存清理。另外peft版本确实有坑,建议换到最新版或退一个稳定版试试,之前我就因为版本问题踩过类似莫名其妙的显存泄漏。
我之前跑13B也遇过一模一样的,不是LoRA的问题,大概率是某个batch里数据长度特别长,导致activation峰值暴涨。你可以试试打开gradient checkpointing的同时,把eval和logging的step设大点,有时候验证集跑一次也会把缓存堆起来。另外看下是不是多卡环境,虽然你单卡但偶尔会有人没设好分布式导致临时buffer叠加。我后来把max_length从2048砍到1024就再没崩过,你可以先拿这个验证下是不是长度触发的峰值。
我之前也踩过这个坑,后来发现是eval时显存没释放,LoRA训练中如果设置了定期验证,评估阶段会额外占一波显存,崩的时间点刚好卡在验证前后。你可以在训练循环里手动清一下cache,或者把eval关掉试试。另外peft版本有时候和transformers版本不匹配也会出这种玄学问题,建议直接换个组合环境,比如peft 0.12+transformers 4.44,我换了之后就没再炸过。还有个小细节,如果你的数据集里某个batch长度特别长,即使batch size=1也可能触发峰值,可以看看是不是数据填充没做均匀。
我之前跑13B也遇到过一模一样的,不是LoRA的问题,八成是激活值峰值或者某个特定batch的序列长度特别长导致的瞬时显存暴涨。你可以试试开一下torch.cuda.empty_cache()定时清缓存,或者用gradient accumulation把batch size压到1但累积步数拉长,这样峰值会小很多。另外peft版本建议升到最新,之前确实有显存泄漏的issue,更新后就好了。
我之前也遇到过一模一样的情况,7B+LoRA跑着跑着突然爆显存,当时排查了半天发现是数据加载那边出了问题,某个batch的seq length特别长,把激活值撑爆了。你可以看看是不是数据里有长度异常长的样本,或者试试在collator里做个按长度排序的bucket,能缓解不少。另外peft版本确实有过显存泄漏的bug,建议升到最新版或者换个分支试试,有时候就是这种玄学问题。
我之前也遇到过一模一样的,不是LoRA的问题,大概率是训练过程中某个batch的序列长度突然变长,导致激活值峰值暴涨,你可以看看是不是数据里有超长样本。另外试试把gradient checkpointing打开,再配合显存碎片清理(比如torch.cuda.empty_cache()),一般能缓解。还有个小技巧,把优化器换成AdamW的8-bit版,能省不少显存,峰值会更平缓。
gradient checkpointing没开对倒是次要的,这个现象更像是在某个step触发了动态计算图的峰值——比如当序列长度或者注意力掩码因为padding变化时,激活值会突然翻倍。你试试把gradient_checkpointing设成True,同时把optimizer换成AdamW + paged_optimizer(bitsandbytes那个),它能把优化器状态临时卸到CPU上。另外peft的lora_dropout如果设了非零值,训练中途dropout mask重新采样也会带来额外显存尖峰,建议先设成0跑一版对比。我上次遇到类似情况,最后发现是DataLoader里collate_fn没固定max_length,不同batch的input长度差异导致显存波动。你可以在训练循环里每50步打印一下torch.cuda.max_memory_allocated(),看看是不是在固定步数附近突然跳高。如果还崩,就试试把model.gradient_checkpointing_enable()放在prepare_model_for_kbit_training之后,别放在LoRA包装之前。
我之前也踩过这个坑,一模一样,7B加LoRA跑到一半突然OOM。后来排查下来,问题往往不在显存总量,而是显存碎片化,尤其是你开了gradient checkpointing的话,它会在反向传播时重新计算激活值,这个峰值可能刚好撞上某个特别长的序列或者attention的计算图,导致瞬间暴涨。你可以用torch.cuda.max_memory_allocated()打印一下峰值到底是多少,我猜实际峰值可能接近20G了。
另外,peft库在某个版本确实有缓存没释放的bug,特别是在多步训练后累积了中间张量,你试试把transformers和peft都升到最新版,或者干脆降到某个稳定版本,比如peft 0.6.2那批。还有个偏方,就是给optimizer加个offload,比如用bitsandbytes的8bit Adam,虽然不能直接解决OOM,但能挤出不少冗余空间。
我最后是用混合精度加torch.cuda.empty_cache()每隔几百步手动清一下,再把max_length限制到512,就稳定跑完了。你那个step挂掉的时间点,也留意下是不是数据里混了超长样本,有些batch的padding特别长,激活值就直接翻倍了。
这情况我也遇到过,十有八九是激活值峰值爆了,试试开gradient checkpointing再配合max_grad_norm限制下。
检查下peft版本,之前有个版本会缓存旧权重导致显存慢慢涨,升级到0.7.2就解决了。
我之前跑13B也撞过一模一样的鬼,后来定位到是eval时把验证集一次性塞进显存了,如果每个step或固定间隔跑评估,峰值会瞬间飙高,LoRA本身反而不会突然多吃多少。你可以开个显存监控工具盯一下是不是正好卡在eval那几个step,或者把evaluation_strategy改成steps并且eval_accumulation_batch_size调小,甚至直接关掉eval跑一轮试试。还有个坑是数据加载器在长序列上会临时缓存attention的中间变量,就算batch=1,如果序列长度分布不均,偶尔来一条超长的样本就会爆,建议看看是不是挂了那几步的loss异常大。peft版本倒不太可能背锅,倒是transformers的版本有时候和flash-attention不兼容会触发缓存泄漏,可以试试把torch的分配器改成plaform或者关掉flash attn。我最后是靠开gradient_checkpointing(记得设use_reentrant=False)加max_length截断解决的,顺手把优化器的state_dict也挪到CPU上,虽然慢点但再也没中途炸过。
这现象我遇到过,大概率不是LoRA本身的问题,而是数据加载或者某个batch里序列特别长导致的峰值显存暴涨。你可以检查下是不是有特别长的样本,或者dataloader的num_workers开太多把内存挤爆了,顺带看看是不是峰值显存刚好卡在某个阈值上。
另外peft版本确实有过诡异的内存泄漏bug,建议直接升级到最新版试试,或者换用transformers自带的神谕微调接口对比一下。还有个小技巧,把optimizer换成AdamW8bit能省不少显存,说不定就熬过那个坎了。
实在不行就开torch.cuda.empty_cache()放在每N步清理下,虽然治标不治本,但能帮你定位是不是缓存累积的问题。
这现象我见过好几次了,大概率不是LoRA本身的问题,而是显存碎片化加峰值叠加。你看到的10G占用是稳态,但训练到中途某个step,比如当某个batch的序列长度恰好触发最长路径,或者当优化器状态和梯度同时驻留时,峰值会突然跳到接近24G。gradient checkpointing如果只开在transformer层,而没覆盖embedding或lora的adaptor,省的那点显存根本不够峰值冲击。
我建议你把torch.cuda.empty_cache()加在每个step结尾试试,先排除缓存堆积。另外查一下你用的peft版本,老版本有些已知bug会在特定step触发额外显存分配,比如当rank为8且target modules设置得比较碎的时候。还有个更隐蔽的点,如果数据里有极端长的样本,即便batch size=1,那个单独样本的激活值也可能爆炸。
我上次遇到过类似情况,最后发现是DataLoader的collate_fn没做padding限制,有个异常长的序列把显存顶穿了。你可以打印一下训练到崩溃前那个batch的输入tensor形状,看看是不是突然比平时大得多。如果真是这样,加个max_length截断或者动态padding就能解决。别急着怀疑库,先把训练循环里每一步的峰值显存打出来,定位到具体是哪个操作导致的飙升,比瞎猜高效多了。
我之前跑13B也遇到过一模一样的,不是LoRA的问题,大概率是激活值或者某些中间变量在特定step触发了峰值。你试试把optimizer换成AdamW8bit,再把gradient_checkpointing那个参数设成use_reentrant=False,很多人都是这么解决的。另外peft版本确实有个老bug会在某个step前额外分配显存,建议升到0.11以上试试。如果还崩的话,盯一下nvidia-smi看是不是别的进程偷偷占了显存,我上次就是被个僵尸进程坑了。
这问题我遇到过,大概率不是LoRA本身的问题,而是数据加载或者某个batch里有个特别长的样本导致的。你可以试着把max_length设短一点,或者检查下是不是有异常长的序列在某个step突然把activation撑爆了。另外peft版本确实有过类似bug,建议先升到最新版试试,顺便开一下torch.cuda.empty_cache()在每个step后手动清一下。我之前就是这么解决的,但也不排除是你某个中间变量没释放,可以开gradient_checkpointing再配合mixed precision看看。
这个现象我遇到过,peft库在某个版本确实有梯度累积相关的显存泄漏问题,尤其当你用了gradient checkpointing但没配合torch的显存清理函数时,每几步就会悄悄多占一点显存。你试过在训练循环里手动调torch.cuda.empty_cache()吗?虽然不治本,但能帮你确认是不是累积性泄漏。另外,LoRA本身不会突然飙升显存,除非你在某个step触发了eval或保存checkpoint的逻辑,这时候如果模型绑定了优化器状态,临时显存会翻倍。我建议你盯一下是不是恰好卡在validation或者logging的step上,把eval和save频率调低或者错开试试。还有个坑是数据加载器的num_workers如果开得高,偶尔会有缓存没释放,导致显存锯齿状增长,你降到0或者2看看。要是还不行,把peft和transformers都升到最新版,之前有个已知bug是lora的weight在反向传播时被重复计算,特定序列长度下会爆。最后,你可以开一下显存日志,每10步打印一次torch.cuda.memory_summary(),定位到具体哪一层开始暴涨,比瞎猜快多了。
我之前跑13B的时候也碰到过一模一样的情况,loss正常但显存在某个step突然冲高然后就崩了。后来排查发现是数据的问题——序列长度不均匀,大部分样本都很短,但每隔几百步就混进来一条特别长的,激活值瞬间爆掉,跟LoRA本身关系不大。你可以在dataloader里按长度排序或者设置max_length硬截断试试。另外gradient checkpointing确实能压峰值,但注意它和微调阶段的某些算子组合可能不生效,尤其是你用了自定义forward的话,得确认checkpoint包住了整个block。还有个小坑,如果你开了flash attention或者用了某些算子融合,它们在长序列下的临时缓冲可能是非确定性的,峰值比预估高很多。peft本身不太可能有这种随机性bug,但你可以试试把gradient accumulation设大一点,batch size保持1,这样至少能定位是不是单step的峰值问题。我之前还遇到过优化器状态里Adam的second moment在某些异常梯度下突然膨胀,导致显存翻倍,加个gradient clip就能缓解。你查一下是不是某个特定batch的loss特别大,如果是,很可能就是数据噪声问题。