最近在搞一个点云相关的项目,自己写了个CUDA扩展算子(就是简单的K近邻查找+特征聚合),forward和backward都单独测过没问题。但一跑完整训练,batch size从8降到4还是OOM,看了下nvidia-smi,显存占用是持续上涨的,不是一开始就爆。怀疑是我在backward里保存了中间变量(比如索引矩阵)没释放?还是说autograd的graph没有正确剪枝?另外想问问大家,自定义算子如果涉及到scatter_add这类操作,反向的时候有没有什么容易踩的坑?有点迷茫,希望有经验的老哥指点一下,谢谢!
用PyTorch写了个自定义算子,训练时显存直接爆掉,求排查思路
全部回复
共 86 条显存持续上涨基本就是graph没释放,试试在backward里把索引矩阵detach掉或者用完置空。
scatter_add反向容易在grad累积时爆显存,用atomicAdd前先检查下梯度shape对不对。
显存持续上涨这个特征其实挺典型的,大概率不是单纯的中间变量没释放,而是你的自定义backward里某个操作导致了autograd graph的节点数随着训练步数线性增长。我遇到过类似情况,最后发现是forward里用了in-place操作或者返回了不可导的临时tensor,导致每次反向都重新构建了一个子图挂在主graph上。你检查下自定义算子的backward输入里有没有包含forward时保存的索引矩阵,如果那个矩阵是动态shape且每次迭代都不同,graph可能没法剪枝,就会一直累积。另外scatter_add反向那个坑我也踩过,就是梯度要scatter回原位置时,如果index有重复,需要手动用atomicAdd或者先做unique再累加,不然梯度会丢或者重复累加导致数值爆炸,但你说单独测过没问题,那就先排除这个。建议你可以在训练循环里每步打印一下torch.cuda.max_memory_allocated(),看是不是稳步增长,同时用torch.autograd.detect_anomaly()跑几步,它会把产生NaN或inf的节点标出来,有时候OOM前其实已经出了数值问题,只是没爆出来。还有个笨办法,把自定义算子换成等价的纯PyTorch实现(哪怕慢点),如果显存曲线变平了,那就肯定是扩展的C++代码里有什么没接好,重点查一下你有没有在C++里手动new了内存但没在析构时释放,尤其是那种每步都调用的临时buffer。如果方便的话,也可以试试把batch size设成1跑几十步,观察显存是不是涨到某个阈值就稳定了,如果一直涨,那和batch size关系不大,就是graph泄漏。
显存持续上涨基本就是graph没释放,检查下backward里有没有把临时tensor赋给了self或者闭包变量,另外自定义算子里如果用了索引矩阵,记得在反向里用完之后手动del加torch.cuda.empty_cache。scatter_add反向其实还好,但要注意梯度累积时atomicAdd的non-determinism,建议先关掉cudnn benchmark试试。我之前也遇到过类似问题,最后发现是保存了K近邻的pairwise distance,那个size是N*K的,小数据没事大数据直接炸,你排查下是不是这个。
显存持续上涨这个现象很典型,基本可以排除算子本身的bug,问题多半出在graph保存的中间张量上。你试试在backward里把那些索引矩阵用detach或者干脆不存,改成在forward里用非持久化buffer,或者直接在反向时重新算一遍。另外scatter_add的反向确实容易踩坑,主要在于梯度要scatter回原位置时,如果索引有重复,梯度累积顺序会影响结果,最好用atomicAdd或者先unique一下。我之前也遇到过类似情况,最后发现是自定义Function里没写materialize_grads,导致梯度graph一直挂着不释放,你可以检查下这个。
大概率是索引矩阵没做detach,反向时梯度流回索引导致graph一直不释放,试试在保存前加个.detach()。
scatter_add反向确实容易踩坑,建议检查下atomic操作的梯度累加是不是用了inplace,不如换成torch.index_add试试。
显存持续上涨这个特征,大概率不是graph没剪枝,而是backward里保存的索引矩阵没在反向计算后手动释放。你试试在自定义Function的backward末尾把中间变量设成None,或者用ctx.set_materialize_grads(True)。scatter_add的反向就是scatter_add本身,但注意要处理重复索引的梯度累加,别用scatter_覆盖了。另外建议你用torch.cuda.memory._record_memory_history()记录一下分配栈,能直接看到是哪个op在涨。
显存持续上涨八成是backward里那个索引矩阵没释放,试试在算子内部用完之后手动置空。
scatter_add反向记得用原子操作或者先算梯度再scatter,不然会重复累加踩坑。
学到了,感谢分享!
显存持续上涨这个现象,我第一反应就是backward里保存了那个索引矩阵没释放,这玩意儿在点云场景下size是N×K,batch一大直接吃满,建议检查一下是不是在自定义Function里用self.save_for_backward存了非必要的大tensor,能改成在反向里重算索引就重算,或者用del手动清掉。另外scatter_add反向的话,注意梯度要scatter回原坐标,很容易出现重复索引累加的问题,最好用atomicAdd或者先unique再分段处理,不然梯度会悄悄出错但数值上不明显。你试试把batch固定到1跑几百步,看显存是不是还在涨,如果还在涨基本就是graph没释放,可以试试torch.cuda.empty_cache()加在step末尾,但治标不治本,根因大概率还是某个缓冲变量在循环里被反复引用。
反向的索引矩阵大概率就是元凶,试试在backward里用完立马置None,别让它留在graph里。
scatter_add反向容易梯度重复累加,记得用atomicAdd或者先unique再处理,不然显存也会莫名涨。
显存持续上涨这个现象挺典型的,基本可以排除“一开始就爆”那种显存碎片问题,更像是graph里某个节点没被释放或者中间tensor被意外保留了。你提到backward里保存索引矩阵,我建议先查一下是不是用了self.xxx去存这些变量,如果是在自定义Function的backward里用ctx.save_for_backward存还好,但要是直接塞到self上,那整个训练周期都不会释放。另外scatter_add这操作确实容易有坑,反向的时候梯度要scatter回原位置,如果你在forward里做了mask或者用了非连续索引,反向时梯度可能没法正确累加,最好用torch.autograd.Function里的mark_dirty和mark_non_differentiable明确标注一下,避免autograd去追踪不该追踪的东西。还有一个排查思路是开torch.cuda.memory_stats(),分步打印每个op前后的allocated和reserved,看是哪个阶段在涨,我之前遇到过类似问题,最后发现是某个中间结果在循环里被反复引用没detach。你单独测forward和backward没问题,但完整训练里optimizer.step()之后graph会重建,如果自定义算子的backward返回了多余的梯度(比如返回了None之外的张量给不需要梯度的输入),也可能会让graph无法剪枝。要不要试试在backward结尾手动del掉不用的变量,再torch.cuda.empty_cache()看下峰值有没有降下来?如果还不行,可以把backward里返回的梯度都打印出来对比下形状,我怀疑是某个维度scatter的时候产生了隐式广播,导致中间tensor膨胀了。
看到“显存持续上涨”这个特征,我第一反应不是中间变量没释放,而是你的backward里可能有个循环或者累积操作在graph上反复挂钩。PyTorch的autograd在自定义算子里经常因为索引操作(比如gather/scatter)自动生成非常长的依赖链,尤其是你在forward里用了非整数倍的采样逻辑,反向时这些临时tensor会被每个step引用,graph就永远剪不干净。我建议先别急着看显存,把backward里的中间变量全部改成不保存索引矩阵,改成在forward里预计算好一个扁平化的邻居列表,反向时用atomicAdd去累加梯度,这样graph会短很多。另外scatter_add的反向确实容易踩坑,主要是梯度要按索引回填,而且不同索引位置重复写入时atomicAdd的顺序不确定,如果你在反向里对同一个位置做了多次累加,建议手动对梯度做一次去重或校验,不然数值上可能不报错但训练loss会飘。还有个笨办法,你在每个iteration手动调torch.cuda.empty_cache()看显存曲线,如果下降但降不到初值,那就是graph的问题;如果完全不降,可能是某个tensor被全局引用住了。我之前遇到过类似情况,最后发现是自定义算子的反向里用了inplace操作把某个buffer的grad_fn给搞断了,导致它一直留在计算图里,你可以用torch.autograd.detect_anomaly()配合输出每层的saved_tensors大小来定位,虽然慢但很直接。
显存持续上涨这个特征很关键,大概率不是graph没剪枝,而是backward里某个tensor被意外retain了。你可以试试在backward里把所有非必要的中间变量都显式del掉,或者用torch.no_grad包住索引生成的部分。scatter_add反向确实容易踩坑,特别是如果forward里用了inplace的index修改,反向时梯度累加的位置可能会对不上,建议检查一下atomicAdd的粒度。另外可以开一下torch.autograd.detect_anomaly,虽然慢但能定位到具体在哪一步开始爆的。
显存持续上涨基本就是graph没释放或者中间变量被引用了,试试在backward里用del手动清一下局部变量。
scatter_add反向就是scatter,注意梯度累加时atomicAdd的重复计算问题。
显存持续上涨这个特征更像是graph没释放而不是单纯变量没清,建议你在backward里把索引矩阵detach掉或者干脆在forward里重算,别存。另外scatter_add反向确实容易踩坑,梯度要手动用index_add回传,不然autograd会默认生成dense mask导致显存爆炸,你可以用torch.profiler看下峰值是哪个节点占的。
显存持续上涨基本就是graph没释放,试试在backward里把索引矩阵detach掉或者用完置空。
scatter_add反向记得用atomicAdd,别自己写循环累加,容易又慢又吃显存。
我之前也遇到过类似情况,持续上涨基本可以排除显存碎片化,大概率是graph里挂了一堆没释放的中间节点。你检查下backward里是不是把那个索引矩阵存成了self的成员变量,这玩意儿会跟着整个训练周期走,得用临时变量或者干脆在forward里重新算一遍。另外scatter_add反向确实容易出问题,建议把梯度检查打开,对比一下手写反向和torch.autograd.gradcheck的结果,定位是不是梯度累积导致的内存泄漏。
反向时把非必要的中间变量detach掉或者干脆重算,scatter_add记得用atomicAdd别直接索引赋值。
显存持续上涨这个特征挺典型的,大概率不是graph没剪枝,而是backward里保存的索引或中间tensor没走detach,导致每次iter都累积在计算图里。我之前写scatter_add反向时踩过坑,grad_output按index回传后,那个index本身得用int64且不能是requires_grad的,否则会额外保存一份用于反向的拷贝。你可以先试试在保存中间变量时统一加.detach(),然后跑几个step看显存曲线是否变平。另外检查下是不是K近邻的索引矩阵维度是(batch, n, k),这个在反向时如果没显式释放,会随着batch累积很夸张。
显存持续上涨多半是graph没释放,试试在backward里把索引矩阵转成bool或者干脆重算一遍。
scatter_add反向确实容易爆,建议用atomicAdd替代或者分块处理,我之前就是这么解决的。