最近在尝试微调一个7B的模型,用的LoRA,rank设的8,batch size调到了2,看显存占用才10G左右(显卡是24G的),按理说应该稳的。但每次训练到大概四五百步的时候,突然就Out of Memory了,训练直接崩掉。我查了日志也没看到什么明显错误,就是显存突然飙升。
我怀疑是不是gradient checkpointing没开对,或者中间缓存没清?也试过把batch size降到1,但还是会在不同步数挂掉。有没有大佬遇到过类似情况?是不是LoRA本身在某个阶段会突然占更多显存?还是我用的peft库版本有bug?
先谢过,真有点被搞懵了。
用LoRA微调7B模型,显存够但训练到一半就OOM了,咋回事?
全部回复
共 165 条我之前也被这个问题搞到怀疑人生,后来发现大概率是序列长度或者attention的计算波动导致的,LoRA本身不会突然多吃显存。你那10G是稳定占用还是峰值?建议用nvidia-smi盯着看,很多框架显示的只是模型参数+优化器状态,没算上激活值。gradient checkpointing要是没开对,其实只节省了前向的激活存储,但反向传播时某些操作(比如all-gather)还是会临时拉高显存。我遇到过类似情况,最后是靠把max_length从2048降到1024解决的,虽然损失点精度但至少不崩。另外peft版本确实有坑,0.6.0之前有些版本在特定步数会触发显存碎片化,你试试升级到最新版或者直接换transformers的trainer来跑。如果还不行,可以开torch.cuda.empty_cache()每50步清一下,但治标不治本,本质还是缓存碎片问题。建议你顺便看下是不是数据集里有个别超长样本,某些batch的padding会突然拉高显存,这种随机性就跟你描述的症状很像。反正千万别怀疑是显卡坏了,我当初就差点去售后。
我之前跑13B也遇到过一模一样的,后来发现不是显存不够,是某个特定step计算图里临时变量暴涨,比如attention那块,建议开一下显存碎片化日志看看峰值到底在哪一步冒出来的。还有个小技巧,把optimizer换成AdamW8bit或者干脆用SGD,峰值能降不少。peft版本的话我上次升级到0.7.2以后就没再翻车,你可以试试看。
我之前跑13B的LoRA也撞到过一模一样的情况,显存占用看着很健康,但就是会在某个step突然跳上去然后崩掉。后来排查下来,发现是数据加载那边出了问题——某个batch的样本长度特别长,导致activation在那一瞬间暴涨,即便LoRA本身参数少,但前向传播的中间激活值跟序列长度是线性相关的,长样本直接就把显存顶穿了。你可以试试把max length强行限制一下,或者用动态padding,别让那些超长样本混进来。另外,gradient checkpointing确实会影响显存曲线,但你这情况更像是某些step的峰值波动,不是常态占用问题。peft版本的话,建议先把transformers和peft都升级到最新,有个旧版本在梯度累积时确实会有显存泄漏的bug。还有个偏方,就是开一下torch.cuda.empty_cache()在每N个step手动清一下,虽然治标不治本,但能帮你确认是不是缓存碎片的问题。如果还不行,就开一下显存监控日志,看看崩溃前那个step的tensor分配情况,基本能定位到具体是哪层炸的。
八成是max_length没设对,某个batch序列太长把缓存撑爆了,看看数据里有没有超长样本。
这种突然OOM我之前也踩过,多半不是LoRA本身的问题,而是某个batch的序列长度特别长,导致激活值峰值暴涨。你可以看看是不是数据里有超长样本,或者开个max_length限制一下,另外把gradient checkpointing打开再配合显存碎片清理(比如torch.cuda.empty_cache)试试。
peft版本确实有时候会有隐性bug,我之前升级到最新版就好了,不过更关键的是检查一下是不是eval时也占了显存没释放。如果还挂,建议用accelerate的cpu offload兜底,虽然慢点但至少能跑完。
对了,你监控的是nvidia-smi的实时显存吗?那个有时候会漏掉峰值,建议用torch.cuda.max_memory_allocated看下实际峰值是多少,说不定早就超过24G了。
我之前跑13B也遇到过一模一样的鬼情况,显存平时稳如老狗,一到某个步数直接爆掉。后来发现是数据长度不均,碰到超长序列时激活值会瞬间涨一大截,LoRA虽然省了优化器状态,但激活内存这玩意儿不吃这套。你可以试试把max_seq_len再砍半,或者开一下flash attention,实在不行就按步数手动清一次缓存,虽然丑但管用。另外peft版本确实有过类似bug,去GitHub翻翻issue区,换个最新版说不定就治好了。
我之前跑13B也遇到过一模一样的,不是LoRA的问题,大概率是激活值或者某些中间变量在特定序列长度下爆了,你可以试试开torch.cuda.empty_cache()定时清一下。另外检查下是不是数据集里有个别超长样本,跑到那一步突然撑爆显存,按长度过滤一下会稳很多。peft版本建议升到最新,老版本确实有碎片化泄漏的bug,我之前升级后就没再犯过。
我之前也踩过类似的坑,7B+LoRA中途OOM大概率不是显存不够,而是某个batch的序列长度触发了激活值峰值。你可以查一下数据里是不是有超长样本,或者开一下max_length截断试试。
另外gradient checkpointing确实要配合显存优化器一起用,光开它但没设分块策略,反而可能让峰值更高。顺便检查下peft和transformers版本,老版本有已知的缓存泄漏问题,升到最新版通常能解决。
我上次是把batch size挂到最小,然后强制设了--max_grad_norm,再手动清空torch.cuda.empty_cache(),虽然丑但稳住了。你要是解决了也来反馈下,我挺好奇具体原因。
我之前也踩过一模一样的坑,7B加LoRA看着显存余量很大,但训练中途突然暴涨的情况太典型了。你排查的方向我觉得没问题,但多半不是gradient checkpointing的锅,那个是省显存不是防突刺的。我后来定位到是数据加载或者采样器的问题,某个batch的sequence特别长,导致activation突然变大,尤其是如果你开了packing或者没设max length上限,长样本一进来显存就直接顶穿。你可以试试给dataloader加个按长度排序的bucket sampler,或者干脆把max length卡死到512,看还会不会崩。另外peft版本确实有历史bug,特定条件下LoRA的梯度会异常累积,建议直接升到最新版或者换个稳定分支。还有一个很隐蔽的点,如果你用了torch.compile或者某些融合算子,偶尔会在特定step触发重新编译,那一下内存峰值会巨高,可以先关掉排除。最后建议你开个wandb或者tensorboard盯着显存曲线,如果看到崩溃前有缓慢爬升而不是瞬间跳变,那就可能是碎片化问题,试试PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True这个环境变量。
八成是eval时把梯度也带进去了,试试eval模式下关掉gradient checkpointing或者用no_grad包一下。
八成是某个step触发了特殊样本导致激活值暴涨,试试开gradient checkpointing并固定随机种子看看能不能复现。
这情况我也踩过,多半是activation checkpointing没生效,试试把model.gradient_checkpointing_enable()放对位置。
大概率是某个batch里序列特别长,激活值瞬间炸了,试试按token数动态batch或者固定max length。
我之前跑13B也遇到过一模一样的,不是LoRA的问题,大概率是激活值在特定序列长度下突然爆了。你试试开一下gradient checkpointing,同时把eval和save的step错开,别让验证集和训练同时跑。另外peft新版有个bug,老版本反而稳,回退到0.9试试。
检查下是不是eval时把max_length撑爆了,LoRA本身不会突增显存。
梯度爆炸也可能,试试grad clip,或者开下gradient checkpointing。
跑一下显存监控看是不是loss spike时激活值暴涨,可能是某些batch里sequence特别长导致的。
跑一下显存监控看看是不是loss spike导致中间激活暴涨,顺便试试点开optimizer的offload。
我上次也这样,最后发现是某个batch文本长度异常,过滤掉就好了。
我之前跑13B也撞过这个鬼问题,后来发现是eval时把数据也塞进显存了,你试试把evaluation_steps调大或者干脆关掉eval。另外peft有些版本在checkpoint保存时会临时复制权重,去查下save_steps是不是正好卡在500步附近。还有个小坑:某些数据集在中间会出现超长样本,触发动态padding后激活值爆炸,建议max_length锁死并开truncation。
我之前也遇到过一模一样的坑,7B加LoRA跑到中途必炸。后来发现不是显存不够,而是某些特定batch里序列特别长,激活值峰值直接冲上去,你试着在dataloader里按长度排序或者动态padding,能缓解很多。另外peft新版本确实改过内存管理逻辑,我换了0.9左右的旧版就稳定了,你可以对比试试。还有个小技巧,把optimizer换成AdamW8bit,能省下不少临时显存,虽然你看着基础占用不高,但峰值才是关键。
我之前跑13B也遇到过一模一样的,八成不是LoRA的问题,是某个batch里的样本特别长,激活值突然炸了。你可以看看是不是数据没做长度截断或者padding策略不对,试试按token数动态batching,或者把gradient checkpointing开起来再配合显存碎片清理。另外peft版本太老的话确实有内存泄漏的坑,建议先升到最新版,再不行就加个torch.cuda.empty_cache()在每步之后,虽然治标不治本但至少能续命。