最近在试着微调一个7B的LLM,用的LoRA,单张4090(24G)。一开始以为LoRA显存占用会很小,结果实际跑起来batch size开到2就OOM了,序列长度也就1024。查了挺多资料,有人说用DeepSpeed ZeRO-3能把优化器状态分片,也有人说直接上两张卡用FSDP更省事。
PyTorch跑大模型总爆显存,换DeepSpeed还是直接上多卡?
全部回复
共 23 条7B模型LoRA在24G上开batch2就OOM,大概率不是优化器状态的问题,先看看是不是梯度检查点没开,或者dataloader那边有额外开销。ZeRO-3确实能省显存,但单卡用ZeRO-3意义不大,它主要是跨卡分片的。如果手里只有一张4090,先把gradient checkpointing和flash attention打开试试,序列1024其实不算长,调完这些batch2应该能跑起来。真要上多卡的话FSDP比DeepSpeed配置简单不少,但两张4090走PCIe通信效率也就那样,不如先压榨一下单卡。
单卡4090跑7B LoRA,seq 1024 batch 2就OOM有点不太对劲,正常情况不该这么吃紧。你可以先检查下是不是跑在fp32上,或者有没有开gradient checkpointing,这两点影响特别大。DeepSpeed ZeRO-3确实能省显存,但它对单卡的收益其实有限,分片主要是多卡场景才明显。如果手上就有两张卡,我更推荐直接上FSDP,PyTorch原生支持,配置比DeepSpeed清爽不少,踩坑也少。
单卡4090跑7B的LoRA按理说不至于这么惨,bs=2 seq=1024就OOM有点不太对劲,建议先检查一下是不是哪里没配好。比如base model是不是用了fp32加载?gradient checkpointing开了没?还有LoRA的target modules如果设成all-linear,参数量其实不小。另外优化器如果用adamw,fp32的optimizer state对7B来说虽然LoRA只更新一小部分,但DeepSpeed ZeRO-3分片的收益在单卡上其实没啥意义,它本来就是为多卡设计的。FSDP两张卡确实能省显存,但通信开销和配置复杂度你得考虑进去,尤其4090没有NVLink,走PCIe会拖速度。我自己的经验是先把gradient checkpointing和8bit adam打开,再把batch size降到1配合gradient accumulation,大概率能跑起来。真要上多卡,FSDP比DeepSpeed在PyTorch生态里更顺滑一些,但调试成本也不低。