最近在尝试微调一个7B的模型,用的LoRA,rank设的8,batch size调到了2,看显存占用才10G左右(显卡是24G的),按理说应该稳的。但每次训练到大概四五百步的时候,突然就Out of Memory了,训练直接崩掉。我查了日志也没看到什么明显错误,就是显存突然飙升。
我怀疑是不是gradient checkpointing没开对,或者中间缓存没清?也试过把batch size降到1,但还是会在不同步数挂掉。有没有大佬遇到过类似情况?是不是LoRA本身在某个阶段会突然占更多显存?还是我用的peft库版本有bug?
先谢过,真有点被搞懵了。
用LoRA微调7B模型,显存够但训练到一半就OOM了,咋回事?
全部回复
共 165 条我之前跑13B也遇到过一模一样的,不是LoRA的问题,是显存碎片化加峰值波动。你试试在optimizer里加一句max_grad_norm=1.0,然后把gradient_checkpointing从True改成use_reentrant=False,这个参数在peft新版里经常出幺蛾子。另外检查下是不是dataloader的num_workers开太高了,worker内存溢出也会算到CUDA头上。我上次是换了torch.cuda.empty_cache()在每个step后手动清一下才稳定,虽然治标不治本但能跑完。
我之前跑13B也遇到过一模一样的状况,后来发现是某个batch里sequence特别长,导致attention的中间激活值瞬间暴涨,跟LoRA本身关系不大。你可以试试开flash attention,或者把max_seq_len显式设短点,能缓解不少。另外peft版本确实有坑,建议换个稳定版或者直接看下是不是和transformers版本不兼容。
我之前跑13B也遇到过一模一样的情况,后来发现是eval时把梯度带进去了,或者某个batch里sequence特别长导致激活值峰值爆炸。你可以试试在training_args里显式关掉eval的梯度,顺便开一下max_grad_norm,再不行就开torch.cuda.empty_cache()手动清下缓存。peft版本倒不太像,这问题更像数据分布不均,偶尔来个超长样本就崩了。你查查是不是某个step的输入长度特别离谱?
我之前跑百川7B也遇到过一模一样的情况,后来发现是数据加载那边出了幺蛾子,某个batch里突然有个超长的序列,attention矩阵直接爆炸。你可以看看是不是输入长度没有统一截断,或者试试在dataloader里加个max_length限制。另外peft版本确实坑多,我之前升级到0.11.0之后类似问题就再没出现过,你可以顺手把transformers和peft都更新下看看。
我之前跑13B也遇到过一模一样的鬼情况,显存看着够但训练中途突然暴涨。后来发现是数据加载那边出了问题,某个batch的序列长度特别长,加上attention的计算峰值直接把显存顶爆了,你查查是不是数据里有长尾样本。还有peft版本确实有坑,建议升到最新版或者换个稳定分支试试,我之前就是更新完就好了。另外你可以开一下gradient checkpointing,再配合torch.cuda.empty_cache()在每步后清一下缓存,虽然治标不治本但能撑久点。
这题我好像也踩过,八成不是LoRA本身的问题,是序列长度或者attention机制在某个step触发了峰值显存,比如长文本片段突然进来。你可以试试用accelerate的--gradient_accumulation_steps把有效batch撑大,同时开max_memory限制,或者干脆监控一下是不是某个特定batch导致激活值爆炸。另外peft确实有过旧版本在训练中途缓存不释放的bug,建议直接升级到最新版再跑一次看看。
这情况我也踩过坑,大概率不是LoRA本身的问题,而是某些batch的序列长度特别长,导致激活值峰值暴涨。你试试在DataLoader里按长度排序或动态padding,把长序列集中处理,能压不少峰值。另外检查下是不是eval时也开了gradient,或者loss缩放没同步,这两个地方容易在特定步数触发隐性显存分配。
这情况我也踩过坑,不是LoRA本身的问题,大概率是激活值或者中间变量在特定batch里刚好触发了峰值。你试过用torch.cuda.max_memory_allocated()打一下峰值吗?我怀疑是某个样本序列特别长,导致attention的临时张量暴增,跟step数没关系,纯粹是数据分布问题。
gradient checkpointing确实能压峰值,但peft有些版本跟它兼容性不好,开了反而可能在某些层重复计算导致内存碎片化。你可以试试不用peft,直接手写LoRA注入,就几行代码的事,有时候反而更稳。
另外检查下是不是dataloader的num_workers太多,子进程预取数据时把CUDA context复制了,这也会让显存慢慢涨到爆。我之前遇到类似问题,最后发现是drop_last没设True,最后一个不完整batch的padding导致异常。
要是实在排查不出来,建议开一下PyTorch的显存快照工具,那个能精确看到分配峰值发生在哪一层。别急着怪库版本,先确认是不是偶然的长序列样本触发的。
我之前跑13B模型也遇到过一模一样的鬼打墙,显存监控看着稳得一批,结果一到某个step就暴毙。后来查了一圈,最可能的原因是数据加载那边出了幺蛾子,比如某个batch的序列特别长,或者padding策略没弄好,导致实际计算图比平时大出一截,虽然LoRA本身参数少,但中间激活值该爆还是会爆。你试着在dataset里加个max_length硬截断,或者用dataloader的collate_fn统一pad到固定长度,看看能不能解决。另外gradient checkpointing确实要确认下是不是真的生效了,有时候peft包版本和transformers版本不匹配,那个参数会被静默忽略。我上次就是升级了peft之后突然好了,虽然没搞懂具体是哪个commit修的。还有个野路子,就是显存快满的时候手动清一下缓存,比如在step间隔里调一下torch.cuda.empty_cache(),虽然治标不治本,但至少能帮你定位是不是碎片化问题。你试过用torch.cuda.set_per_process_memory_fraction限制硬上限吗?这样至少能提前看到OOM的风险,而不是整个进程被干掉。
我之前跑13B也遇到过一模一样的情况,尤其用peft的时候,问题多半不在LoRA本身,而是数据加载那边出了幺蛾子。你试试把DataLoader的num_workers设成0,或者检查下是不是有某个batch的序列特别长,导致attention的中间变量爆炸。另外gradient checkpointing要确认是True不是默认的False,有时候模型加载时会把它重置掉。实在不行就把optimizer换成AdamW8bit,能省不少临时显存。
检查下是不是某个batch的数据长度特别长,padding后序列变长导致激活值暴涨,给dataset加个按长度分桶的采样器试试。
这情况我也踩过坑,不是LoRA本身的问题,多半是显存碎片化或者某个batch的激活值突然暴涨。你可以试试开一下torch.cuda.empty_cache(),或者把gradient_checkpointing打开,再不行就检查下数据加载时有没有意外把整个数据集塞进显存。另外peft版本确实有过类似bug,建议升级到最新版,之前我换了个版本就稳定了。
八成是某个batch的activation峰值问题,试试把gradient_checkpointing设成True然后显存碎片化看看。
我之前也这样,后来用torch.cuda.empty_cache()加在每步后面就稳了。
试试关掉gradient checkpointing再开,有时候这玩意反而会触发显存峰值,我之前就这么解决的。
八成是loss突然炸了导致激活值暴涨,试试梯度裁剪或者看下是不是数据集里混了超长样本。
peft版本问题我也踩过坑,换个0.6.2试试,顺便把optimizer换成adamw8bit能省不少。
我之前跑13B也遇到过一模一样的情况,后来发现是PyTorch的缓存分配器在某个step触发了碎片化,显存看着够但实际可用块不连续,你可以试试在训练循环里定期调torch.cuda.empty_cache(),或者把环境变量PYTORCH_CUDA_ALLOC_CONF设成max_split_size_mb:128。另外peft版本确实有老bug,建议升到0.10以上,顺手把transformers也更新下,这俩组合容易在特定长度序列上突然爆显存。
我之前也踩过类似的坑,不是LoRA本身的问题,大概率是激活值缓存或者某些layer的中间变量在特定batch size下突然炸了。你可以试试开一下torch.cuda.empty_cache()在每N步手动清一下,或者看看是不是dataloader的最后一个batch形状不匹配导致的隐式padding。另外peft最近几个版本确实有内存泄漏的issue,建议直接升到最新版或者干脆换个分支试试。
我之前跑13B也遇到过一模一样的,LoRA显存曲线不是线性的,某个step突然爆掉多半是激活值峰值,跟rank和bs关系不大。你可以试试在peft里开disable_grad_scan或者直接用unsloth改内核,能省不少。另外检查下是不是数据长度不统一,偶尔来一条超长序列把缓存打爆了,用max_length截断或者packing能解决。我后来换了gradient_accumulation_steps=4加微batch,稳得一批,你可以参考下。
显存突然飙升大概率是某个batch的序列长度特别长,查下数据里有没有 outlier 吧。
我之前跑13B也碰到过这鬼情况,后来发现是某个batch里序列特别长,导致激活值突然爆炸,跟LoRA本身关系不大。你可以看看是不是数据里混了超长文本,或者在dataloader里按长度sort一下。另外peft版本确实有过显存泄漏的issue,建议换个稳定版试试,gradient checkpointing开了的话注意下是不是跟模型并行冲突了。
我也遇到过,不过是在优化器状态上栽的跟头。你确认下是不是用了paged_adamw之类的,有些优化器会在特定步数触发碎片化分配,显存直接飙上去。可以先换个adamw_simple跑跑看,或者把max_length限制一下,有时候是动态padding搞的鬼。
这情况八成是缓存没释放,尤其是attention的中间张量。你试试在训练循环里手动清一下torch.cuda.empty_cache(),虽然治标不治本但能定位是不是累积问题。另外检查下是不是验证集在某个step被加载了,eval模式下的显存峰值往往被忽略。
我猜是gradient accumulation的时序问题,LoRA在反向传播时对低秩矩阵的梯度计算偶尔会触发额外显存分配。你可以把gradient_accumulation_steps设小点,或者直接用deepspeed的zero stage2,省下的显存能兜住这个波动。我上次就是靠这个解决的。