最近在试着用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来说长序列的activation特别吃显存,先砍到1024试试,很多教程默认512不是没道理的。另外transformers 4.31配peft确实有已知的显存泄漏问题,建议升到4.36以上,或者直接换用unsloth这个库,同样配置能省一半显存。还有个小trick是开flash attention,3090支持的话能显著降峰值,代码里加一句attn_implementation="flash_attention_2"就行。
max_length设2048确实是显存杀手,7B模型光attention的KV cache就得吃好几个G,可以先砍到512试试。transformers 4.31问题不大,关键是看下是不是pytorch的显存碎片化太严重,试试开PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128。另外LoRA的target_modules别全上,只改q和v能省不少,还有optimizer换成AdamW的8bit版或者干脆用SGD。我自己的配置是batch_size=1,grad_accum=8,max_len=1024,跑13B都稳,你参考下。
max_length砍到512试试,我当初2048也爆,降到1024立马稳了。
transformers升到4.35+,老版本对LLaMA的显存优化差不少。
max_length砍到1024试试,3090跑7B只开gradient checkpoint不太够,还得配合gradient accumulation。
我最近也踩过这个坑,max_length设2048确实太狠了,7B模型光attention就是平方级的开销,降到1024或者512立马就不一样。另外你可以试试用gradient accumulation模拟大batch,配合梯度裁剪,别小看这个,有时候比开checkpoint还管用。transformers 4.31有点老,升到4.35以上对LLaMA的显存优化做了不少改进,尤其是flash attention支持,能省不少。还有个小技巧,把optimizer换成8-bit AdamW或者用bitsandbytes的4bit量化基座,显存占用能再压一截。
max_length砍到1024,再换上flash attention试试,显存能省不少。
max_length降到1024试试,很多教程没提这个,加上gradient checkpoint要配合output_attentions=False才有效。
24G跑7B LoRA按理说够用了,问题大概率出在max_length上,2048太长,试试512或768,显存占用直接砍半。另外transformers 4.31有点老,换4.38以上版本,很多显存优化是后加的。还有个小技巧,把gradient_checkpointing和optimizer并行加载开起来,能省不少峰值内存。我上次跑13B也是这么调下来的,你先试试这个组合。
max_length砍到1024试试,LoRA加gradient checkpointing在24G上极限也就这样了。
3090跑7B按理说24G够用,但max_length设2048确实有点狠,LoRA虽然省了trainable参数,但activation和attention的显存还是跟序列长度直接挂钩的,可以先砍到512试试。另外transformers 4.31对LLaMA的attention实现有点老,建议升到4.35+,顺便检查下是不是意外加载了完整base model的optimizer state。我自己的经验是加上gradient_accumulation_steps=8配合max_length=1024,再开个flash_attention(如果环境支持),基本能稳在19G左右。你那个OOM是出现在step刚开始还是跑到一半?如果是后者,可能还有hidden_state缓存没清干净的小bug。
max_length砍到1024立马见效,transformer版本升到4.36以上,很多坑都是版本太老挖的。
我之前也遇到过一模一样的情况,后来发现是max_length的问题,2048对7B来说太奢侈了,降到1024甚至512会好很多。另外你可以试试优化器用adamw_8bit或者干脆用SGD,显存占用能降一截。还有个小技巧,把attention的flash attention打开,transformers 4.31应该支持,能省不少内存。版本倒不觉得是主因,但建议升到4.35以上,有些显存优化是后加的。我刚跑通的时候也是各种爆,调了两天才稳,别急。
24G跑7B LoRA按理说是够的,但max_length=2048确实是个隐形杀手,序列长度对显存的影响是平方级的,你试着把max_length砍到1024甚至512,显存占用能掉一大截。另外你检查下是否真的把梯度检查点应用到了所有层,某些教程只开个开关但没在模型配置里设use_gradient_checkpointing=True,等于白开。还有一个容易忽略的地方是optimizer状态,AdamW的动量项在bf16下也吃显存,试试8bit的AdamW或者干脆用SGD with momentum,虽然收敛慢但省得多。transformers 4.31有点旧了,4.36之后对LLaMA的attention实现做了优化,建议升到4.40+,顺便把flash attention开起来,这玩意儿能省不少显存。如果你用的是HuggingFace的Trainer,记得把gradient_accumulation_steps设大一点,比如8,这样虽然单步batch小但整体步数不变,实际显存压力会小很多。最后检查下是不是dataloader里把整个数据集都加载进内存了,有时候不是模型爆而是数据预处理阶段爆的。我刚跑通7B的时候也踩过这些坑,调完这些基本24G能稳跑。
max_length设2048确实挺吃显存的,LoRA虽然省了 optimizer 状态,但激活值还是按序列长度线性涨的,你可以试试把 max_length 砍到1024,或者用 gradient accumulation 模拟更大 batch 但保持序列短点。另外 transformers 4.31 的 attention 实现可能没走 SDPA,换 4.36+ 或者手动开 torch.compile 说不定能省不少。我自己的经验是,7B 模型在 24G 上跑 LoRA,序列长度1280 左右比较稳妥,再高就得用 sequence packing 或者 offload 了。还有个小 trick,把 flash attention 打开,显存能降个两三G,你检查下是不是没装 flash-attn 库。
24G跑7B LoRA按理说是够的,问题大概率出在max_length上,2048的序列长度对显存压力比batch size还大,你可以先砍到1024试试。另外transformers 4.31确实有点老,4.38之后对LLaMA的attention实现优化了不少,升级一下说不定就稳了。还有个冷门trick是关掉flash attention,某些版本下它反而更吃显存,换成sdpa模式能省不少。我上次跑类似的配置是max_length 1024,batch size 2,gradient checkpoint开着,加上zero3 offload,勉强能塞进24G,你可以参考下。
把max_length砍到512试试,bf16下长序列的attention缓存特别吃显存,另外transformers升到4.35以上版本,旧版对梯度检查点支持有坑。
max_length砍到1024试试,bf16在3090上有时不如fp16稳,transformers升到4.35+也有奇效。
同样配置跑不起来太正常了,transformers 4.31对LLaMA的支持还不算成熟,很多显存优化是在后续版本才补上的,建议先升到4.38以上再试。max_length设2048确实是个坎,LoRA虽然只训练 adapter 参数,但前向/反向的激活值还是按完整序列长度算的,24G卡跑7B加2048上下文,稍不留神就爆。你可以试试把 max_length 砍到1024看看能不能稳定跑完一个step,如果能跑通那就确认是序列长度的问题。另外检查一下 gradient_checkpointing 是不是真的生效了,有时候和 model parallel 或者某些缓存机制冲突会静默失效,可以打印模型显存占用对比一下。还有个容易被忽略的点是 optimizer 的状态,AdamW 的动量项在 LoRA 里虽然小,但如果你没对非 trainable 参数做冻结处理,它们还是会被算进 optimizer 里。我自己的经验是,用 bitsandbytes 的 4bit 量化加载 base model,LoRA 用 8bit 的 adapter,配合 paged_adamw,24G 跑 2048 长度基本稳。另外检查下是否开了 flash attention,如果你的环境支持直接用,能省不少内存。版本别停在4.31,至少换到4.36以上,有些显存泄漏的修复就在那几个版本里。
我最近也刚跑完类似的活儿,7B+LoRA,卡是4090,但一开始也爆。你提到max_length设2048,这个确实是个大头,序列长度对显存是二次方影响,不是线性的,尤其attention的中间激活值。我试过把max_length砍到1024,显存占用直接掉了快40%,如果任务不是特别需要长上下文,可以先压到1024试试。另外检查一下你是不是把gradient_checkpointing开在了model.enable_input_require_grads()之前,顺序不对有时候不生效,还有optimizer选的是AdamW的话,它的状态量也吃不少显存,可以换8bit的adam或干脆用SGD(效果略差但省显存)。transformers 4.31应该没问题,我用的4.35,但感觉不是版本坑。还有个容易被忽略的点:LoRA的target_modules如果设了全部线性层,比如q,k,v,o,gate,up,down,那参数和激活都会涨,只设q,k,v能省不少。最后实在不行就开flash-attention,虽然要编译,但显存能再降一截。你那个3090跑7B理论上是够的,多半是序列长度加target_modules太肥了。
max_length2048确实顶不住,试试1024加梯度累积,或者换flash-attention能省不少。