最近在试着用LoRA微调7B的LLaMA模型,参考了网上一些教程,batch size设到1,gradient checkpoint也开了,用的还是bf16,结果3090(24G)还是跑着跑着就OOM了。我看别人同样的配置能跑起来,难道是我tokenizer的max_length设太长(2048)?还是说attention机制里有些隐藏的显存占用我没注意到?另外,我用的transformers库版本是4.31,不知道是不是版本问题。有没有老哥能分享下实际能跑的配置,或者推荐一些更省显存的trick?提前谢谢了,刚入门大模型,真的有点懵。
用PyTorch跑LLaMA微调,显存总爆,到底是哪里设置不对?
全部回复
共 170 条max_length设2048确实有点猛,7B模型本身吃显存就厉害,可以试试先砍到1024或者用梯度累积来变相降batch。另外transformers 4.31有个已知的attention计算冗余问题,升到4.35以上能省不少显存。我自己跑的时候还把gradient checkpoint再开一层,配合deepspeed的zero2,24G跑7B loRA基本稳住了。
max_length设2048确实有点高,7B模型在24G显存上跑这个长度挺吃紧的,建议先降到1024或者512试试看。另外transformers 4.31有个已知的attention缓存问题,升到4.35以上能省不少显存,我当时也是升级后直接跑起来的。还有个trick是试试gradient accumulation,虽然batch size=1但多累积几步也能等效大batch,显存压力小很多。
max_length砍到1024试试,或者检查下tokenizer是不是把pad设成eos了,这俩坑我踩过。
max_length砍到1024试试,seq_len对显存影响是二次方的,24G跑2048确实悬。
max_length降到1024试试,我3070跑7B就是这么救回来的,flash attention也能省不少。
我之前也卡在同样的问题上,3090跑7B LoRA按理说应该够用,但OOM往往不是单一原因。你max_length设2048确实偏激进,训练时序列长度会直接决定激活内存峰值,试下把max_length砍到1024甚至512,显存占用能掉一大截。另外transformers 4.31有个已知问题,就是Llama的attention实现里会额外缓存past_key_values,哪怕你只做微调不生成,这个缓存也会占空间,升级到4.36以上或者直接用peft的官方示例代码会好很多。还有个隐藏坑是gradient checkpointing和bf16的交互,有些版本下checkpoint会失效,你可以在loss.backward()前打印一下内存变化,看是否真的在省显存。最省事的方案是直接换unsloth优化版的LoRA,它重写了attention和计算图,同样24G能跑batch size 2甚至4,速度还快不少。你要是想继续用原生transformers,可以把优化器换成adamw_8bit,再把gradient_accumulation_steps设小点,省下来的显存足够撑住2048长度了。我自己的配置是batch size=1,seq_len=1024,8bit优化器加gradient checkpoint,稳定跑完3万步没炸过。
24G跑7B LoRA理论上是够的,但max_length 2048确实有点激进,我试过降到1024显存占用直接少一半。另外transformers 4.31有个已知的attention mask显存泄漏问题,建议升到4.35以上,顺便把gradient_checkpointing的use_reentrant设成False试试。还有个小技巧,把优化器换成Adafactor,能省不少显存,虽然收敛慢点但能跑起来最重要。
max_length砍到1024试试,7B全量微调本来就不轻松,LoRA也扛不住长序列。
max_length 2048确实挺吃显存的,7B模型就算LoRA,seq len翻倍显存占用也接近指数涨,可以先砍到1024试试。另外你用的transformers 4.31有点旧了,换到4.35+对LLaMA的attention实现优化不少,能省一截内存。还有个冷门技巧是给模型加个gradient_checkpointing_kwargs把use_reentrant设成False,有时候能省出几个G。实在不行就把LoRA的r值降到8,或者用QLoRA直接4bit加载,24G跑7B完全够。
max_length砍到1024试试,顺手把attn实现换成flash-attention,显存能省一大截。
max_length降到512试试,八成是序列长度把显存吃满了,很多教程都没提这坑。
max_length砍到1024试试,bf16下7B用24G本来就紧,开gradient checkpoint还得把batch size压到1以下。
max_length砍到1024基本能救,4.31确实有显存泄漏的坑,换4.36试试。
我之前也卡在这过,后来发现max_length从2048砍到1024,显存占用直接掉了快6G,你这设置大概率是罪魁祸首。另外transformers 4.31确实有点老,有些算子没优化,升到4.38以上能省不少显存。还有个冷门trick是把attention的dropout关掉,推理时无所谓,但训练时能挤出一小块缓冲。别太迷信别人的配置,3090跑7B LoRA本来就紧巴巴,实在不行就用8bit的QLoRA,效果差不了太多。
max_length砍到512试试,LoRA微调根本不用那么长,省下的显存立竿见影。
我最近也在折腾这个,7B加LoRA按理说24G应该够的,但你这情况我也遇到过。max_length设2048确实有点激进,序列长度对显存的影响是二次方的,试试砍到1024或者512,效果可能立竿见影。另外transformers 4.31有个已知问题,就是attention的中间激活没被正确释放,尤其配合gradient checkpointing的时候,建议直接升到4.35以上,我换了之后明显稳很多。还有个容易被忽略的点,你检查下是否真的把model parallel或者device_map设对了,有时候默认会把模型均匀塞到多个设备上,反而产生额外的显存碎片。再就是LoRA的target_modules别一股脑全加上,只改q_proj和v_proj能省不少,虽然效果略降但能跑起来才是关键。最后建议你监控一下nvidia-smi,看是不是显存碎片化严重,如果是,试试在训练循环前加个torch.cuda.empty_cache(),虽然治标不治本,但有时候能撑过峰值。我这边实测7B用2048长度,batch size=1,开gradient checkpoint,不加其他优化,峰值大概在20G左右,你如果还爆,可能就是代码里有什么地方不小心把张量复制了一份,比如在loss计算时用了detach后没回收。
我之前也卡在这步,折腾了快一周才明白问题多半不在LoRA本身。你max_length设2048确实偏长了,7B模型在bf16下光输入序列的激活值就够吃满好几个G,就算开了gradient checkpoint,attention的KV cache还是会随长度线性涨,建议先砍到1024试试,显存能立刻松一大截。另外transformers 4.31有个已知问题,就是跟新版peft搭配时,base model的forward会重复计算某些中间张量,LoRA层虽然只训练一小部分参数,但推理路径上的显存峰值没降下来,换个4.35以上的版本可能就顺了。还有个冷门trick是检查tokenizer的padding侧,如果没有设成左侧,batch内长度差异大时,右侧padding会导致计算图里多出一堆无意义的attention位置,白白占显存。我自己的跑法是batch=1,max_length=1024,gradient checkpoint开,再加一个4-bit量化加载模型,把LoRA的r设成8,alpha=16,这样3090上能稳定跑完一个小epoch。你如果还爆,可以试试在dataloader里用pin_memory=False,有时候内存碎片也会间接影响显存分配。版本问题优先排查,真的能省不少事。
这配置看着没啥大问题,但我赌五毛是你max_length的锅,2048对7B来说太狠了,LoRA虽然省了优化器显存,但激活值照样吃满。你可以试试把max_length砍到1024,或者用gradient accumulation模拟更大batch,显存曲线会平缓很多。另外transfomers版本确实可能影响显存分配,4.31有点旧了,建议升到4.38+,新版本对attention的显存优化明显。我自己的经验是,24G跑7B微调,max_length限制在1536以内,加上unsloth那套动态padding,基本能稳定跑完。
我最近也踩过这个坑,问题很可能出在max_length=2048上,7B模型配24G卡跑这个长度确实极限,先砍到1024试试,显存占用直接掉一截。另外transformers 4.31有个已知的attention mask显存浪费问题,建议升到4.35+,LoRA的target_modules记得把q_proj和v_proj都加上,别只改一个。还有个冷门trick:把gradient_checkpointing配合use_cache=False一起用,能省不少峰值显存。我最后是batch size=1、max_length=1024、8-bit AdamW才稳定跑完的,你可以试试。
看到你提到max_length设2048,这个基本就是显存杀手了,7B模型在bf16下光激活值就很吃紧,LoRA虽然省了优化器状态但attention的中间张量还是按序列长度平方算的,建议先砍到512试试,很多教程为了省事没提这个细节。另外transformers 4.31确实有点老,后面版本对flash attention和显存碎片优化了不少,但升级前最好确认下和peft的兼容性。还有个容易忽略的点是gradient checkpointing要配合input batch内动态padding,如果你没做padding的话,短序列也会按2048分配空间,白白浪费显存。我自己的经验是哪怕batch size为1,长序列下也要把gradient_accumulation_steps调大,用时间换空间,而且可以试试torch.compile,有时候能压掉不少峰值显存。最后建议你用nvidia-smi监控一下是哪个时刻显存突然飙高,多半是forward时把整条序列的kv cache都存下来了,如果开了gradient checkpoint理论上不该这样,可能你踩了某个模型的bug,换成最新的accelerate和transformers组合说不定就解决了。