最近在搞一个点云相关的项目,自己写了个CUDA扩展算子(就是简单的K近邻查找+特征聚合),forward和backward都单独测过没问题。但一跑完整训练,batch size从8降到4还是OOM,看了下nvidia-smi,显存占用是持续上涨的,不是一开始就爆。怀疑是我在backward里保存了中间变量(比如索引矩阵)没释放?还是说autograd的graph没有正确剪枝?另外想问问大家,自定义算子如果涉及到scatter_add这类操作,反向的时候有没有什么容易踩的坑?有点迷茫,希望有经验的老哥指点一下,谢谢!
用PyTorch写了个自定义算子,训练时显存直接爆掉,求排查思路
全部回复
共 86 条显存持续上涨这个特征其实挺典型的,大概率不是索引矩阵没释放,而是你backward里那个自定义的K近邻索引在反向传播时被autograd当成graph的叶子节点保存了,导致每步训练都累积一份。你可以试试在forward里把索引用.detach()包一下,或者干脆用ctx.mark_non_differentiable标注,这样graph就不会追踪它了。scatter_add反向确实容易踩坑,因为它的梯度要按index做gather,你得确保index在CPU上或者用torch.cuda.synchronize检查一下是不是有隐式同步导致显存碎片。还有个小建议,把nvidia-smi改成每100ms采样一次,看是不是某个特定层峰值涨,而不是整体缓涨,能帮你缩小范围。
显存持续上涨大概率不是graph剪枝的问题,更像是backward里保存的索引矩阵没被释放,试试在自定义Function的backward里把不需要的中间变量用del手动清掉,或者改用save_for_backward只存必要的。scatter_add反向确实容易踩坑,我之前遇到过梯度累加顺序不对导致数值错误,建议你用torch.autograd.gradcheck先验证一下梯度,同时监控一下每步的显存峰值,看看是不是某个临时tensor没释放。
显存持续上涨这个现象其实挺典型的,大概率不是autograd graph剪枝的问题,而是你的自定义backward里某个tensor被意外保留了下来。我之前写索引类算子也遇到过,尤其你提到保存了索引矩阵,那个如果是int64的,shape又是N×K的话,占用比想象中大得多,而且如果它被绑定在graph里,每个step都不会释放。你可以试试在backward里把不需要的中间变量显式置None,或者用del加torch.cuda.empty_cache()看下曲线有没有掉下来,但别依赖这个,治标不治本。
另外scatter_add的反向确实容易踩坑,主要在于它的梯度是scatter_add本身,但如果你在forward里用了index的复用,反向时index的梯度是None,但autograd可能会因为index参与运算而试图构建一个巨大的稀疏矩阵,间接导致显存爆炸。我建议你检查下是不是用了类似torch.zeros_like再scatter的模式,这个在点云里特别常见,但反向时梯度会变成一个dense的累加缓冲,非常吃显存。可以试试把scatter改成index_add_或者用segment_csr这类更底层的方式,能省不少内存。
还有个排查思路,就是直接在训练循环里对每个step打点,看是forward阶段涨还是backward阶段涨。如果是forward涨,那就是临时变量没释放,如果是backward涨,那就是graph累积或者梯度缓冲的问题。你可以用torch.autograd.detect_anomaly()跑一下,虽然慢,但能定位到具体是哪一行触发了问题。另外确认下你的自定义算子是不是每次调用都新建了CUDA context,这个有时候也会导致显存碎片化,看着像泄漏。
反向时把索引矩阵转成bool mask或者干脆重算,别硬存,省下的显存能救你命。
scatter_add反向就是scatter,记得把梯度accumulate到正确位置,不然会静默丢数据。
显存持续上涨而不是直接爆,大概率不是graph剪枝的问题,更像是backward里保存的索引矩阵没被释放,KNN那个索引在batch大时非常占显存,建议用non_blocking或者直接在forward里重算索引。scatter_add反向确实容易踩坑,它的梯度是scatter_add本身,但如果你用了原地操作,得确保grad累积时没有race condition,可以试试用torch.scatter_reduce的reduce="add"来替代。另外你确认一下自定义算子是否实现了setuptools的cuda caching allocator,没实现的话每次调用都可能额外分配内存。
显存持续上涨更像是graph没释放,试试在backward里用del手动清下中间变量。
显存持续上涨大概率是backward里保存了索引矩阵没释放,试试在反向计算完后手动del掉。另外scatter_add反向记得用原子操作,不然梯度会丢。
显存持续上涨基本就是graph没释放,试试在backward里把索引矩阵detach或者用完就清,别留着。
scatter_add反向确实容易炸,建议检查下atomic操作有没有加锁,或者改成segment_sum试试。
显存持续上涨基本就是graph没释放,试试在backward里用no_grad包一下索引操作。
scatter_add反向确实容易踩坑,记得用put_代替,能省不少显存。
显存持续上涨这个特征很关键,我猜大概率是backward里那个索引矩阵被autograd当成leaf tensor存进graph了,试试在自定义Function的backward里把不需要梯度的中间变量用.detach()或者干脆不保存,改用重建的方式算一遍。另外scatter_add反向确实容易出问题,我记得它对应的是gather操作,但如果你在backward里又用scatter_add去回传梯度,可能会有重复累加导致的显存泄漏,建议先检查一下是不是每次step都在往一个变量上累加梯度。还有个小技巧,把batch size调到1跑几个step看看显存曲线,如果还是线性涨,基本就是graph没释放,可以考虑每步手动清一下缓存。
显存持续上涨这个特征其实挺典型的,大概率不是graph没剪枝,而是backward里保存的索引矩阵在batch循环里没被释放,或者被autograd当成leaf tensor缓存了。你可以试试在backward里把不需要梯度的中间量用detach()包一下,或者干脆用non_blocking=True的临时tensor覆盖。另外scatter_add反向确实容易踩坑,特别注意一下index的维度对齐和重复索引的梯度累加逻辑,我之前就是没处理重复点导致梯度爆炸,跟你这个OOM一起出现的话,建议先加个torch.cuda.empty_cache()观察每步显存曲线。
大概率是索引矩阵没做detach或者反向时被graph持有,试试在保存前clone并detach一下。
scatter_add反向确实容易爆,建议把中间变量转成稀疏表示或者直接重算,别存完整张量。
显存持续上涨这个特征确实很像graph没释放,但我觉得更可能是你在backward里保存的索引矩阵被autograd当成需要梯度的节点了,试试在保存的时候detach一下或者干脆用non_blocking的临时tensor。另外scatter_add反向就是scatter,坑在于索引的梯度是累加的,你如果没处理好重复索引的梯度累积,很容易在batch变大时把显存堆上去。我建议你用torch.autograd.detect_anomaly跑一下,它能帮你定位到具体是哪一行爆的,比盲猜快多了。顺便问下你用的是不是原子操作的scatter_add?那个在反向时如果没加锁,显存波动会很奇怪。
显存持续上涨而不是直接爆,大概率不是graph没剪枝,而是你backward里保存的索引矩阵每个step都在累积,试试在自定义Function的backward里显式del掉那些大tensor,或者用ctx.save_for_backward但反向结束后把ctx置空。另外scatter_add反向确实容易踩坑,如果是inplace操作记得加torch.cuda.synchronize(),不然异步执行下显存峰值会虚高。你测单独forward/backward时是固定输入size吗?训练时点云数量会不会变化,导致每次分配的workspace大小不一致?建议先python自带tracemalloc定位下是哪一行申请的内存,比猜效率高。
显存持续上涨而不是直接爆,大概率不是graph剪枝的问题,而是backward里保存的索引矩阵在反向传播时又被复制了一份,而且CUDA context本身也有缓存,可以试试torch.cuda.empty_cache()看下曲线是不是锯齿状。scatter_add反向的话,注意梯度要按index做累加而不是赋值,之前我在这踩过坑,用torch.autograd.Function的话记得在backward里把grad_output按index维度展开再scatter_add回去。另外你那个索引矩阵如果是整形的话,可以考虑用torch.ops.aten.index_put_的inplace版本或者直接存到self里但用完后手动释放,不知道你具体怎么写的,方便贴下backward代码吗?
显存持续上涨而不是一开始就爆,这个特征其实挺关键的,大概率不是graph没剪枝,因为PyTorch的autograd在反向传播结束后会自动释放中间节点的buffer,除非你在forward里显式retain_graph或者把中间结果存成了self.xxx。我怀疑你是在backward里构造了新的tensor参与梯度计算但忘了手动释放,或者是在循环里累积了某些引用没清掉,建议用torch.cuda.reset_peak_memory_stats()配合torch.profiler看一下峰值到底涨在哪个step。scatter_add这个操作确实容易踩坑,尤其是反向的时候,如果你对同一个索引位置多次累加,梯度会自然累积,这在数学上没问题,但如果你在backward里又用scatter_add去回传梯度,得注意atomicAdd的随机性可能导致非确定性结果,影响调试。另外建议检查一下自定义算子的backward函数里,是不是隐式创建了和输入shape相关的临时变量,比如索引矩阵如果用了torch.nonzero或者repeat_interleave,这些操作在反向时会产生巨大的中间梯度图,即使你最终没保存它们。还有个野路子,你可以试着把batch size降到1,如果显存不再涨了,那基本就是中间变量和batch size线性相关,这时候就得考虑用torch.utils.checkpoint或者手动重计算来省显存。
显存持续上涨这个特征基本可以排除临时buffer的问题,更像是有节点没从计算图里摘掉。建议在backward里把保存的索引转成bool mask或者干脆在forward里重算一遍,省得图一直挂着。scatter_add反向确实容易踩坑,我记得grad得用index_add_回传,直接scatter的话梯度会覆盖而不是累加,你可以先检查下是不是这里写岔了。另外你试试开torch.cuda.empty_cache()看下曲线,能帮你区分是碎片还是真泄漏。
显存持续上涨而不是一开始爆,这个特征其实挺关键的,大概率不是简单保存了中间变量,而是graph里某个节点引用了不该引用的tensor,导致autograd把整个计算图都留着不放。我之前写过类似的自定义scatter_add反向,踩过一个坑是:如果你在forward里用了torch.cuda.graphs或者显式保存了input的索引,但是backward里又对这些索引做了in-place操作,graph就永远剪不掉。建议你先用torch.autograd.detect_anomaly()跑一下,它能定位到具体是哪个op导致的graph保留,顺便检查一下你的K近邻索引是不是用了long tensor还转成了cuda,如果索引本身没被detach,反向的时候它会被当成叶子节点参与梯度计算,那内存就炸了。另外scatter_add反向有个经典问题,就是梯度要scatter回去的时候,如果index里有重复值,你直接grad.scatter_add_会累加,但如果你用了grad.index_add_或者直接赋值,可能会覆盖导致梯度错误,这个和内存倒是关系不大,不过建议你统一用scatter_add_保证正确性。还有个排查技巧,把batch size设成1跑一个step,然后看显存峰值和结束后的占用差,如果结束还占着几百MB,那就是graph没释放,如果只占一点,那就是forward里分配了太多临时buffer没清。你提到batch从8降到4还是涨,那我觉得可能是你的算子内部用了动态shape的临时tensor,比如根据点数分配了最大尺寸的buffer,这个buffer在每次迭代都会重新申请而旧的不释放,试试在算子内部用torch.cuda.memory_format或者显式del掉中间量,或者干脆把最大点数固定下来预分配。最后问一下,你的backward里有没有返回给forward的input的梯度?如果input本身不需要梯度,记得return None,不然graph会强制保留input的引用,这个特别容易忽略。
显存持续上涨大概率是backward里保存的索引没detach,试试在保存前加.clone().detach()。
scatter_add反向确实容易踩坑,建议检查一下梯度是否在重复索引位置累加错了。
显存持续上涨这个现象很关键,大概率不是graph剪枝的问题,而是你backward里那个索引矩阵被autograd当成叶子节点保存了,试试在保存前用detach或者干脆转成普通tensor。另外scatter_add反向确实容易踩坑,你确认下梯度是不是只回传到参与聚合的那个点上,不然会有重复累加导致显存翻倍的情况。我之前也遇到过类似问题,最后是用torch.utils.cpp_extension的debug模式跑了一遍,能看到每个op的显存分配,你可以试试。