最近在试着用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砍到1024试试,bf16在3090上可能没生效,检查下transformers版本和flash attention。
max_length设2048确实挺吃显存的,7B模型光KV cache就得占好几个G,你可以先试试把max_length砍到1024或者512,看看OOM是不是立刻缓解。另外transformers 4.31对LLaMA的attention实现有点老,建议升到4.35以上,新版用flash attention能省不少显存。我自己的经验是LoRA的r设小点(比如8),target modules别全选,只挑q和v,能再挤出一部分空间。还有个小技巧,把gradient_checkpointing和optimizer的offload一起开,虽然慢点但稳定不少。
24G跑7B LoRA按理说是够的,问题八成出在max_length=2048上,序列一长激活值直接爆炸,试试把长度砍到1024甚至512,显存能省一大截。另外transformers 4.31对LLaMA的attention实现有点老,建议升到4.35以上,顺手把gradient_checkpointing_kwargs里的use_reentrant改成False,有时候这个会影响显存释放。还有个冷门技巧是给优化器加个offload参数,或者用paged_adamw,能再挤点空间出来。
max_length砍到1024试试,LoRA用4bit量化能省不少,24G跑7B很稳的。
gradient checkpointing开了的话,attention的缓存也得清,transformers升到4.35以上试试。
24G跑7B LoRA按理说够用,但max_length拉到2048确实有点狠,sequence length对显存的影响是平方级的,试试砍到1024或者512,能省出一大块。另外transformers 4.31对LLaMA的支持不算太好,建议升到4.35以上,很多显存优化是后加的。我怀疑你OOM不是单点问题,而是attention的中间激活值没被gradient checkpoint完全覆盖,可以检查下是不是只开了model的checkpoint,没管input的缓存。再不行就上unsloth吧,它把LoRA的前向过程重写了,显存占用能再降三四成。
我之前也遇到过一模一样的情况,3090跑7B LoRA按理说够用,问题大概率出在max_length=2048上,序列一长中间激活值直接爆炸,试试把长度砍到1024或者512,显存立刻松快很多。另外transformers 4.31的attention实现确实有点吃显存,建议升到4.38以上,或者手动开一下use_flash_attention_2,能省不少。还有个野路子,把LoRA的r值从8降到4,效果损失不大但显存压力小一截,你可以先跑通再慢慢调回去。
max_length设2048确实是显存大头,7B模型在bf16下光输入embedding就要占不少,但更关键的是attention的中间激活值,这个长度下batch size=1也可能爆。你试试把max_length砍到1024甚至512,先跑通再说,长度对LoRA微调效果影响没那么大,后面可以再渐进式加回去。另外gradient checkpointing开了的话,记得确认一下是否真的生效了,有时候transformers版本或者模型配置里没正确传参,它会静默失效,你可以在训练日志里看有没有gradient checkpointing相关的提示。我怀疑你那个transformers 4.31可能有bug,之前4.28到4.30之间有些版本的attention实现会额外缓存,建议直接升级到4.35以上,或者试试用peft库自带的prepare_model_for_kbit_training,它会把模型转成更省显存的格式。还有个trick,你可以把优化器换成分片优化器,比如bitsandbytes的8位Adam,能省下大概2-3G的优化器状态显存,但要注意梯度累积步数得调大一点。最后,如果还是爆,看看是不是pytorch的缓存碎片问题,可以在每个step之后清一下torch.cuda.empty_cache(),虽然治标不治本但有时候能多撑几个step。我自己的经验是7B模型在24G卡上用LoRA,seq_len=1024,batch=1,显存占用大概在16-18G左右,你对比下这个基线,如果超太多就肯定是设置问题。
说实话你这个配置理论上真能跑起来,我怀疑问题出在max_length=2048上,7B模型在bf16下光激活值就吃得很凶,LoRA虽然省了梯度但attention的中间张量一点没少,24G卡跑2048长度确实极限了。你试试把max_length砍到1024或者512,如果只是实验的话,效果差距不会特别大,但显存能宽裕很多。另外transformers 4.31有个已知问题,就是past_key_values的缓存管理在长序列下会多占不少内存,建议直接升到4.36以上,那个版本改进了SDPA的显存分配。还有个冷门trick,你可以把gradient_checkpointing的use_reentrant参数改成False,有些时候默认的True反而会让峰值显存变高,虽然这听起来反直觉但我实测过能降低几百MB。还有就是你检查下有没有不小心把model parallel或者device_map设成auto了,有时候它会把层分散到不同设备上反而增加额外开销,手动把模型全放cuda:0可能更稳定。我自己的配置是3090跑7B,seq_len=1024,batch=1,8张卡用deepspeed zero2才敢上2048,单卡就别太贪长了。最后建议你开个显存监控脚本,比如用pynvml每步打印峰值占用,定位到底是前向还是反向爆的,这样改起来更有方向。
max_length 2048确实挺吃显存的,你先降到512试试,很多时候数据根本用不了那么长。另外transformers 4.31对LLaMA的LoRA支持有点坑,建议升到4.36以上,配合peft最新版会稳很多。还有个容易忽略的点:优化器状态,如果你用AdamW,7B模型光优化器就占不少,换成bitsandbytes的8bit adam能省一大截。3090跑7B LoRA batch size 1理论上是够的,检查下是不是dataloader的num_workers或者pin_memory偷偷占用了。
max_length 2048 这个点确实值得先查一下,7B 模型在 3090 上跑 LoRA,序列长度对显存的影响比 batch size 还猛,你把它降到 512 或 1024 试试,很多时候 OOM 就是长序列那几批触发的。另外 gradient checkpoint 和 LoRA 一起用的时候,如果 LoRA 没正确加到所有 linear 层,反而可能没省多少,建议确认下 target_modules 是不是覆盖了 q_proj、v_proj 这些。transformers 4.31 本身问题不大,但 flash attention 如果没装或者没启用,attention 那块显存占用会高不少,可以看看能不能上 flash-attn 或者 xformers。还有个小坑是 tokenizer 的 padding 策略,如果 padding 到 max_length 而不是 longest,那显存基本就是按 2048 在吃。我之前在 3090 上跑 7B LoRA,max_length 1024、batch 1、gradient accumulation 8,bf16 加 checkpoint,大概占 18G 左右,你可以照着这个基准调。实在不行就上 QLoRA,4bit 量化之后显存直接砍半,微调效果差距也没想象中那么大。