最近在调一个图像分割模型,用的是PyTorch,之前跑得好好的,但加了几个数据增强操作后,训练到第5个epoch显存直接爆了(12G不够用)。我怀疑是某个transform或者自定义的Dataset里内存没释放,但代码写得很乱,一时找不到具体哪一行有问题。试过torch.cuda.memory_summary(),但输出信息太杂,看不懂。有没有什么工具或方法能定位到是哪一行代码导致显存泄漏或突然暴涨?比如能不能像Python的cProfile那样,能按行看显存占用?先谢谢各位老哥了。
用PyTorch训练模型时,显存突然暴涨,怎么排查具体是哪一行代码导致的?
全部回复
共 179 条我之前也踩过类似的坑,后来发现用torch.cuda.set_per_process_memory_fraction限制一下最大显存能快速定位到爆显存的代码段。还有就是可以把transform里的操作逐个注释掉跑一遍,大概率是某个增强操作创建了临时变量没释放。另外推荐试试pytorch的autograd.detect_anomaly,它能在反向传播时提示异常梯度,有时候显存暴涨跟梯度爆炸也有关系。
试试PyTorch的torch.cuda.set_per_process_memory_fraction限制最大显存,或者用torch.autograd.detect_anomaly定位梯度爆炸的地方。
遇到这种显存突然暴涨的情况确实挺头疼的,特别是代码写乱的时候。除了memory_summary,你可以试试torch.cuda.set_per_process_memory_fraction先给显存设个上限,这样爆了会直接报错,能帮你定位到具体epoch或batch。另外推荐用torch.utils.checkpoint,它能把中间激活值扔掉换时间省显存,但本质还是得找出泄漏点。我一般会开个torch.autograd.set_detect_anomaly(True)结合print去打印每个tensor的shape和device,或者用tracemalloc跟踪Python对象的内存分配。数据增强那块最容易出问题的是transforms里用了ToTensor后没及时释放中间变量,比如PIL图像转Tensor后原图没删干净,或者自定义Dataset里__getitem__返回了不需要的缓存。你也可以试试用torch.cuda.memory_snapshot把当前显存分配情况dump下来,再用第三方工具比如pytorch_memlab或者memray做可视化分析,虽然麻烦点但比瞎猜快。对了,确认下你的DataLoader是不是设了pin_memory和num_workers,有时候多进程反而会卡住显存不释放。
这种情况我也遇到过,建议试试torch.cuda.set_per_process_memory_fraction限制最大显存,配合torch.autograd.detect_anomaly能快速定位梯度爆炸的位置。如果怀疑是数据增强的问题,可以用torch.utils.data.DataLoader的worker_init_fn在每个worker里单独加个gc.collect(),或者把transform拆开逐个跑一遍看看哪个峰值异常。还有一个取巧的方法:在训练循环里每步打印torch.cuda.memory_allocated(),对比前后差值最大的那段代码基本就是罪魁祸首。
我之前也遇到过类似的情况,后来发现是transforms里某个操作把中间结果留在了GPU上,用torch.cuda.empty_cache()试了一下没解决,最后是靠torch.cuda.set_per_process_memory_fraction限制显存上限才定位到具体模块的。你可以试试在关键函数前后加print(torch.cuda.memory_allocated())对比一下,或者用torch.utils.checkpoint把梯度检查点打开,有时候能临时绕过暴涨问题。
这种情况确实挺头疼的,尤其是代码改乱之后很难定位。我建议可以先试试用torch.cuda.set_per_process_memory_fraction限制最大显存,这样爆显存时会直接报错而不是卡死,配合torch.autograd.set_detect_anomaly(True)能更快抓到反向传播里的异常。另外你提到怀疑是transform的问题,可以单独把数据增强部分拆出来跑一个循环,用torch.cuda.memory_allocated()每步打印显存变化,如果某一步突然飙升就是那行代码的问题。还有个小技巧是检查下DataLoader的num_workers,有些数据增强操作(比如随机裁剪)在多进程下可能会复制多次张量导致显存暴涨,尤其是用了worker_init_fn没处理好随机种子的时候。如果还不行,可以试试pytorch_memlab这个库,它能像cProfile那样逐行记录显存,不过需要稍微改改代码。最后提醒下,自定义Dataset里如果用了Python原生的list或者dict来存中间结果,记得用完后手动del或者用gc.collect(),有时候PyTorch的缓存机制不会自动释放这些临时变量。
试试torch.cuda.set_per_process_memory_fraction设个上限,爆了直接报错能快速定位。另外可以hook一下DataLoader的worker,打印每个batch的显存变化,毕竟增强操作里容易出问题的就是随机resize或者cutmix这种动态张量。我之前遇到过torchvision的RandomResizedCrop会缓存中间结果,换成自己写的就解决了。
试试用torch.cuda.set_per_process_memory_fraction限制显存,配合torch.autograd.detect_anomaly能快速定位异常增大的地方。
可以用torch.cuda.set_per_process_memory_fraction限制显存上限,配合try-except逐步缩小范围。
我最近也遇到过类似的问题,后来发现是DataLoader的num_workers开太多,加上transform里用了类似RandomCrop这种会保留中间张量的操作,导致显存一直没释放。你可以试试torch.cuda.empty_cache()在每个epoch结束调用一下,或者用torch.utils.checkpoint把中间变量丢掉。另外有个叫pytorch_memlab的工具可以按行追踪显存分配,比memory_summary直观很多。
我之前也遇到过类似问题,后来发现是DataLoader的num_workers设太高,加上某些transform里用了原地操作,导致显存一直堆积。建议试试torch.cuda.set_per_process_memory_fraction来限制单次分配,或者干脆在transform里加个torch.cuda.empty_cache(),虽然不优雅但能快速定位。另外可以写个简单脚本,只跑数据加载部分,看显存会不会涨,这样就能排除是模型还是数据的问题。
我最近也碰到过类似问题,后来发现是某个数据增强里用了torchvision的RandomResizedCrop,没加torch.no_grad导致梯度图一直挂着。你可以试试给每个transform加个with torch.no_grad(),或者用torch.autograd.set_detect_anomaly(True)然后看回溯,虽然慢但能抓出是哪一行爆的。另外pytorch的torch.cuda.memory._record_memory_history()配合snapshot工具也挺好用,能按行看分配栈。
这种情况我也踩过坑,尤其是加了复杂的数据增强之后,显存爆炸往往不是模型本身的问题,而是数据加载时某些操作把中间结果留在了计算图里。你试试在transform里面加一句torch.no_grad(),或者把每个增强操作单独封装成函数,然后在每个函数里显式调用del和torch.cuda.empty_cache()来测试,看哪个环节显存不回缩。另外,可以用torch.cuda.set_per_process_memory_fraction(0.8)先限制上限,这样爆显存时会直接报错,配合torch.cuda.memory._record_memory_history()和torch.cuda.memory._dump_snapshot()能生成一个可视化内存快照,在chrome://tracing里看,比那个summary直观很多。还有个笨办法:把数据增强一步步注释掉,跑一个mini batch,用nvidia-smi盯着显存变化,虽然麻烦但最直接。对了,检查下你的Dataset里的__getitem__有没有把增强后的tensor无意间留在属性里,比如self.last_image = transform(img)这种,很容易造成累积。
这种情况我太熟了,pytorch的显存问题有时候真的让人头大。你加的可能是像RandomResizedCrop或者RandomRotation这种操作吧?它们如果在__getitem__里频繁创建临时tensor,确实容易把显存堆爆。我建议你先试一下torch.cuda.set_per_process_memory_fraction(0.9),强行限制最大显存占用,这样爆了会直接报OOM错误,比突然崩掉好定位。然后可以装个pytorch_memlab,它能按行记录显存分配,在可疑的transform前后打上@profile装饰器,跑一两个batch就能看到哪一步涨得最离谱。另外,你检查下数据增强里有没有用transforms.ToTensor()之后又手动做了torch.from_numpy,这种重复转换会多占一份显存。如果自定义Dataset里有torch.no_grad()或torch.inference_mode()没用好,也可能导致梯度图残留。最后一个小技巧:在训练循环里每隔几步打印torch.cuda.max_memory_allocated(),对比不同epoch的峰值,如果持续上涨就是有泄漏。
试试给每个transform前后打点记录cuda memory,二分法定位很快,别靠猜。
试试把transform里的操作拆开逐个跑一遍,多半是某个增强函数返回了没detach的图。
我上次就是被RandomResizedCrop坑了,加个with torch.no_grad()包一下就好了。
这个情况我也踩过坑,加了transform之后显存炸了大概率不是泄漏,而是某个操作把整张图或者中间变量留在计算图里了。你可以试试在训练循环里用torch.autograd.detect_anomaly(),它虽然慢但会直接报出异常张量产生的代码位置。另外如果怀疑是Dataset的问题,可以单独跑一遍数据加载流程,把num_workers设成0看显存还涨不涨,这样能快速定位是不是数据侧的问题。我之前就是被一个随机裁剪的边界检查坑了,生成了一张超大图塞进batch里。
我之前也遇到过类似情况,加了个随机裁剪之后显存直接翻倍,后来发现是transform里用了太多临时tensor没及时del,而且PyTorch的缓存分配器会把显存留着不还,看着像泄漏其实不是。你可以试试torch.cuda.reset_peak_memory_stats()配合torch.profiler,那个能按操作符级别看内存分配,比memory_summary直观很多。不过最粗暴有效的办法是二分法,先把新加的数据增强全注释掉,确认基线没问题,再逐个加回来跑一个step看显存峰值,基本两三轮就能定位到嫌疑代码。另外检查下Dataset的__getitem__里是不是把整张图或者大数组存成了实例变量,Python的GC有时候不会立刻回收,累积几个epoch就爆了。还有个小技巧,在训练循环里定期打印torch.cuda.memory_allocated()和torch.cuda.memory_reserved(),如果allocated稳定但reserved涨,那就是缓存碎片问题,可以试试torch.cuda.empty_cache()临时缓解。真要按行看显存,目前没有现成工具,但你可以把可疑函数拆成小块,每块之间插一行峰值打印,比自己瞎猜快多了。
试试用torch.autograd.detect_anomaly()配合torch.cuda.set_per_process_memory_fraction把显存限制死,跑崩的时候看堆栈能定位到具体操作。另外查一下是不是数据增强里用了torch.no_grad没包住,或者ToTensor之前没做copy,有些transform会保留计算图导致显存累积。我上次就是这么排查出来的,比直接看memory_summary直观多了。
试试用torch.profiler带memory profiling,能按操作符看显存占用,比summary直观多了。
我一般用torch.cuda.set_per_process_memory_fraction限制显存,爆了会直接报错定位到具体行。