最近在试着用LoRA微调7B的LLaMA模型,参考了网上一些教程,batch size设到1,gradient checkpoint也开了,用的还是bf16,结果3090(24G)还是跑着跑着就OOM了。我看别人同样的配置能跑起来,难道是我tokenizer的max_length设太长(2048)?还是说attention机制里有些隐藏的显存占用我没注意到?另外,我用的transformers库版本是4.31,不知道是不是版本问题。有没有老哥能分享下实际能跑的配置,或者推荐一些更省显存的trick?提前谢谢了,刚入门大模型,真的有点懵。
用PyTorch跑LLaMA微调,显存总爆,到底是哪里设置不对?
全部回复
共 7 条max_length 2048对7B模型来说确实挺吃显存的,尤其attention的kv cache会随序列长度平方增长。我建议你试试把max_length先降到1024或者512,同时检查下是否启用了flash attention(transformers 4.31好像还不直接支持,可以手动换xformers)。另外LoRA的r值设到8或16就行,别太大,还有记得关掉模型本身的缓存(use_cache=False),能省不少显存。最后实在不行就上deepspeed stage 2或者ZeRO-3,24G跑7B应该是够的。
max_length设2048确实容易爆,试试把gradient checkpoint加上并调小max_length到1024,同时检查下transformers版本,4.31可能有点旧。
2048的max_length确实有点高了,很多人跑7B LoRA其实用1024甚至512就能出不错的效果,你可以先降下来试试。另外transformers 4.31有个已知的attention计算冗余问题,升到4.35以上能省不少显存。还有个trick是把model parallel或者device map设成auto,有时候能自动分块避免峰值。我自己的经验是把LoRA的r值降到8,target modules只选q_proj和v_proj,24G跑2048长度也没爆过。
max_length设2048确实有点吃紧,7B模型用LoRA的话,建议先降到1024试试,能省不少显存。另外transformers 4.31有个已知的attention缓存问题,升到4.35以上或者手动清理下past_key_values会好很多。还有个trick是model.enable_input_require_grads()和gradient_checkpointing_enable()配合用,偶尔能多挤出几百M。你试试看还爆不爆?
24G跑7B LoRA按理说是够的,我试过把max_length降到1024,同时用gradient_accumulation_steps=2来模拟大batch,这样显存能省不少。另外检查下你是不是加载了完整模型而不是quantized版本,4.31的transformers我记得有个attention的显存优化开关,叫use_flash_attention_2或者类似的名字,开一下试试。还有,确认下你的tokenizer是不是自动加上了特殊token导致序列变长,我上次就是被这个坑了。
max_length设2048确实有点高,7B模型用bf16单卡24G的话,1280左右比较稳,可以试试先降到1024跑通再说。transformers 4.31有个已知的attention缓存问题,升级到4.35以上能省不少显存。另外检查下是不是把model parallel或者device map设成了auto,有时候它会自动把层分散到多卡,单卡反而更吃内存。
老实说2048的max_length在7B上确实挺吃显存的,我试过把max_length降到1024或者甚至512,batch size就能提到2或者4,训练效率反而更高。另外你可以试试gradient accumulation,配合梯度检查点能进一步压显存。还有transformers 4.31有个已知的attention显存泄漏问题,升级到4.35以上会好很多。如果还不行,检查下是不是加载了完整的模型权重,用device_map='auto'或者bitsandbytes的8bit量化能省不少。