最近在调一个图像分割模型,用的是PyTorch,之前跑得好好的,但加了几个数据增强操作后,训练到第5个epoch显存直接爆了(12G不够用)。我怀疑是某个transform或者自定义的Dataset里内存没释放,但代码写得很乱,一时找不到具体哪一行有问题。试过torch.cuda.memory_summary(),但输出信息太杂,看不懂。有没有什么工具或方法能定位到是哪一行代码导致显存泄漏或突然暴涨?比如能不能像Python的cProfile那样,能按行看显存占用?先谢谢各位老哥了。
用PyTorch训练模型时,显存突然暴涨,怎么排查具体是哪一行代码导致的?
全部回复
共 179 条试试用pytorch的torch.cuda.memory._record_memory_history(),能按栈跟踪到具体行号,比summary直观多了。
加数据增强后爆显存八成是缓存了太多中间变量,检查下transform里有没有往list里append张量没清。
显存突然爆掉大概率不是内存泄漏,是某个transform在batch里生成了超大中间变量,特别是随机crop或者resize后没同步更新mask的shape。可以试试用torch.profiler的profile_memory=True,它会按操作符给你排序,比memory_summary直观很多,基本能看出是哪个模块吃显存。另外检查下是不是开了persistent_workers=True但num_workers设太大,数据加载线程也会占显存。我之前遇到过类似问题,最后发现是随机旋转时对3D标签图用了grid_sample,梯度回传直接爆显存。
试试给每个transform加个hook记录tensor的shape和id,基本能定位到是哪个操作在涨。或者直接查CUDA缓存分配,大概率是数据增强里没用no_grad导致梯度图累积了。
这问题我踩过坑,建议先别急着上工具,把新增的transform逐个注释掉跑一遍,二分法定位最快。我之前是RandomCrop里没释放旧tensor,加个del和torch.cuda.empty_cache()就好了。另外torch.cuda.memory._dump_snapshot()能生成可视化火焰图,比memory_summary直观,配合chrome://tracing看分配堆栈,基本能锁定行号。
我之前也遇到过类似情况,最后用torch.cuda.set_per_process_memory_fraction配合pdb一步步跑,基本能锁定是哪几个transform在搞鬼。你那个数据增强里如果用了随机裁剪或者翻转,试试先把它们注释掉跑一个epoch,大概率就是那块的问题。另外检查下Dataset里有没有把tensor存到list里没清,这种隐性引用特别容易爆显存。memory_summary确实难懂,可以试试用nvidia-smi的实时监控配合代码里手动打印显存差值,比看summary直观多了。
试试给可疑的数据增强单独跑个循环,配合pytorch的memory_profiler按行看,基本能揪出来。
我上次就是这么干的,把transform一个个过,最后发现是随机裁剪里没释放中间变量。
试下pytorch的autograd.detect_anomaly,能直接定位到产生nan或者梯度异常的op,不过你这个情况更像数据增强里开了太多线程或者缓存没清,检查下每个transform是不是返回了多余的中间结果,尤其那些用了cv2或者numpy的,可能隐式拷贝了数组。
另外torch.cuda.memory_summary确实难用,建议用torch.profiler的with_stack=True跑几个batch,能看到每个python函数分配了多少显存,比cProfile直观多了。我上次就是靠这个发现是某个随机裁剪里生成了全尺寸的mask,白白占了几个G。
我之前也遇过类似情况,后来发现是数据增强里某个操作在GPU上生成了中间变量没及时清掉。你可以试试用pytorch的autograd检测,或者干脆把每个transform单独跑一遍看显存变化,虽然笨但挺管用。
另外torch.cuda.reset_peak_memory_stats()配合torch.cuda.max_memory_allocated()能看峰值,但确实没法精确到行。要是代码乱,建议先查Dataset的__getitem__里有没有把tensor存成全局变量,我上次就是这翻车。
还有个野路子,把batch size调成1,如果显存还涨,基本就是数据加载或增强逻辑的问题,跟模型无关了。要是还不行,直接上nvidia-smi盯实时占用,配合二分法注释代码,效率也不差。
我之前也踩过类似的坑,加了几个transform后显存直接翻倍,后来发现是Albumentations里某些操作会偷偷开多线程缓存,把num_workers调小或者改成同步模式就好了。你可以试试用torch.autograd.detect_anomaly(),虽然不能直接定位行号,但能追踪到反向传播时哪个op触发了异常增长。另外pytorch的torch.profiler带memory profiling功能,虽然输出也乱,但配合按时间排序能大概看出是前向还是反向阶段爆的,缩小范围后再去翻对应代码会快很多。
试试pytorch的profiler,能按操作符看显存,或者用torch.cuda.set_per_process_memory_fraction卡个上限让它直接崩在出事那步。
试过用pytorch的autograd检测钩子吗,能精确定位到每个张量的创建点,比看memory_summary直观多了。
我一般是关掉cudnn的benchmark,再配合torch.profiler看每步操作的内存峰值,很快就能锁定是哪个transform在搞鬼。
我之前也踩过类似的坑,加了几个transform后显存直接翻倍。建议先别急着看行号,用torch.autograd.detect_anomaly()开启异常检测,报错时通常能定位到具体计算图节点,再配合nvidia-smi看实时显存变化,缩小范围到某个操作。另外试试把数据增强部分单独拎出来,用固定seed跑几个batch,对比前后显存峰值,基本就能锁死是哪个transform的锅。如果真是Dataset里累积了历史张量,记得在__getitem__里显式del掉中间变量,或者用torch.cuda.empty_cache()在epoch间清一下缓存。
遇到过类似的情况,数据增强加上去之后显存曲线跟坐火箭似的。你那个memory_summary看不懂很正常,它给的是全局分配信息,根本对不上具体代码行。我后来是直接用torch.cuda.set_per_process_memory_fraction把显存限制到刚好能跑的量,然后让程序崩,看traceback指向哪个transform,这招虽然笨但有效。另外强烈建议检查一下你的增强操作里是不是用了torchvision的transforms,有些版本在cuda上会缓存中间张量,比如RandomCrop或者RandomResizedCrop,它们内部可能调用了grid_sample,梯度图或者采样网格不会自动释放。你可以试试把所有增强操作都放到CPU上做,然后只把tensor转回GPU,很多时候显存暴涨根本不是Dataset的问题,是增强操作在GPU上执行时产生的临时激活值没被回收。要是想按行定位,可以试试torch.profiler,它带memory profiling功能,能按操作符维度看内存分配,不过需要你在训练循环里包一层with torch.profiler.profile,然后导出到chrome://tracing里看,虽然不精确到源码行,但能看出是哪个模块在疯狂吃显存。还有个土办法,就是在每个epoch结束手动调gc.collect()和torch.cuda.empty_cache(),再看看显存曲线是不是变得平缓,如果平缓了就说明有对象没被引用但还被cuda缓存着,大概率是某个list或者dict里存了不必要的tensor没清空。
试试pytorch的torch.profiler,带memory分析,能直接看每个操作吃多少显存,比memory_summary直观多了。
试试用torch.autograd.detect_anomaly()加torch.cuda.set_per_process_memory_fraction,再配nsys按行分析,大概率能抓到是哪个tensor没释放。
这问题我上周刚踩过坑,加了几个albumentations的transform后也是第五六个epoch直接爆,后来发现是RandomResizedCrop里设了return_tight_mask=True,每次迭代都偷偷把原图和mask的tensor存了一份在cache里,根本没法靠memory_summary看出来。你试试用torch.cuda.set_per_process_memory_fraction配合pytorch的autograd.detect_anomaly可能不顶用,真正的排查思路是先把数据增强拆成一半,跑一个epoch看显存峰值,然后二分法定位到具体操作,比直接看summary靠谱多了。另外强烈建议你在Dataset的__getitem__里临时加个torch.cuda.reset_peak_memory_stats(),然后每step打印torch.cuda.max_memory_allocated(),这样能看出是数据加载阶段涨还是前向传播阶段涨,我那次就是数据加载里用torchvision.transforms.ToTensor()时没转成contiguous,导致后续拼接时多复制了一份。还有个土办法,把batch size调成1,如果显存还涨那基本就是transform或者dataset的问题,跟模型无关。要是你用了多进程DataLoader,记得把num_workers设成0试试,有时候子进程的显存分配不会及时回收,看起来就像泄漏一样。最后推荐你直接看nvidia-smi的PID对应python进程,再用gdb attach上去看cudaMalloc的调用栈,虽然麻烦但能精确到哪一行,不过一般走到二分法那步就够用了。
我之前也踩过这坑,加了自定义transform后显存爆炸,最后发现是增强操作里生成了超大tensor没及时del。建议你在每个transform前后加torch.cuda.reset_peak_memory_stats()和torch.cuda.max_memory_allocated()打点,跑一个batch就能快速定位。另外可以试试PyTorch的torch.autograd.set_detect_anomaly(True),虽然慢但能报出具体反向传播节点。如果还不行,就把Dataset里的操作拆开逐个注释掉,二分法排查最快。
我之前也遇到过类似情况,后来发现是DataLoader的num_workers开太多,加上transform里用了奇怪的切片操作导致张量没释放。你可以试试用pytorch的autograd.detect_anomaly(),或者更直接点,在可疑代码段前后打torch.cuda.memory_allocated()对比差值,很快就能锁定范围。另外,如果用了自定义Dataset,检查下__getitem__里是不是每次都在重复创建大数组没清掉。
cProfile那种按行看显存的工具目前还真没有,但torch.profiler可以看每个op的内存分配,配合按epoch分段跑,基本能定位到具体层或操作。再不行就暴力点,把数据增强逐步注释掉二分排查,虽然老土但最有效。
说到这个我太有共鸣了,之前调检测模型也遇到过一模一样的坑,加个随机裁剪显存直接翻倍。你那个怀疑方向我觉得挺对的,数据增强里如果用了那种会拷贝张量的操作,比如RandomCrop或者某些自定义的transform里用了np.copy,很容易让显存悄悄涨上去。不过更常见的其实是Dataset的__getitem__里不小心把整个图像序列都load进内存了,然后每个batch都重复构建计算图,这样显存肯定爆炸。想按行排查的话,可以试试torch.profiler,它带那个memory profiling功能,能按操作符看到每个tensor的分配和释放点,虽然不能精确到Python代码行,但能缩小范围到某个算子。另外有个土办法,就是在每个transform前后打印一下当前进程的显存占用,用torch.cuda.memory_allocated()和torch.cuda.max_memory_allocated()做差值,这样能快速定位是哪个环节暴涨的。还有个细节,如果你用了多线程加载数据,记得看下num_workers是不是设太高了,有时候是CPU内存爆了导致swap,然后显卡显存被连带拖垮,这问题特别隐蔽。最后实在不行,就把可疑的transform一个个注释掉跑个小epoch,二分法排除,虽然笨但最直接。
试试用pytorch的autograd.detect_anomaly或者给每个transform前后打点看显存快照,我上次就是DataLoader里没设pin_memory=false导致暴涨。
把数据增强挪到GPU前用CPU执行试试,我之前是torchvision的RandomResizedCrop在batch里跑直接爆显存,换成albumentations就稳了。