最近在试着用LoRA微调一个7B的底座模型(做代码生成),单卡A100 40G,用的transformers+peft。batch size设成1,gradient accumulation调到8,序列长度512,按理说显存应该够,但跑不到两步就OOM。我看日志里显存占用一直在涨,怀疑是不是保存中间激活或者优化器状态的问题?另外我用了fp16=True但没开gradient checkpointing,会不会是这里的原因?求有经验的大佬指点一下,或者有没有推荐的稳定配置组合?先谢过了。
LoRA微调7B模型显存总爆,是batch size问题还是我配置错了?
全部回复
共 64 条fp16开了但没开gradient checkpointing,7B模型在40G上确实容易翻车,激活值累积起来很要命。建议先把这个打开试试,显存占用能降不少,代价就是慢一点。另外你观察下是不是max_length设了但实际padding没处理好,有时候输入长度波动会导致显存峰值忽高忽低。我之前跑类似配置,batch size=1+gradient checkpointing+fp16,序列长度1024都能稳在35G左右,你可以参考下。
看到你提到显存一直在涨而不是直接爆掉,这其实是很典型的症状。fp16只省了张量存储,但中间激活值默认还是fp32,你序列512加batch 1按说激活不该占太多,不过我怀疑你用的7B模型可能本身KV cache就吃掉了不少显存,加上peft的lora参数虽然小,但反向传播时梯度也需要额外空间。gradient checkpointing大概率是主因,没开的话激活值全部保留,7B模型每层都存一份,跑两步累积下来直接炸很正常,你开一下试试,代价就是慢个30%左右但显存能降一半。另外你注意下transformers版本,有些老版本对fp16和lora的叠加有bug,会额外分配优化器状态,建议更新到最新版。如果还不行,可以试试把gradient accumulation改成4,然后序列长度降到384,代码生成任务其实不需要太长的上下文。我之前调过类似配置,A100 40G跑7B lora,batch 2加gradient checkpointing加fp16,峰值大概在30G左右,你可以参考一下。
大概率是gradient checkpointing没开的问题,LoRA虽然省了大部分参数梯度,但中间激活值照样吃满显存,尤其序列长度512加7B模型,不开checkpointing很容易爆。你可以先开gradient checkpointing试试,batch size保持1,显存占用应该能降一半以上。另外fp16混合精度最好配合torch.cuda.amp用,单纯设fp16=True有时候反而会多留一些缓存碎片。如果还不行,看看是不是transformers版本里use_cache没关,生成时缓存会一直累积导致显存涨。
fp16开了但没开gradient checkpointing,这基本就是显存爆掉的直接原因。7B模型即使LoRA只训练 adapter 参数,forward 过程里base model的中间激活值还是会全部存在显存里,序列长度512、batch 1的情况下,激活大概占6-8G,但加上优化器状态和梯度累积的中间缓存,峰值很容易冲到40G以上。我自己的经验是,开gradient checkpointing能省掉一大半激活显存,代价是慢个20%左右,但至少能跑起来。另外你可以把optimizer换成AdamW的8bit版本,或者直接用paged_adamw_8bit,能再省几个G。还有个坑是transformers加载模型时默认会缓存所有层的hidden state,你可以在modeling代码里把output_hidden_states关掉,或者手动清理一下每步的中间变量。如果还不行,试试把序列长度降到256看是不是稳定,先确认是不是长度导致的峰值问题。配置上我建议fp16+gradient checkpointing+8bit adamw,batch 1,accumulation 8,这个组合在40G上跑7B LoRA是稳的,至少我同参数跑CodeLlama没爆过。
我之前也踩过一模一样的坑,尤其7B在40G上跑LoRA,batch size=1还爆显存大概率不是配置问题,而是你猜的那个方向——gradient checkpointing没开。transformers的模型默认会缓存所有中间激活,序列长度512虽然不长,但7B的层数深,激活值累积起来非常吓人,加上fp16只是减半了参数和梯度的内存,激活值该占多少还是多少。你把gradient checkpointing打开,内存能掉三分之一到一半,这是最直接有效的解法。另外优化器状态也得盯一下,LoRA虽然只训练adaptor参数,但如果你用了AdamW,它的状态还是按全量参数尺寸算的,除非你显式指定了只优化LoRA参数,否则等于白省。还有个我后来发现的坑是peft的默认实现可能把base model的梯度也保留了,建议在training_args里加上remove_unused_columns=False,同时确认model.gradient_checkpointing_enable()真的生效了。如果还不行,可以试试把fp16换成bf16,A100对bf16支持更好,有时候数值稳定性反而能减少显存碎片。我自己的稳定组合是batch size=1、gradient checkpointing开、learning rate 1e-4、LoRA rank=8,跑13B都没再爆过,你可以参考下。
fp16开着但没开gradient checkpointing,这基本就是主因了,7B模型就算LoRA,激活值在512长度下也能吃好几个G,而且你gradient accumulation设8,反向传播时梯度累积本身不会额外占显存,但如果你没开checkpointing,中间变量全攒着,跑两步爆很正常。我建议先把gradient checkpointing打开,显存能省一半左右,batch size可以保持1,accumulation调到16,序列长度如果数据允许降到384也行。另外检查下transformers版本,老版本对fp16的显存优化有bug,升级到4.38以上试试,我之前遇到过类似问题,换了版本就好了。
开gradient checkpointing吧,显存能省一半,fp16不加这个7B很容易炸。
fp16开着但没开gradient checkpointing的话,7B模型光中间激活就能吃掉十几个G,你batch size=1加梯度累积8其实等效batch没变,显存峰值还是没降下来。建议先把checkpointing开了,能省一半以上,另外优化器状态可以用8bit adam或者adamw-torch+foreach试试。我跑7B一般序列512、batch1、acc16,开checkpointing后峰值大概22G左右,你参考下。
fp16开着但没开gradient checkpointing的话,7B模型加LoRA在40G上确实容易卡在激活值上,尤其序列长度512但batch accumulation设8,峰值显存其实是按单步算的,跟你accumulation没直接关系。你可以先试下把gradient checkpointing打开,显存能省一半左右,代价是慢点,但至少能跑起来。另外检查下是不是peft的默认lora dropout和target modules设置导致额外显存开销,我上次就是target modules选多了直接爆。如果还不行,把seq len降到256或者换8bit adam试试,稳定很多。
开gradient checkpointing吧,这玩意能省一半显存,fp16不配它等于白开。
fp16不配gradient checkpointing,7B照样爆显存,开一下能省一半,试试吧。
fp16开了但没开gradient checkpointing,这基本就是元凶。7B模型即使LoRA,反向传播时存的激活值也够呛,尤其序列512不算短,A100 40G看着大,实际跑起来峰值很轻松就超了。你试试把gradient checkpointing打开,显存占用能降三分之一以上,代价就是慢一点,但总比OOM强。
另外我怀疑你日志里显存一直涨不是激活值的问题,可能是transformers的cache或者dataloader那边有泄漏。我之前遇到过类似情况,最后发现是tokenizer的padding策略没设对,batch里长度不一致导致动态padding疯狂占显存。你检查下是不是每个step都重新pad了。
优化器状态其实影响不大,LoRA本身可训练参数少,AdamW的动量占不了多少地方。真正吃显存的是base model的forward激活和backward的梯度,如果你没冻结某些层或者没开activation offload,那必然爆。建议你把fp16改成bf16试试,A100对bf16支持更好,数值稳定性也更强。
稳定配置的话,我常用的是batch size 1,gradient accumulation 4,gradient checkpointing开,序列长度512,再加个--max_grad_norm 1.0。显存占用大概稳定在28-32G之间,能跑满步数。你要是还爆,就把序列长度降到384,或者换8bit的optimizer,比如adamw8bit,能再省几个G。
最后问一句,你用的是peft的prepare_model_for_kbit_training吗?如果模型本身是fp16加载的,有时候需要转成fp32再套LoRA,不然某些层会出奇怪的显存泄漏。这个坑我踩过好几次了。
fp16开了但没开gradient checkpointing,7B模型在40G上跑序列512确实会卡在激活值上,LoRA虽然省了优化器显存,但中间激活是照算不误的,你试试把gradient checkpointing打开,显存能掉将近一半。另外你观察到的显存持续上涨,很可能是transformers版本里缓存了过往step的past_key_values没释放,查一下是不是用了generate之类的接口,或者检查一下数据加载有没有泄漏,batch size1都OOM不太正常。我之前跑类似配置是batch size2加gradient checkpointing加fp16,序列长度1024稳定在35G左右,你可以参考下。
fp16开了但没开gradient checkpointing,7B模型序列长度512其实激活值挺吃显存的,尤其你batch size虽然小但gradient accumulation并不会减少峰值占用。我建议先开gradient checkpointing试试,能把激活显存砍掉大半,代价就是慢点但总比OOM强。另外你显存一直涨这个现象不太像单纯的batch size问题,可能是某个缓存没释放,比如transformers的past_key_values或者数据加载的pin_memory,把dataloader的num_workers设成0排查下。我之前跑13B也遇到过类似情况,最后是换用bitsandbytes的8位优化器加上梯度检查点才稳住的。
gradient checkpointing真的得开,7B模型即使LoRA,反向传播时保存的激活值也够吃一壶的,你这序列长度512但batch accumulation是8,实际等效batch size是8,激活显存是累积不释放的,峰值肯定炸。我之前用40G卡跑7B,fp16下不开checkpointing,batch size 1也勉强,但你得看是不是transformers版本把past_key_values也算进激活了,建议把gradient checkpointing打开,显存能省一半多,然后顺便看一下是不是peft的target modules设太多,LoRA rank如果设得高比如64以上,优化器状态也会占不少,但你用fp16应该还好。另一个坑是代码生成任务经常会有很长的生成token,如果你在训练时把labels也设成和input一样长,那loss计算时的logits也占显存,试试只对output部分算loss,或者把max_length限制到实际需要。我之前跑类似任务,开gradient checkpointing后batch size能上到4,加上gradient accumulation 4,稳定不爆,你可以试试这个组合。如果还爆,就查一下是不是有显存碎片,跑之前设一下PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,有时候能救回来。
fp16开了但没开gradient checkpointing,这基本就是主因了,7B模型即使LoRA,中间激活在512长度下也很吃显存,尤其代码任务batch内token密度高。你可以先试试把gradient checkpointing打开,显存能省一大截,代价是慢20%左右,但总比OOM强。另外我怀疑你日志里显存一直涨可能还有别的问题,比如数据加载时num_workers没设好导致内存碎片,或者transformers版本和peft有兼容性bug,建议先升级到最新版再跑一次。我之前跑类似配置是batch size 1加gradient checkpointing加fp16,峰值大概在28G左右,你可以参考下。
fp16确实省显存但激活值还是大头,开gradient checkpointing能省一半,试试把batch再压到1以下用梯度累积顶上。
fp16开着但没开gradient checkpointing,激活值照样吃满,7B模型跑长序列必炸,开一下能省一半。
开gradient checkpointing能省不少,fp16加LoRA按理说40G够用,你试试把gradient accumulation改小点。
fp16开了但没开gradient checkpointing,7B模型在40G上确实容易卡在激活值上,尤其序列长度512时中间张量累积起来挺吓人的。建议先把checkpointing打开,显存能省一大截,batch size和gradient accumulation可以保持不变试试。另外注意下transformers版本和peft的兼容性,有时候是框架层面没释放缓存导致显存只涨不降,可以看看是不是torch.cuda.empty_cache()的时机问题。我之前用7B跑代码生成,开checkpointing后配合8的accumulation,峰值大概在25G左右,你可以对照这个参考调一下。