最近在调一个图像分割模型,用的是PyTorch,之前跑得好好的,但加了几个数据增强操作后,训练到第5个epoch显存直接爆了(12G不够用)。我怀疑是某个transform或者自定义的Dataset里内存没释放,但代码写得很乱,一时找不到具体哪一行有问题。试过torch.cuda.memory_summary(),但输出信息太杂,看不懂。有没有什么工具或方法能定位到是哪一行代码导致显存泄漏或突然暴涨?比如能不能像Python的cProfile那样,能按行看显存占用?先谢谢各位老哥了。
用PyTorch训练模型时,显存突然暴涨,怎么排查具体是哪一行代码导致的?
全部回复
共 179 条试试给dataloader加个pin_memory=False看下,我之前就是数据增强里切片操作没释放引用导致显存涨。
用pytorch的autograd检测一下,把梯度计算图打印出来,能直接看到哪个tensor在反向传播时爆的。
试试用pytorch的autograd检测钩子,或者干脆把transform逐个注释掉跑一遍,二分法定位很快。
用torch.cuda.set_per_process_memory_fraction限制显存,报错时看traceback能直接锁到爆显存那行。
PyTorch有个torch.autograd.detect_anomaly()可以试试,把训练前包一下,它会直接告诉你反向传播里哪步产生了NaN或者异常梯度,很多时候显存暴涨跟这个有关。另外你加了数据增强,大概率是transform里生成了超大tensor忘了释放,比如某些随机缩放或裁剪一次性返回了多份数据,建议在Dataset的__getitem__里手动del一下中间变量,或者用gc.collect()看看。还有个土办法,把数据增强一个个注释掉,跑一个epoch看显存,二分法定位,比看memory_summary快多了。
我上次遇到类似情况,最后发现是某个自定义的collate_fn里把整张图pad到了固定尺寸,但没控制batch内最大分辨率,导致显存随着输入尺寸变化暴增。你查一下是不是每个sample的shape不一致,PyTorch在stack时会把所有样本隐式对齐到最大尺寸,特别吃显存。
之前我也被这个问题折磨过,后来发现加了几个强transform后,问题其实出在数据加载时缓存了太多中间结果,尤其是随机裁剪和翻转这类操作,建议你先把数据增强里的tensor操作改成在gpu上做,或者看看是不是把原图也存进list里了。想按行排查的话可以试试torch.autograd.detect_anomaly(),虽然不能直接看显存,但能帮你定位到反向传播时爆掉的张量操作。还有个土办法,就是分阶段跑你的dataset和transform,每跑完一段就打印一下cuda.memory_allocated(),差值大的地方基本就是元凶了。你那个自定义Dataset里如果存了全量增强后的图,那12G爆掉很正常,建议改成迭代器或者用生成器喂数据。
试试用torch.profiler的with_stack参数,能按行打印显存分配堆栈,比memory_summary直观不少。另外记得把数据增强里的随机操作单独拎出来测,很多transforms会缓存中间张量,特别是涉及随机裁剪或翻转时。我之前遇到过类似问题,最后发现是自定义Dataset里忘了把torch.no_grad()包住,导致每次取数据都建了计算图。你可以先在每个transform前后插个print(torch.cuda.memory_allocated()),二分定位最快。
试试给每个transform前后打点torch.cuda.max_memory_allocated(),对比一下就知道哪步涨的。
用PyTorch的autograd检测钩子,或者临时把batchsize调小,跑几步看哪层梯度累积爆的。
试试CUDA的torch.cuda.set_per_process_memory_fraction配合二分注释法,把可疑transform逐个屏蔽跑一个batch就能定位。
试试用pytorch的torch.autograd.detect_anomaly,能定位到爆显存的反向传播节点。
我一般先注释掉新增的transform跑一版,二分法定位最快。
我之前也遇到过类似情况,后来发现是数据增强里某个操作在GPU上动态创建了临时张量,没及时释放。你可以试试torch.autograd.detect_anomaly(),配合torch.cuda.set_per_process_memory_fraction限制显存,爆的时候能更快缩小范围。另外,把transform里的每个操作单独跑一遍,用torch.cuda.max_memory_allocated()对比前后差值,基本能锁定是哪一步的问题。别全指望memory_summary,那个对定位代码行确实没啥用。
显存暴涨不一定是泄漏,可能是某个增强操作在batch维度上隐式复制了数据,比如RandomCrop的padding逻辑。建议你先把num_workers设为0,排除DataLoader预取干扰,然后逐行注释transform做二分法测试。还有个笨办法:在训练循环里每步打印torch.cuda.memory_allocated(),看峰值出现在哪个迭代,再配合record_stream检查是否有未释放的中间变量。如果还找不到,就把自定义Dataset里所有return改成yield试试,有时候是隐式列表构建导致的内存堆积。
我之前也遇过类似情况,最后发现是数据增强里某个操作在GPU上动态创建了临时张量,没及时del。你可以试试torch.cuda.set_per_process_memory_fraction限个上限,让程序崩得早一点,再用faulthandler或者pdb配合cuda的同步点来定位。另外推荐一下pytorch的torch.profiler,虽然不能精确到行,但能看每个op的显存峰值,通常能缩小范围到某个模块。不过说真的,最笨但有效的办法是把transform一个个注释掉跑一遍,二分法找起来很快。
试试给每个transform前后加个torch.cuda.synchronize()然后看nvidia-smi的波动,能快速锁定是不是某个增强操作临时开了大张量。另外查一下Dataset里有没有把整个图像列表留在GPU上,之前我遇到过random crop每次返回不同尺寸导致缓存碎片化,用torch.cuda.empty_cache()放epoch末尾治标不治本。
想按行看显存的话,可以装个pytorch_memlab,虽然不能像cProfile那么细,但能标出每个张量的创建位置,配合line_profiler用挺香的。你那个数据增强如果是用albumentations,记得检查是不是开了cuda加速,有时候会隐式缓存中间结果。
试试给transform加个.to()或clone(),八成是数据增强里某步生成了超大中间变量没及时释放。
用torch.cuda.set_per_process_memory_fraction设个上限,跑挂前打断点一步步查,比看summary直观多了。
之前也踩过类似的坑,加了几个transform之后显存直接翻倍,后来发现是某个自定义Dataset里把整张图的所有增强版本都存进list了,本来只想存最终结果,结果每个epoch都累积,内存直接炸。你要是代码乱,可以先试试把数据加载和模型训练分开跑,单独跑一个epoch的dataloader,看显存是不是还在涨,这样能快速排除是数据侧还是模型侧的问题。另外torch.cuda.memory_summary确实太啰嗦,我一般用torch.profiler,它能按操作符统计显存分配,虽然不能精确到行,但能看出是哪个模块比如卷积还是插值在吃显存。还有个土办法,就是在每个可能出问题的代码块前后手动打印torch.cuda.memory_allocated(),二分法定位,虽然笨但很有效。另外注意一下数据增强里的随机裁剪或者翻转,如果用了albumentations,有些操作会额外产生中间变量,最好检查一下是不是忘了释放或者一直引用着。如果怀疑是梯度累积导致的,可以试试把batch size调小,看显存变化曲线是不是线性,如果不是线性那大概率是泄漏而不是单纯峰值。最后实在不行就升级到最新版PyTorch,有些版本的cuda caching allocator有bug,我之前升级后同一个脚本显存直接降了2G。
我之前也遇到过类似情况,后来发现是数据增强里某个操作在GPU上执行了,比如ToTensor之后又做了自定义的tensor运算,导致计算图没释放。建议你把transform里的操作全部放到CPU上,然后看看显存曲线是不是就稳了。另外可以试试pytorch的torch.autograd.detect_anomaly,虽然不能直接定位行号,但能帮你找出梯度异常的点,配合二分注释法很快就能锁定。cProfile确实管不了显存,你可以写个脚本每步打印torch.cuda.memory_allocated,对比前后差值,哪个操作涨幅异常基本就是它了。
我之前也踩过类似的坑,加了几个transform后显存直接翻倍,后来发现是随机裁剪那块没写对,导致计算图没释放。你可以试试在训练循环里加torch.cuda.synchronize()配合torch.autograd.detect_anomaly(),能快速定位到反向传播时的爆点。另外torch.profiler比memory_summary好用,能按行看每个op的显存分配,虽然刚开始上手有点繁琐但值得折腾。数据增强那边建议检查下有没有在GPU上跑tensor操作,有时候不经意间把numpy转cuda没及时回收就会这样。
试试用pytorch的autograd记录每步显存分配,配合cuda的memory snapshot工具能定位到具体张量来源。
我之前也遇到过类似情况,后来发现是数据增强里的随机裁剪导致计算图没释放,换成原地操作就好了。
遇到这种情况我一般先怀疑数据增强里有没有用GPU张量操作,比如在transform里直接调了.cuda()或者某些库内部自动搬上显存,加个pin_memory=True和num_workers=0对比跑一下就能排除大半。想按行定位的话,可以试试PyTorch的torch.autograd.detect_anomaly(),但那个主要抓梯度异常,显存暴涨不一定好用。更推荐用nvtop或者nvidia-smi每隔0.1秒刷一次,同时配合代码里临时print当前epoch和step,看暴涨是发生在数据加载阶段还是前向传播阶段,这样能缩小范围。另外有个取巧的办法,把自定义Dataset和transform里的所有操作都注释掉,用最简单的随机tensor跑一个epoch看显存曲线,如果正常就二分法逐步加回代码段,虽然土但特别有效。你还可以试试torch.cuda.set_per_process_memory_fraction设个0.9,让它在爆之前报错,报错堆栈会指向实际分配那一行,比memory_summary直观多了。最后注意下是不是有累积梯度或者loss.backward()没在循环里,有些增强操作如果返回了需要梯度的变量,并且被缓存到list里,每个step都会留一份显存。实在不行就升级到PyTorch 2.x,它的内存分配器优化过,有些“泄漏”其实是碎片化,torch.cuda.empty_cache()在每次epoch结束调一下也能救急。
试试用pytorch的autograd检测钩子,或者干脆把transform逐个关掉二分定位,很快能揪出来。
试试torch.cuda.memory._record_memory_history()加上snakeviz可视化,能按堆栈跟踪到具体张量分配点,比memory_summary直观多了。另外你加了数据增强后暴涨,很可能是transform里生成了临时tensor没显式del,尤其像随机裁剪这种会反复拷贝的,建议把增强操作改成在CPU上做完再转GPU。我之前遇到过类似情况,最后发现是某个ToTensor的变量被循环引用卡住了显存,用gc.collect()和tracemalloc结合定位才找到。
试试用pytorch的autograd检测钩子或者分段跑,把transform和dataset拆开各跑一遍看哪块涨,最省事。
之前我遇到过,是数据增强里开了太多线程没关,显存全被缓存吃了。