最近在搞一个点云相关的项目,自己写了个CUDA扩展算子(就是简单的K近邻查找+特征聚合),forward和backward都单独测过没问题。但一跑完整训练,batch size从8降到4还是OOM,看了下nvidia-smi,显存占用是持续上涨的,不是一开始就爆。怀疑是我在backward里保存了中间变量(比如索引矩阵)没释放?还是说autograd的graph没有正确剪枝?另外想问问大家,自定义算子如果涉及到scatter_add这类操作,反向的时候有没有什么容易踩的坑?有点迷茫,希望有经验的老哥指点一下,谢谢!
用PyTorch写了个自定义算子,训练时显存直接爆掉,求排查思路
全部回复
共 86 条显存持续上涨这个特征基本可以排除单纯是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。