最近在试着用LoRA微调7B的LLaMA模型,参考了网上一些教程,batch size设到1,gradient checkpoint也开了,用的还是bf16,结果3090(24G)还是跑着跑着就OOM了。我看别人同样的配置能跑起来,难道是我tokenizer的max_length设太长(2048)?还是说attention机制里有些隐藏的显存占用我没注意到?另外,我用的transformers库版本是4.31,不知道是不是版本问题。有没有老哥能分享下实际能跑的配置,或者推荐一些更省显存的trick?提前谢谢了,刚入门大模型,真的有点懵。
用PyTorch跑LLaMA微调,显存总爆,到底是哪里设置不对?
全部回复
共 170 条24G跑7B LoRA其实挺极限的,你max_length设2048确实是最大的显存杀手之一。我试过把max_length砍到1024,同样配置下显存占用直接从23G降到16G,效果其实差不了太多,毕竟微调时序列长度对生成质量的影响没那么夸张。另外transformers 4.31有个已知的attention mask实现问题,会额外分配一部分临时张量,建议你升到4.35以上,或者干脆用peft库自带的prepare_model_for_kbit_training把梯度检查点做得更细一点。还有个容易忽略的地方是optimizer状态,就算用AdamW,LoRA参数本身虽然少,但如果你把bias和layer_norm也设成可训练,那显存还是会被这些全参部分吃满,建议把trainable_params只锁在LoRA的A和B矩阵上。最后实在不行就上gradient accumulation配合更小的per_device_train_batch_size,比如0.5这种伪批处理,虽然训练慢点但能稳定跑完。我刚入门那会儿也卡在这,后来发现是dataloader的pin_memory默认开了,多线程预加载也会临时占几个G,关掉又能省一截。
我之前也遇到过一模一样的情况,后来发现是max_length=2048太吃显存了,尤其7B模型加上attention的KV cache,24G真的扛不住。你可以试试把max_length降到1024或者512,显存占用能降一大截。另外transformers 4.31确实有点老,有些显存优化是后面版本才加的,建议升到4.38以上,说不定直接就好了。还有个冷门trick,把flash attention打开,能省不少显存,就是得装对应依赖。要是还不行,看看是不是pytorch版本和CUDA不匹配,有时候这个也会导致显存分配异常。
24G跑7B的LoRA按理说够用,但max_length=2048确实有点顶,我之前用1024才稳,你可以先砍到512试试,顺便确认下是不是数据集里真的有那么多长样本。transformers 4.31我记得有个版本对LLaMA的attention缓存处理有bug,换4.35+或者直接上最新版,能省不少显存。另外你检查下是不是把gradient_checkpointing开在model.enable_input_require_grads()之前了,顺序不对等于没开。实在不行就上8bit优化器加CPU offload,虽然慢点但至少不爆。
max_length改成512试试,7B用2048就算LoRA也吃紧,我这么调完稳得很。
我最近也碰到过类似的情况,24G跑7B按理说够用,但max_length拉到2048确实容易爆,尤其是长序列下attention的中间激活值会翻倍涨。你可以试试把max_length先砍到1024,或者用gradient_checkpointing配合显存清理,另外transformers 4.31有个已知的显存泄漏问题,升级到4.35+可能就稳了。还有个土办法是开torch.utils.checkpoint把input_embeds也包进去,能省不少。
我之前也卡在这过,后来发现max_length设成2048在7B上确实挺吃紧的,尤其LoRA虽然省了主参数梯度,但激活值还是按全量算的。建议先降到1024试试,或者用gradient accumulation把batch size拆成更小,比如显存只够跑一步就accumulate个8步。另外transformers 4.31有个已知的attention mask问题,会多占显存,升级到4.35+能好不少。我最后是开了flash attention(如果显卡支持)才稳住的,你可以查一下你那个3090的驱动版本,老驱动可能不支持。
max_length砍到1024试试,3090跑7B LoRA这个长度挺极限的,4.31版本也有点旧了。
我最近也在折腾这个,7B+LoRA在24G上确实有点极限。max_length拉到2048会直接让激活值爆炸,我降到1024之后明显稳了很多,你可以先试试这个。另外transformers 4.31的attention实现有点老,换4.36以上版本能吃到更好的flash attention优化,显存能省不少。如果还爆,试试把LoRA的rank降到8,或者用8bit的AdamW优化器,这几个组合下来基本能跑通。
max_length设2048确实是个大头,7B模型在bf16下光KV cache就要吃掉不少显存,你试试把max_length砍到1024或者512,很多情况下OOM直接就消失了。另外transformers 4.31版本对LLaMA的支持其实有点尴尬,建议升到4.35以上,我之前用旧版跑llama2也遇到过奇怪的显存泄漏,换版本后就好了。gradient checkpointing开了的话,记得把input和output的显存占用也估算进去,有时候forward过程里activations反而比backward更吃显存。还有个冷门trick,把attention的dropout关掉,训练时能省一点显存,虽然影响很小但聊胜于无。另外你确认下是不是真的用了LoRA而不是全量微调,有些教程代码里没写对,导致模型参数还被冻结在GPU上。最后实在不行就试试8bit量化加载模型,配合LoRA能再省一半,24G跑7B应该很宽松,就是速度会慢些。我自己的经验是,batch size=1+max_length=1024+bf16+gradient checkpoint,3090跑7B完全没问题,你先把这几个条件固定下来再排查别的。
我最近也遇到过类似问题,最后发现是max_length惹的祸,2048对7B来说真的有点顶,降到1024或者用动态padding能省不少。另外你试试把flash-attention装上,显存占用能掉一截,transformers 4.31对flash attn支持不太好,升到4.35+会稳很多。还有个冷门技巧是关掉gradient checkpointing里那个use_reentrant参数,有时候反而更省。实在不行就把LoRA的r值调小到8或者4,效果差不了太多但能救回来。
3090跑7B LoRA按理说24G是够的,但max_length拉到2048确实挺吃显存,我一般习惯设1024,微调时对大多数任务影响不大。另外你可以试试4-bit量化加载base model,配合peft的prepare_model_for_kbit_training,显存能再省一截。transformers 4.31有点旧了,换到最新版试试,有些版本的attention实现确实有额外缓存。还有个小技巧是优化器用8-bit Adam,或者干脆用SGD,能省下不少状态显存。你跑的时候留意下是爆在forward还是backward阶段,可以用torch.cuda.max_memory_allocated()看下峰值到底在哪一步。
max_length砍到1024试试,LoRA加gradient checkpointing应该能压下来,transformers升到4.35以上有不少显存优化。
max_length设2048确实挺吃显存的,尤其是7B模型即使LoRA也要预留不少激活值。你试试把max_length砍到1024,或者用gradient_accumulation_steps配合小batch,体感能省下2-3G。另外transformers 4.31有个已知的attention mask显存泄漏问题,升到4.35+或者换peft最新版可能直接解决。我跑13B时还开了torch.utils.checkpoint配合gradient_checkpointing,不然真扛不住。
说实话我第一反应也是max_length,2048确实挺夸张的,LoRA虽然冻结了原模型,但attention的中间激活值跟序列长度是二次方关系,你这等于直接把显存大头锁死了。我之前拿A6000跑7B,max_length设1024,batch 1,开gradient checkpoint,峰值能到20G左右,你3090要是再跑个eval或者累积梯度,爆掉太正常了。建议先把max_length砍到512试试,如果任务允许,甚至256都行,很多教程默认128也能收敛。另外transformers 4.31有点老了,新版本对flash attention的支持更好,你可以试试4.38以上,开attn_implementation="flash_attn_2",显存能再省一截。还有个坑是gradient_checkpointing要配合input_ids的requires_grad用,有些教程没提这个,导致checkpoint没生效,你检查下训练循环里有没有把model.gradient_checkpointing_enable()放在正确位置。最后如果还爆,试试optimi的8位优化器,或者干脆用torch.utils.checkpoint手动包住attention层,能榨出不少空间。
max_length设到2048确实是个大头,7B模型在bf16下光attention的中间激活值就很吃显存,LoRA虽然省了优化器状态,但激活值一点没少。你可以试试把max_length砍到1024或者512,很多场景其实用不到那么长的上下文,尤其微调指令数据的话512基本够了。另外transformers 4.31的LLaMA实现确实有些已知的显存效率问题,建议直接升到4.38以上,新版对SDPA的支持能省不少内存。还有个小技巧,gradient checkpointing开了之后记得把input batch size再压小一点,比如设成0或者1,虽然会慢一些但能稳定跑完。我之前也遇到过类似情况,最后发现是DataLoader的pin_memory和num_workers开太多导致额外的显存碎片,关掉之后立刻好了。你要是还爆,就试试accelerate的cpu offload,把部分层放到内存里,速度损失可以接受。另外检查一下是不是用了model.enable_input_require_grads(),这个会强制保留所有输入梯度,对LoRA来说是多余的。
max_length降到1024试试,LoRA的target_modules得选对,不然白省显存。
max_length 2048确实有点猛,7B模型光attention的KV cache就够吃几个G了,你可以先砍到512试试,跑通再慢慢加。另外transformers 4.31对llama的支持确实有点老,换4.36以上版本,很多显存优化是自动启用的。还有个小技巧,把gradient checkpointing配合torch.utils.checkpoint一起用,能再省一截。我自己的经验是,3090跑7B LoRA,batch size 1加max_length 1024,大概能压到15G左右,留点余量给反向传播。
24G跑7B LoRA按理说够用,你先查下是不是max_length设2048太长,很多教程默认512或1024,长文本的激活值吃显存特别狠。另外transformers 4.31有点老,换4.38以上试试,新版本对attention的显存优化明显。我自己的经验是,把lora的target modules只放q和v,别全加,能省不少,还有optimizer换成adafactor,比adamw省一半显存。你跑起来的时候看下nvidia-smi,是峰值爆了还是一直涨,一直涨可能是数据加载泄漏,那就不关模型事儿了。
我之前也卡在同样的问题上,24G跑7B LoRA真的没有想象中那么宽裕。你max_length设2048确实是个大头,很多教程默认是512或1024,长序列下attention的显存占用是平方增长的,这个影响比batch size还明显。另外transformers 4.31对LLaMA的attention实现确实比较老,建议升到4.35以上,新版用flash attention能省不少,而且对LoRA兼容性更好。
还有个容易忽略的点是你有没有把gradient checkpointing真正生效,光开那个flag不够,得确认模型在forward时确实用了checkpoint函数,有时候包装层没包对就白开了。我自己的经验是把max_length降到1024,然后加上gradient accumulation,虽然速度慢点但稳得很。另外你用的LoRA rank是多少?如果设到16或32,那个增量矩阵也会吃显存,降到8试试,效果差别不大但省很多。
还有个小技巧,可以把optimizer换成AdamW的8bit版,bitsandbytes库那个,能省2-3G。你跑之前看一眼nvidia-smi,如果显存从开始就接近满了,那大概率是模型加载阶段就超了,可以试试用device_map="auto"让部分层跑到CPU,或者直接load_in_8bit量化基座模型,LoRA本身对精度影响不大。我最后就是用8bit基座加rank=8加max_len=1024,3090跑得挺稳,你可以试试这个组合。
说实话max_length 2048在7B上确实挺吃紧的,LoRA虽然省了训练参数但激活值照样按序列长度算,试试把max_length砍到1024甚至512,显存能掉一大截。另外transformers 4.31的attention实现确实有点老,建议升到4.38以上,flash attention能自动启用,省不少显存。我跑7B一般开gradient checkpointing + paged optimizer,batch size 1,max_length 1024,24G稳稳的,你可以先照这个调调看。