最近在搞一个点云相关的项目,自己写了个CUDA扩展算子(就是简单的K近邻查找+特征聚合),forward和backward都单独测过没问题。但一跑完整训练,batch size从8降到4还是OOM,看了下nvidia-smi,显存占用是持续上涨的,不是一开始就爆。怀疑是我在backward里保存了中间变量(比如索引矩阵)没释放?还是说autograd的graph没有正确剪枝?另外想问问大家,自定义算子如果涉及到scatter_add这类操作,反向的时候有没有什么容易踩的坑?有点迷茫,希望有经验的老哥指点一下,谢谢!
用PyTorch写了个自定义算子,训练时显存直接爆掉,求排查思路
全部回复
共 86 条显存持续上涨这个特征很关键,基本可以排除单次forward/backward的临时分配问题,更像是graph里某个节点被反复引用导致内存无法回收。你提到保存了索引矩阵,这个确实是大坑,尤其是K近邻的索引如果用了int64,在batch size=4时可能就占了几百MB,而且autograd的graph会一直持有它直到backward结束,如果这个索引还是从输入动态生成的,那基本上每个step都会累积一次。我建议你先在backward里手动del掉那些大tensor,再调一下torch.cuda.empty_cache()看曲线是否平缓,如果还是涨,那就是graph本身没剪干净。另一个思路是用torch.autograd.detect_anomaly()结合显存快照,能定位到具体是哪个节点持有内存。scatter_add的反向确实容易出问题,因为它的梯度是scatter_add本身,但你要注意index的梯度应该是None,如果误把index也当叶子节点参与反向,graph就会异常臃肿。我之前遇到过类似情况,最后是改成用torch.index_add的原子操作版本,并且把索引转成int32,显存直接降了40%。另外你确认一下自定义算子的backward里有没有返回和forward输入数量一致的梯度,多返回一个None都会让graph多留一个引用。如果实在排查不出来,可以用torch.profiler看内存分配的时间点,基本能锁定是哪一行代码在持续申请显存。
显存持续上涨基本就是graph没释放,试试在backward里把索引矩阵detach掉或者用完手动清。
scatter_add反向记得用atomicAdd,不然梯度会丢,我之前就栽在这上面。
大概率是中间变量被graph持有没释放,试试在backward里用完索引就detach或清零。scatter_add反向记得用put_避免梯度累积爆炸。
显存持续上涨这个特征其实挺典型的,大概率不是graph剪枝的问题,而是你的自定义backward里某个tensor被意外保留在了计算图里。比如你保存索引矩阵的container是python list而不是tensor,autograd就不会自动释放它,得手动清。另外scatter_add反向确实容易踩坑,建议检查一下是不是在反向里对同一个grad_output做了多次in-place操作,这会导致显存翻倍增长。你可以试试在backward里把中间变量改成不保存,或者用torch.no_grad()包一下非必要计算,看显存曲线能不能稳住。
显存持续上涨这个特征挺典型的,八成不是一次性分配的问题,而是每次迭代都在累积。你提到backward里保存了索引矩阵,这个嫌疑很大,因为点云的K近邻索引如果没转成稀疏表示,batch size 4情况下那玩意儿可能轻松占掉几个G,而且它会被autograd当成graph的一部分一直挂着。建议你查一下是不是在forward里把索引当成了非tensor的list存下来,或者用detach()切开,不然每次反向都会新建节点。
另外scatter_add的反向确实容易踩坑,因为你得把梯度再scatter回去,这里很容易不小心在原地操作或者产生广播导致额外显存。我之前写过类似的,发现用torch.autograd.Function的时候,如果ctx.save_for_backward里存了太大的张量,即便后续不再需要,graph也不会自动释放,得手动在backward里把不需要的变量置None。
你可以试试用torch.cuda.memory_stats()打点,看是reserved bytes在涨还是allocated bytes在涨,能区分是缓存碎片还是真泄漏。还有个小技巧,把backward里的索引矩阵转成int16或者直接存邻接表,省个四倍显存。如果实在排查不出来,干脆把自定义算子在训练循环里单独跑几个step,用torch.autograd.detect_anomaly()看看能不能定位到具体哪一步。
显存持续上涨这个特征挺典型的,大概率不是autograd剪枝的问题,而是你的自定义backward里某个tensor被graph持有没释放。我之前遇到过类似情况,查了下发现是scatter_add的反向里,如果对index做了gather操作且没detach,那个索引矩阵会被当成leaf node保存下来。建议你在backward里临时变量用完都手动置None,或者试试torch.cuda.empty_cache看能不能定位到具体是哪一步在涨。另外点云场景batch size小但点数多,索引矩阵的shape很容易忽略,可以打印一下保存变量的内存占用,确认是不是索引矩阵本身太大。
显存持续上涨像是有节点没释放,试试在backward里用no_grad包住索引矩阵,或者干脆重算别存。
scatter_add反向一般就是scatter一次,但记得把grad用index加回去,不然容易梯度错位。
显存持续上涨这个特征其实挺典型的,大概率不是单次爆掉而是graph没释放。你可以试试在训练循环里加torch.cuda.empty_cache()看曲线有没有回落,或者用torch.autograd.detect_anomaly()定位具体是哪一步累积的。另外scatter_add反向确实容易出问题,特别是索引重复的时候梯度会累加,如果你在backward里手动构造了sparse gradient矩阵,记得要确保形状和dtype跟原tensor一致,不然会隐式广播导致显存膨胀。
我之前写类似算子时踩过坑,保存索引矩阵本身没问题,但如果你把整个输入x也存下来当中间变量,那batch一多就特别吃显存。建议只保存必要的shape信息和索引,反向时重新计算或者用轻量级缓存。还有个小技巧,检查一下你的K近邻查找是不是在forward里返回了完整距离矩阵,那个东西复杂度是O(N²),点云稍大一点就直接炸了,改成分块计算能省不少。
显存持续上涨这个特征太典型了,基本可以排除单次峰值问题,就是graph上挂着的东西没被释放。你提到backward里保存了索引矩阵,这个嫌疑最大,尤其是K近邻这种数量级的索引,batch一叠起来就是天文数字。我建议你先在backward的入口和出口各打一次torch.cuda.max_memory_allocated(),看看峰值差多少,如果差值远大于你预期的临时变量,那基本实锤是graph没剪干净。另外autograd的graph剪枝有个坑,就是如果你的自定义op在backward里返回了None给不需要梯度的输入,但某些中间节点还强引用着输出,graph就永远缩不掉,这个用torch.autograd.graph.saved_tensors_hooks可以查得比较清楚。
scatter_add反向这个我踩过类似的,主要是要小心grad在重复索引上的累加行为,别自己手动去重或者用torch.unique,那样梯度就错了。正确做法是直接用index_add_,或者干脆用scatter_add的数学性质,forward的scatter对应backward的gather,但要注意维度对齐。你既然单独测过backward没问题,那大概率不是数值问题,而是生命周期问题。还有个思路,你可以试试在训练循环里每隔几步手动调一下torch.cuda.empty_cache()看显存会不会回落,如果会,说明是碎片化而不是泄漏,那方向就得转向减少临时tensor的分配次数。实在不行就开torch.autograd.detect_anomaly跑几个step,虽然慢,但能定位到具体是哪个op的backward在搞事。
显存持续上涨这个特征太典型了,基本可以锁定是graph没释放或者中间变量被hold住了。你单独测backward没问题是因为只跑一次,graph会立刻释放,但训练循环里每个step的graph会累积,除非你手动调了retain_graph,不然肯定不是这里的问题。我怀疑更可能是你在自定义Function里把索引矩阵存成了self的成员变量,而且没有在backward里置空,这样每个iteration都会把旧数据留在显存里,叠加起来就爆了。建议你用torch.cuda.memory_record记录一下分配历史,看看具体是哪一行代码在持续申请内存。至于scatter_add的反向,确实有个大坑就是它本身是不可导的,你需要在backward里手动构造对grad_output的gather操作,而且要注意index的维度匹配,否则很容易隐式广播产生巨大的中间tensor。另外一个小建议,如果你K近邻的K值比较大,可以考虑用torch.nn.functional.embedding_bag或者segment_csr这类现成算子代替手写scatter,显存效率会好很多。我上次写类似算子也遇到过,最后发现是forward里为加速查询多存了一份稠密距离矩阵,那个才是大头。
显存持续上涨这个特征太典型了,基本可以排除单纯显存碎片化,大概率就是graph里某个节点没被释放。你检查下backward里是不是把索引矩阵转成long tensor存了,这玩意特别吃显存,能转int32就转int32。scatter_add反向确实容易踩坑,我建议你手动实现下梯度回传,别依赖autograd对scatter的自动推导,那个经常把grad accumulate到同一个位置导致内存暴涨。另外可以试试在backward里显式del掉中间变量再调torch.cuda.empty_cache(),虽然治标不治本但能定位问题。
显存持续上涨这个特征,大概率不是中间变量没释放,而是backward里保存的索引矩阵在autograd看来是“需要梯度”的,导致graph一直没被剪掉。你试试在保存索引的时候加.detach(),或者干脆用save_for_backward,这样能强制切断追踪。
另外scatter_add反向确实容易踩坑,它反向是gather操作,但要注意索引的维度匹配,特别是batch维度,很多OOM都是因为这里不小心广播了巨大张量。你可以先单独把backward里每个临时tensor的shape打出来,找找看有没有意外膨胀的。
我之前也遇到过类似情况,最后发现是自定义算子里用了python的list存中间结果,导致graph无法释放。建议你尽量用torch的tensor操作,避免python对象参与反向传播。
显存持续上涨这个特征基本可以排除单纯显存峰值问题,很可能是backward里保存了不该存的张量,比如索引矩阵如果没显式detach或者转成long之后还留在计算图里,graph会一直挂着。我之前遇到过类似情况,用torch.autograd.gradcheck配合torch.cuda.memory._record_memory_history()去抓具体哪一行分配了内存,比瞎猜快很多。scatter_add反向的话,注意梯度要按index做gather而不是直接回传,另外如果用了inplace操作记得检查版本兼容性,建议写个最小复现脚本把K近邻的K设小一点跑几个step看内存曲线,基本能定位到是算子内部泄漏还是graph累积。
显存持续上涨这个特征基本可以排除单纯graph没剪枝的问题,大概率是backward里某个tensor被意外retain了。你检查下自定义Function里是不是把self.save_for_backward和普通self.xxx混用了,后者会导致整个graph生命周期被拉长。另外scatter_add反向时注意grad_output在index重复位置上的累加,最好用atomicAdd或者先unique再分段处理,不然容易出数值问题而且debug起来很头疼。
显存持续上涨这个特征挺典型的,大概率不是单次峰值问题,而是graph里某个节点把tensor引用挂住了。你可以试着在backward里把索引矩阵转成bool mask或者干脆用index_put_,少存一个int64的大张量能省不少。另外scatter_add反向确实容易踩坑,因为梯度要回传到源位置,建议你检查下是不是在反向时对同一地址做了累加操作,导致autograd的buffer像滚雪球一样越滚越大,试试用torch.autograd.detect_anomaly()定位一下具体是哪一步在涨。
显存持续上涨而不是直接爆,大概率不是单次分配的问题,更像是graph里某个节点没被释放,或者你保存的索引矩阵在每次iter里不断累积。我之前写过类似的聚合算子,反向里如果用scatter_add,记得把gradient的accumulation显式置零,不然它会一直叠加。另外你可以试试torch.cuda.memory_summary()看下分配趋势,能定位是哪一层在涨。还有个思路,把backward里的中间变量改成不保存,用重计算的方式,虽然慢点但显存稳。
显存持续上涨这个现象挺典型的,大概率不是autograd没剪枝,而是你那个索引矩阵在backward里被反复引用,导致graph一直没释放。我之前写类似算子时踩过同样的坑,建议你把需要保存的中间变量改成显式调用detach()或者在自定义Function里用nonlocal存,而不是一股脑塞进ctx.save_for_backward。另外scatter_add反向确实容易出问题,它本质上是gather操作,容易产生梯度累积,你最好检查一下是不是在backward里对同一位置重复写入了,导致显存碎片化。可以先试试在训练循环里加torch.cuda.empty_cache()看能不能缓解,但根治还得从变量生命周期入手。
显存持续上涨这个特征基本能排除单纯graph没剪枝的问题,更像是backward里某个操作在循环中累积了buffer。你查一下KNN索引是不是用了Tensor.clone()或者detach()后又参与了反向,这种最容易悄悄占显存。scatter_add反向的坑主要是grad累加时索引重复会覆盖,建议用atomicAdd或者把索引展开成one-hot再乘,不过后者更费显存。另外确认下自定义算子的backward有没有显式释放非必要保存的变量,尤其是那些大shape的int64索引。我之前写类似算子时试过在backward末尾手动清空缓存,效果挺明显。
显存持续上涨这个特征挺典型的,大概率不是graph没剪枝,而是backward里保存的索引矩阵没被释放,试试在自定义Function的backward里把不需要的中间变量用del手动清掉,或者改成不保存索引、在反向时重新计算。scatter_add的反向确实容易踩坑,主要是要确保对grad_output做gather时索引维度对齐,不然会静默出错,建议用torch.autograd.gradcheck单独验证一下梯度。另外你batch size降一半还涨,看看是不是哪里不小心创建了新的计算图,比如在循环里用了requires_grad的变量做累积。
显存持续上涨这个特征挺典型的,大概率不是算子本身泄漏,而是backward里保存的索引矩阵在batch维度上累积了,试试在反向函数里把不需要的中间变量显式置None,或者用ctx.mark_non_differentiable标记那些索引张量。scatter_add反向确实容易踩坑,我遇到过梯度重复累加的问题,建议用at::index_add_反向替代手写循环,反正先查一下是不是每个step的graph都没释放。