最近在调一个图像分割模型,用的是PyTorch,之前跑得好好的,但加了几个数据增强操作后,训练到第5个epoch显存直接爆了(12G不够用)。我怀疑是某个transform或者自定义的Dataset里内存没释放,但代码写得很乱,一时找不到具体哪一行有问题。试过torch.cuda.memory_summary(),但输出信息太杂,看不懂。有没有什么工具或方法能定位到是哪一行代码导致显存泄漏或突然暴涨?比如能不能像Python的cProfile那样,能按行看显存占用?先谢谢各位老哥了。
用PyTorch训练模型时,显存突然暴涨,怎么排查具体是哪一行代码导致的?
全部回复
共 179 条碰到过类似的,加了几个augmentation之后显存直接翻倍,最后发现是随机裁剪那块在GPU上做了太多张量拼接,每次迭代都留了中间变量。你可以试试用pytorch的autograd检测机制,开detect_anomaly()跑一个epoch,它会直接给出反向传播时哪块计算图出了问题,虽然是报错形式但能定位到具体操作。还有个土办法,把batch size调成1,然后把transform逐个注释掉,二分法测哪个操作吃显存,虽然慢但绝对有效。另外你说的cProfile那种按行看显存,目前没有现成工具,但可以用torch.profiler的memory profiling功能,它会记录每个op的分配和释放,输出结果里搜allocated bytes,能看出是哪个层或者哪个transform在持续累积。注意一下你的自定义Dataset,如果里面有在__getitem__里做重计算但又没释放numpy数组,或者用了num_workers但没设persistent_workers,也可能导致显存逐步涨。最后建议盯一下数据增强里有没有用.to(device)或者.cuda(),有些增强库会自动把tensor搬GPU,你后面再搬一次就容易爆。
遇到这种问题别慌,我上次也被增强操作坑过,后来发现是某个transforms在GPU上做张量操作忘了加non_blocking=True,导致中间变量一直留在显存里。你可以试试torch.autograd.detect_anomaly()配合torch.cuda.set_per_process_memory_fraction把显存限到接近爆的阈值,这样报错时能直接定位到创建那个张量的代码行。另外如果自定义Dataset里用了torch.from_numpy,记得加.pin_memory()不是必须的,但每次迭代完要手动del掉大数组,或者干脆把数据预处理挪到__getitem__外面,提前缓存好。实在不行就拆batch跑二分法,把transform一个个禁掉,基本能锁定嫌疑犯。
我之前也遇到过类似的情况,加了几个transform之后显存直接爆掉,后来发现是随机裁剪的时候没处理好边界,导致生成了超大tensor。建议你试试用torch.autograd.detect_anomaly()配合断点,或者干脆在DataLoader里加个hook打印每个batch的tensor形状,基本能锁定是哪个操作在搞鬼。memory_summary确实难懂,我一般直接看allocated_bytes的峰值和当前值的差值,暴涨基本就是有中间变量没释放。还有个笨办法,把数据增强一步步注释掉跑一个batch看显存变化,虽然麻烦但最直观,比瞎猜效率高。
试试用torch.profiler的with_stack参数,能直接看到每个操作的显存分配和Python调用栈,比memory_summary直观多了。另外排查transform时,先单独跑一遍预处理流程,看数据加载前后的显存变化,重点检查有没有在__getitem__里把tensor留在GPU上,或者用了torch.no_grad但没释放中间变量。我之前遇到类似情况,最后发现是某个自定义transforms里顺手存了batch到self,改成局部变量就好了。
试试给每个transform前后加个torch.cuda.synchronize再打印显存,二分法定位很快的。
建议用pytorch的autograd检测钩子,给可疑层挂上,看哪层前向开始暴涨。
试试用torch.profiler,能按操作符看显存分配,比memory_summary直观多了。
或者干脆把transform逐个注释跑一遍,二分法定位,五分钟就能找到元凶。
我之前也踩过这个坑,加了几个transform之后显存直接翻倍,后来发现是随机裁剪那步没释放中间变量。你可以试试torch.autograd.detect_anomaly(),它能定位到反向传播里哪块计算图出问题,虽然慢点但很准。另外建议检查下有没有在循环里反复创建tensor,比如把transforms对象移到__init__里而不是每次__getitem__都new一个,这招对我挺管用。
之前也踩过这坑,加transform后显存暴涨大概率是数据加载时把整个batch的中间结果都留在计算图里了,试试把transform里那些tensor操作都用@torch.no_grad()包一下,或者直接用torch.utils.data.DataLoader的pin_memory和num_workers配合看看。真要定位到行号的话,可以用torch.autograd.detect_anomaly()开异常检测,虽然会慢点,但能直接炸出是哪次前向传播出的问题,比看memory_summary直观多了。另外也可以试试pytorch_memlab这个库,能按行打印每个tensor的分配点,亲测比自带的summary好用。
看到你说加了数据增强才爆显存,我第一反应就是问题大概率出在transform里,而不是模型本身。PyTorch的DataLoader默认会缓存一部分张量,如果transform里有用到像RandomCrop或者Resize这种会改变tensor形状的操作,有些实现会在GPU上做计算,导致中间变量没及时释放,尤其是你开了pin_memory或者num_workers>0的时候,问题会被放大。你可以试试把所有的transform都改成在CPU上执行,然后在Dataset的__getitem__末尾强制del掉中间变量,看看显存曲线会不会平缓下来。至于定位具体行,有个取巧的办法:在训练循环里每隔几步手动调用torch.cuda.empty_cache(),同时用nvidia-smi实时监控,如果显存是阶梯式上涨而不是平稳波动,那就能确定是某个batch处理完没释放,这时候可以在Dataset里加个print输出当前样本索引,二分法定位到具体是哪个操作。另外推荐试下torch.profiler,它能按操作名和调用栈显示显存分配,比memory_summary直观很多,虽然有点学习成本但值得花半小时折腾。最后提个醒,如果自定义Dataset里用了全局变量存中间结果,检查一下是不是每个epoch结束后还保留着引用,这个坑我踩过。
试试给每个transform前后打个显存快照,用torch.cuda.reset_peak_memory_stats()配合peak_stats对比,基本能锁定是哪个操作在涨。另外检查下数据增强里有没有拼接或重复调用tensor.to(device)的,我之前就是RandomCrop里忘释放中间变量,加了del和gc.collect就好了。memory_summary确实难用,可以看看pytorch的memray,虽然主要针对CPU,但配合cuda的events能定位到具体行。
这问题我上周刚踩过坑,加了几个随机裁剪和亮度抖动的transform,到第三个epoch直接OOM。memory_summary()那玩意儿确实能看,但更像在翻垃圾桶找针,信息量太大反而不容易定位。我建议你先别急着改代码,把batch size调成1跑一个step,如果显存还涨那就是数据链路的问题,这时候去Dataset里逐行注释掉transform,二分法排查看哪个操作触发增长,比瞎猜快得多。另外有个土办法,直接在训练循环里每10个step打印一次torch.cuda.max_memory_allocated(),如果数值是阶梯式上升而不是稳定在高位,大概率是某个变量被意外留在了计算图里,比如你在循环里用了loss.backward()但没写optimizer.zero_grad(),或者某个中间tensor被赋给了self属性。至于按行看显存,PyTorch官方没这工具,但可以试试torch.profiler的with_stack参数,它能把每个op的显存分配和对应的Python调用栈打出来,虽然输出还是有点乱,但至少能指向具体函数名。还有个小技巧,把transform里所有操作换成in-place版本,比如x = x[:, :, ::2, ::2]这种切片不会复制内存,但torchvision.transforms.Resize返回的是新tensor,如果累积在list里就会炸。实在不行就在每个transform前后手动torch.cuda.synchronize()然后看memory_allocated()差值,虽然笨但绝对准。
试试用pytorch的torch.profiler,能按行看显存分配,比memory_summary直观多了。
之前也遇到过,八成是数据增强里某个操作在GPU上执行没释放,排查时把transform逐个禁掉对比试试。
试试给每个transform前后加个torch.cuda.synchronize然后打印allocated_memory,很快就能锁定是哪步涨的。
或者用pytorch的memory_profiler配合逐行执行脚本,能精确到行号,比summary直观多了。
我之前也遇到过这种加了transform之后显存炸了的情况,后来发现是某个RandomCrop在每次迭代时都保留了原图副本没释放。建议先用torch.autograd.detect_anomaly()看看是不是反向传播时爆的,如果没报错就写个简单脚本把transform一个个单独跑,配合gpustat实时盯显存变化,基本能锁到具体操作。
另外也可以试试pytorch_memlab这个库,能按行输出张量分配点,比memory_summary直观多了。还有个小技巧,把Dataset里的__getitem__返回的tensor用.cpu()强制搬回内存,看看会不会缓解,如果有效那就是GPU上累积了太多中间变量。
要是还不行,直接二分法注释代码,先固定住数据增强,把训练循环里每轮print一次torch.cuda.memory_allocated(),对比哪一轮开始突增,范围就很小了。反正别盯着memory_summary看,那个真不是给人读的。
我之前也遇到过类似情况,最后发现不是transform的问题,是Dataset里把增强后的图像存成了list,每个epoch都在往里append,显存自然就爆了。你这个情况可以试试用torch.utils.bottleneck,虽然它主要管CPU,但能帮你看到数据加载那块的耗时和内存变化,至少能缩小范围。另外torch.cuda.memory._record_memory_history()这个API能开启堆栈跟踪,配合torch.cuda.memory._snapshot()可以导出成json,用chrome的tracing工具看,能精确到每行Python代码的分配点,比memory_summary直观多了。不过说实话,如果代码很乱,建议先把所有transform抽出来单独测,比如跑一个epoch的纯加载(不训练),看显存是否稳定;如果还涨,就二分法注释掉一部分增强操作,定位到具体那个操作。还有个小技巧,用gc.collect()和torch.cuda.empty_cache()放在每个batch结束后,虽然治标不治本,但能确认是不是缓存没释放。最后提醒下,有些in-place操作比如x += 1在autograd下会保留计算图,也可能造成显存累积,检查下有没有类似写法。
试试给每个transform前后加个torch.cuda.synchronize再打印allocated_memory,二分定位很快的,之前我这么干过。
加个memory_profiler配合pytorch的钩子,能按行标内存,不过小batch跑几轮就够看了。
我之前也遇到过这种,加了几个transform之后显存直接翻倍,后来发现是RandomCrop这类操作在每次调用时都会保留中间变量,试试用torch.utils.checkpoint或者把transform里的临时tensor显式del一下。另外推荐用pytorch_memlab这个库,能按行输出内存分配,比memory_summary直观多了。不过也要注意是不是数据加载时worker数开太多,导致CPU和GPU之间缓存堆积,我之前就是卡在这。
用torch.cuda.memory._record_memory_history()配合tracemalloc试试,能把分配栈打到具体行号,比summary直观多了。
我之前也遇到过类似的情况,最后发现是数据增强里某个操作在GPU上执行了,比如把tensor从CPU转到GPU没转回来,显存就只增不减。你可以试试用torch.autograd.detect_anomaly()或者给每个transform加个打印,看哪个前后显存变化最大,比看memory_summary直观多了。另外,别忽略num_workers设太高也可能导致缓存堆积,先把worker降到0跑一跑排除掉。
试试pytorch的torch.cuda.set_per_process_memory_fraction配合tracemalloc,把显存分配和python内存分配绑一起看,能缩小范围。另外检查下你的Dataset里是不是在__getitem__里对CUDA tensor做了切片或clone,这玩意儿容易隐式保留计算图。我之前遇到过类似问题,最后发现是transforms里的随机裁剪返回了不同shape的tensor,导致显存碎片化,加个torch.cuda.empty_cache()在epoch间调用就稳多了。