最近在调一个图像分割模型,用的是PyTorch,之前跑得好好的,但加了几个数据增强操作后,训练到第5个epoch显存直接爆了(12G不够用)。我怀疑是某个transform或者自定义的Dataset里内存没释放,但代码写得很乱,一时找不到具体哪一行有问题。试过torch.cuda.memory_summary(),但输出信息太杂,看不懂。有没有什么工具或方法能定位到是哪一行代码导致显存泄漏或突然暴涨?比如能不能像Python的cProfile那样,能按行看显存占用?先谢谢各位老哥了。
用PyTorch训练模型时,显存突然暴涨,怎么排查具体是哪一行代码导致的?
全部回复
共 179 条试试跑一个小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逐个注释掉跑一次循环,定位到具体操作后再看源码,比自己硬翻代码快很多。
遇到过类似的情况,当时是用了太多的随机裁剪和翻转叠加,导致中间变量没及时释放。可以试试把数据增强部分单独拎出来跑一次迭代,然后看torch.cuda.max_memory_allocated()对比一下,基本能锁定是哪个transform在吃显存。另外检查一下dataloader的num_workers和pin_memory,有时候多进程也会让显存莫名其妙涨上去。
试试PyTorch自带的torch.cuda.set_per_process_memory_fraction限制显存上限,报错时能直接定位到爆显存的张量位置。另外推荐用torch.cuda.memory._record_memory_history()结合memory_summary()看具体分配堆栈,比直接看摘要清晰很多。之前我也遇到过类似问题,最后发现是transform里忘了对tensor用.clone()导致计算图累积。
这种问题我也遇到过,建议试试torch.cuda.set_per_process_memory_fraction先限制下显存上限,这样爆的时候能更快定位到异常代码块。另外可以配合torch.autograd.set_detect_anomaly(True)打开梯度异常检测,有时候数据增强里的操作会导致计算图异常放大。我上次就是自定义transform里用了torch.where没注意广播维度,显存直接炸了。
这种问题我也踩过坑,尤其是数据增强堆多了以后,显存突然炸掉真的很让人头大。建议你先试试torch.cuda.set_per_process_memory_fraction限制最大显存,这样至少能快速定位到是哪个阶段爆的。然后结合torch.autograd.set_detect_anomaly(True),虽然会拖慢速度,但能帮你追踪到产生NaN或者梯度异常的ops。不过你说的按行看显存占用,目前PyTorch官方没有像cProfile那样细粒度的工具,但你可以用torch.cuda.memory._record_memory_history()配合torch.cuda.memory._snapshot()导出内存快照,再用chrome://tracing打开看,那个图比memory_summary直观多了,能精确到每个Tensor的分配位置。另外我个人经验是,很多“泄漏”其实是数据增强里用了transforms.ToTensor()或者Normalize时没处理好batch维度,比如在__getitem__里重复加载了原始图像没释放引用,或者用了torchvision.transforms.RandomCrop这类会缓存中间结果的op。你检查下自定义Dataset里是不是在__init__阶段就把所有图片读到内存了?还有增强操作里如果有RandomApply或者RandomOrder,注意它们内部可能维护了状态,多进程DataLoader下会累积显存。
试试用 torch.cuda.set_per_process_memory_fraction 限制显存,再配合 torch.autograd.set_detect_anomaly(True) 定位异常梯度。
试试PyTorch的torch.cuda.set_per_process_memory_fraction限制上限,或者用torch.cuda.memory._record_memory_history能按调用栈看分配。
试过用torch.cuda.set_per_process_memory_fraction限制最大显存吗?这样起码爆的时候能更快定位到出问题的epoch。另外推荐pytorch_memlab这个库,可以装饰一下可疑函数,每次调用完自动打印显存变化,比手动打日志好用。我之前也遇到过类似情况,最后发现是自定义Dataset里有个to(device)操作忘了移除,每个batch都在拷贝数据。
这问题我也踩过坑,显存暴涨确实恶心。之前我搞分割模型时也遇到过类似情况,最后发现是某个自定义transform里开了太多临时变量,而且没有及时del掉,导致中间结果一直卡在显存里。torch.cuda.memory_summary()确实信息爆炸,但你可以配合torch.cuda.set_per_process_memory_fraction先限制最大显存,让它提前爆,报错时就能看到具体堆栈了。另外推荐用torch.cuda.memory._dump_snapshot生成内存快照,然后扔到chrome://tracing里可视化,能看到每个Tensor在哪个时间点被分配,定位到具体代码行。还有一招比较暴力,就是写个装饰器或者用py-spy在训练循环里逐行打印显存增量,虽然慢但很准。建议你重点检查数据增强里有没有在原地修改张量,或者用到了torchvision的某些ops,它们有时会意外保留梯度。最后提醒下,PyTorch的DataLoader如果num_workers开太多,子进程的显存泄漏也会累积,试试降到2看看。
这种问题我也遇到过,数据增强堆多了显存直接炸,排查起来确实头大。torch.cuda.memory_summary()那块报告说实话对新手不太友好,一堆缓存池和碎片信息,看多了容易懵。我个人比较推荐两个土办法:第一个就是分段注释法,先注释掉所有transform,如果显存不爆了就逐个加回来,每加一个跑一个mini-batch,看哪个transform把显存拉高了;第二个是试试torch.cuda.memory_allocated()和torch.cuda.max_memory_allocated()配合着在每个transform前后打log,能比较直观地看出哪一步在持续累积显存。另外注意下你的DataLoader的num_workers是不是设太大了,以及pin_memory=True有时也会让显存占得偏高。还有一个容易被忽略的点——如果你在transform里用了像transforms.RandomCrop这种,它可能每次都会创建新张量并且没被及时回收,可以检查下有没有在循环里重复创建dataset实例或者没有用del清理中间变量。实在不行就上pytorch的profiler,虽然配置麻烦点,但能抓到每一步分配的内存。
试试用torch.cuda.set_per_process_memory_fraction限制最大显存,配合pdb在可疑代码段逐行跟踪。
试试torch.cuda.set_per_process_memory_fraction加上memory_profiler一行行跑,或者用torch.utils.checkpoint手动查中间变量。
这种问题我太懂了,之前搞检测模型的时候也被类似的问题折磨过。你这种情况大概率不是显存泄漏,而是某个transform或Dataset里不小心保留了计算图,导致反向传播后梯度没释放干净。建议你先试试torch.cuda.empty_cache()加在每次epoch结束,如果显存能降下来,那就不是真泄漏,而是临时缓存占着茅坑不拉屎。要定位具体行的话,可以试试用torch.autograd.set_detect_anomaly(True),它能帮你找出反向传播中异常的节点,虽然不直接显示显存,但至少能锁定可疑的计算链路。另外我强烈建议你给每个transform加个torch.no_grad()上下文,尤其是那些不需要梯度的数据增强操作,不然它们也会参与计算图构建。还有一个野路子:把代码分段跑,比如在Dataset的__getitem__里每处理一步就打印一下当前显存,配合nvidia-smi -l 1实时监控,很快就能看出哪一步在疯狂吃显存。要是还搞不定,可以试试torch.cuda.memory._record_memory_history(),它能记录每个张量的分配位置,虽然输出很啰嗦但确实能追溯到具体行号。
试过用torch.cuda.set_per_process_memory_fraction限制最大显存吗?虽然不能精确定位,但至少能防止爆掉让你慢慢排查。另外推荐你用pytorch的autograd.detect_anomaly()配合torch.no_grad去逐块注释代码,我上次就是靠把transform一个个取消才发现是某个随机翻转忘了释放内存。
我之前也遇到过类似情况,后来发现是自定义Dataset里忘了在__getitem__里对图像做深拷贝,导致多个batch共享同一份内存,显存越堆越高。可以试试用torch.utils.checkpoint或者给每个transform加个显存监控,比如在可疑操作前后打一下torch.cuda.memory_allocated(),差值大的地方基本就是问题代码。另外torch.cuda.set_per_process_memory_fraction也能帮忙卡个阈值,让报错更早出现。