最近开始尝试用Llama 3.1 8B做领域微调,跟着教程写了LoRA,batch size设到1,gradient checkpointing也开了,结果3090(24G)还是OOM。我的输入长度大概2k tokens,是不是跟序列长度有关?看到有人说用DeepSpeed ZeRO-3或者Flash Attention能省显存,但配置起来有点复杂,不太确定是哪里出了问题。另外,torch.compile会有帮助吗?现在用的是PyTorch 2.1,cuda 12.1。求有经验的大佬指点下常见坑,谢谢!
刚转大模型方向,用PyTorch跑LLM微调,显存总爆掉怎么优化?
全部回复
共 178 条试试4bit量化加unsloth,24G跑8B绰绰有余,2k长度别开ZeRO,纯属给自己找麻烦。
你这配置按理说跑LoRA不该爆,先查下是不是max_length没设对,另外Flash Attention真能救急,装上能省不少。
torch.compile也建议开,但得先确认下和DeepSpeed的兼容性,不然容易白折腾。
说实话看到这个配置我第一反应是怀疑你的tokenizer是不是把特殊token算太多了,2k长度加上attention mask的padding,实际显存占用可能比你想的翻倍。我之前用7B模型跑2k输入,LoRA开gradient checkpointing,batch size=1在24G卡上勉强能过,但你要是用了动态padding或者没开padding_side="left",那多出来的计算量真的会压垮显存。DeepSpeed ZeRO-3确实能救急,但配置起来容易踩坑,比如offload到CPU后训练速度会慢到怀疑人生,而且跟gradient checkpointing叠加有时候会出奇怪的bug。我觉得你不如先试试Flash Attention,这个改动最小,直接在模型配置里把attn_implementation="flash_attention_2"打开就行,2.1版本应该支持。torch.compile的话,对显存帮助不大,主要是加速,而且首次编译会卡很久,建议先别折腾。还有个坑是优化器状态,如果用了AdamW,8B模型光优化器状态就占好几个G,可以试试Adafactor或者8bit优化器,能省不少。最后检查下max_seq_len是不是被模型默认值撑大了,Llama 3.1的config里可能设了8192,你输入2k但模型会按最大长度预分配显存,手动改成2048试试。
24G跑8B LoRA按理说够的,你batch=1还爆大概率是2k序列长度在作祟,attention这块占的显存随序列长度平方涨。Flash Attention绝对值得配,能省不少,而且现在transformers直接传attn_implementation="flash_attention_2"就行,不用手动改代码。ZeRO-3对单卡意义不大,先试试offload optimizer到CPU。torch.compile这版本收益不明显,建议先把input长度截断到1k或者用gradient accumulation模拟大batch,看看峰值显存降没降。另外检查下是不是把eval也开在训练循环里了,那个也吃显存。
这配置跑8B LoRA按理说不该爆的,2k长度确实吃显存,但24G卡开gradient checkpointing是够的。你检查下是不是把LoRA加在了所有linear层上,或者用了默认的fp32训练?换成bf16能省一半。Flash Attention可以试试,装起来不复杂,torch.compile对显存帮助不大但能提速,建议先把attention scaling和padding那边优化下。
Flash Attention先安排上,2k长度下显存大头基本都在attention矩阵上,这步能省不少。
8B全参微调24G本来就极限,试试把max_seq_len砍到1k,flash-attn必开,ZeRO-3配LoRA反而有坑。
3090跑8B LoRA还爆显存,大概率就是序列长度惹的祸,2k tokens塞进去activation footprint很夸张,flash attention基本是必开的,直接能省一半左右。DeepSpeed ZeRO-3对单卡LoRA其实帮助不大,反而ZeRO-2或者干脆关掉offload更省心,你可以先试试把max_seq_len砍到1k看还爆不爆。torch.compile在2.1上对Llama这类模型收益一般,而且容易踩动态shape的坑,建议先把前两个搞定再折腾。另外记得检查下是不是把LoRA加到了所有linear层,有时候target modules设少点也能省不少显存。
24G跑2k长度的8B LoRA确实紧,先试下把序列截到1k或换ZeRO-2,Flash Attention能省不少。
24G跑8B LoRA按理说应该够,2k序列确实是个坎,但你这配置全开还爆,我赌多半是优化器状态和中间激活在打架。我之前用7B模型试过,光靠gradient checkpointing不够,得把LoRA的target modules选少点,别一股脑全怼上,尤其别动embedding和lm_head,那俩的梯度会额外吃显存。DeepSpeed ZeRO-3确实能救急,但配起来容易跟LoRA的param group冲突,我建议你先试ZeRO-2加offload optimizer,改动小见效快。Flash Attention对长序列提升非常明显,2k长度下能省掉一大块attention矩阵的显存,torch.compile反而没那么关键,PyTorch 2.1上compile对动态shape支持一般,容易踩坑。还有个冷门但实用的点——检查下你的dataloader有没有把input_ids和attention_mask同时送进模型,有时候重复计算attention mask也会偷偷占显存。另外,你确定用的是bf16而不是fp16吧?fp16在8B上跑很容易因为梯度下溢导致loss不稳,但不会直接OOM,如果显存还差一口气,试试把max_length临时砍到1536,先跑通再慢慢加。
24G跑8B LoRA还爆显存,多半不是显存不够,而是你踩了序列长度和激活内存的坑。2k tokens对于8B模型来说,即使batch size=1,激活值也会吃掉很大一块,gradient checkpointing能省但别指望它把峰值降一半以上。我建议你先看一眼NVIDIA-SMI,如果显存占满但利用率很低,那大概率是activation memory在作祟,这时候Flash Attention确实能救急,因为它把注意力矩阵的中间态从显存挪到重计算里了。DeepSpeed ZeRO-3在这个场景下其实不太对症,那是为多卡分布式设计的,单卡上跑反而会引入额外通信开销,除非你后续要上多卡,否则别折腾。torch.compile可以试试,它对显存优化有限,但能降低显存碎片化,有时候能挤出10%-15%的空间,代价是编译时间很长,而且PyTorch 2.1配CUDA 12.1可能会有算子兼容问题,建议先升级到2.3+。我遇到过更隐蔽的坑是tokenizer的padding策略,如果你没设padding_side='left'且固定长度,模型内部会动态生成超长序列,导致显存忽高忽低。另外,LoRA的target_modules别全加上,比如把lm_head和embedding也设成可训练,那几乎等于全量微调了,显存直接翻倍。最后实在不行就把max_seq_len临时砍到1024验证一下,如果能跑通,就能确认是长度的问题,再考虑用梯度累积来模拟更大batch。
3090跑8B LoRA,2k长度确实紧,但你这配置不该直接爆,先查下是不是max_seq_len没对齐导致padding浪费,把tokenizer的padding设成side=left或者直接truncate到实际长度试试。Flash Attention能省不少,主要是把attention的中间矩阵省掉了,PyTorch 2.1里可以直接用SDPA的flash内核,不用额外配DeepSpeed。torch.compile对显存帮助不大,主要是提速,而且跟gradient checkpointing有时候会冲突,建议先不开。我之前踩过坑是optimizer状态占太多,用AdamW的话8B模型光优化器就要吃好几个G,换个Adafactor或者8bit优化器能立省一大截。
24G跑8B LoRA还爆显存,大概率是序列长度在作祟,2k tokens对attention来说吃显存很凶,可以先试下把max_length砍到1k或者512看看降幅。Flash Attention确实值得装,能省不少内存而且速度还快,torch.compile对显存优化帮助不大,主要是提速。如果你不想折腾DeepSpeed,可以先试试把LoRA的rank调低到8或者4,再配合gradient checkpointing,基本能稳住。另外确认下是不是把embedding和lm_head也设成可训练了,那俩参数占大头,冻结掉能省一截。
你这配置跑2k长度8B确实紧,先把flash attention开了试试,能省不少显存,torch.compile对显存帮助不大。
我刚开始搞的时候也总OOM,后来发现光开gradient checkpointing不够,还得配合显存碎片优化,比如PyTorch里设PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True能立省不少。Flash Attention确实值得搞,尤其你2k长度,直接换掉多头注意力那块能砍掉大半激活显存,torch.compile在2.1上对训练加速有限,但偶尔能省点内存,可以试试但别抱太大期望。另外你LoRA是不是加到所有线性层了?有时候只加qkv能再压一截,还有检查下有没有把eval模式下的buffer也塞进显存,或者优化器状态用8bit,这几个坑我踩完才勉强跑起来。
2k序列确实吃显存,试试flash attention加bf16,能省不少。
2k序列长度下8B模型光激活值就很吃显存,LoRA本身省的是优化器状态和梯度,激活这块还得靠Flash Attention或者梯度检查点配合。你gradient checkpointing开了还OOM,可能是batch里padding没处理好,或者optimizer用的还是fp32,试试bitsandbytes的8bit Adam。ZeRO-3在单卡上收益不大,反而通信开销拖后腿,不如先把max_seq_len降到1k验证下是不是长度问题。torch.compile在2.1上对LLM支持还不稳,别折腾了,先升到2.3+再说。
2k序列对8B模型确实容易爆,试试Flash Attention加bf16,能省不少。