最近在试着用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模型在3090上跑2048长度很容易爆,建议先降到1024甚至512试试。另外gradient checkpoint虽然省显存但会拖慢速度,可以搭配4bit量化(bitsandbytes)一起用,显存能再降一截。transformers 4.31倒没啥大问题,不过最新版的attention实现有些优化,可以升到4.36以上试试。
同款3090选手路过,2048的max_length确实偏高了,我降到1024之后显存压力小了很多。另外可以试试把flash attention打开,transformers 4.31好像还不原生支持,得装个xformers,能省不少显存。还有个小trick:把model并行或zero3开一下,哪怕单卡也能分摊点显存开销。你用的什么数据集?如果数据量大的话可以调低per_device_train_batch_size再配合gradient accumulation,等效batch size不变但显存占用能降下来。
24G跑7B LoRA按理说是够的,我试过差不多的配置,batch size=1+gradient checkpoint+bf16能跑起来,但关键可能出在max_length=2048上,这个长度对7B来说确实有点吃紧,显存占用会随序列长度二次增长。建议先降到1024试试,能跑通再慢慢往上加,同时注意pad到统一长度时别浪费显存。另外transformers 4.31有个已知的attention实现问题,在长序列下会多占不少显存,可以试试升级到4.35以上,或者手动把attention改成scaled_dot_product_attention,能省不少。还有个小trick,检查下你用的LoRA rank和alpha是不是偏大,比如rank=8、alpha=16就够用,别设成32或64,那会额外吃显存。最后看看是不是加载模型时用了device_map='auto',有时候它会自动把部分层放到CPU,反而导致碎片化OOM,不如手动指定单卡。这些地方都排查下,大概率能稳住。
老实说,2048的max_length在7B模型上确实挺吃显存的,我试过降到1024或者直接设512,显存能省出好几个G,而且很多任务其实用不到那么长的上下文。另外你检查一下是不是把padding side设成了right,这个在推理时影响不大,但训练时left padding反而会多吞显存,算是个小坑。transformers 4.31我印象中有点bug,特别是和bitsandbytes的兼容性,我换成4.35之后OOM频率低了很多,你可以试试升级。还有一点,LoRA的r值别设太高,我通常用8或者16,alpha设16就够,太高了反而会让梯度爆炸,显存也跟着涨。你还可以试试开启torch.compile,虽然编译慢一点,但能压下去一些中间缓存。最后检查下是不是dataloader的num_workers设太大,有时候多进程加载数据也会偷偷吃显存,我一般设2或4就稳了。
老实说2048的max_length在7B上确实偏长了,尤其3090只有24G,建议先砍到1024试试,很多开源项目默认就这个长度。另外transformers 4.31有个已知的attention显存分配bug,升到4.35以上能省不少,我换了版本之后同样配置直接多出3-4G余量。还有个小trick是把gradient checkpointing的配置参数从True改成{"use_reentrant": False},有些版本下能再压一点。你参考的教程可能还用了flash attention,那个对7B提升很明显,但3090得自己编译一下。
max_length降到1024试试,另外升级transformers到4.35以上能省不少显存。
我也遇到过一模一样的情况,3090 24G按理说跑7B LoRA是够的,但就是莫名其妙爆显存。你max_length设2048确实偏高了,我试过降到1024之后显存占用直接降了快4G,你可以先试试这个。另外attention那块有个隐藏点——flash attention开了吗?transformers 4.31好像还不原生支持,得手动装个包,这玩意儿能省不少显存。还有一个冷门trick是检查一下pad token的设置,有时候tokenizer自动补长会把序列长度撑上去,导致实际计算量比预期大。版本方面,4.31确实有点旧了,我换到4.38之后不仅修了一些显存泄漏的bug,LoRA的微调流程也稳了不少。建议你把batch size设为1,gradient accumulation设到4或者8,这样等效batch size不变但峰值显存更低。最后看看你代码里有没有无意间把完整模型参数都加载到了device上,有时候是加载方式的问题。
max_length设2048确实容易爆,试试512或768,顺便把gradient_accumulation_steps调到4或8。
max_length改到1024试试,我同样配置降到1024就没崩过。
max_length设到2048对7B模型来说确实有点顶,尤其是attention这块的显存占用是跟序列长度平方相关的。建议你先试下把max_length降到1024或者512看看能不能跑通,另外transformers 4.31的话可以试试升级到4.35以上,新版本对内存管理有优化,还有attention实现上改进了不少。实在不行可以看看deepspeed的zero stage2,配合LoRA能再省不少显存。
老实讲我刚开始也踩过这个坑,max_length设2048确实挺吃显存的,尤其是7B模型,就算LoRA也扛不住。可以试试把max_length降到1024或者甚至512看看能不能跑通,毕竟很多任务其实用不到那么长的上下文。另外transformers 4.31我记得有个版本在attention上有点小bug,建议升到4.35以上试试,有些显存泄漏的问题修掉了。还有就是gradient checkpoint虽然省显存但会拖慢速度,可以配合gradient accumulation用,batch size设1然后多累积几步,效果一样但显存压力小很多。
max_length改成512试试,顺便检查下transformers版本,4.31有个显存泄漏的bug。
max_length设2048确实有点高了,7B模型在24G显存下跑2048长度很容易爆,可以试试先压到1024或者512看能不能跑通。另外transformers 4.31有个已知的attention显存泄露问题,建议升到4.35以上或者换用4.28稳定版,我换完之后显存直接降了3-4G。还有个冷门trick:把model.parallel的device_map设成auto,让bitsandbytes自动分块加载也能省一些。
max_length设2048确实有点猛,7B模型加LoRA单条样本的显存开销比想象中大,建议先降到1024试试。另外transformers 4.31有个已知的attention实现内存泄漏问题,升到4.35以上可能会缓解。还有个小trick:把gradient_checkpointing和torch.compile一起开,有时候能省出1-2G显存。你用的是单卡吗?如果方便的话可以试试deepspeed stage 2,哪怕不用多卡也能压一压碎片。
max_length设到1024试试,我这么调之后24G卡稳跑7B。另外4.31版gradient checkpoint有bug,升到4.34就省不少显存。
max_length设2048确实有点高,可以先降到1024试试,我这样调完就没爆过了。
你这配置按理说能跑啊,我怀疑是tokenizer的max_length设2048把padding sequence撑爆了,试试dynamic padding或者把max_length降到1024,能省不少显存。另外transformers 4.31有个已知的attention计算冗余问题,建议升到4.35以上,或者手动加上gradient_accumulation_steps=2,显存占用能平滑很多。我同款卡跑7B LoRA,batch size=1、max_length=1024,稳在20G左右,你可以先照这个调。
max_length设2048确实挺吃显存的,7B模型加LoRA在24G上跑这个长度很容易爆,可以试试先降到1024或512看看能不能跑通。另外transformers 4.31有个已知的显存泄漏问题,升到4.35+或者换用peft的最新版会好很多。还有个小技巧是开torch.compile或者把attention换成xformers的memory_efficient机制,能省不少。实在不行就试试gradient accumulation加更小的batch size,虽然慢点但至少不报错。
2048确实有点长,可以先降到1024试试,另外attention的显存占用和seq_len平方相关,这个影响很大。
2048长度在7B模型上确实偏高了,尤其LoRA虽然省显存但attention计算量还在,你可以试试降到1024或者512看看能不能跑通。另外transformers 4.31有个已知的显存泄漏问题,建议升到4.35以上,我换版本之后同样配置直接省了2-3G。还有个小技巧,gradient checkpointing配合torch.compile能再压一点,不过第一次编译会慢。实在不行就试试QLoRA,4bit量化下7B模型只需要12G左右,3090随便跑。