最近在试着用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模型来说确实挺吃显存的,尤其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量化能省不少。
你这配置按理说应该能跑起来,我怀疑问题可能出在几个地方。第一,max_length=2048对7B模型来说确实有点长,尤其是3090的24G显存,建议先降到1024试试,很多教程实际用的是512或768,长序列的显存开销是平方级增长的。第二,transformers 4.31确实有点老,后面几个版本对LLaMA的显存优化改进不少,比如4.35之后支持了更好的flash attention集成,建议升到4.38以上。另外,你检查过gradient checkpointing是不是真的生效了吗?有时候代码里虽然调用了,但模型某些层可能没被包裹进去,导致实际没开启。还有一个冷门trick:可以试试把模型加载时用device_map="auto"配合bitsandbytes的4bit量化,这样即使不降max_length也能省下将近一半显存,虽然精度会降一点,但LoRA微调影响不大。另外,你训练时是不是把optimizer states也塞进显存了?用bitsandbytes的8-bit Adam能再省几G。最后,可以看看是不是dataloader的num_workers设太高,或者pin_memory=True导致显存碎片,有时候这些细节反而会引发奇怪的OOM。
max_length降到1024试试,我这么调完显存直接降了小一半。
3090跑7B LoRA按理说24G是够的,max_length设2048确实有点激进,我之前试过,光是输入输出padding到2048,单条样本的显存占用就比1024翻倍不止。建议你先试试把max_length降到1024甚至512,看看能不能跑通,然后再逐步往上加。另外attention的显存占用跟序列长度是平方关系,这个挺容易忽略的,你如果开了gradient checkpoint但没做flash attention,那长序列还是会爆。可以试试装一下xformers或者直接用transformers自带的flash attention选项,能省不少。还有一个小trick:把optimizer换成AdamW的8-bit版本,或者干脆用SGD,显存压力会小很多。transformers 4.31我印象里有点老,LoRA相关的显存优化可能没跟上,建议升到4.35以上,有些bug修了。最后确认一下你是不是真的把模型加载成了bf16,有时候代码里写的是torch.bfloat16但没设device_map='auto',结果还是默认float32在跑,那肯定爆。
2048的max_length确实有点长了,7B模型在这个长度下即便是LoRA也很吃显存,很多人跑7B都是用512或1024,你可以先降到1024试试。另外transformers 4.31有个已知的小bug,某些版本的attention实现会多占显存,建议升到4.35以上,同时开启use_flash_attention_2,能省不少。还有个小trick:把gradient_accumulation_steps设成2或4,配合梯度检查点,batch size不用动但能稳住显存。
max_length设2048确实高了,7B模型光attention的KV cache就吃不少,先降到1024试试。另外LoRA的r值和target modules也会影响显存,r设8或者16就够了,别贪大。transformers 4.31有个已知的显存泄漏bug,升到4.35以上能解决。我自己跑的时候还加了--gradient_accumulation_steps 4配合batch size 1,算是省显存又稳定。
max_length设2048确实挺吃显存的,可以试着砍到1024或者512看看,LoRA本身省不了多少attention的缓存。另外transformers 4.31有个已知的显存泄漏bug,建议升到4.35以上,我升完直接省了3-4G。还有个trick是把model并行拆开用device_map=“auto”,让bitsandbytes自动分配层,能挤出点空间。实在不行就降一下lora的r值到8,效果差不太多但显存压力小很多。
这配置按理说应该能跑,我猜问题可能出在max_length上,2048对于7B模型确实挺吃显存的,试试改到1024或者512看看能不能稳住。另外transformers 4.31有个已知的attention显存泄漏bug,升到4.35以上应该能改善不少。还有个trick是把model parallel或者zero3打开,虽然慢点但能省不少显存,我刚入门那会儿也被折磨过,慢慢调就好了。
max_length设2048确实挺吃显存的,7B模型加LoRA在24G上跑这个长度很容易爆,建议先降到1024试试。另外transformers 4.31有个已知的attention缓存问题,升级到4.35以上能省不少显存。我自己之前也是3090,开了gradient checkpoint后还把batch size设成1,但把max_length砍到512才稳下来,你可以先按这个配置跑通了再慢慢加长度。
max_length设到2048确实容易爆,7B模型加LoRA在24G卡上跑2048长度挺极限的,我试过把长度降到1024或者用动态padding能省不少显存。另外transformers 4.31有个已知的attention计算冗余问题,试试升级到4.35以上或者换用flash attention,显存占用能降一截。还有个小trick是检查下是不是把基座模型本身也加载到了optimizer状态里,LoRA微调时冻结原始参数可以省掉那些梯度缓存。
我最近也遇到过一模一样的情况,3090 24G按理说跑7B LoRA应该够的,但确实容易翻车。你max_length设2048其实不算夸张,问题可能出在attention计算的中间变量上——特别是长序列下,QK^T的显存占用是平方级的,就算加LoRA也躲不开这个。建议你试一下在训练时把gradient checkpointing换成更激进的版本,比如transformers里那个use_reentrant=True的选项,或者手动把torch.compile打开,虽然编译慢点但能省不少显存。另外transformers 4.31有个已知的显存泄漏问题,升到4.35以上或者降到4.28试试看,很多人更推荐4.28。还有就是检查一下是否加载了模型全部权重,LoRA只更新部分参数,但推理时基座权重还是全量载入的,如果没用low_cpu_mem_usage=True或者device_map="auto",无关层的缓存也会占地方。我个人实际能跑通的配置是:batch_size=1, max_length=1024, bf16, 开gradient checkpoint, 用bitsandbytes的4bit量化加载基座模型,这样LoRA在3090上大概能剩下3-4G余量。你可以先降max_length到1024跑一轮试试,如果还不稳定,大概率是版本或者加载方式的问题。
max_length设2048确实容易爆,试试把batch size降到1的同时把max_length砍到1024,或者换8bit adam。
max_length设2048确实会吃不少显存,7B模型用LoRA的话建议先降到1024试试,很多教程默认512都能跑。另外transformers 4.31有个已知的attention显存泄漏问题,升到4.35以上应该能缓解。还有个trick是加上--gradient_accumulation_steps 2或者4,这样batch size虽然设1但等效batch size不变,显存压力会小很多。我16G的卡跑7B就是靠这几个方法稳住的。
max_length降到1024试试,很多人忽略这个,其实占不少显存。