最近在微调一个7B的Llama模型,单卡A100 80G,尝试用DeepSpeed的ZeRO-3跑,结果一启动就显存爆炸,直接OOM。我查了文档,把offload参数也打开了,optimizer和param都offload到了CPU,batch size降到1,梯度累积也调了,但还是在第一步就崩。
我怀疑是不是我的模型加载方式有问题?或者ZeRO-3需要特殊的模型并行配置?看网上有人说ZeRO-2就够用,但我怕显存不够。有没有大佬遇到过类似情况?是不是我漏了什么关键参数?真诚求教,实在不想为了省钱白嫖半天还跑不起来……
用DeepSpeed跑Llama微调,ZeRO-3总是OOM,是我配置姿势不对吗?
全部回复
共 177 条我之前也遇到过类似问题,后来换了方案。
A100 80G跑7B用ZeRO-3确实容易OOM,我遇到过类似情况,后来发现是模型加载时默认用了全精度,改成bf16或者int8能省不少显存。另外offload到CPU后记得设一下cpu_offload_use_pin_memory=True,不然CPU和GPU之间传输太慢也可能触发奇怪的内存问题。你试试ZeRO-2加offload,7B单卡其实够用,我跑过几个微调任务没崩过。
我也遇到过类似的情况,A100 80G跑7B按理说ZeRO-3不该第一步就炸。检查下模型加载是不是用了from_pretrained的默认float32,改成bf16或者把dtype显式设成half能省不少显存。另外offload到CPU后别忘了调一下cpu_offload_eff_bucket_size,有时候默认值反而会卡住。
你这情况我遇到过,7B模型在单卡A100上硬上ZeRO-3确实容易翻车,尤其offload没配好时CPU内存也可能成瓶颈。试试把offload改成只offload optimizer,param留在GPU上,或者直接切到ZeRO-2配gradient checkpointing,我实测单卡A100跑7B完全没问题。另外检查下是不是用了torch的默认加载方式,用from_pretrained加device_map="auto"有时反而会干扰DeepSpeed的显存分配。
说实话ZeRO-3在单卡A100上搞7B确实容易OOM,我踩过类似的坑。你offload都开了还崩,很可能是模型加载时没有用from_pretrained的device_map="auto"或者手动指定model.to('cpu')再转,因为ZeRO-3默认会把参数分布到各个device上,单卡环境下反而可能因为初始化阶段显存分配策略出问题。我之前试过一个trick:先用ZeRO-2配合cpu_offload跑,batch size设1,gradient accumulation设8,其实80G足够塞下7B的梯度+优化器状态,显存占用大概在60G左右,反而比ZeRO-3更稳。另外检查下deepspeed_config.json里是不是忘了设zero_force_ds_cpu_optimizer: false,这个参数不关的话即使offload到CPU也可能因为优化器实现问题爆显存。还有个小细节,试试把fp16换成bf16,A100对bf16支持更好,能省点显存。如果还是崩,建议把model_parallel_size设为1,或者直接用Hugging Face的PEFT库搞LoRA微调,7B用LoRA的话单卡80G完全无压力,ZeRO-2都不用开。
说实话,ZeRO-3在7B模型上单卡A100 80G确实有点冒险,尤其offload到CPU后如果内存带宽不够,第一步爆显存挺常见的。我建议你先试试ZeRO-2加cpu offload,7B模型显存占用大概在50-60G左右,单卡应该能撑住,实在不行再调低seq length。另外检查下transformers版本和模型加载时是不是用了from_pretrained(..., torch_dtype=torch.float16),这个能省不少显存。
说实话我最近也踩过类似的坑,7B模型在单卡A100上开ZeRO-3确实容易翻车,尤其第一次加载时候显存会有一个峰值。你offload都开了还崩,可能是模型加载时没有用from_pretrained(..., torch_dtype=torch.float16)或者device_map="auto",因为ZeRO-3默认会把参数均匀分到所有设备,单卡情况下反而多了一层冗余的中间缓存。我建议你先试试ZeRO-2加上offload_optimizer,7B模型在80G上完全够用,batch size调到4都没问题,毕竟ZeRO-3的通信开销在单卡场景下其实是负优化。另外检查下你的deepspeed_config.json里有没有设zero_optimization.stage=3的同时还开了cpu_offload的pin_memory,这个选项在某些CUDA版本下会导致显存预分配翻倍。如果你坚持要ZeRO-3,试试在from_pretrained前先手动torch.cuda.empty_cache(),或者在脚本开头显式设置CUDA_VISIBLE_DEVICES=0,有时候多卡环境变量没清也会占显存。实在不行就改ZeRO-2吧,我跑13B模型都这么搞的,7B真的没必要硬上ZeRO-3。
试试把offload全关了,单卡A100跑7B用ZeRO-2就行,ZeRO-3反而容易炸。
试试把cpu_offload改成只offload optimizer,param留gpu上,ZeRO-3对7B单卡确实容易崩。
A100 80G跑7B其实ZeRO-2就够了,ZeRO-3的通讯开销和碎片化内存反而容易踩坑。建议先试ZeRO-2+offload optimizer,batch size设1观察显存峰值是不是在30-40G左右。另外检查下模型加载时是不是用了huggingface的from_pretrained,那个默认会先加载到0号卡再分片,容易瞬间爆显存。
同用A100 80G跑7B,ZeRO-3确实容易炸,我试过把offload的pin_memory关掉,再设个--zero_force_ds_cpu_optimizer false反而稳了。你检查下transformers版本是不是太新,有些版本加载模型时会在GPU上多占一块buffer,降回4.30.2试试。另外可以用deepspeed的auto模式先跑个profile看看峰值在哪,我那次发现是hidden_states没进offload列表。
我最近也刚踩过这个坑,7B模型在单卡A100上跑ZeRO-3确实有点极限,特别是Llama这种大词汇量的模型,光是embedding层就能吃掉不少显存。你开了offload还崩,大概率不是参数配置的问题,而是模型本身的初始化方式——很多框架在加载模型时会先把完整参数放到显存里,然后再offload,这一瞬间就爆了。你可以试试用zero.Init()上下文管理器来加载模型,这样参数一开始就在CPU上,ZeRO-3的分片才能生效。另外你有没有设stage3_gather_16bit_weights_on_model_save这个参数?有时候它默认行为会额外吃显存。说实话,对单卡来说ZeRO-2确实更省心,7B用ZeRO-2配合gradient checkpointing和4-bit量化,80G显存跑batch size 1甚至2是没问题的,我后来索性换回ZeRO-2了,稳定很多。你如果非要死磕ZeRO-3,可以试试把stage3_max_live_parameters和stage3_max_reuse_distance调小一点,限制同时驻留在GPU上的参数数量。最后想问下你是用transformers的from_pretrained加载的吗?那个接口有时会和DeepSpeed的offload打架。
我遇到过类似情况,A100 80G跑7B用ZeRO-3确实容易踩坑。试试把offload先关掉,只保留optimizer offload,param offload反而会增加CPU-GPU通信开销,导致显存碎片化。另外检查下模型加载是不是用了from_pretrained的device_map="auto",这个和ZeRO-3有时会冲突,改成手动分配可能好点。实在不行切ZeRO-2吧,7B模型在80G上完全够用,我跑过很多次了。
我也遇到过类似的情况,后来发现是模型加载时默认用了float32,显存直接翻倍。你试试在from_pretrained里加上torch_dtype=torch.float16或者bfloat16,一般能省不少。另外ZeRO-3对7B模型其实ZeRO-2就够了,80G跑7B单卡完全没压力,你先把offload关掉试试看。
说实话我最近也踩过这个坑,感觉ZeRO-3对7B模型在单卡A100上确实有点过于激进了。80G显存理论上够用,但ZeRO-3本身会保留模型参数的完整副本用于前向,加上优化器状态和梯度,第一步就崩很可能是CPU offload的带宽瓶颈导致内存爆炸,而不是真正的显存不足。你可以试试先把optimizer offload关掉,只offload param,或者干脆换ZeRO-2加offload——我实际试下来ZeRO-2配CPU offload在7B上完全够用,batch size甚至能撑到4左右。另外检查下你有没有用deepspeed的zero.init上下文加载模型,如果直接torch.load的话确实会在第一步爆显存。还有个小细节,检查下你gradient checkpointing有没有开,这个对显存优化很关键。最后问一句,你用的是transformers的Trainer还是自己写的训练循环?后者容易漏掉一些自动的内存管理逻辑。
说实话看到你这个配置我第一反应是ZeRO-3开offload之后还OOM确实有点诡异。我猜可能不是显存瓶颈,而是CPU内存或者NVMe交换那块儿出问题了——有时候offload的默认路径会往/tmp写,如果系统盘空间不够直接卡住。另外你检查过模型加载是不是用了meta device吗?如果用原生的transformers加载,即使offload了参数还是在初始化阶段就占满显存。我自己的经验是7B模型在单卡A100上用ZeRO-2加fp16其实能跑,batch size设到2甚至4都没问题,除非你用了特别长的序列。ZeRO-3的显存节省主要靠切分参数,但通信开销和CPU offload的延迟反而可能让第一步更吃资源。你可以试试先关掉offload,只用ZeRO-2跑一个极小的batch看能不能过,如果还OOM那就检查下是不是代码里有什么隐藏的显存泄漏,比如没清理的梯度累积buffer。还有个小技巧是显存监控加个torch.cuda.empty_cache()在step之前手动释放一下,有时能救急。
试试把offload的pin_memory关了,或者换ZeRO-2加activation checkpointing,7B单卡其实不用硬上ZeRO-3。
试试把cpu_offload改成nvme,或者先检查下模型是不是完整加载到显存了,ZeRO-3本身不会导致OOM这么离谱。
ZeRO-3的offload默认可能没完全生效,检查下zero_optimization.stage3_gather_16bit_weights_on_model_save是不是开了?
ZeRO-3对7B模型来说开销确实大,试试ZeRO-2加CPU offload,说不定能跑起来。