最近在调一个图像分割模型,用的是PyTorch,之前跑得好好的,但加了几个数据增强操作后,训练到第5个epoch显存直接爆了(12G不够用)。我怀疑是某个transform或者自定义的Dataset里内存没释放,但代码写得很乱,一时找不到具体哪一行有问题。试过torch.cuda.memory_summary(),但输出信息太杂,看不懂。有没有什么工具或方法能定位到是哪一行代码导致显存泄漏或突然暴涨?比如能不能像Python的cProfile那样,能按行看显存占用?先谢谢各位老哥了。
用PyTorch训练模型时,显存突然暴涨,怎么排查具体是哪一行代码导致的?
全部回复
共 6 条试试跑一个小batch然后用torch.cuda.set_per_process_memory_fraction限制显存,异常时会直接报错定位到具体行。
试试PyTorch自带的torch.cuda.set_per_process_memory_fraction或者torch.autograd.set_detect_anomaly(True),后者能帮你定位到具体产生NaN或梯度异常的地方,虽然不一定直接显存泄漏,但有时暴涨跟梯度爆炸有关。另外强烈推荐torch.utils.checkpoint分段检查点,或者用memory_profiler结合nvidia-smi循环采样,在数据增强函数前后加个torch.cuda.synchronize()和torch.cuda.max_memory_allocated()对比,基本能圈定是哪个transform在搞鬼。
这种情况我遇到过,排查起来确实头疼。可以先试试用torch.cuda.set_per_process_memory_fraction限制显存上限,这样爆显存时会直接报错,再配合torch.autograd.detect_anomaly()看梯度异常。另外推荐pytorch_memlab这个库,能按行打印每行代码分配的张量大小,比memory_summary直观很多。我上次就是用它发现某个transform里to(device)没加non_blocking=True导致累积了中间变量。
遇到过类似情况,后来发现是某个transform里把中间结果存到了列表里没清掉,导致每个batch都在累积。建议可以先试着在每个epoch结束手动调一下torch.cuda.empty_cache(),看看峰值有没有降下来,如果降了八成是哪个变量一直留在显存里。另外可以试试用torch.utils.checkpoint,把一些中间激活值扔掉,吃显存能少很多。
我之前也踩过类似的坑,尤其是复杂的数据增强链里,有些transform会缓存中间结果或者创建不释放的临时张量。可以试试把训练代码拆开,用torch.cuda.empty_cache()配合torch.cuda.memory_allocated()在每个关键步骤前后打点,比如每个transform执行完都打印一次,基本能定位到暴涨点。另外torch.utils.bottleneck这个工具也能看内存分配热点,但得花点时间读输出,比直接看memory_summary清楚多了。
我也遇到过类似问题,后来发现是某个数据增强操作里用了torch.no_grad()但忘了把中间变量detach,导致计算图一直挂着。可以用torch.cuda.set_per_process_memory_fraction限制一下显存,配合torch.cuda.memory_snapshot和snapshot_to_chrome可视化,能比较直观地看到哪一步分配了最多内存。另外建议把transform逐个注释掉跑一次循环,定位到具体操作后再看源码,比自己硬翻代码快很多。