最近在试着用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 条这问题我太有共鸣了,之前我调bloom的时候也卡在类似的地方。你那个demo能跑通但自己代码不行,我赌八成不是DeepSpeed和自定义forward不兼容,而是你的dataloader在fetch数据时偷偷占了显存,尤其如果用了pin_memory=True再加上num_workers开很大,内存和显存之间会疯狂搬运,A100看起来还有空间但实际碎片化严重。另外你说offload到CPU了,但ZeRO-3的offload是分层的,你确认optimizer states和gradients都真正offload了吗?有时候只offload了参数,梯度还是留在GPU上,那照样爆。还有一招你可以试试,在trainer里显式设个max_length,把padding截断到实际长度,别让batch里最长那条拖着所有样本的显存。我之前就是有个样本特别长,结果整个batch的显存占用直接翻倍。至于model.to('device')那行,如果写死了cuda:0而DeepSpeed又自己做了device placement,可能会有冲突,建议直接删掉让DeepSpeed管理。最后实在不行,把batch size降到8跑一遍,看显存曲线是不是线性降,能定位是不是数据形状的问题。
我之前也踩过类似的坑,A100 40G跑7B按理说很宽裕,但问题往往出在自定义forward里的中间变量没释放,尤其是attention或norm层里显式创建的张量,DeepSpeed的显存统计和实际峰值是两码事。建议你把batch size降到1,先确认是不是数据collate阶段把padding搞太多了,我之前就是sequence长度没对齐导致激活值爆炸。另外,试试在trainer里加个torch.cuda.empty_cache()看能不能缓解,或者干脆把model.to('cuda')这行去掉,让DeepSpeed自己管理设备分配。如果还不行,大概率是你代码里某个子模块用了固定device的buffer,和ZeRO-3的partition逻辑冲突了。
我之前也踩过类似的坑,问题往往不在显存总量,而在碎片化或者activation峰值。你试试把gradient checkpointing和offload同时打开时,确认一下stage3的pin_memory参数,有时候这个默认开True反而会爆。另外,你自定义forward里如果有大tensor的临时变量,记得del掉再手动清一下cache,我上次就是有个中间变量没释放导致OOM。还有,对比demo脚本时,注意看它是不是用了gradient_accumulation_steps,你batch16如果等价于4步累加,峰值显存会差很多。
我之前也踩过类似的坑,显存算得刚刚好但实际跑起来就是爆。你试试把batch size再减半,或者看看是不是dataloader的num_workers开太多,CPU内存交换到显存也会挤爆。另外自定义forward里如果有创建临时tensor没及时释放,ZeRO-3的显存统计会不准,建议用torch.cuda.memory_summary()看一下具体峰值在哪。还有个偏方,把offload的cpu_device改成nvme试试,虽然慢点但能跑通。
八成是你代码里没把input_ids那些tensor显式放到cuda上,或者混用了CPU和GPU的tensor,查查dataloader的collate_fn。
这问题我也踩过,把model.to(device)改成model.to('cuda:0'),再检查下是不是有变量在循环里被重复创建了。
我之前也踩过类似的坑,A100 40G跑7B理论上确实够,但问题多半不在显存总量,而是峰值碎片化。你试试把batch size降到8,然后把gradient accumulation设成2,看还OOM不。另外,自定义forward里如果有动态shape或者中间变量没释放,ZeRO-3的显存规划会特别容易炸,建议查一下有没有把不需要的tensor手动del掉。至于model.to('cuda')那边,DeepSpeed其实会接管device,你显式调用反而可能跟它的offload逻辑冲突,改成让它自己管理试试。
我之前也踩过类似的坑,而且最后发现根本不是显存不够,是代码里某个地方偷偷把模型又复制了一份。你检查下是不是在自定义forward里用了什么会导致参数被同步或者重新实例化的操作,比如在模型内部又调了一遍model.to(device),或者用了torch.no_grad但没处理好梯度流。另一个常见问题是数据加载时num_workers设太大,每个worker会额外留一份显存缓存,A100 40G看着够用,但被这些隐性开销吃掉了。至于DeepSpeed和自定义forward不兼容这事,确实存在,特别是如果你用了动态shape或者自定义的梯度累积逻辑,ZeRO-3的partition参数和你的forward里某些tensor操作会对不上。我建议你先用最小的batch size跑通,然后把gradient_checkpointing和offload都关掉,一步步加回去,看是哪一步触发OOM。还有,确认下你的trainer有没有重复调用optimizer.zero_grad,有时候梯度累积会导致activation峰值叠加。如果demo脚本能跑,你就拿你的数据和demo的dataloader对比下,看是不是batch里max_length不一致,或者padding策略不同,导致实际sequence长度超了预期。最后,可以试试在代码里打印每层显存占用,用torch.cuda.memory_summary()看看峰值到底出现在哪。
说实话你这个现象我太熟了,A100 40G跑7B加ZeRO-3按理说16的batch真不算离谱,但OOM往往不是显存总量的问题,是碎片化或者瞬时峰值爆了。你提到自定义forward,我第一反应就是activation的显存分配可能没被DeepSpeed的显存规划器统计进去,尤其是如果你在forward里用了额外的中间张量或者非标准的attention实现,那offload和checkpointing可能根本覆盖不到那部分峰值。我之前碰到过类似情况,最后发现是我在loss计算里偷偷存了logits用于指标统计,那玩意在ZeRO-3下会强制gather全量参数,直接撑爆显存。建议你先把trainer里的compute_metrics关掉,或者把logits改成在step结束后单独用eval模式算,看看能不能跑通。另外你那个demo脚本能跑,大概率是它用的模型类有DeepSpeed官方适配的注册逻辑,而你自己的代码可能没走HookedTransformer或者没有正确设置tie_word_embeddings,导致embeddings和lm_head被重复切分。你可以试试在deepspeed config里把zero_force_ds_cpu_offload设成true,然后看看是不是某些buffer没有被正确移到device上——我之前就是忘了把position_ids的buffer注册成persistent,结果它在CPU上被反复访问,虽然不占显存但会触发同步,反而让显存峰值变得不可预测。还有个小技巧,你可以在训练循环开头手动跑一个dummy forward,把各层的显存占用打印出来,对比一下demo脚本的差异,基本能定位到是哪个模块在偷显存。
大概率是你代码里手动调了model.to(),和deepspeed的device管理打架了,试试把显式转移删掉。
数据加载器如果有多线程预取,也可能爆显存,把num_workers调成0验证下。
查一下是不是自定义forward里用了绝对位置编码或者动态shape,ZeRO-3对这两类操作会额外分配显存。
我上次就是数据collator里多留了个token,结果显存直接翻倍,去掉就好了。
我之前也踩过类似的坑,A100 40G跑7B按理说ZeRO-3加offload是够的,但你提到demo能跑通自己的不行,大概率问题出在自定义forward里有没有保留一些中间变量没释放,或者你数据collate的时候把整个序列pad到固定长度了,显存瞬间爆炸。另外检查下是不是在model.to('cuda')之后又用了accelerate或deepspeed的初始化,导致重复分配了参数。建议先关掉offload,把batch size降到4试跑一下,然后逐步加回去定位是哪一步触发的OOM,比盲调config靠谱。
我之前也栽在过这上面,后来发现是数据collator返回的labels维度没对齐,DeepSpeed的ZeRO-3对shape特别敏感,一不对就假装OOM。你可以试试在forward里print一下input_ids和labels的shape,跟官方demo对比下。另外自定义forward如果没用strict模式,有些tensor会被意外传到GPU上,建议用torch.cuda.memory_summary()看看是不是有隐藏的缓存没清。我上次就是漏了个.detach().cpu(),折腾了两天才发现。
我之前也踩过类似的坑,ZeRO-3开offload之后显存是降了,但CPU和GPU之间来回搬运反而容易触发碎片化,尤其自定义forward里有中间变量没释放的话。建议你开一下NVIDIA的nsys看看实际显存峰值在哪个阶段爆的,八成不是batch size的问题,是某个tensor没detach或者梯度累积没清。还有,你那demo脚本是不是没开混合精度?fp16和bf16对显存占用影响挺大的,A100上bf16能再省一截。对了,检查下data collator是不是把pad token搞成动态padding了,这玩意儿有时候能占掉好几个G。
我之前也踩过类似的坑,多半不是显存不够,而是你自定义forward里有些中间变量没释放,或者用了不该用的detach。你可以试试在训练循环里手动清一下cache,或者把offload的cpu_offload改成nvme_offload看看,虽然慢点但能定位问题。另外检查下dataloader的num_workers,有时候数据预取会占一堆显存,调成0试试。还有那个model.to('cuda')后面是不是漏了device_map,最好用deepspeed的initialize帮你处理。
八成是自定义forward里有些中间变量没释放,试试用torch.cuda.empty_cache加内存分析钩子查一下。
我之前也卡在过这种魔幻OOM上,最后发现是自定义forward里有个中间变量没做detach,导致计算图没释放。你可以试试在数据加载那步直接构造好input_ids和attention_mask,别让trainer每次现算,再把model.to('cuda')改成先to('cpu')再to('cuda'),有时候顺序真能影响显存碎片。另外查一下是不是dataloader的num_workers开太多,每个worker会复制一份显存上下文,A100 40G看着大但ZeRO-3的partition buffer也挺吃紧的。
我上次是改了优化器的offload策略,把optimizer和param都offload,只留gradient在GPU,瞬间就稳了。你那个demo能跑可能因为它用的是官方示例的固定配置,而你的脚本里可能某个参数没对齐,比如zero_optimization里的stage3_gather_16bit_weights_on_model_save,这个开关开着会额外占显存。建议你直接跑一下官方demo,然后改一行你的代码就对比一下显存变化,这样定位快很多。
之前我也踩过类似的坑,主要问题不一定在显存总量,而是碎片化或者某个中间变量突然暴涨。你试试把batch size直接砍到1,然后把gradient accumulation设大一点,如果这样能跑通,基本就是峰值显存没算对。另外自定义forward里如果有大tensor的临时拼接或者切片操作,很容易触发OOM,尤其ZeRO-3对张量形状变化很敏感,建议把那些操作挪到forward外面做。还有检查下dataloader是不是在GPU上直接做padding,有时候collate_fn里偷偷把数据搬到cuda就会占掉显存。你那个demo能跑通的话,对比下两边的trainer配置,重点看model_parallelism和zero_optimization.stage3_gather_16bit_weights_on_model_save这两个参数,大概率是某个细节不一样。
我之前也踩过类似的坑,最后发现是dataloader里num_workers设太高,加上自定义collate_fn里偷偷把序列pad到固定长度,显存直接翻倍了。你检查下是不是数据预处理时有个隐藏的tensor复制操作?另外ZeRO-3对模型结构里的动态shape特别敏感,试试把自定义forward里所有tensor都显式声明device,或者干脆用Peft包一层lora再套DeepSpeed,绕过原生forward的兼容问题。
试试把trainer的batch size调成1,梯度累积拉满,我之前这么干就好了,感觉是显存碎片问题。
八成是自定义forward里有些临时变量没释放,检查下有没有显式del或者用torch.no_grad包一下。
我之前也踩过类似的坑,而且最后发现根本不是显存的事。你提到同样配置跑demo没问题,但自己代码就OOM,那大概率是某个中间变量被意外保留在了计算图里。比如自定义forward里如果有Python list存了多个tensor,或者用了不该有的detach位置,都会导致激活显存峰值暴涨,ZeRO-3虽然省了参数和梯度,但激活值是不offload的。另外你说offload到CPU了,但有没有确认optimizer state和param都真正offload了?有时候只开了配置项,实际没生效,尤其版本不匹配时。还有一个容易忽略的点:你的数据加载是不是用了pin_memory=True?如果DataLoader的num_workers开太大,加上pin_memory,会在每个worker里复制一份CUDA context,显存直接翻倍。你可以试着把batch size降到1,然后把model.to('cuda')之后立刻print一下torch.cuda.memory_summary(),看看是哪个分配峰值爆了。我之前就是靠这个发现有个reshape操作把序列维度给平方了,瞬间变相加大了batch。至于DeepSpeed和自定义forward不兼容,我怀疑你用了Python的@torch.jit.script装饰器或者某些动态控制流,ZeRO-3的partitioning会跟它冲突,建议你先用ZeRO-2跑一遍对比,能通就基本锁定了。