最近在调一个图像分割模型,用的是PyTorch,之前跑得好好的,但加了几个数据增强操作后,训练到第5个epoch显存直接爆了(12G不够用)。我怀疑是某个transform或者自定义的Dataset里内存没释放,但代码写得很乱,一时找不到具体哪一行有问题。试过torch.cuda.memory_summary(),但输出信息太杂,看不懂。有没有什么工具或方法能定位到是哪一行代码导致显存泄漏或突然暴涨?比如能不能像Python的cProfile那样,能按行看显存占用?先谢谢各位老哥了。
用PyTorch训练模型时,显存突然暴涨,怎么排查具体是哪一行代码导致的?
全部回复
共 179 条试试用 torch.autograd.detect_anomaly() 或者给每个transform加个hook,跑起来看哪一步的tensor没释放就行。
当时我排查泄漏就是靠给Dataset里每个操作单独打日志,爆的那步直接锁定,比看memory_summary省心多了。
之前也遇过类似情况,后来发现是数据增强里某个操作返回了CUDA tensor,没转回CPU,导致累积在显存里。可以试试给每个transform加个torch.cuda.memory_allocated()打印,二分定位到具体模块,比memory_summary直观多了。另外torch.profiler的with_stack选项能按行看分配点,虽然跑起来慢点,但排查这种问题很管用。如果还不行,检查下Dataset的__getitem__里有没有把中间结果存成成员变量,那个最容易不知不觉吃显存。
试试用torch.profiler,能按行看显存分配,或者把transform拆开逐个跑一遍,问题大概率在某个数据增强上。
建议把batchsize调小点跑几个step,配合nvidia-smi盯着看,哪个transform加进去显存涨了就是哪个。
试试用pytorch的autograd检测钩子,在loss.backward()之前把每层的activations存下来对比,不过简单点可以直接用torch.cuda.set_per_process_memory_fraction把显存限制到6G,跑起来报错会直接定位到爆的那一步。我之前碰到类似情况是数据增强里用了随机crop,每次返回的tensor没detach,梯度图越积越大,你重点查一下transform里有没有对tensor做in-place操作或者把非叶子节点传进了网络。用memory_summary确实很乱,可以配合tracemalloc看CPU内存,但显存的话先用二分法把transform逐个注释掉跑前几个batch,很快就能缩小范围。
试试PyTorch的torch.autograd.detect_anomaly(),虽然它主要查梯度问题,但有时能顺藤摸瓜找到爆显存的op。另外可以用torch.cuda.set_per_process_memory_fraction把显存限制死,让OOM提前触发,配合CUDA_LAUNCH_BLOCKING=1看回溯,比看memory_summary直观多了。我之前遇到过类似情况,最后发现是某个自定义transform里把整个batch的tensor切片后没detach,导致计算图越滚越大。
试试用pytorch的autograd记录hook,在每个tensor的backward里打个标记,配合torch.cuda.memory._record_memory_history()能看到分配栈,比memory_summary直观很多。我上次遇到类似问题就是这么定位到是某个transform里repeat操作把中间结果攒着没释放。另外注意下数据增强的worker进程,num_workers设太大也可能显存被多个进程各自拷贝一份。
我之前也踩过这个坑,后来发现是自定义Dataset里没把numpy数组转成torch tensor,导致每次取batch都隐式复制一份到GPU。你可以先检查下增强操作里有没有生成大数组后忘记del,或者用gc.collect()强制回收看看显存变化。如果还不行,就把训练循环拆成几步,每步打印torch.cuda.memory_allocated(),二分法定位到具体模块。
楼上说的record_memory_history确实好用,不过老版本PyTorch不支持。另一个办法是直接用tracemalloc跟踪CPU内存,因为很多显存暴涨其实源头是CPU端数据累积,比如transform里把整张图存进list没释放。我之前排查时是先把num_workers调到0,排除多进程干扰,再逐行注释transform,很快就能锁定。记得把pin_memory也关掉试试。
说实话,我之前也踩过类似的坑,加了几个transform之后显存直接翻倍,最后发现是某个随机裁剪操作在每次迭代时都保留了整张图的梯度图引用,导致显存只增不减。memory_summary()那个输出确实太劝退了,我后来是改用torch.cuda.set_per_process_memory_fraction把显存限制在比如8G,这样爆掉的时候报错栈会直接指向分配那一行的Python调用点,比看summary直观得多。另外可以试试pytorch_memlab这个库,它的LineProfiler能按行输出每个tensor的分配大小和堆栈,配合tracemalloc一起用,基本能定位到是Dataset里哪一步在偷偷累积显存。还有个土办法,在训练循环里每隔几步打印一下torch.cuda.memory_allocated()和torch.cuda.memory_reserved()的差值,如果差值持续变大,基本就是有tensor没释放,这时候把transform挨个注释掉,二分法定位很快。对了,你检查过num_workers吗?如果数据加载线程数开太高,有时候每个worker的缓存也会占显存,虽然这听起来像是内存问题,但某些场景下PyTorch会把这个缓存搬到GPU上。最后,如果代码里有loss.backward()之后没做optimizer.zero_grad(),那每个batch的梯度会累加,显存暴涨也是瞬间的事,虽然你说之前跑得好好的,但新加的数据增强可能改变了loss计算路径,值得再确认一下。
我之前也踩过类似的坑,加了几个transform后显存暴涨,最后发现是随机裁剪里某个张量没detach,一直在计算图里累积。建议你把每个transform单独拆出来跑一遍,配合torch.autograd.detect_anomaly()看它报不报错,虽然慢但能定位到具体操作。另外可以用tracemalloc查CPU内存,显存的话nvidia-smi dmon实时看哪个进程涨,再配合pdb在训练循环里打断点,一步步排除。数据增强里如果用了torchvision.transforms,注意RandomCrop返回的坐标是int,但某些自定义操作会不小心保留梯度,检查下有没有no_grad包裹。
我之前也遇到过类似情况,最后发现是数据增强里某个操作在GPU上做了张量拼接,梯度图没释放。你可以试试用torch.autograd.detect_anomaly(),虽然慢点但能直接报出错操作的位置,比memory_summary直观多了。
另外如果怀疑是Dataset里累积了历史batch,检查下transform里有没有把输入存成了类属性,或者用了list.append没清空。我之前用albumentations的某个版本就有这坑,升级后就好了。
实在不行就按epoch拆小批跑,配合nvidia-smi -l 1实时盯显存曲线,哪个epoch开头暴涨就断点排查那个阶段,比盲猜快。
之前我也被这个折磨过,加了几个transform之后显存直接螺旋起飞。后来我发现大概率不是数据增强本身的问题,而是你的Dataset在__getitem__里存了太多中间变量,或者某些transform返回了不同shape的tensor导致后续计算图没有释放。你可以试试给每个transform前后加个torch.cuda.reset_peak_memory_stats(),然后手动打印torch.cuda.max_memory_allocated(),分块定位是哪个操作把峰值拉起来的。
另外有个土办法,但挺管用:把batch size调成1,如果显存还是暴涨,那就肯定不是数据批量的问题,而是单张图处理过程中有东西在累积。再不行就开torch.autograd.detect_anomaly(),虽然慢,但它会在反向传播异常时直接告诉你哪一行参与了梯度计算,配合torch.cuda.set_per_process_memory_fraction把显存限制住,爆的时候看堆栈就能抓到真凶。
至于按行看显存,PyTorch官方没有现成的,但你可以用torch.profiler,里面有个with torch.profiler.profile(activities=[ProfilerActivity.CUDA]),跑完用print(prof.key_averages().table(sort_by="cuda_time_total")),虽然它按算子分类,不是严格按行,但能看出是哪个模块在疯狂申请显存。我上次就是这么找到的,结果是某次我忘了把输入tensor从.cuda()转回.cpu(),导致整个transform链都跑在GPU上,白占了一大块。你检查下是不是也有类似操作,尤其是自定义的collate_fn里,那地方最容易藏雷。
试试给dataloader加个num_workers=0跑一遍,如果没问题就八成是数据加载进程的显存没回收,以前我也被这坑过。真要定位到行号的话,可以看看pytorch的torch.autograd.detect_anomaly,但那个只能查梯度异常。还有个土办法,把数据增强一个个注释掉跑小批量,二分法定位,比看memory_summary直观多了。
我之前也遇到过类似情况,后来发现是数据增强里某个操作在GPU上动态生成了超大tensor,没及时del。你可以试试在训练循环里用torch.cuda.reset_peak_memory_stats()配合torch.cuda.max_memory_allocated(),在每个epoch前后打点,再二分法注释掉可疑的transform,基本能锁定范围。
另外推荐个工具叫pytorch_memlab,能按行输出每个tensor的分配位置,比memory_summary直观多了。不过它需要你稍微改下代码,用它的line_forward包装一下模型。还有个土办法,把batch size调成1跑一两个step,如果显存还涨,那肯定不是数据量的问题,就是某个操作本身有bug。
我上次就是靠这个定位到是RandomResizedCrop里一个参数没设对,导致每次迭代都重新分配缓存。你检查下你有没有设置torch.backends.cudnn.benchmark=False,有时候这个开关也会让显存波动。
试试用torch.profiler看每个op的显存分配,或者直接把transform逐行注释跑一遍,二分定位很快。
要不先查查数据增强里有没有用detach或clone,不然多半是计算图没释放。
试试pytorch的autograd检测,或者用torch.profiler按行打表,很快能定位到爆显存的op。
之前我遇到类似情况,把transform里的cache清掉就好了,你看看是不是开了persistent_workers。
试试给每个transform前后加个torch.cuda.synchronize()然后打印allocated memory,配和nvidia-smi的实时监控看是哪个阶段涨的,比直接看summary直观多了。另外自定义Dataset里如果用了list存tensor,记得做完增强后del掉再gc.collect(),我之前就是有个归一化临时变量没清,一个batch多占了几百M,跑几个epoch就爆了。
试试用pytorch的autograd检测钩子,或者直接二分法注释掉新增的transform,跑几个step看哪个爆。
之前遇到过是自定义Dataset里把整个tensor存了list,切片时没释放,换成索引就没事了。
试试给每个transform加个显存快照对比一下,或者检查下数据增强里是不是创建了临时变量没释放。
试试torch.profiler,带record_shapes=True,跑几个step后看profiler的table,能按operator和调用栈拆,基本能定位到是哪个模块在涨。另外,你加了数据增强,大概率是num_workers开太多,或者transform里有个别操作把中间变量留在显存了,比如转成float后再算,注意下torch.no_grad有没有漏。还有个笨办法,每轮迭代后手动print(torch.cuda.max_memory_allocated()),配合二分注释代码,比看summary快。
我之前也遇到过类似情况,后来发现是数据增强里某个操作在GPU上执行了,比如ToTensor之后又做了不必要的cuda转换,显存就默默涨上去了。你可以试试把transform里的操作全部移到CPU上,然后给DataLoader加个pin_memory=False,很多时候问题就出在这里。另外torch.cuda.set_per_process_memory_fraction能设个上限,爆了会直接报错,比看summary直观些。要是还找不到,就按batch二分法,把train循环里每段代码拆开跑,配合nvidia-smi监控,基本能锁死是哪一行。
试试给每个transform前后都打一下torch.cuda.memory_allocated(),用二分法缩小范围,比看summary直观多了。另外检查下是不是开了persistent_workers=True但num_workers设太大,数据增强里如果有随机生成的大tensor没及时del,也可能累积爆显存。我之前遇到过类似情况,最后发现是某个自定义Dataset里把整张图的所有增强结果都存进了列表,改成yield就好。
推荐你装个pytorch_memlab,它能用装饰器直接输出函数级显存占用,定位到具体方法很快。不过要注意它跟torch.compile有时不兼容,如果开了编译就先关掉再测。还有个小技巧,把torch.cuda.empty_cache()临时插到每个epoch开头,如果显存能回收,说明不是泄漏而是峰值问题,那就重点查batch里是不是有异常大的样本。
显存暴涨最烦的就是这种偶发性,我一般直接改batch_size=1跑一个step,如果还爆,就基本锁定在数据加载或transform里。你可以试着把数据增强单独抽出来跑一遍,用nvidia-smi盯着显存,哪个操作让内存曲线往上跳就是它。另外别忽略albumentations这类库,有些版本有已知的显存不释放bug,升级或降级版本可能就好了。