最近在试着用LoRA微调7B的LLaMA模型,参考了网上一些教程,batch size设到1,gradient checkpoint也开了,用的还是bf16,结果3090(24G)还是跑着跑着就OOM了。我看别人同样的配置能跑起来,难道是我tokenizer的max_length设太长(2048)?还是说attention机制里有些隐藏的显存占用我没注意到?另外,我用的transformers库版本是4.31,不知道是不是版本问题。有没有老哥能分享下实际能跑的配置,或者推荐一些更省显存的trick?提前谢谢了,刚入门大模型,真的有点懵。
用PyTorch跑LLaMA微调,显存总爆,到底是哪里设置不对?
全部回复
共 170 条3090 24G跑7B LoRA理论上确实够,但你这配置里最可疑的就是max_length 2048,序列一长,激活值呈平方增长,哪怕batch size是1也容易爆。我之前用7B模型,max_length从2048降到1024,显存直接省了快6G,你可以先试试这个,代价只是训练样本得截断一下。
另外transformers 4.31版本有点旧,LoRA这块的显存优化在后续版本里改了不少,建议升到4.38以上,有些隐藏的缓存机制会更省。还有一个你大概率没注意的坑:attention的key/value cache在推理时是好事,但训练时如果你没显式关掉,它会额外占一份显存,检查下model.config里use_cache是不是False。
gradient checkpoint开了的话,记得同时把optimizer换成AdamW的8bit版本,或者用paged_adamw_8bit,这部分省下的显存很可观。
最后,别死磕单卡,试试deepspeed的zero stage 2,哪怕不开offload也能把优化器状态拆开,我实测能再省20%左右。
如果这些都调完还爆,那大概率是你数据加载时pin_memory或者num_workers设太高了,导致CPU显存碎片化,把这两个调低点。
同款卡同款模型,我之前也是这么爆的,后来发现是max_length背锅,2048对7B来说长序列的中间激活值太恐怖了,先砍到1024试试。另外transformers 4.31有个已知的attention mask分配问题,升到4.35+能省不少显存。你还可以把LoRA的target modules限定在q和v上,别碰k和o,实测能省将近2G。最后检查下是不是把padding设成了max_length,改成动态padding到batch内最长那条,效果立竿见影。
我最近也遇到差不多的情况,最后发现是max_length背锅,2048对7B来说太狠了,拉到1024立马稳了。另外transformers 4.31确实有点老,升级到4.38+之后显存占用明显下降,你可以先试试这个。还有一个冷门trick是关掉attention里那个flash_attn的开关,有时候新版库默认开了反而更吃显存。要是还不行,就检查下是不是dataloader里num_workers设太高,内存和显存打架也会导致假OOM。
max_length砍到1024试试,bf16下7B全参微调24G本来就很紧,LoRA也得留足激活值。
max_length砍到1024试试,3090跑7B LoRA这长度确实悬,flash-attention也装上能省不少。
24G跑7B LoRA按理说是够的,问题八成出在max_length=2048上,序列一长,中间激活值直接爆炸。你可以先砍到1024试试,或者用gradient_accumulation_steps把batch补回来,另外检查下是不是flash-attention没装上,transformers 4.31对它的支持有点坑,换个4.36+版本能省不少显存。我自己的经验是开bf16时把optimizer换成8bit adamw,能再挤出一两个G,你试试看。
我之前也卡在这步,max_length 2048对7B来说确实很吃紧,尤其LoRA虽然省了全量微调的显存,但激活值还是按最长序列算的。你试试把max_length砍到1024,或者用梯度累积来模拟更大batch,效果会明显很多。另外transformers 4.31对LLaMA的attention实现有点老,建议升到4.36+,有官方优化的flash attention路径,能省不少显存。我3090跑7B,LoRA rank=8,max_length 1024,batch 1加gradient checkpoint,峰值大概18G,你可以参考下。
max_length设2048确实挺吃显存的,LoRA虽然省了优化器状态,但激活值还是按序列长度算的,你可以先砍到1024试试,效果差不了太多。另外transformers 4.31有个已知问题,Llama的attention实现会多缓存一部分中间张量,升到4.35+或者直接换flash-attention能省不少。我自己的经验是gradient checkpointing要配合显存碎片优化,设一下PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,有时比调别的都管用。还有个小坑,如果你用了gradient accumulation,记得确认一下它是不是真的生效了,有时候会静默回退到普通累积,那显存压力直接翻倍。
24G跑7B LoRA按理说是够的,但你这配置我一眼就看出问题大概率不在max_length上。2048确实偏长,可bf16+gradient checkpoint下纯推理都不至于爆,问题出在optimizer states和activations的叠加——LoRA虽然只训adapters,但AdamW的momentum和variance还是按全参数规模算的,你试试把optimizer换成8bit的adamw,或者干脆用paged_adamw,能省出3-4G。另外transformers 4.31有个已知坑,就是LLaMA的attention实现会默认缓存past_key_values,你在generate时得显式设use_cache=False,但训练时反而要检查模型是否把hidden_states的梯度保留下来了——好多教程会用gradient_checkpointing_enable(),但忘了同时调model.config.use_cache=False,这俩一冲突会额外吃显存。还有个冷门技巧:把input_ids切成长度为1024的多个segment,用sequence packing方式拼接,但注意要加attention_mask隔离,能极大降低峰值显存。我自己的3070跑7B是batch_size=1、max_len=1536、lora_rank=8、8bit adamw,峰值大概17G,你照这个调肯定能稳住。系统swap记得开大点,真OOM了能抢救一下。
我之前也卡在这过,后来发现多半是max_length的锅,2048对7B来说太狠了,LoRA虽然省了优化器显存,但激活值照样吃满。你试试把max_length压到1024,或者用gradient_accumulation把有效batch补回来,显存能降一个量级。另外transformers 4.31有个已知的attention mask bug,会多占缓存,升到4.35+基本能解决。我3090跑7B,loRA rank=16,max_len=1024,batch=1,checkpoint开着,峰值大概19G,稳得很。你先把长度砍半,基本就不爆了。
3090跑7B LoRA按理说24G是够的,但你max_length拉到2048确实挺吃显存,序列长度对attention的占用是平方级的,我建议先降到1024试试,很多教程默认512也能跑。另外transformers 4.31有个已知的显存泄漏问题,升到4.35+或者换peft最新版能改善不少。你还可以看看是不是把model parallel或者device_map设成了auto,有时候它会额外预留显存,手动指定单卡反而更稳。我自己的经验是把gradient checkpointing的input caching关掉,再配合unsloth那个库,能省下快3G。
我之前也遇到过类似问题,后来换了方案。
max_length砍到1024试试,bf16下7B的activation峰值很吃显存,实在不行用4bit量化加paged optimizers。
max_length设2048确实挺吃显存的,7B模型光attention的KV cache就够喝一壶了,你可以先试试把max_length砍到512或者用梯度累积来模拟更大batch,我跑13B都没这么容易爆。另外transformers 4.31的LLaMA实现有些老,建议升到4.35以上,里面修了不少内存碎片问题,还有那个use_flash_attention_2选项开了能省不少。实在不行就上QLoRA,4bit量化后24G稳稳的,就是训练速度会慢点。
max_length 2048加上bf16其实还好,但你是不是忘了关gradient_checkpointing的use_reentrant参数?新版transformers默认不开这个会导致显存翻倍。我上次也是卡在这,换成gradient_checkpointing_kwargs={"use_reentrant": True}之后立刻从爆显存变成只用15G。另外检查下是不是dataloader的pin_memory和num_workers设置太高,把CPU内存占满也会间接导致CUDA OOM。
我遇到过类似情况,最后发现是attention里torch.nn.functional.scaled_dot_product_attention的memory-efficient模式没生效,你可以在model.config里显式设attn_implementation="flash_attention_2",前提是装好flash
max_length设2048确实有点顶,7B模型光是KV cache就吃不少,你可以先砍到1024试试,很多任务其实用不了那么长。另外transformers 4.31的Llama实现有个已知问题,就是attention的显存分配比新版浪费不少,建议升到4.36以上。我自己的经验是还得把gradient_checkpointing和input_embeds的缓存一起配合调,有时候光开开关不够。如果还爆,试试用unsloth那个库,同样的LoRA设置能再省一半显存,我3090跑7B到2048长度都没问题。
max_length2048确实太狠了,我降到1024再配合gradient checkpoint就能跑,你试试。
max_length砍到1024试试,我上次降到512立马稳了,显存瞬间少好几个G。
同款配置我也踩过这坑,max_length 2048在7B上确实太激进了,LoRA虽然省了主干的梯度,但attention的KV cache还是按全量长度算的,试试把max_length砍到1024甚至512,显存直接少一大截。另外transformers 4.31的attention实现有点老,建议升到4.38以上,新版本对SDPA的显存优化明显。还有一个冷门技巧,把gradient_checkpointing配合use_reentrant=False用,能再省点。我最后是batch size 1 + max_length 512 + 8bit量化才稳跑24G,你参考下。
max_length调到1024试试,LoRA用qlora直接4bit量化,24G稳稳的。
max_length拉到1024试试,bf16下7B的激活值真没那么友好,另外4.31的flash attention记得手动开一下。