最近在试着用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,求大佬指点
全部回复
共 134 条我之前也踩过类似的坑,最后发现是dataloader没设pin_memory,加上自己写的collate_fn里有隐式转换,导致显存碎片化。建议你先把数据加载改成最朴素的写法,排除掉这个变量再试。另外,自定义forward里如果有Python list或者动态shape操作,DeepSpeed的显存规划会失效,可以试着把输入都pad到固定长度。还有个笨办法,把ZeRO-3的stage3_max_live_parameters调小一点,强制更频繁地释放参数,虽然慢点但能跑通。
八成是自定义forward里没用model的device,试试把input显式.to(device)再传。
之前也踩过类似的坑,后来发现是自定义forward里用了绝对position embedding或者中间变量没走model parallel的接口,ZeRO-3会把这些张量当成普通参数来切分,反而爆显存。你可以试试把trainer里的model.to('cuda')去掉,让DeepSpeed自己管理device,另外检查一下dataloader有没有在batch里带额外的CUDA tensor,这种最容易隐性OOM。还有个排查技巧,把offload先全关掉,batch size降到2跑一次,如果还爆就是代码问题,不爆就慢慢加回去找临界点。
我之前也踩过类似的坑,后来发现多半不是显存不够,而是显存碎片化或者临时tensor没释放。你试试在trainer里加个empty_cache,或者把batch size调到8看看能不能跑通,如果8能跑16不行,那就是峰值内存的问题。另外,自定义forward里如果有大中间变量,记得用with torch.no_grad()包一下,或者检查下是不是有变量被意外存到了GPU上。我上次就是有个list在循环里一直append tensor,结果全堆显存里了。
我之前也踩过类似的坑,症状几乎一模一样,最后发现不是显存的问题,而是显存碎片化导致的。ZeRO-3会把参数、梯度、优化器状态都切分,如果你的自定义forward里有临时大tensor(比如中间激活没及时释放),或者数据加载时每个batch的shape不一致,就特别容易触发碎片化OOM。你可以试试在训练循环里手动加torch.cuda.empty_cache(),或者把batch size调小一半先跑通,看是不是瞬间就好了。
另外你提到offload到CPU了,但要注意offload通常是配合NVMe用的,纯CPU offload反而会拖慢速度,而且如果CPU内存不够或者swap配置没调好,也会报奇怪的OOM。我建议你检查一下DeepSpeed的zero_optimization配置里有没有设“reduce_scatter”和“allgather”的bucket大小,有时候默认值在7B模型上会爆。
至于demo脚本能跑通而你自己的不行,我赌五毛是trainer里的data collator或者模型输入格式有问题。比如你的模型forward里如果用了attention_mask和position_ids,但数据没传全,DeepSpeed的封装层会额外申请内存来补默认值。你可以把自定义forward简化成最朴素的输入输出,先排除模型本身的问题。最后问一下,你用的DeepSpeed版本是多少?之前有个版本和transformers新接口有冲突,升到0.14以上可能就解决了。
大概率是你自定义forward里没用model的embedding那些层,导致ZeRO-3切分参数时没被感知到,试试给自定义层也包上deepspeed的hook。
这问题多半出在自定义forward上,ZeRO-3对动态图和显存分配很敏感。你把offload改成ZeRO-2试试,大概率能跑通。
之前也踩过类似的坑,多半不是显存算错的问题,而是你代码里某个张量悄悄被复制了一份,比如在自定义forward里用了.cpu()或者.detach()之后又传回GPU,这样ZeRO-3的显存规划就完全失效了。建议你用torch.cuda.memory_summary()在报错前打一下,看是不是有非模型参数的缓存涨得特别离谱。另外,HuggingFace那个demo能跑通是因为它可能默认把input_ids和labels都放在同一设备上,而你自己的trainer如果手动做了to('cuda:0')之类的操作,反而会打乱DeepSpeed的partition逻辑。可以试试把整个trainer换成Trainer类,只改model_parallel相关参数,别的都保持默认,大概率能定位出问题。
之前也踩过类似的坑,排查下来多半不是DeepSpeed本身的问题,而是自定义dataset或collator里偷偷把tensor留在了GPU上没释放。试着在trainer里加个显存追踪,每个step打印一下reserved和allocated,能很快定位到是哪一步暴涨。另外你提到demo脚本能跑通,建议直接diff两边data和model的差异,很多时候是padding策略不同导致序列长度远超预期。如果自定义forward里有中间变量没清理,ZeRO-3的显存优化会失效,可以试试把offload从CPU换成NVMe看报错是否变化。
我之前也栽在过类似的坑里,排查到最后是自定义forward里有个中间变量没做detach,导致计算图一直没释放。建议你先把trainer的batch_size调到1跑一次,如果还OOM就直接看堆栈,大概率是某个激活值没被checkpoint包住。另外你说的demo能跑通,很可能它用的model.forward是原生的,而你改了输入输出结构,ZeRO-3的partition逻辑会在自定义层上出问题。可以先试试不offload,纯ZeRO-3配gradient checkpointing,看显存峰值差多少,这样能快速定位是不是数据加载那块额外占了你没注意到的显存。
检查下自定义forward里是不是多了没必要的中间变量,那玩意儿不释放显存也扛不住。
先别急着怀疑DeepSpeed,大概率是你自己代码里某个tensor没detach或者数据没清干净,建议从数据加载那块排查。
ZeRO-3下显存够用但OOM,八成不是模型本身的问题,而是offload和checkpointing的配置没对齐。你提到HF demo能跑通、自己的代码不行,这其实是个很关键的线索,说明瓶颈大概率在数据侧或者trainer的组装方式上。比如自定义dataset如果返回的tensor没有pin_memory,或者collate_fn里偷偷保留了计算图引用,stage3分片参数时就会多占一份显存。另外ZeRO-3对forward里临时创建的参数副本很敏感,如果你在model内部手动做了weight拼接或clone,分片机制会失效,瞬间把40G吃满。建议先关掉offload单独跑一个step,用torch.cuda.memory_summary看看到底是参数、梯度还是激活占了大头。还有个小坑是batch size 16在ZeRO-3下如果没开contiguous_gradients和overlap_comm,通信缓冲区也会吃掉不少显存。可以先把batch降到4,确认能跑通后再逐步往上加,同时对比HF demo和你自己代码在trainer初始化时的deepspeed config差异,通常问题就藏在那几行里。
ZeRO-3下OOM跟显存“够不够用”经常不是一回事,参数被切分之后通信buffer、临时all-gather出来的完整层、还有优化器状态peak都会叠在一起,40G跑7B bs16其实挺悬的。你说demo脚本能跑、自己的代码不行,这个信息量很大,八成是训练循环里有些细节不一样,比如loss是不是做了mean、有没有在forward里额外保留中间激活、labels的shift位置对不对。另外你offload到CPU之后如果pin_memory没配好,或者stage3_gather_16bit_weights_on_model_save这类设置没开,反而会在某一步把整层参数拉回GPU导致炸显存。建议先把batch size降到1跑通,再用ds_report看实际配置,然后逐步往上加,顺便确认你的自定义forward有没有在torch.no_grad外面偷偷建计算图。还有个坑是数据加载器返回的tensor如果默认在GPU上,或者collate_fn里做了padding到很长,也会让显存悄悄涨上去,可以先打印一下每个micro-batch的真实seq长度。