最近在尝试微调一个7B的模型,用的LoRA,rank设的8,batch size调到了2,看显存占用才10G左右(显卡是24G的),按理说应该稳的。但每次训练到大概四五百步的时候,突然就Out of Memory了,训练直接崩掉。我查了日志也没看到什么明显错误,就是显存突然飙升。
我怀疑是不是gradient checkpointing没开对,或者中间缓存没清?也试过把batch size降到1,但还是会在不同步数挂掉。有没有大佬遇到过类似情况?是不是LoRA本身在某个阶段会突然占更多显存?还是我用的peft库版本有bug?
先谢过,真有点被搞懵了。
用LoRA微调7B模型,显存够但训练到一半就OOM了,咋回事?
全部回复
共 165 条这种跑几百步突然OOM的情况我也碰过,大概率不是LoRA本身的问题,而是显存碎片化攒到一定程度炸了。你可以试试设置PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,或者把optimizer换成8bit的,能省不少。另外检查下是不是eval或者save的时候触发了额外显存分配,有时候是checkpoint保存那一下顶爆的。
你这种跑到几百步才炸的情况,八成是碎片化显存攒出来的。我遇到过类似的,前面看着占用不高,但PyTorch缓存分配器一直没释放干净,到某个点突然要一块连续大显存就崩了。可以试试设一下PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,或者每几百步手动清一下cache。另外检查下有没有开evaluation或者save时的临时峰值,那一下经常是压死骆驼的最后一根稻草。
可能是内存碎片攒的,试试 torch.cuda.empty_cache() 加定期清理,或者设 PYTORCH_CUDA_ALLOC_CONF 环境变量。
我之前也踩过这个坑,大概率不是LoRA本身的问题,而是训练到中途触发了某种动态显存增长。你留意下是不是在eval或者save checkpoint的时候崩的,HF的Trainer默认会保留一堆中间变量。另外gradient checkpointing如果和某些peft版本搭配不当,反而会在反向传播时多占显存,建议确认下是不是每步都在重算。可以试试设个save_steps和eval_steps错开,再开个PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True看看能不能缓解。
这个现象挺典型的,我之前用LoRA微调7B的时候也踩过类似的坑,不一定是显存本身不够,而是训练到中途某些操作触发了峰值。你可以先看看是不是在eval或者save的时候炸的,有时候是保存checkpoint或者跑验证集那一下显存突然上去了,尤其是你如果没设置eval_steps和save_steps错开的话。另一个常见原因是token长度分布不均,前面几百步都是短样本,显存看着很稳,后面突然来了一批长序列,attention那块直接把显存顶爆了。gradient checkpointing如果开对了确实能压不少,但有些模型和peft版本搭配会有bug,你可以打印一下每步的max_seq_len确认下。还有个小细节是LoRA的optimizer state,如果你用的是adamw,它会在某些步数累积动量,虽然理论上不会暴涨,但配合长序列就容易出事。建议试试把max_seq_len卡死,或者用dynamic padding改成固定长度,另外检查下dataloader有没有把整个batch的label都算loss,那个也会偷偷吃显存。