最近在试着用DeepSpeed ZeRO-3来微调一个7B的Llama模型,我的显卡是A100 40G。理论上两个batch size跑16应该没问题吧?但每次一跑就报CUDA out of memory。已经设置了offload到CPU,也用gradient checkpointing了。更奇怪的是,用同样的配置跑HuggingFace上的一个demo脚本反而能跑通,换成我自己的代码就不行。我怀疑是不是自己的数据加载或者trainer配置有问题,但查了一下午也没找到原因。有没有老哥遇到过类似情况?是不是DeepSpeed和某些自定义的forward函数不兼容?或者我应该在model.to('cuda')之前做点什么?救救孩子,调了两天快崩溃了。
显存明明够用,但用DeepSpeed跑Llama微调总是OOM,求大佬指点
全部回复
共 118 条我之前也踩过这个坑,多半不是DeepSpeed的锅,而是你自己的trainer里某个参数跟ZeRO-3冲突了。比如你手动调用了model.to('cuda')或者把embedding单独挪到GPU上,这会让ZeRO-3的分片逻辑直接崩掉。建议你把自定义forward里所有显式device操作都去掉,让模型统一走accelerate的device管理。另外数据加载那边,如果用了collator里做了pad到固定长度,也可能导致每个batch的显存峰值暴涨,试试动态padding看能不能缓解。最后实在不行就开个profile看看是哪一步爆的,别瞎猜了。
跑过类似配置,感觉你问题可能出在梯度累积或者优化器状态上。ZeRO-3本身就会把参数分片到各设备,但你如果同时开了offload到CPU,又用了gradient checkpointing,这俩叠加会反复在CPU和GPU间搬数据,显存占用反而变高。你可以试试把offload关掉,只用ZeRO-3的分片,然后把batch size降到8看看,如果还OOM那就是你代码里有什么地方偷偷创建了大的中间张量,比如在loss计算里用了one-hot或者重复了logits。demo脚本能跑大概率是因为它走的是最标准的流程,没动那些犄角旮旯的API。
这问题我太熟了,之前用ZeRO-3跑CodeLlama也踩过一模一样的坑。你那个“demo能跑通自己代码就炸”的现象,十有八九是卡在模型输入的数据结构上——比如自定义的attention mask或者position ids形状没对齐,导致显存碎片化,ZeRO-3在分片时反而把临时buffer撑爆了。建议你把trainer的batch size先降到1,然后把offload改成offload_optimizer和offload_param分开开,看看报错是不是变成具体的tensor size mismatch。另外检查一下是不是你的forward里有任何显式调用了.cuda()或者.to('cuda')的操作,DeepSpeed会自己管理device,你硬指定的话容易跟它的分片逻辑冲突。还有个骚操作是直接在ds_config里把zero_force_ds_cpu_optimizer设为false,有时候默认的CPU optimizer会跟A100的特定驱动版本打架。如果还不行,开个nsys profile看一眼peaked memory到底在哪一步飙升,比我在这瞎猜准。
说实话你这情况我太熟了,上次我调一个13B模型也卡在类似问题上,最后发现根本不是显存的事。你那个demo能跑通但自己代码不行,我赌八成是huggingface的Trainer在背后悄悄做了些你没注意到的处理,比如自动padding或者把labels给移到了正确设备上,你手写的训练循环可能没管这些。另外你说model.to('这行没写完,我猜你是不是想offload到CPU或者meta device?ZeRO-3下每个layer都是按需加载的,如果你提前把整个model放到cuda上,反而会打破它的分片逻辑,导致每张卡上都存一份完整参数。还有个坑是数据加载,如果你自己写dataloader时没设pin_memory或者num_workers太高,CPU和GPU之间的搬运也可能把显存碎片化,看起来明明够用但就是分配不出连续块。我建议你先用torch.cuda.memory_summary()看一眼内存分配情况,是不是有大量reserved但未使用的碎片。再有就是检查一下你的optimizer状态,ZeRO-3默认会把optimizer状态也分片,但如果你用了Adam的偏置修正项或者自定义的梯度裁剪,可能让它在某个节点上突然复制全量状态。最后实在不行,把offload_optimizer和offload_param都打开,然后gradient_accumulation_steps设成2,虽然慢点但至少能跑通,先确定是不是代码逻辑问题再优化性能。
我之前也踩过类似的坑,问题往往不在显存总量,而是CUDA内存碎片或者某个中间变量峰值爆炸。建议你开一下DeepSpeed的wall_clock_breakdown和memory_breakdown日志,看看具体是哪一步分配失败。另外,自定义forward里如果用了Python list存tensor,或者有动态shape操作,ZeRO-3的显存规划很容易失效,试试把输入pad到固定长度,或者用torch.cuda.empty_cache()在step间手动清一下。还有,你的model.to('cuda')是不是在deepspeed.initialize之后又执行了?那个会破坏offload的device_map,导致tensor又回显存了。
我之前也踩过类似的坑,后来发现是自定义forward里用了绝对位置编码的缓存,导致显存碎片化严重,ZeRO-3分片后反而放大了这个问题。你可以试着把输入序列长度固定成8的倍数,或者关掉eager attention试试。另外检查下dataloader有没有把label也放到GPU上,有时候多卡数据并行会隐式复制一份。你那个demo能跑通,大概率是序列长度或者padding策略不同,建议对比一下两边实际的max_length和batch内的tensor shape。
我之前也踩过类似的坑,最后发现问题根本不在显存够不够,而在你代码里某个张量悄悄被复制了一份。A100 40G跑7B加batch 16理论确实没问题,但DeepSpeed的ZeRO-3会把参数切到各卡上,如果你自定义forward里用了类似torch.Tensor的显式索引或者直接对模型权重做了操作,比如model.layers[i].weight.data += something,这就会触发全量参数物化,瞬间把显存撑爆。你那个HuggingFace demo能跑通,多半是它里面没有这些“小动作”,数据流很干净。
我建议你查两件事:第一,把zero3的reduce_scatter和gather相关日志打开,看看是不是有层在forward时被整个gather到单卡上;第二,检查你的dataloader有没有在__getitem__里做GPU上的张量操作,比如把label直接搬到cuda,这会让每个batch额外占一块显存。另外,offload到CPU其实是个双刃剑,如果你的offload_optimizer和offload_param设置得太激进,CPU和GPU之间频繁搬运反而会卡住,甚至报OOM因为临时缓冲没释放。
我怀疑更可能是你trainer配置里max_length或者padding策略不一致,导致实际序列长度比预期长很多。你可以试着在训练循环里打印loss.backward()之前torch.cuda.max_memory_allocated(),看峰值到底出现在哪个阶段。如果峰值就在forward里,那八成是某个自定义模块用了tensor.detach().clone()或者cat操作,这些都会打破ZeRO-3的分片假设。我上次就因为在loss里加了个torch.norm对全参数计算,结果直接OOM,去掉就好了。
大概率是自定义forward里没用model.xxx而是直接操作了tensor,导致ZeRO-3的显存视图失效,试试把输入显式传到device看看。
之前也踩过类似的坑,最后发现是data collator里自己写了padding逻辑,导致每个batch的sequence长度波动特别大,显存分配直接炸了。你可以先打印一下每个batch的max_length看看,或者干脆把padding统一到固定长度试试。另外自定义forward里如果有动态shape的tensor,比如用了torch.where或者mask,ZeRO-3的显存优化可能失效,建议用torch.profiler看下具体是哪一层爆的。demo能跑通很可能是因为它数据预处理写得很规整,你的代码里如果用了自定义dataset,检查下有没有在训练循环里偷偷保留中间变量。
我之前也踩过类似的坑,A100 40G跑7B理论上确实够,但OOM往往不是显存总量问题,而是碎片化或者某个中间激活值突然爆了。你试试把batch size降到1,然后梯度累积开大点,先排除是不是峰值显存的问题。另外,自定义forward里如果用了变量长度的padding或者动态shape,DeepSpeed的显存规划和缓存可能没法正确预估,容易在某个节点炸掉。你对比下demo脚本和你代码的输入tensor形状是不是完全一致,有时候就是多了一个维度或者mask没对齐的事儿。还有,model.to('cpu')那句是不是没写完?我猜你是想先加载再offload,顺序反了也可能导致OOM。
试试把dataloader的num_workers调成0,我之前就是多进程加载数据导致的显存虚高。
我之前也踩过类似的坑,最后发现是dataloader的num_workers开太高,CPU内存爆了反而拖累显存分配。你试试把batch size调到1跑一遍,如果还OOM,大概率是模型本身有显存碎片或者某些层没走ZeRO分区。另外自定义forward里如果有显式创建的临时tensor没清理,也可能干扰显存统计,建议用torch.cuda.empty_cache()手动清一下看看。
还有个思路:你对比下demo脚本和自己的代码,看看是否用了不同的mixed precision策略,比如demo可能是bf16而你是fp16,A100上bf16的显存占用会低不少。如果你开启了offload,确认下optimizer和gradient都真的挪到CPU了,有时候模型参数分区了但中间激活值没管住。实在不行,试试把ZeRO-3降到ZeRO-2,省下来的通信开销可能反而让显存更够用。
我之前也踩过类似的坑,A100 40G跑7B按理说ZeRO-3加offload是够的,但问题多半出在你自己的trainer配置上。建议先检查一下dataloader的num_workers和pin_memory,有时候数据预取会悄悄占掉显存。另外你说的自定义forward,如果是用了梯度累积或者中间变量没释放,DeepSpeed的显存规划确实会失灵,可以试试在forward里手动清一下缓存。还有个笨办法,把batch size降到8,先排除是不是显存碎片化的问题,再用torch.cuda.max_memory_allocated看峰值到底在哪一步爆的。
试试把trainer的batch size调成1再开梯度累积,我之前这么干直接好了,八成是你自定义forward里显存没释放。
我之前也踩过类似的坑,最后发现是自定义forward里用了太多中间变量没释放,导致activation峰值爆了,哪怕gradient checkpointing也只管backward时重算,管不了你手动保留的tensor。你试试在loss.backward()之前显式del掉那些大tensor,或者看看是不是dataloader里每个batch的padding都特别长,把序列长度统一到固定值(比如512)再测一下,大概率能解决。另外,model.to('cuda')没问题,但确认一下你的optimizer state是不是也被offload到CPU了,ZeRO-3有时候会忽略某些参数组的offload配置,得看下实际日志里每张卡的显存占用变化。
我之前也踩过类似的坑,多半问题不在DeepSpeed本身,而是你自定义的forward里有些中间变量没被显式释放,ZeRO-3对这部分特别敏感。建议你先把batch size降到1跑通,再逐步排查是不是某个tensor的shape在跨设备时出了问题。另外,offload到CPU不代表万事大吉,如果CPU内存也吃紧或者数据加载有瓶颈,照样会报OOM。可以试试在trainer里加个torch.cuda.empty_cache()的钩子,或者在每个step后手动清一下缓存,有时候能救急。
我之前也踩过类似的坑,后来发现问题是出在自定义的forward里用了绝对位置编码或者中间变量没释放,导致显存峰值比理论高不少。你可以试试把batch size降到8跑一下,如果能过就说明是峰值问题,再逐步排查是不是某个tensor没detach。另外检查一下你的trainer有没有把model.to(device)和deepseed的引擎初始化顺序搞反,这个顺序错了会直接导致offload失效。还有个小建议,把dataloader的pin_memory关掉试试,有时候这个会和ZeRO-3的通信冲突。
跑通demo但自己代码不行,大概率是数据维度的坑。你看下是不是自己把input_ids和attention_mask拼在一起传了,或者用了动态padding导致每个batch的shape不一致,DeepSpeed的显存规划器遇到这种会直接疯掉。我以前就是self-attention里多算了一次mask,显存直接翻倍。建议先在CPU上把forward逻辑跑一遍,打印每层tensor的shape和内存占用,对比官方demo的差异。
这情况我熟啊,八成是gradient checkpointing和ZeRO-3的partition没配合好。你检查下是不是用了自定义的激活函数或者自定义loss,这类东西有时候会阻止DeepSpeed对中间激活做分区。我之前用了个swish变体就炸了,换回原版gelu立马好了。另外,你试试把
我之前也踩过类似的坑,A100 40G跑7B按理说ZeRO-3加offload是够的,但你提到demo能跑通自己的不行,八成问题出在数据collator或者模型输入上。比如有些自定义forward里如果用了额外的中间变量没释放,或者pad到固定长度,显存占用会翻倍。建议你用torch.cuda.max_memory_allocated()打一下峰值,对比demo和自己代码的差距,再检查下dataloader是不是有意外保留的图。另外,你model.to('后面是写device_map还是直接to('cuda')?如果是DeepSpeed的话,手动调用model.to可能反而会干扰它的显存管理,试试完全交给trainer处理。
碰到过类似的,最后发现是DataLoader的num_workers和pin_memory在搞鬼,尤其是pin_memory=True的时候,每张卡都会额外预留一部分锁页内存,ZeRO-3又把参数切得七零八落,这部分显存碎片叠加起来比想象中夸张得多。你试试把pin_memory关掉,或者把num_workers降到0,有时候这个比offload配置更影响实际显存占用。
另外你说HF demo能跑通但自己代码不行,我怀疑问题出在model.to()的时机上。ZeRO-3要求模型先通过deepspeed.initialize()包装,之后才能移动设备,如果顺序反了,有些参数会残留在默认设备上,导致显存被悄悄吃掉。你可以检查一下是不是在初始化之前调用了model.cuda()或者model.to('cuda'),这个坑我栽过两次。
还有个思路,你自定义的forward里如果有中间变量没删干净,比如attention的权重矩阵或者embedding的临时结果,在ZeRO-3下这些可能不会自动释放,因为分片机制会让内存管理变得没那么及时。建议在forward结束前手动del一下大tensor,再torch.cuda.empty_cache()看看。
最后想确认下,你offload到CPU是用的optimizer offload还是param offload?如果只开了前者,参数还在显存里,7B模型光参数就要14G,加上梯度和优化器状态,16batch跑16000序列长度的话,40G确实可能不够。可以试试把offload_optimizer和offload_param都打开,虽然慢点但至少能跑起来。
八成是你数据collator没pad到同一长度,动态padding搞一下显存直接降一半。
offload设了但batch size 16还是爆,八成是CPU offload没生效,查下ZeRO-3的stage3_params和gradient分区配置对不对。