最近在试着用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 条八成是你自定义forward里用了绝对位置编码或动态shape,ZeRO-3对这类操作会强制同步全量参数,显存瞬间爆掉。
看到你说同样的配置跑官方demo没问题,自己代码就崩,我第一反应就是数据这块的batch size没对齐。DeepSpeed ZeRO-3的显存计算是按全局batch size来的,但如果你在DataLoader里用了自己的sampler或者collate_fn,实际塞进GPU的tensor形状可能跟你想的不一样,尤其是padding到固定长度那种,一个样本顶三个样本的显存。
另外你说offload到CPU了,但得确认一下是不是所有参数都真正offload了,比如optimizer states和梯度,有时候模型参数offload了但activation还在显存里,gradient checkpointing只是省了中间激活的存储,可如果forward里用了自定义的attention mask或者额外的buffer,这些不会自动offload。
我遇到过类似的情况是自己在forward里创建了临时tensor没删,Python的GC有时候不及时,显存被碎片占着但看起来free memory够,实际训练时一申请大的连续内存就爆。你可以试试在训练循环里手动清一下缓存,或者用torch.cuda.max_memory_allocated看看峰值到底在哪一步涨上去的。
还有个坑是DeepSpeed的zero_optimization配置里,如果offload_optimizer的pin_memory设了True,但你的数据加载是异步的,可能会有隐形的临时副本在显存里。建议把offload的buffer大小调小一点,或者干脆先关掉offload,纯ZeRO-3跑一个很小的batch看能不能过,排除配置冲突。
最后怀疑一下你的model.to('cuda')之后,是不是某些子模块还留在CPU上没同步,ZeRO-3对设备一致性要求很高,混合设备会触发隐式同步导致显存峰值异常。你可以在模型加载后打印一下所有参数的device,看看有没有漏网的。
实在不行就开个profile看看显存曲线,别用nvidia-smi看,那个不准,直接用torch.profiler记录每个op的分配,基本能定位到是embedding还是某个linear层爆的。这种问题通常不是玄学,就是某个细节没对上。
我之前也被这个坑过,A100 40G跑7B按理说ZeRO-3加上offload应该绰绰有余,但问题往往不在显存总量,而在碎片化或者峰值分配。你提到同样的配置跑demo没问题,那大概率是你自己的代码里有某些操作偷偷把activation或者中间tensor留在了GPU上,比如自定义的loss计算或者数据collate里用了detach但没清缓存。建议你先把batch size降到1跑通,然后逐步加,看是哪一步开始爆的,同时用torch.cuda.max_memory_allocated()打一下峰值,对比下是不是有异常高的瞬时占用。
另外你说的自定义forward函数不兼容,这个确实有可能,特别是如果你在forward里用了python原生循环或者动态shape,DeepSpeed的partitioning逻辑会对tensor的metadata做额外处理,容易导致临时buffer没释放。我上次就是在一个自定义attention mask的处理上卡了半天,最后发现是用了list comprehension生成mask,导致每次迭代都创建新的tensor,而ZeRO-3的offload机制没法及时回收这些临时对象。你试试把forward里所有中间结果都显式赋值给局部变量,或者用torch.no_grad()包裹一些不需要梯度的计算。
还有个小细节,你的model.to('cuda')之后,有没有调用过model.train()?如果模型处于eval模式,但某些batch norm或者dropout层的行为不一致,也可能导致显存分配异常。我建议你直接用deepspeed.initialize()来构建engine,别手动做model.to(),让DS自己管理设备放置,这样能避免很多隐性问题。最后,检查下你的optimizer是不是用了AdamW,它对显存的要求比SGD高不少,如果开了offload optimizer,记得确认下offload_device是cpu,而不是nvme,否则会有额外的IO开销。
我之前也踩过类似的坑,A100 40G跑7B按理说ZeRO-3加offload是够的,但你提到同样的demo能跑通自己的不行,那八成不是显存计算问题,而是你代码里某个隐性张量在偷偷占显存。比如数据加载时如果没做pin_memory=False,或者自定义的collate_fn里把整个batch的label也送上了GPU,都会导致显存峰值暴涨。另外你说怀疑自定义forward不兼容,这个方向其实挺对的——DeepSpeed对动态控制流或者带Python list操作的forward支持很差,它会尝试把整个计算图切分,一旦有非张量操作就可能退回全量显存。你可以试试在forward里加torch.cuda.synchronize()然后打印每一层的activation显存,看看是不是某个中间变量没释放。还有就是trainer配置里如果用了gradient_accumulation_steps,加上ZeRO-3的partition梯度,有时候会额外多出一份临时buffer,建议把optimizer的offload和param的offload分开设置,别一股脑全丢CPU。我最后是自己写了个简单的training loop,不用HuggingFace的Trainer,才彻底解决,因为Trainer内部有些默认行为会覆盖DeepSpeed的配置,比如自动把model.half()或者强制no_sync,这些都可能引发OOM。你具体用的是Trainer还是自定义loop?如果是Trainer,试试把deepspeed_config里的zero_force_ds_cpu_optimizer设为false,这个选项经常被忽略。
我之前也踩过类似的坑,后来发现是自己写了个custom loss把中间变量存下来了,导致显存峰值直接翻倍。你可以试试用torch.cuda.max_memory_allocated()看下峰值是不是在backward那一步爆的,或者干脆把batch size压到8跑一下对比下显存曲线。另外既然demo能跑通,建议直接diff一下两边trainer的配置,特别是model_parallel和zero_optimization里的stage3_gather_16bit_weights_on_model_save这些参数,我上次就是栽在offload的optimizer和params没分开配。
我之前也踩过类似的坑,最后发现是dataloader的num_workers开太多,每个worker都会复制一份显存里的缓存,直接顶爆了。你试试把num_workers调成0或者关掉pin_memory,看会不会好点。另外你检查过forward函数里有没有显式创建大tensor吗?比如某些mask矩阵,这种在ZeRO-3下会绕过分片直接占显存。
还有个思路,你对比下demo脚本和你的代码,是不是自己写了trainer的梯度累加或者loss缩放?有时候这些细节会和DeepSpeed的optimizer冲突,导致显存分配异常。建议先用DeepSpeed的默认trainer跑你的数据,排除模型本身的问题。
我之前也踩过类似的坑,A100 40G跑7B按理说ZeRO-3加offload是够的,但你提到同样的配置demo能跑通,自己代码不行,那问题多半出在data collator或者模型输入上了。我上次就是自定义了一个forward返回了额外的tuple,结果DeepSpeed的梯度同步炸了,显存直接翻倍。建议你先关掉offload,把batch size降到1跑一次,看看是不是模型本身就有冗余张量在反向传播时没释放。另外检查下dataloader有没有num_workers泄露显存,或者pad到固定长度导致序列特别长,这俩都是隐形杀手。
之前跑bloom也踩过类似的坑,多半不是显存不够而是显存碎片化,试试在代码里加上torch.cuda.empty_cache(),或者把batch size先降到4看看是不是立马能跑。另外自定义forward里如果有中间变量没释放,ZeRO-3的显存统计会失真,你可以用torch.cuda.memory_summary()看下实际峰值。还有别用model.to('cuda'),ZeRO-3下得让DeepSpeed自己管理设备分配,不然容易双重占用。
offload到CPU还OOM的话,八成是你自定义forward里显存没释放,查下中间变量吧。
之前也踩过类似的坑,ZeRO-3把参数切片后,如果你的模型里有任何自定义的forward里做了跨层操作或者用了静态缓存,很容易让显存计算失效。建议先别急着怀疑数据加载,试试把trainer里的model.to('cuda')去掉,让DeepSpeed自己管理设备分配,有时候手动转移反而会破坏它的分片逻辑。另外,你把gradient checkpointing和offload同时开的时候,注意一下CPU内存是不是也爆了,有时候OOM报的是CUDA但实际是卡在数据传输上。可以的话,把demo脚本和你的代码逐行diff一下,重点看optimizer和lr scheduler的配置,我上次就是多设了个warmup step导致显存峰值暴涨。
我之前也撞过一模一样的墙,最后发现是自己dataloader里num_workers开太多,显存被CPU缓存撑爆了,把worker降到4就没事了。你那个demo能跑通,很可能就是因为它用的默认collator,而你自定义的batch里藏着没释放的中间变量。另外检查一下是不是模型里某些层没走DeepSpeed的包装,比如自己写的loss函数里额外forward了一次。
我之前也踩过类似的坑,A100 40G跑7B按理说ZeRO-3加offload确实够,但问题多半不在显存总量,而是碎片化或者临时buffer的峰值。你那个demo能跑通很关键,建议直接把你的trainer配置和demo的逐项对比,特别是save_steps和logging的间隔,有时候评估或保存模型瞬间会额外申请一大块显存。另外自定义forward里如果有动态shape或者用了大tensor的中间变量,DeepSpeed的显存规划会失效,试试把输入pad到固定长度,或者关掉offload只看纯ZeRO-3的峰值占用。数据加载那边也检查一下collator是不是把整个batch都搬到GPU了,用CPU pin memory有时候反而会加剧瞬时压力。
我之前也踩过类似的坑,最后发现是dataloader里num_workers设太高,加上自定义collate_fn里偷偷做了to('cuda'),显存被多进程复制吃掉了。你试试把num_workers调成0,或者检查一下数据加载那部分有没有隐式的设备转移。另外,DeepSpeed和自定义forward确实容易有兼容问题,特别是如果你在forward里用了Python list comprehension或者动态shape,ZeRO-3的partition逻辑会炸,建议把forward里所有张量操作都统一成tensor,别混着list。还有个笨办法,先把你自己的代码里所有model相关的部分换成HuggingFace原版Trainer,逐步二分定位到底是你哪段逻辑触发的OOM。
我之前也踩过类似的坑,最后发现是自定义forward里有个中间变量没显式释放,ZeRO-3对显存管理特别敏感,你试试把非必要的中间tensor用del删掉再清下缓存。另外你确认一下dataloader的collate_fn是不是把整个batch都pad到相同长度了,如果某个样本特别长,显存峰值可能远超预期。还有,demo脚本能跑不代表你的配置没问题,建议把batch size降到1跑一次,看是不是还OOM,能快速定位是数据还是代码的问题。
八成是自定义forward里没用DeepSpeed的engine包装,或者input_ids没走device_map。试试把model.to('cuda:0')去掉,让ZeRO接管。
试试把自定义forward里的中间变量清一下,之前我遇到过activation累积不释放的情况。
offload到CPU后记得调大NVMe缓存,不然可能卡在swap上。
我之前也踩过类似的坑,A100 40G跑7B按理说真够用,但问题往往不在显存总量,而在碎片化和activation峰值。你那句“换成demo脚本能跑,自己代码不行”特别关键,十有八九是自定义forward里某个中间tensor没释放,或者data collator把序列padding到了超长,导致单步显存暴涨。建议先老老实实把batch size调到1,开trace看看每一步峰值显存分配在哪个模块,顺便查下是不是有hidden state被重复保留没detach。另外,DeepSpeed ZeRO-3和自定义forward确实偶尔有兼容问题,特别是如果你在forward里用了python list存tensor,它没法正确切分参数,试试把那部分改成tensor操作或者换个ZeRO-2跑一下对比看看。
八成是自定义forward里没用deepspeed的wrapper包tensor,或者数据没pad到同一长度导致动态shape炸了显存碎片。