最近在搞一个点云相关的项目,自己写了个CUDA扩展算子(就是简单的K近邻查找+特征聚合),forward和backward都单独测过没问题。但一跑完整训练,batch size从8降到4还是OOM,看了下nvidia-smi,显存占用是持续上涨的,不是一开始就爆。怀疑是我在backward里保存了中间变量(比如索引矩阵)没释放?还是说autograd的graph没有正确剪枝?另外想问问大家,自定义算子如果涉及到scatter_add这类操作,反向的时候有没有什么容易踩的坑?有点迷茫,希望有经验的老哥指点一下,谢谢!
用PyTorch写了个自定义算子,训练时显存直接爆掉,求排查思路
全部回复
共 86 条显存持续上涨这个现象很典型,大概率不是graph剪枝的问题,而是backward里保存的索引矩阵没被释放,尤其是K近邻这种index tensor,尺寸跟batch size挂钩,建议在自定义Function里用ctx.mark_non_differentiable标记索引,然后手动把不需要反向传播的变量用detach或者不保存。另外scatter_add的反向确实容易踩坑,如果你是直接对index做gather的话,记得处理重复索引的梯度累加,别用inplace操作,我之前就是在这卡了好久。
显存持续上涨八成是backward里索引没释放,试试在反向函数里手动清一下中间缓存。
scatter_add反向容易梯度累加错位,建议用atomicAdd或者检查一下index的维度对齐。
显存持续上涨而不是直接爆,大概率不是graph没剪枝,而是backward里那个索引矩阵被当成叶子节点保存了,每步都累积。你可以试试在自定义Function里把不需要梯度的中间变量显式写成non_tensor或者用packed结构,别让它进autograd的saved_tensors。另外scatter_add反向时,如果索引有重复,梯度累加的顺序可能和预期不一致,建议检查一下atomicAdd的粒度,或者干脆用torch的index_add_代替手写CUDA,排除一下是不是算子本身的问题。
显存持续上涨这个现象其实挺典型的,不一定是graph没剪枝,更大概率是你在backward里保存了那个索引矩阵没释放,因为K近邻的索引维度是BNK,点云项目里N动不动上万,这个中间量在训练时每个step都会累积在计算图里直到step结束才清。你可以试试在自定义Function的backward里用完索引后手动del掉,或者干脆用non-blocking的方式在forward里就把它转换成稀疏表示,能省不少显存。另外scatter_add反向的话,我猜你可能是对梯度做scatter_add回到原特征位置,这里有个坑是梯度累加的时候会重复计算,如果索引有重复元素,梯度会被放大,最好用atomicAdd或者确保索引唯一,不然数值上会出问题但不会爆显存。还有个排查技巧,你可以把batch size设成1跑几个step,同时监控显存变化曲线,如果还是持续增长,那就是代码里有变量在跨step累积,比如把临时tensor存到了self里。我之前遇到类似情况是发现反向里用了in-place操作改了某个buffer,导致autograd的graph没法释放,你可以检查下有没有对叶子tensor做in-place修改。你forward和backward单独测没问题但合起来爆,很可能就是中间变量生命周期的问题,建议用torch.cuda.memory_summary()看具体哪个op分配了最多内存,一目了然。
显存持续上涨这个现象很关键,大概率不是autograd没剪枝的问题,而是你在backward里保存的索引矩阵被当成叶子节点参与了反向传播,导致每次迭代都累积计算图。你可以试试在保存索引时加上.detach(),或者用torch.no_grad()包一下那部分操作。另外scatter_add反向确实容易踩坑,如果你是用index_put_这类inplace操作,记得在forward里把输入clone一份,否则梯度会写进原始buffer。我之前遇到过类似问题,最后是把K近邻的索引改成int16存储,显存直接省了一半,你可以试试。
显存持续上涨这个特征太典型了,基本可以排除单纯的中间变量没释放,更像是backward里某个操作在graph上累积了历史节点。你试试在训练循环里加torch.cuda.empty_cache()看峰值有没有回落,能区分是缓存碎片还是真泄漏。scatter_add反向确实容易踩坑,我之前遇到过grad在索引重复位置累加导致显存翻倍的问题,建议把索引矩阵转成int32或者干脆用torch.utils.checkpoint重算一遍,省显存效果立竿见影。另外你检查下自定义算子有没有正确实现set_context,有时候默认的materialize_grads=True会把不需要的梯度也存下来。
显存持续上涨这个特征挺典型的,我猜大概率不是graph没剪枝,而是你在backward里保存的索引矩阵没释放,试试在自定义Function的backward里用del显式删掉大tensor,或者改成不保存索引、forward里现算。另外scatter_add反向确实容易踩坑,主要得注意梯度要scatter回原始位置,别用inplace操作,不然autograd的buffer会越积越多。你单独测backward没问题但训练爆,也可能是loss backward之后没手动清空中间缓存,检查下有没有什么地方把临时tensor挂到了self上。
显存持续上涨基本就是graph没释放,试试在backward里把不需要的中间变量detach掉,或者用del手动清一下。
显存持续上涨这个特征很关键,我赌八成不是graph没剪枝,而是backward里那个索引矩阵被autograd当成需要梯度的叶子节点了。你试试在保存index的时候用detach(),或者干脆用register_buffer,不然每次反向都会累计一份新的引用。另外scatter_add的反向确实是经典坑,尤其是当同一个index对应多个梯度时,pytorch的atomicAdd行为在不同架构上表现不一致,我建议你手动检查一下梯度累加的顺序,或者用segment_csr那套替代方案。还有个小技巧,你可以用torch.cuda.memory._record_memory_history()抓一下分配栈,看看到底是哪一行触发了峰值,比盲猜快很多。顺便问下,你的K近邻是用暴力法还是grid加速?如果是暴力法,那中间那个N×K的distance矩阵就算不保存,可能在kernel内部也占了一波显存,试试把这个矩阵改成half或者直接复用内存。最后,如果batch size降到4还爆,建议先跑一次单batch的完整训练,确认不是优化器状态或者batch norm buffer在累积,我之前被这个坑过。
显存持续上涨多半是graph没释放,试试在backward里用no_grad包一下索引计算。
显存持续上涨这个特征其实挺典型的,大概率不是graph没剪枝,而是backward里那个索引矩阵被autograd当成需要梯度的叶子节点保存了,试试在自定义Function的backward里把不需要梯度的tensor用.detach()或者直接转成long再返回。另外scatter_add反向时要注意梯度累加的顺序,如果同一个位置被多个点scatter,梯度应该sum而不是覆盖,我之前就在这里踩过坑。你可以在训练循环里每隔几步打印一下torch.cuda.memory_allocated()看是不是线性增长,如果是那就是有节点没释放。
显存持续上涨这个特征挺典型的,大概率不是graph没剪枝,而是backward里那个索引矩阵被当成叶子节点保存了,试试在自定义Function里用ctx.mark_non_differentiable标记一下,或者干脆把索引换成int32的tensor存到CPU上。另外scatter_add反向确实容易踩坑,你得手动实现这个op的grad,注意一下梯度要scatter回原始位置而不是累加到目标位置,不然梯度会越传越大。我之前也遇到过类似问题,最后是给中间变量加了个detach才解决,你可以先检查一下是不是有变量意外参与到了反向图里。
显存持续上涨这个特征,大概率不是graph剪枝的问题,更像是你保存的索引矩阵在训练循环里被某个外部list引用了,导致autograd的buffer没法释放。我之前遇到过类似情况,排查时可以先试试在backward里把非必要的中间变量用del手动删掉,再配合torch.cuda.empty_cache()看曲线有没有变化。另外scatter_add反向确实容易踩坑,主要得确认你的atomicAdd在梯度累加时没被重复计算,尤其是当索引有重复值时,最好在CPU上用小case对比一下数值。你试试把batch size调回8,但把点云数量减半跑一下,能更快定位是显存泄漏还是计算量问题。
显存持续上涨这个现象很关键,说明大概率不是单次前向/反向的峰值问题,而是graph里某个节点一直在被引用导致内存无法回收。你提到backward保存索引矩阵,这个很可疑——试试在自定义Function里把非必要的中间量改成用index_add_之类的重计算代替存张量,或者干脆在反向最后手动删掉临时变量。另外scatter_add的反向确实容易踩坑,因为梯度要scatter回原位置,如果索引有重复,记得grad用atomicAdd的语义,别用简单的赋值。你检查过cuda caching allocator的碎片吗?有时候不是真爆了,是碎片太多导致分配不出连续块,可以试试环境变量PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True。
我之前也遇到过类似的情况,最后发现是backward里保存的索引矩阵没做detach,导致整个计算图一直挂着,显存自然就只涨不跌。你可以试试在保存中间变量时用detach或者只保存必要的索引,别把整个张量都留下来。另外scatter_add反向确实容易踩坑,特别是如果索引有重复,梯度累加的顺序会影响结果,建议用torch.autograd.gradcheck先验证一下反向的正确性。还有个小技巧,可以用torch.cuda.memory_summary()看看具体是哪一层分配的内存,能帮你快速定位问题。
显存持续上涨这个特征其实挺典型的,大概率不是graph没剪枝,而是backward里保存的索引矩阵在训练循环里被反复累积了。你试试在自定义Function里把不需要梯度的中间量用non_blocking或者直接不存,改成在forward里重新算一遍,虽然慢点但能救急。scatter_add反向的话,注意梯度要按index做atomicAdd,别直接用scatter回传,容易丢数据。你检查下是不是batch内每个点的K近邻索引长度不一致,导致padding的部分也被当成有效梯度累积了?
大概率就是索引矩阵没释放,建议查下backward里有没有detach或者把非必要梯度变量缓存清掉。scatter_add反向记得用put_代替,能省不少显存。
这个显存持续上涨的曲线很典型,基本就是graph没释放或者中间变量被retain了。你查一下backward里有没有把索引矩阵存成self.xxx,这样每个step都会累积。另外scatter_add反向确实容易踩坑,建议用torch.autograd.Function的ctx.save_for_backward只保存必要的,别顺手把大tensor也存进去。
我之前写过类似的算子,发现是forward里用了non_blocking=True但没同步,导致显存碎片越来越多。可以试试在backward开头加个torch.cuda.synchronize(),顺便用torch.cuda.max_memory_allocated()定位一下是哪一步峰值最高,这样排查更快。
显存持续上涨这个特征挺典型的,我怀疑不是中间变量没释放的问题,因为Python的引用计数一般会自动回收,更像是你的自定义算子每次迭代都在graph上累积了额外节点。你既然forward和backward单独测过没问题,那大概率是反向时返回的梯度张量形状或者device跟你预期的不一致,导致PyTorch没法正确剪枝,每次迭代都重新构建了完整的反向计算图。关于scatter_add,我踩过一个坑就是反向时需要用gather把梯度映射回去,但如果你在forward里用了in-place操作,哪怕是很隐蔽的,也可能导致autograd的版本计数失效,graph就永远不释放了。建议你试试在训练循环里加上torch.cuda.synchronize()然后打印当前graph的节点数,或者用torch.autograd.detect_anomaly()跑一下,能定位到是哪个op在累积。另外确认下你的K近邻索引是不是用了torch.long类型的tensor,如果是的话,反向时计算梯度跟索引类型没关系,但要是你把它当成buffer存进module里了,那它就会一直留在显存里。最后问下,你的点云数量是动态变化的吗?如果每次batch的点数不一样,那索引矩阵大小变化也可能导致缓存碎片,显存波动上涨不一定全是graph的问题。
八成是索引在反向时被autograd留着不撒手,试试在backward里手动置空或者用non_blocking释放。