最近在调一个图像分割模型,用的是PyTorch,之前跑得好好的,但加了几个数据增强操作后,训练到第5个epoch显存直接爆了(12G不够用)。我怀疑是某个transform或者自定义的Dataset里内存没释放,但代码写得很乱,一时找不到具体哪一行有问题。试过torch.cuda.memory_summary(),但输出信息太杂,看不懂。有没有什么工具或方法能定位到是哪一行代码导致显存泄漏或突然暴涨?比如能不能像Python的cProfile那样,能按行看显存占用?先谢谢各位老哥了。
用PyTorch训练模型时,显存突然暴涨,怎么排查具体是哪一行代码导致的?
全部回复
共 179 条之前也遇到过类似的情况,后面发现是数据增强里某个操作返回了不该有的tensor,比如把numpy转cuda后忘了detach,导致计算图越积越大。建议先试试用torch.autograd.detect_anomaly()配合pdb,或者干脆在每个transform前后打一下torch.cuda.memory_allocated(),二分法定位很快。另外也可以看看是不是Dataset的__getitem__里用了全局变量或者缓存没清,这种隐性引用挺坑的。
我之前也遇到过类似情况,后来发现是数据增强里某个操作在GPU上动态生成了超大中间张量,根本没进Dataset的锅。你可以试试用torch.cuda.set_per_process_memory_fraction配合tracemalloc去看CPU侧内存,但更直接的办法是给每个transform单独跑一次前向,观察显存曲线变化。另外PyTorch有个torch.autograd.detect_anomaly,虽然主要查梯度问题,但有时能顺带暴露异常分配。如果代码太乱,不如先把增强操作逐个注释掉做二分排查,比看summary快多了。
试试pytorch的torch.cuda.memory._record_memory_history,能按行回溯分配点,比summary直观多了。
试试给每个transform单独跑一遍前向,用torch.cuda.reset_peak_memory_stats()包着,看哪个峰值涨得离谱,基本能锁定。另外查查是不是开了num_workers=0但Dataset里存了太多中间变量,有些增强库比如albumentations会缓存结果。我之前遇到类似情况是随机crop的边界计算里有bug,导致tensor被反复复制。如果还不行,建议用pytorch的autograd检测一下是否有叶子节点没detach,这个坑经常很隐蔽,但显存暴涨多半跟它有关。
试试用pytorch的autograd检测钩子,在backward之前打印每层tensor的size,配合torch.cuda.reset_peak_memory_stats()和torch.cuda.max_memory_allocated()分段跑,比如把每个transform单独拎出来过一遍,基本能锁定是哪个操作在涨。我之前遇到过类似情况,最后发现是自定义Dataset里把整张图的所有增强结果都存了list没清,你可以重点查下有没有往成员变量里塞临时tensor。另外torch.cuda.memory._dump_snapshot那个可视化工具比summary直观,能看内存分配栈,就是配置稍微麻烦点。
这问题我太熟了,之前调分割模型也遇到过一模一样的,加了几个增强后第3个epoch直接OOM。你怀疑transform或者Dataset里内存没释放,方向大概率是对的,因为很多增强操作比如随机裁剪、翻转会在每次迭代时创建临时张量,如果没及时del或者被计算图引用住,显存就会像滚雪球一样。不过想定位到具体行,torch.cuda.memory_summary()确实太笼统,我建议你试试torch.profiler的with torch.profiler.profile(profile_memory=True),它能把每个操作的内存分配和释放都列出来,按self_cuda_memory_usage排序,基本一眼就能看出哪个算子把显存吃掉了。另外还有个土办法,就是把你新加的数据增强一个个注释掉跑一个epoch,二分法定位,虽然笨但特别直观,我上次就是这么找到的——结果是我在自定义Dataset里把增强后的图像append到了list里忘了清空,相当于每张图都存了个副本,显存不爆才怪。还有个小技巧,可以在每个epoch结束打印torch.cuda.max_memory_allocated(),如果这个值随着训练单调递增,那就说明有累积泄漏,如果是突然跳高,那大概率就是单次操作的问题。你要是想按行看,也可以试试torch.autograd.detect_anomaly()配合with torch.cuda.amp.autocast(),但那个主要是查梯度异常的,对纯显存占用帮助有限。总之别慌,先把你加的增强代码块逐段隔离,再结合profiler看,基本半小时就能揪出来。
试试给dataloader加个pin_memory=False,大概率是数据增强里张量没释放,我之前也栽在这过。
你试试用torch.autograd.detect_anomaly配合二分注释代码,或者开个nsys看看内存曲线,比看memory_summary直观多了。
试试用pytorch的autograd检测钩子,把tensor的grad_fn打印出来,配合torch.cuda.memory._record_memory_history()能抓到分配栈,比memory_summary直观多了。我之前也遇到过类似情况,最后发现是数据增强里某个op开了retain_graph=True没关,导致计算图越积越大。另外检查下是不是在循环里反复创建了优化器或loss函数,这些对象会持有中间变量。实在不行就二分法,把transform逐个注释掉跑几个step看显存变化,虽然笨但有效。
试试用pytorch_memlab的MemReporter,能按行看增量,比memory_summary直观多了。
显存爆一般不是泄漏,多半是某个transform在batch上返回了不同shape,检查下自定义Dataset的__getitem__里是不是把整个tensor都return了。
说到这个我太有同感了,之前跑检测头的时候也遇到过类似情况,加了几个随机裁剪直接爆显存,后来用pytorch的autograd检测hook才定位到是某个tensor在反向传播时被保留了引用。你那个情况,如果怀疑是数据增强的问题,可以先试试把transform逐个注释掉跑一个step,用二分法缩小范围,比直接看memory_summary直观多了。另外torch.cuda.set_per_process_memory_fraction设个上限,让程序在爆显存前直接报错,配合faulthandler或者pdb就能抓到具体调用栈。还有个土办法,在Dataset的__getitem__里加torch.cuda.synchronize()和打印当前显存占用,看哪个样本的缓存突然涨得离谱,基本就能锁定是某个操作搞的鬼。至于按行看显存,目前没有cProfile那么成熟的工具,但可以试试torch.profiler的memory profiler,它按操作符维度给内存分配统计,虽然不能精确到py行,但能看出是卷积还是拼接这类操作在吃显存。要是急用,直接给transform加个no_grad或者把增强放到CPU上做,再用pin_memory预取,大概率能缓解峰值。最后提醒下,别迷信显存泄漏,有时候就是输入尺寸没对齐,某个分支悄悄生成了超大中间特征图。
试试用pytorch的autograd检测钩子,或者用gputil配合分段打印显存,基本能锁定到具体操作。
可以查下是不是数据增强里用了detach或者梯度没清零,我之前就是crop操作里忘了释放旧变量。
试试用pytorch的内存分析器,能精确到算子级别的显存占用,比summary直观多了。
之前遇到过类似问题,最后发现是增强操作里把tensor留在计算图上了,记得加detach。
试试用pytorch的autograd记录hook,在backward前打印每层显存,或者直接二分注释代码,最快能定位。
显存暴涨多半是数据增强里没释放中间变量,你把每个transform单独跑一遍看哪个峰值高就明白了。
显存这种东西最坑的就是它不像CPU内存那样能直接给你报个错告诉你哪一行炸了,尤其是加了transform之后,很可能是某个操作在计算图里被保留了下来。我之前遇到过类似情况,最后发现是自定义Dataset里把增强后的图像tensor存成了list成员变量,每个epoch叠加导致显存只增不减。你可以试试用torch.profiler,它比memory_summary直观得多,能在profiling的时候按操作符和调用栈看allocated memory,配合record_shapes=True基本能定位到是哪个层的输出在累积。另外有个小技巧,如果你用的是albumentations或者torchvision的transforms,检查一下是不是有概率性的操作(比如cutout)在某个batch里生成了异常大的中间变量,这种偶发暴涨经常是输入尺寸没对齐导致的。实在不行就开个子进程,每隔几步打印torch.cuda.memory_allocated(),用二分法注释掉部分数据增强逻辑,虽然土但最有效。还有一个容易忽略的点,如果开了num_workers>0,DataLoader的prefetch机制会把下一批数据提前塞到显存里,配合pin_memory有时候会瞬间多出几百MB,但这通常不会持续累积,更像你说的“暴涨”而不是“泄漏”。建议把saved_tensors_hooks也用上,能追踪到反向传播时哪些中间激活被意外保留,PyTorch 1.10+自带这个功能。最后提醒一下,记得关掉cudnn的benchmark模式试试,某些卷积配置在输入尺寸变化时会重新选算法,临时显存峰值可能翻好几倍。
巧了,我之前也踩过这坑,加了几个增强之后显存直接翻倍。你试试torch.cuda.set_per_process_memory_fraction设个上限,让它提前崩,然后用faulthandler或者pdb配合catch Exception去抓,但更实用的是把transform拆开逐个跑一遍看每个操作前后的显存增量,这比直接看memory_summary直观多了。另外查一下你的Dataset是不是在__getitem__里存了中间变量没清,特别是那种把整张图或者mask的list挂在self上不释放的写法,很容易爆。如果用了随机crop或resize,注意一下是不是每个iter都new了新的tensor但没detach,累积到图里了。还有个骚操作,就是给每个transform包一层hook,打印当前allocated和reserved的差值,定位基本秒级。最后实在不行就开--reload_dataloader_every_n_epochs,把缓存周期打散,不过这只是治标。真要看每行显存,可以试试pytorch_memlab,它的line_profiler能按行报增量,虽然偶尔不准,但比肉眼找强多了。
试试用pytorch的profiler带record_shapes,能追踪到具体op的显存峰值,比原生的清楚多了。
建议试试PyTorch的torch.autograd.detect_anomaly()或torch.cuda.set_per_process_memory_fraction配合二分法注释代码,能快速缩小范围。我之前遇到过类似问题,最后发现是自定义Dataset里存储了所有增强后的图像没释放,改成在__getitem__里动态生成就好了。也可以用pytorch_memlab这个库,它的LineProfiler能按行输出显存分配。另外数据增强操作如果用了torchvision.transforms,注意有些函数会缓存中间结果,把num_workers调成0跑几个step试试,能排除是不是数据加载的并行问题。
我之前也碰到过类似的,最后发现是数据增强里有个随机裁剪搞出了超大尺寸的中间张量,训练到后面才触发。你可以试试 torch.cuda.memory_snapshot() 配合 pycharm 的调试器,在可疑的 transform 前后各打一次快照对比。另外用 torch.utils.checkpoint 能省不少显存,但得确认下是不是累积了计算图没释放。
显存爆在第五个epoch挺典型的,大概率是某个transform每次把新tensor挂到了计算图上,或者Dataset里缓存没清。可以试试torch.cuda.memory_allocated()在训练循环里打点,配合line_profiler按行看谁在涨。另外用torch.utils.checkpoint或者把transform放到DataLoader的worker里,别在主进程做。实在不行就二分法注释代码,先砍掉一半增强跑一轮看看。