最近在尝试用DeepSpeed微调一个7B的LLaMA模型,卡是两张A100 40G。用了ZeRO-3,offload也开了,但每次跑起来没几分钟就报OOM。我看了一些教程说7B模型用ZeRO-3应该能跑,但我设置zero_optimization.stage=3之后,显存占用直接冲到35G以上,训练步数一多就崩。我试过调小per_device_train_batch_size到1,还是不行。是不是我offload_param和offload_optimizer的device没设对?或者ZeRO-3本身对多卡通信有额外显存开销?有没有大佬分享下实际成功跑7B微调的config配置?先谢过了。
用DeepSpeed跑LLaMA微调,ZeRO-3总报显存不足,是我配置有问题吗?
全部回复
共 112 条说实话你这个现象我太熟了,之前用ZeRO-3跑6.7B也踩过一模一样的坑。关键点其实不在offload本身,而是你把offload_param和offload_optimizer都设成cpu之后,通信开销会变得特别离谱,每步都要在GPU和CPU之间来回倒腾权重和梯度,显存反而被临时buffer给占满了。我建议你先别急着全offload,试试只offload optimizer,param留在GPU上,然后打开zero_force_ds_cpu_optimizer=false,这样能省不少显存。另外你两张卡互联如果是PCIe而不是NVLink,ZeRO-3的all-gather和reduce-scatter会放大通信显存占用,我遇到过直接多出4-5G临时显存的情况。还有个容易忽略的坑是,你数据加载那边如果num_workers开太高,CPU内存占满会反过来挤压页锁内存,间接影响显存分配。我后来把per_device_train_batch_size调到1,gradient_accumulation_steps加到16,然后offload_optimizer的device设成nvme(虽然慢但稳定),勉强能跑完一个epoch。你可以看看nvidia-smi里是不是有碎片化显存,有时候不是总量不够,是分配不了连续块。最后建议你开一下zero_quantized_weights,虽然会掉一点点精度,但显存能省15%左右。
offload只开param没用,optimizer也得全扔CPU,另外试试把gradient_checkpointing打开能省不少。
offload设cpu试试,另外A100 40G两张跑7B本来就紧,梯度检查点开了没?
offload设cpu试试,另外梯度检查点开了没?这俩才是省显存关键。
说实话你这情况我太熟了,之前用ZeRO-3跑13B也踩过一样的坑。你提到offload开了但显存还是冲到35G,我怀疑你offload_param的device可能设的是cpu,但offload_optimizer没跟着设,或者两个都设了但pin_memory没开,导致优化器状态还是留在GPU上。另一个很隐蔽的点是ZeRO-3在做forward和backward的时候会触发全量参数收集,这个过程本身会临时把当前层的参数广播到所有卡上,通信缓冲区会吃掉不少显存,尤其两张卡之间NVLink带宽不够的时候,这个临时buffer可能比你想的大得多。我之前试过把zero_force_ds_cpu_optimizer设为false,然后配合stage3_gather_16bit_weights_on_model_save,能省出不少显存。还有个小技巧是把gradient_checkpointing打开,它虽然会慢一点,但能把激活内存压得很低,配合offload基本能稳。你per_device_train_batch_size已经调到1了还崩,那大概率不是batch size的问题,试着把zero3_max_live_parameters和zero3_max_reuse_distance调小一点,限制同时驻留在GPU上的参数数量。最后确认下你的DeepSpeed版本,老版本对LLaMA的某些算子支持有问题,建议升到0.9以上再试。
说实话你这个问题我当初也踩过,7B用ZeRO-3双卡A100 40G理论上是能跑的,但很多人忽略了一个关键点:offload_param和offload_optimizer默认的device是cpu,但如果你没设offload_device或者没配nvme,反而会把大量参数塞回显存做临时缓冲,尤其是梯度checkpoint没开的话,激活值直接爆掉。你可以试试把activation checkpointing打开,这能省不少内存,代价是慢一点。另外ZeRO-3在多卡通信时确实有额外的all-gather开销,每层计算前都要把参数广播回来,这部分峰值显存经常被任务管理器忽略,我建议你把zero_force_disable_cpu_offload设成false,同时把stage3_gather_16bit_weights_on_model_save关掉,这个会在保存时额外拉全量权重。还有个坑是optimizer的offload,如果你只offload了param没offloadoptimizer,那Adam状态照样吃满你的显存,两个必须同时开。我自己的配置是per_device_train_batch_size设成2,gradient_accumulation_steps设成8,加上offload到cpu,跑7B勉强稳在36G左右,但偶尔还是会抖一下。你可以先开着nvidia-smi监控,看是不是在某一步突然跳涨,如果是,大概率是通信峰值和checkpoint叠加导致的,试着把stage3_prefetch_bucket_size调小一点,默认值太大了。如果还不行,干脆降到ZeRO-2加offload,7B微调其实也能扛住,没必要死磕stage3。
40G双卡跑7B按理说ZeRO-3+offload是够的,但你这显存冲到35G+大概率是offload没生效,或者把offload参数也留在GPU上了。我之前踩过坑,offload_param和offload_optimizer的device都得明确写cpu,而且pin_memory最好设成false,不然反而会增加显存碎片。另外你检查下zero_force_ds_cpu_optimizer是不是被默认设成false了,这会导致优化器状态还在GPU上,我改成true之后显存直接降了10G。通信开销确实有,但你batch都1了还崩,更像是配置里没把stage3_gather_16bit_weights_on_model_save关掉,试试在训练时禁用这个,应该能再省一点。
每步之间用torch.cuda.empty_cache()手动清一下缓存会有帮助,不过更关键的可能是你数据加载那边,num_workers设太大也容易吃显存。我之前是把这个加到4,加上把gradient_checkpointing打开,7B在单张40G上都能跑起来,双卡其实更宽松。你试试把offload的cpu_offload改成nvme看看,虽然慢但能救命,实在不行就降到6B模型吧。
我最近也刚用DeepSpeed跑过7B微调,跟你情况挺像的,两张A100 40G,一开始也是开offload就崩。后来我发现一个坑是offload_param和offload_optimizer的device如果都设成cpu,反而会把PCIe带宽占满,导致通信瓶颈,显存看着没满但训练直接卡死或者OOM。你可以试试只offload optimizer,param留在GPU上,或者把zero_force_opt_offload关掉,让ZeRO自己决定哪些放哪。另外,ZeRO-3确实有额外的显存开销,主要是每层all-gather的临时buffer,你可以在config里把zero_3_use_jit开起来,同时把reduce_scatter改成allgather试试,有时候能省不少。我最后是把per_device_train_batch_size设成1,gradient_accumulation_steps加到16,然后offload_optimizer的pin_memory打开,才勉强跑稳。还有个建议,你检查下zero_3_round_robin这个参数,如果没设,默认可能每个rank都存全量参数分片,反而更占显存。最后问下,你用的是DeepSpeed还是DeepSpeed+HF的Trainer?如果是后者,model_parallel那块也可能有冲突,我上次就是被这个坑了。
我之前也卡在这,7B用ZeRO-3其实能跑但40G很极限,你试试把offload_param的device设成nvme,offload_optimizer设成cpu,同时关掉zero_force_disable_cpu_offload,显存能省不少。另外ZeRO-3确实会为每层参数收集多一份通信buffer,你把zero_optimization.allgather_bucket_size调小到5e7,reduce_bucket_size调到2e7,能压住峰值。还有个小坑,checkpoint保存时也会触发全量参数收集,容易瞬间爆,记得开zero_optimization.stage3_gather_16bit_weights_on_model_save=false。我最后是batch_size=1,gradient_accumulation_steps=16才稳住的,你可以参考下。
offload设cpu试试,还有gradient_checkpointing开了没,这俩能省不少。
说实话你这情况我太熟了,之前用ZeRO-3跑13B也撞过一模一样的墙。你提到offload_param和offload_optimizer,我猜大概率是device设成了cpu,但没留意NVMe offload的路径和缓存大小,这两个一没配好,反而会把显存占满,因为通信缓冲区和临时张量全挤在卡上。另外ZeRO-3确实有额外的显存开销,主要是all-gather和reduce-scatter时每个layer的临时权重会暂存,两张卡40G算下来,7B模型理论能塞下,但实际峰值很容易超,尤其你开了offload后CPU-GPU传输慢,导致GPU侧累积的激活值没及时释放。我建议你先把zero_optimization里的reduce_bucket_size和allgather_bucket_size调小一点,比如设到5e7或者1e8,这能显著降峰值。还有,stage3_prefetch_bucket_size也别默认,设0或者特别小,不然预取机制会额外吃显存。再一个坑是activation checkpointing必须开,不开的话激活值直接炸,你batch size调到1也没用。最后检查下gradient_accumulation_steps,如果设很大,反向传播时梯度会累积在显存里,跟offload的优化器状态打架。我最后跑通是用的是offload_optimizer到cpu,offload_param到nvme,加上上面那三个参数全改小,batch size=1,gradient accumulation=16,稳得一批。你要是还崩,把完整config贴出来,咱们对着看下是不是pin_memory或者通信后端的问题。
offload设cpu了吗?还有试试gradient_checkpointing,能省不少显存。
之前也踩过这个坑,后来发现ZeRO-3的partition size和NVMe offload的配置比想象中敏感,光开offload不够,zero3_offload_param的pin_memory和nvme_path如果没设好,反而会多占一份显存做buffer。另外A100 40G跑7B全参微调本来就紧,我后来是配合gradient checkpointing + ZeRO-3才稳住的,batch size=1也不一定救得了,因为激活值峰值在forward阶段。你试试把zero_optimization.reduce_bucket_size和allgather_bucket_size调小到5e8,再把stage3_gather_16bit_weights_on_model_save关掉,显存能掉下来不少。多卡通信确实有临时显存开销,但一般不会到几个G,感觉还是offload配置的问题。
我之前也被这个坑过,ZeRO-3的显存占用看着特别吓人,其实有一部分是通信缓冲区和临时activation的锅。你可以试试把zero_force_ds_cpu_offload设为true,同时确认下offload的device是不是都写成了cpu,另外pin_memory关掉有时候能省不少。还有个偏方,把gradient_checkpointing打开,虽然慢点但显存能压下去不少,我当时就是靠这个才跑通的。
我之前也踩过这个坑,ZeRO-3的显存开销不只是模型参数,还有每步的梯度同步和通信缓冲,两张40G其实挺紧的。offload确实要设对,param和optimizer都建议指到cpu,但更关键的是把zero_force_ds_cpu_optimizer设为false,不然容易有隐藏的显存碎片。另外你可以试试把zero_allow_untested_optimizer打开,关掉offload_optimizer的pin_memory,有时候这个反而会多占一块固定显存。我后来是开了ZeRO-3+offload,batch size=1,梯度累积调到8,才稳定跑完的,但速度慢得让人想摔键盘。
offload开全但device设cpu试试,另外检查下allgather的buffer大小,A100 40G跑7B按理够。
我之前也踩过这个坑,ZeRO-3的显存开销不只是模型参数,还有每步通信的buffer和碎片化问题,40G两张卡跑7B确实很紧。你试过把zero_force_ds_cpu_offload设成true,或者手动给offload_param指定nvme路径吗?另外检查下zero3_init是不是默认把参数全塞进显存了,可以搭配pin_memory和stage3_prefetch_bucket_size调小点试试。
offload设cpu试试,另外检查下allgather的buffer大小,默认配置吃显存很凶。
ZeRO-3的显存开销确实不只是模型参数,通信buffer和gradient partitioning也会吃不少,尤其两张卡通信效率不高时,峰值可能比理论值高很多。你试试把zero_optimization.overlap_comm和round_robin_gradients都打开,再把offload_param的pin_memory设成true,有时候这几个开关能省出好几个G。另外确认下你offload的device是cpu还是nvme,如果nvme没配好反而会加重显存压力,我之前就栽在这上面。还有个小技巧,把optimizer的offload单独开,param的offload可以先关掉,7B模型参数也就14G,两张卡硬扛应该行。
offload只开param不行,optimizer也得全扔CPU,而且梯度检查点没开的话照样爆。
ZeRO-3的通信峰值也挺吃显存,试试把partition size调小点,或者换ZeRO-2加offload更稳。