最近在搞一个点云相关的项目,自己写了个CUDA扩展算子(就是简单的K近邻查找+特征聚合),forward和backward都单独测过没问题。但一跑完整训练,batch size从8降到4还是OOM,看了下nvidia-smi,显存占用是持续上涨的,不是一开始就爆。怀疑是我在backward里保存了中间变量(比如索引矩阵)没释放?还是说autograd的graph没有正确剪枝?另外想问问大家,自定义算子如果涉及到scatter_add这类操作,反向的时候有没有什么容易踩的坑?有点迷茫,希望有经验的老哥指点一下,谢谢!
用PyTorch写了个自定义算子,训练时显存直接爆掉,求排查思路
全部回复
共 97 条显存持续上涨这个特征基本可以排除单纯是graph没剪枝,更像是backward里保存的索引矩阵在训练循环里被反复累加引用没释放,建议在backward里用非持久化的临时变量或者及时del掉。另外scatter_add反向确实容易踩坑,如果用了atomicAdd要确认梯度累加的顺序和确定性,我之前写过类似算子,最后改成先unique再segment sum才稳定。你检查下是不是每个step的中间变量都绑在了loss上没解绑?
显存持续上涨这个特征基本可以排除单纯中间变量没释放的问题,更像是有节点没从计算图里 detach 掉,导致反向时 graph 越积越长。你试试在 backward 里用 ctx.mark_non_differentiable 标记那些索引矩阵,这能让 autograd 不追踪它们。至于 scatter_add 反向,坑主要在重复索引的梯度累加顺序上,建议用 atomicAdd 而不是直接写回,同时把 forward 里的邻居索引存成 int32 而不是 int64,能省不少显存。你可以先用 torch.profiler 看下具体哪一行分配了显存,比瞎猜快。
大概率是backward里那个索引矩阵没释放,试试在反向计算完后手动del再清下缓存。
另外scatter_add反向记得用原子操作,不然梯度累加会出问题。
大概率是backward里保存的索引和中间结果没做detach,试试在保存前detach一下,或者用临时变量别存进graph。scatter_add反向坑很多,建议检查下atomic加法的梯度累积顺序。
八成是索引矩阵没释放,试试在backward里用non_blocking=True或者显式del一下。scatter_add反向记得用index的grad累积,容易踩重复索引的坑。
显存持续上涨基本可以排除graph剪枝的问题,更像是backward里某个tensor被意外保留在了计算图里,比如索引矩阵如果用了non_blocking或者没detach,很容易被autograd盯上。建议在backward结束的地方手动del掉中间变量,或者用torch.cuda.empty_cache()在每步训练后看显存曲线是否变平。scatter_add反向确实容易踩坑,我遇到过梯度重复累加的问题,记得在CUDA里对atomicAdd做线程同步,或者直接改用scatter_reduce试试。你方便贴一下backward里return的梯度shape吗?有时候广播维度不一致也会导致隐式保存大tensor。
大概率是索引矩阵没detach或者没在backward里置空,试试把不需要梯度的变量显式detach一下。
显存持续上涨这个特征,挺像graph没释放的,建议在backward里别把索引矩阵直接存成self.xxx,试着用ctx.save_for_backward并且确保自定义Function的backward返回的梯度数量和forward输入对得上,不然autograd会一直保着整张计算图。scatter_add反向我踩过坑,梯度要scatter回原位置时容易产生重复累加,记得用index_put_或者atomicAdd,而且最好先确认下是不是K近邻索引里大量重复点导致梯度爆炸,你可以在每个step打一下torch.cuda.max_memory_allocated看峰值是哪个op涨的。另外如果持续上涨而不是瞬间爆,也有可能是训练循环里某个tensor没detach,比如loss里拼接了graph相关的量,检查下有没有把中间特征不小心return出来参与loss计算。
显存持续上涨这个特征其实挺典型的,如果forward/backward单测都没问题,大概率不是算子本身泄漏,而是训练循环里每次迭代都有新的graph节点没被释放,比如把中间tensor存到了self或者list里。你可以试试在loss.backward()后手动把optimizer.zero_grad()改成set_to_none=True,顺便查一下有没有把不需要梯度的tensor误传进autograd。另外scatter_add反向就是scatter本身,但索引矩阵如果没转成int64或者没加detach,很容易把整个graph撑住,建议把索引相关的变量在backward里用完立刻置None。
显存持续上涨更像是graph没释放或变量被graph引用,试试backward后手动清下缓存,再查下索引矩阵有没有detach。
显存持续上涨这个特征,大概率不是graph没剪枝,而是你的backward里确实有张量在跨step累积。检查下索引矩阵是不是用了.detach()或者干脆在forward里就转成int64的non-differentiable tensor再存,这样autograd不会管它但显存会一直占着。另外scatter_add反向很容易踩的坑是梯度要回传到src的每个位置而不是只回传index,建议手写一个index_add_的对称反向逻辑,别直接用pytorch的gather去拼。我之前遇到过类似问题,最后发现是保存的indices没转long,导致隐式转换又建了个临时张量。你可以试试在backward入口先torch.cuda.synchronize()然后看torch.cuda.memory_summary(),能定位到具体是哪一行分配的内存。
显存持续上涨这个现象基本可以排除graph没剪枝的问题,更像是反向时张量没释放或者被重复保留。你试试在backward里把索引矩阵转成long之前先detach一下,或者干脆用index_put_这类不保存完整索引的方式。另外scatter_add反向确实容易踩坑,grad输出会自动sparse,如果你手动实现了反向,记得把grad_output先to_dense再操作,不然显存会莫名其妙翻倍。建议用torch.autograd.detect_anomaly开一下,能定位到具体是哪一行爆的,比盲猜快多了。
显存持续上涨更像是graph没释放或者索引没清,试试在backward里把不需要的中间量手动置空,顺便检查下scatter_add的反向是不是产生了重复梯度累加。
显存持续涨而不是一开始就爆,基本可以锁定是backward里保存的中间变量没释放,或者autograd graph被意外持有了。你检查下backward里是不是把索引矩阵、邻接表这类东西存成了成员变量,或者有python list一直在append。scatter_add反向确实坑多,如果用了in-place操作或者对同一索引重复累加,很容易让梯度图变得异常庞大。建议先拿torch.cuda.memory_summary()看看是哪个阶段在涨,比盲猜快很多。
显存持续上涨这点挺关键的,基本可以排除那种一次性大分配导致的OOM,更像是某些tensor被graph一直引用着没释放。你说backward单独测没问题,但那通常是单次迭代,循环跑起来才会暴露引用泄漏,建议用torch.cuda.memory_summary()对比几个step之间的allocated和reserved变化,看是活跃内存涨还是缓存涨。自定义算子如果backward里把索引矩阵、临时buffer存成成员变量或者挂在ctx上没清理,训练循环里就会一直累积。scatter_add那块确实容易踩坑,反向本质是gather,但如果你forward里对索引做了去重或排序,backward的梯度回传顺序和原始输入对不上就会静默出错,同时index的形状如果没对齐还会触发隐式的broadcast分配。还有个隐蔽的点是你在CUDA里如果用了at::empty又没及时释放,autograd的saved_tensors会把它钉住,哪怕后面不用了也等不到free。可以试试在backward里手动del掉不用的中间变量,或者用torch.autograd.set_detect_anomaly跑一小段看有没有异常引用。另外确认下你的算子是不是每次forward都新建了cuda stream或者workspace,那种也会让显存只增不减。
显存持续涨而不是一开始就爆,大概率是反向图里存了东西没释放。你确认下自定义Function的backward里是不是把forward的索引tensor直接存成成员变量了,那种会一直被graph引用着。scatter_add反向本身不背锅,但它对应的索引如果被autograd retain住,每个iteration都会累积。建议用torch.cuda.memory_summary看看是哪块在涨,再单独跑几个step对比下。
显存持续涨大概率是保存了索引没释放,试试用checkpoint或者手动detach索引。scatter_add反向容易重复累加,注意去重。