最近在尝试微调一个7B的模型,用的LoRA,rank设的8,batch size调到了2,看显存占用才10G左右(显卡是24G的),按理说应该稳的。但每次训练到大概四五百步的时候,突然就Out of Memory了,训练直接崩掉。我查了日志也没看到什么明显错误,就是显存突然飙升。
我怀疑是不是gradient checkpointing没开对,或者中间缓存没清?也试过把batch size降到1,但还是会在不同步数挂掉。有没有大佬遇到过类似情况?是不是LoRA本身在某个阶段会突然占更多显存?还是我用的peft库版本有bug?
先谢过,真有点被搞懵了。
用LoRA微调7B模型,显存够但训练到一半就OOM了,咋回事?
全部回复
共 165 条这个情况我其实也踩过坑,感觉不完全是显存不够的问题。7B模型用LoRA正常跑,显存占10G其实已经留了不少余量,但突然OOM更像是某个中间变量或者缓存没被及时释放导致的。我建议你检查一下dataloader里num_workers的设置,如果设得太大,有时候子进程会偷偷累积显存碎片,跑着跑着就炸了。另外你提到的gradient checkpointing,可以确认一下是不是只开了模型层的,而LoRA那部分没被包含进去——有些peft版本对checkpointing的兼容性确实有问题,会导致某些模块的中间结果不释放。batch size降到1还崩的话,可以试试把gradient accumulation steps设大一点,比如4或8,这样等效batch size不变但单步压力更小。还有就是看看是不是用了transformers的cache机制,那个在长序列训练时偶尔会爆。最后,如果你用的是较新的peft版本,可以回退到0.7.0左右试试,我遇到过某个版本对梯度裁剪和显存管理有bug,降版本就稳了。
我遇到过类似的情况,感觉不一定是LoRA本身的问题,更像是某个中间变量或者优化器状态在训练过程中累积导致显存爆炸。可以试试在优化器里开个torch.cuda.empty_cache()手动清一下缓存,或者检查下dataset的padding策略是不是随着训练步数变了。另外peft库确实有过一些显存泄漏的旧版本bug,更新到最新版说不定就解决了。
这情况我也遇到过,大概率不是LoRA本身的问题,而是PyTorch的显存分配机制在搞鬼。你观察得很准,显存不是一开始就爆,而是训练到某个步数突然飙升,这其实是计算图缓存或者中间变量没释放导致的。就算batch size调到1,如果gradient checkpointing没正确生效,反向传播时临时张量还是会堆积。我建议你检查一下是否在peft的配置里明确设置了gradient_checkpointing=True,并且确保模型调用了model.enable_input_require_grads(),否则这个开关可能只是摆设。另外,可以试试在训练循环里手动清空缓存,比如每N步调用torch.cuda.empty_cache(),虽然治标不治本,但能临时绕过去。还有一点,检查下dataloader的num_workers是不是设得太高了,有时候多进程加载数据也会在某个时间点突然占用额外显存。你用的peft版本是0.10还是更新的?之前0.9.x确实有个bug,LoRA的权重更新时会临时创建过大的中间矩阵。如果这些都不管用,建议直接用transformers的Trainer并开启deepseed stage 2,对7B模型来说效果很稳定。
我之前跑13B也遇过一模一样的状况,后来发现是PyTorch的缓存分配器在搞鬼,显存占用到峰值不代表真的用完,但训练中途某个大张量突然申请连续内存就会爆。你可以试试在训练循环里加个torch.cuda.empty_cache(),或者设PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,这样能让分配更碎一点,防止碎片化导致突然OOM。另外LoRA本身不会突然涨显存,但如果你在某个step恰好触发了评估或者梯度累积清零,那一下的峰值可能比平时高不少,建议开一下gradient checkpointing(不只是LoRA层,整个模型都开),同时把optimizer的momentum缓存也检查下,有时候AdamW的state会随着训练慢慢膨胀。peft库最近确实有几个版本改过内存管理逻辑,如果方便的话可以试试0.6.2或者直接更新到最新,顺便看下transformers版本是不是太老。还有个土办法,把数据加载的num_workers设成0,有时候dataloader异步预取会额外占显存。我最后是靠减少保存checkpoint的频率解决的,因为保存时会复制模型权重,正好撞上某个大激活值释放的时机就崩了。
我之前跑13B也遇到过一模一样的,不是LoRA的锅,大概率是某些token序列特别长导致激活值突增。你可以试试在dataloader里按长度排序,或者开max_seq_length截断,能稳定很多。另外peft版本确实有坑,建议锁到0.9.0试试,新版有时候会偷偷缓存梯度。
我之前也踩过这个坑,7B加LoRA正常不该中途爆显存,问题多半不在rank和batch上。你试试把gradient checkpointing打开后,再确认下是不是evaluation时也在算梯度,或者dataloader里有没有什么动态padding导致序列长度突然拉长。另外peft版本确实有老bug会在step到一定数量后累积激活值,建议直接升到最新版,顺便把optimizer的momentum缓存清一下试试。如果还崩,就开一下显存日志(比如torch.cuda.memory_summary),看看到底是哪块峰值涨的,多半是某个特定长度的样本触发的。
我之前也碰到过一模一样的,7B + LoRA跑着跑着显存突然暴涨。后来发现是某个batch里sequence特别长,导致激活值峰值远超平时,gradient checkpointing只对前向有效,但反向时的临时buffer还是会被撑爆。你可以试试按token数动态batch,或者把max_seq_len硬性截断一下,大概率能解决。另外peft版本确实有过类似bug,升级到最新版再看看。
我之前跑13B也遇到过一模一样的,症状就是loss正常、显存曲线平稳,但到了某个固定步数突然暴毙。后来查出来是数据集里某个特别长的样本在特定轮次被采样到了,激活值瞬间冲高,跟LoRA本身关系不大,更像是输入序列长度的分布问题。你可以试试在dataloader里按长度排序或者设个max_length硬截断,看是不是能稳定跑过去。另外gradient checkpointing其实对这类突刺帮助有限,它省的是常规显存,但解决不了极端长序列的峰值。peft这边的话,建议确认下是不是最新版,之前有过版本在梯度累积步数切换时缓存没释放的issue,不过现在应该修了。还有个土办法,就是给PyTorch设个环境变量PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,能减少碎片化导致的空间利用率下降,虽然治标不治本。你试试看把训练步数打点,看OOM前一步的loss和输入长度是不是异常,大概率能定位到是数据问题。
我之前跑7B也遇到过一模一样的,后来发现是eval的时候把数据pad到最大长度了,显存峰值直接翻倍。你试试把evaluation关了或者缩短max length,应该能撑过去。另外peft版本确实有点玄学,更新到最新版或者换个commit试试。
之前跑13B也遇到过一模一样的,显存看着稳但中途突然爆,后来发现是某些batch里序列特别长,激活值peak直接翻倍。你查查训练数据是不是有长文本分布不均的情况,把max_seq_len卡死或者开个梯度累积试试。另外peft有的版本对某些模型会缓存中间态不释放,换个0.11.0或者直接上最新版看看。
我那次是checkpoint保存的时候触发的,因为会额外算一次inference的激活值,正好卡在峰值上。你观察下崩的时间点是不是跟save_steps对得上,把save_strategy改成epoch或者关掉eval试试。
顺便说一句,我开gradient_checkpointing之后反而更容易崩,感觉它跟flash-attention有的组合会出内存泄漏,你如果没开flash-attn的话可以试试开一下,显存占用反而更平滑。
我遇到过一模一样的,也是7B+LoRA,24G卡跑着跑着就爆。后来发现是数据加载那边的问题,某个batch的序列长度特别长,导致激活值突然飙上去,跟LoRA本身关系不大。你可以看看是不是有超长样本,或者试试在dataloader里按长度排序,另外把gradient checkpointing开了再配合显存碎片清理,大概率能解决。
显存突然飙升八成是激活值峰值爆了,试试把gradient checkpointing打开再配合显存碎片清理看看。
我之前也踩过类似的坑,7B配24G显存跑LoRA按理说很宽裕,但中途OOM多半不是rank或batch size的锅。你试试看是不是序列长度的问题,有些数据集里偶尔会出现特别长的样本,到那一步激活值会突然暴涨,显存曲线看起来就像“飙升”一样。我当初就是没做max length的截断,结果每跑几百步就炸一次,后来统一截到1024就再没出现过。另外gradient checkpointing建议确认下是不是真的生效了,光在配置里开了但没调model.gradient_checkpointing_enable()就容易白开,你可以在训练循环里打印一下model.gradient_checkpointing看看。还有一个容易被忽略的点是optimizer的状态,AdamW的动量项会随着训练步数慢慢累积,虽然LoRA只更新少量参数,但如果优化器没冻结基座模型参数(比如没设requires_grad=False),它照样会把梯度缓存都算进去,显存就会像温水煮青蛙一样涨上去。peft库版本倒不太可能有这种bug,我更怀疑是你数据加载时有个别batch特别“畸形”,比如padding没对齐导致计算图特别长。你可以试着在DataLoader里加个按长度排序的sampler,或者干脆把max_length设死,应该能解决。要是还不行,就开个显存监控脚本(比如pynvml)盯着每个step的显存峰值,看看是不是每步都在缓慢增长,如果是那就是泄漏,得查查是不是有tensor被意外绑到了计算图上。
我之前跑13B也遇到过一模一样的,不是LoRA的问题,是你某个batch里序列长度突然变长导致的峰值显存暴涨。可以试试在DataLoader里对sequence length做个max_length截断,或者把gradient checkpointing打开确认一下是不是真的生效了(看下显存曲线有没有台阶式下降)。另外peft版本建议锁到0.7.1左右,新版有些时候会偷偷把base model的gradient也开了,挺坑的。
我之前跑13B的LoRA也撞过这情况,后来发现是某个特定batch里序列长度暴涨,padding没控制住,显存峰值直接翻倍。你可以看看是不是数据里有超长样本,或者试试在collator里显式max_length截断。另外gradient checkpointing确实得配合显存优化器用,单独开有时候反而触发碎片问题,建议把torch的缓存清理和max_split_size_mb调一下。peft版本的话,我之前用0.6.2有类似毛病,升到最新版就好了,你顺手看看transformers版本对不对得上。
我之前跑13B也遇到过一模一样的坑,不是LoRA本身的问题,大概率是某个中间步骤的激活值或者梯度累积导致的峰值显存暴涨。你rank8、batch2看着稳,但训练到中途loss下降后,某些层输出的分布变化可能让临时buffer分配突然变大,这种情况在量化或者混合精度下特别常见。建议你先把gradient checkpointing确认开对,用model.gradient_checkpointing_enable(),同时检查一下accelerate或者transformers版本,最近peft有几次更新改过内存管理逻辑。还有个偏方是给PyTorch设PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,能缓解碎片化导致的假OOM。另外你试过用torch.cuda.empty_cache()在每次step后手动清一下吗?虽然治标不治本,但能帮你确认到底是不是缓存堆积。如果还崩,开CUDA_LAUNCH_BLOCKING=1跑一次,看能不能抓到真正的报错栈。
我遇到过一模一样的,不是LoRA的问题,是显存碎片化加中间激活值峰值。你试试把gradient checkpointing打开,然后设个环境变量PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,基本能解决。另外检查下是不是dataloader的num_workers设太高了,有时候数据加载也会挤占显存,尤其到某个epoch边界的时候。
我上次也被这个坑了,后来发现是evaluation的时候把dev集也塞进显存了,训练中途跑验证就爆了。你可以在训练脚本里确认下是不是每个多少步会跑一次eval,如果是的话,把eval的batch size调小或者干脆关掉看看。peft库版本倒是不太可能,这问题太普遍了。
你用的啥优化器?AdamW的话偶尔会有那种momentum相关的临时张量暴涨,换个8bit优化器或者把optimizer的eps调大点试试。我之前换成bnb的AdamW8bit之后就没再炸过,不过那次是12G的卡,情况可能不完全一样。
我之前也踩过类似的坑,后来发现是数据加载那边的问题,不是LoRA本身,你试试把dataloader的num_workers调成0或者pin_memory关掉,有时候缓存不释放会突然爆一波。另外检查下是不是某个batch的序列特别长,导致激活值突增,peft在长序列下动态图算得比预想多很多。我后来加了max_length截断和梯度累积,反而稳了,你可以对比下loss曲线是不是在长样本附近崩的。
遇到过一模一样的坑,最后排查下来是数据集里个别样本特别长导致的。LoRA本身不会突然涨显存,但如果你没按sequence length做bucket或者固定max length,某个超长batch会把activation footprint瞬间拉爆,尤其是深层的attention计算,7B模型在长序列下的中间激活值比你想的夸张得多。建议你检查一下训练数据的长度分布,把超过某个阈值的样本截断或过滤掉,或者干脆在collator里统一padding到固定长度试试。
另外gradient checkpointing这个确实要确认下,peft默认不会自动帮你开,得在transformers的TrainingArguments里显式设gradient_checkpointing=True,而且最好配合input_names之类的设置让重计算生效。如果开着但显存还是骤升,可以试试把optimizer换成AdamW的8bit版本,或者关掉cache相关的选项,有时候peft的forward里会缓存一些中间tensor用来加速,但反而会在长序列上累积。
我上次还发现一个隐蔽问题:如果用了packing或者concat数据集,某些样本拼接后长度远超预期,这时候即使batch size很小,单条样本的显存占用也会直接翻倍。建议你在训练循环里加个hook打印每步的max_memory_allocated,定位到具体是哪一步爆的,然后对比那一步的输入长度和前面正常步数有什么区别,基本就能锁定了。
至于peft库版本,我建议先升级到最新版,老版本确实有关于lora_dropout和scaling在特定batch下内存泄漏的issue,不过你这种情况更像是数据长度波动不是bug。总之先查长度分布,再确认checkpointing,最后再考虑版本问题,应该能解决。
八成是某个batch里sequence特别长,attention缓存突然爆了,试试max_length截断或者按长度分桶。