最近在尝试微调一个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的锅,试试把优化器状态和中间激活显存盯一下,峰值容易爆。
八成是数据加载或者loss计算里藏了动态shape,跑几步才爆显存,试试固定max length。
我上次这么崩是tokenizer padding没设对,你查查是不是某个batch特别长。
八成是显存碎片化,或者某个batch的序列长度突然变长导致峰值暴涨,试试固定max length。
我之前跑13B也碰到过一模一样的情况,不是显存不够,而是显存碎片化加临时峰值叠加的结果。你观察到的“突然飙升”大概率不是LoRA本身的问题,而是某些特定batch里序列长度不均匀,导致attention矩阵计算时显存瞬间暴涨——7B模型就算只算单层,长序列下的激活值也能轻松吃掉几个G。建议你查一下数据加载时有没有做padding到固定长度,或者尝试开一下flash-attention,它能显著降低峰值显存。另一个可能被你忽略的地方是优化器状态,AdamW的momentum和variance在训练中途会逐渐膨胀,如果用了8-bit优化器但没正确配置,也会出现后期显存缓慢爬升然后崩掉的情况。梯度检查点建议确认是开在模型层上,而不是只包了LoRA适配器,否则反向传播时还是会把中间激活全算一遍。最后可以试试在训练循环里手动清一下torch.cuda.empty_cache(),虽然治标不治本,但能帮你确认是不是缓存没释放的问题。如果排除了这些,建议把peft和transformers都升到最新版,老版本确实有过中途显存泄漏的bug。
我之前也踩过类似的坑,LoRA本身显存曲线应该是平稳的,问题大概率出在数据加载或者优化器状态上。你试试把gradient_accumulation_steps加上,虽然总batch不变,但能降低瞬时显存峰值。另外检查下是不是有eval时把模型切回全参数了,或者某个batch里序列特别长,你可以log一下每步的input长度看看。peft版本我记得之前有个版本跑7B会在特定步数触发重算导致显存暴涨,升级到最新版试试,我换完就好了。
你这情况多半是激活值峰值爆了,试试把max_length调小或者开个torch.cuda.empty_cache定时清下缓存。
我之前跑13B也遇到过一模一样的鬼情况,loss还在正常降呢,结果某个step显存直接拉满然后崩掉。你怀疑gradient checkpointing没开对,但这个其实不是罪魁祸首,因为LoRA本身不会在某个阶段突然变胖,问题更可能出在数据加载或者日志记录上。我当时查了半天,最后发现是evaluation时把验证集整个塞进显存去算指标了,那个瞬间峰值比训练时高出一大截。你可以看看是不是设了eval_steps,并且验证集没有用dataloader分批,或者验证时忘了关gradient。另外peft库确实有过一些版本在特定模型上内存泄漏的issue,建议你直接升级到最新版,或者换个分支试试,我换完就再没崩过。还有个歪招,就是给训练脚本包一层try-except,在OOM前自动保存checkpoint然后重启,虽然不治本但至少不浪费时间。你最好也监控一下是不是某个特定batch的数据特别长,比如padding策略没设好导致序列长度突增,这个在文本类任务里很常见。
多半是某个batch的sequence特别长,峰值显存炸了,试试max_length截断或者按长度排序batch。
我之前跑13B也遇到过一模一样的,不是LoRA的问题,大概率是激活值峰值或者某个特定step的数据batch特别不均匀导致的。你可以试试把gradient checkpointing打开,再把优化器换AdamW的8bit版,显存占用能压下去不少。另外查一下是不是dataloader在某个step加载了特别长的序列,长序列的attention计算会突然爆显存,这个在日志里看不出来。
这现象我见过好几次,不是LoRA本身的问题,更像是数据集里混了超长样本。你想想,虽然平均长度看着不高,但每隔几百步突然来一条特别长的序列,激活值和计算图会瞬间爆掉,显存曲线就是那种锯齿状突然拉满的。我之前在7B上跑QLoRA也踩过这个坑,后来用grad accumulation配合动态padding,把每个batch内的序列长度对齐到接近最大长度,OOM概率立马降了八成。你检查下dataloader是不是有shuffle,有时候撞上长样本纯属运气问题,所以步数才不固定。另外peft的版本确实有老坑,某些版本在特定rank下会缓存额外的梯度张量,建议直接升到0.11以上,顺便把disable_gradient_checkpointing这个参数确认下,很多人开了gradient checkpointing但忘了把模型输入也设为requires_grad=False。最后给你个偏方,训练循环里每隔200步手动清一次torch.cuda.empty_cache(),虽然治标不治本,但能帮你定位是不是碎片化导致的缓慢泄漏。如果试完还崩,就把transformers的modeling_utils.py里的_gradient_checkpointing_func打点日志,看崩之前到底走了哪个模块。
我之前也踩过这个坑,不是LoRA本身的问题,大概率是激活值或者临时变量在长序列下突然爆了。你可以试试在dataloader里加个max_length限制,或者把optimizer换成AdamW8bit,显存能省不少。另外peft版本建议升到最新,有个旧版bug是训练中途会缓存某些梯度导致显存泄漏。如果还不行,开一下torch的memory_stats看看具体哪一步涨的,比盲猜快。
大概率是某个batch里序列长度暴长,把激活值撑爆了,试试按token数动态batch或者max length截断看看。
多半是loss突然炸了导致激活值激增,试试gradient clipping或者看看是不是数据集里有异常长序列。
也可能是peft的lora dropout在前向随机激活了更多显存,换个seed跑跑看?
我之前也踩过类似的坑,最后发现是PyTorch的缓存分配器在搞鬼。显存占用看着只有10G,但训练到中途某个step时,某些中间激活值或者梯度checkpoint的临时张量会突然申请一大块连续显存,而这时候allocator手里虽然有碎片但凑不出一整块,就直接OOM了。你可以试试在训练循环里每N步手动调一下torch.cuda.empty_cache(),或者干脆设置PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,把分配粒度调小,让碎片能被利用起来。另外,你说batch size降到1也会挂,那基本可以排除batch维度的问题,更像是某个特定长度的序列触发了显存峰值——比如数据里偶尔有一条特别长的样本,即便LoRA只训adaptor,但forward/backward还是要过完整base model,长序列的attention激活值会指数级涨。我建议你查一下dataloader里每条样本的input长度分布,把超过某个阈值的样本截断或者过滤掉试试。至于peft库的bug,可能性不大,但你要是用了比较老的版本,可以升级到最新看看,因为之前确实有过gradient accumulation状态下显存泄漏的issue。最后,如果实在不行,就把gradient checkpointing显式打开,配合gradient_accumulation_steps=2,牺牲点速度换稳定。
我之前也遇到过一模一样的,LoRA显存按理说很稳,但训练中途突然爆掉大概率不是rank或batch的问题,更像是某个step的梯度或激活值异常导致临时峰值。你可以试试把gradient checkpointing显式打开,同时用torch.cuda.empty_cache()在每步后清一下缓存,虽然治标不治本但能确认是不是碎片化累积。另外peft版本确实有坑,我之前在0.6左右某个版本遇到过类似问题,升级到最新版或者换个稳定版试试看。要是还崩,建议开一下mixed precision训练,有时候fp16的loss缩放会搞出瞬时显存飙升,虽然奇怪但确实存在。
我之前也踩过类似的坑,跟你显卡大小没关系,纯粹是显存碎片化问题。LoRA虽然整体占用不高,但训练中某个step如果刚好遇到激活值峰值,加上PyTorch缓存分配不及时释放,就会瞬间爆掉。你可以试试在训练循环里加个torch.cuda.empty_cache(),或者把环境变量PYTORCH_CUDA_ALLOC_CONF设成max_split_size_mb:128,大概率能缓解。
另外检查下是不是dataloader的num_workers开太多,有时候数据加载也会突然多吃显存。peft版本的话,我上次升级到0.7.0就正常了,之前0.5.0确实有过类似问题。你用的哪个版本?
我之前跑13B也踩过这个坑,LoRA虽然省显存但峰值不一定稳,问题多半出在loss震荡时激活值突然暴涨,可以试试把gradient checkpointing打开同时调小max_length,或者用torch.cuda.empty_cache()在每步后手动清一下缓存。另外peft版本确实有老bug,建议直接升到最新版,顺便看下transformers版本是不是太旧,我之前升级完就再没崩过。还有一个思路是冻结某些层或者用gradient accumulation模拟小batch,但你这个显存余量按理说够,先查数据加载是不是有内存泄漏吧。
我之前跑13B的LoRA也踩过类似的坑,显存看着够但训练中途突然爆掉,后来发现不是LoRA本身的问题,而是数据加载器在某个epoch边缘会触发额外的缓存分配,尤其是当你的数据集长度不是batch size整数倍的时候,最后那几条样本会临时撑大计算图。你试试在dataloader那边加个drop_last=True,或者把shuffle关了看看能不能复现,这样能排除数据侧的问题。另外gradient checkpointing确实会影响显存峰值,但按理说开了之后应该更稳才对,你确认一下是不是真的在模型配置里生效了,有时候peft包和transformers版本不匹配会导致checkpointing被静默覆盖。还有一个常见坑是loss scaling或者梯度累积相关的临时tensor,在你这个rank=8的情况下不太明显,但可以试着把optimizer换成AdamW8bit,顺便开一下optimizer的offload,这样能把峰值显存摊到CPU上。如果还崩,就直接在训练循环里手动清一下torch.cuda.empty_cache(),虽然治标不治本,但能帮你确认是不是缓存碎片化的问题。最后peft库最近更新挺勤的,你要是用的老版本,建议升到0.10以上,之前有issue提到过某些rank下forward时动态padding会突然多分配显存,正好对应你这种“固定步数后飙升”的现象。
跑一下显存监控脚本看看峰值,八成是激活值在某个batch突然炸了,和LoRA本身关系不大。
试试开gradient checkpointing,再把优化器状态换成分片加载,应该能压住峰值。
我之前跑13B也遇到过一模一样的,不是LoRA的问题,大概率是某个batch里序列特别长,activation突然爆炸。你试试在dataloader里按长度sort或者设个max_length截断,能稳很多。
另外peft版本确实坑多,建议直接升级到最新版,顺便把transformers也升了,之前有个旧版本会在特定step触发缓存泄漏,跟你的现象完全吻合。
还有个土办法,把optimizer换成AdamW8bit或者Sophia,显存峰值能压下去一截,我这么干之后再没崩过。