最近在调一个图像分割模型,用的是PyTorch,之前跑得好好的,但加了几个数据增强操作后,训练到第5个epoch显存直接爆了(12G不够用)。我怀疑是某个transform或者自定义的Dataset里内存没释放,但代码写得很乱,一时找不到具体哪一行有问题。试过torch.cuda.memory_summary(),但输出信息太杂,看不懂。有没有什么工具或方法能定位到是哪一行代码导致显存泄漏或突然暴涨?比如能不能像Python的cProfile那样,能按行看显存占用?先谢谢各位老哥了。
用PyTorch训练模型时,显存突然暴涨,怎么排查具体是哪一行代码导致的?
全部回复
共 179 条我之前也踩过类似的坑,加了几个transforms之后显存直接失控。可以试试用torch.autograd.detect_anomaly(),它能在反向传播时精确报出是哪一行计算图出了问题,比memory_summary直观多了。
另外建议把数据增强单独拎出来跑一遍,用torch.utils.data.DataLoader的pin_memory和num_workers调成0,排除多进程缓存占用。如果还是涨,就在每个transform后面插个torch.cuda.synchronize()和print(gpu_memory_allocated()),二分法定位特别快。
试试给每个transform前后加个torch.cuda.synchronize然后配合nvidia-smi的实时刷新看显存曲线,虽然土但能快速锁定是哪个操作涨的。另外可以检查下数据增强里有没有不小心把图像slice或者mask保留成了计算图的一部分,用del和torch.cuda.empty_cache手动清一下验证。如果代码实在乱,建议把Dataset里返回的样本改成只返回tensor的copy,排除共享内存的坑。之前我遇到过类似问题,最后发现是随机裁剪里用了可变长度导致pin_memory疯狂申请空间,换成固定尺寸就好了。
我之前也踩过这个坑,后来发现多半不是数据增强本身,而是transform里不小心把Tensor留在计算图上没detach,或者Dataset的__getitem__里反复创建了没释放的临时变量。你可以试试用pytorch的autograd记录hook,或者干脆把每个transform单独拆出来跑一遍,看哪个step前后显存跳变最明显,比直接看memory_summary直观多了。另外如果加了随机裁剪之类的操作,记得检查一下有没有在循环里累积了图,导致反向传播时显存翻倍。
试试用pytorch的autograd.detect_anomaly,再配合cuda的NVTX手动标一下关键代码段,基本能锁死问题。
建议直接二分法注释代码,或者用torch.profiler看每步内存变化,比看memory_summary直观多了。
这问题我踩过坑,先别急着上高级工具,试试把transform里的随机操作换成确定性版本跑一个step,看显存还涨不涨,大概率是某个增强算子(比如随机crop或mixup)在反向传播时把计算图撑爆了。另外你可以在每个epoch结束手动调torch.cuda.empty_cache(),但注意这只能清缓存,真正泄漏得靠tracemalloc配合torch.autograd.detect_anomaly(),后者能直接告诉你哪步反向传播爆的。如果还定位不了,就写个最小复现脚本,把数据增强一个个加回去二分测,比看memory_summary直观多了。
试试用torch.cuda.set_per_process_memory_fraction限制上限,配合二分注释transform代码,比看summary直观多了。
给每个transform加个计数器打印显存变化,跑一个batch就能看到是哪个操作在涨,之前我这么定位到过问题。
我之前也遇到过类似的,数据增强加的多了之后显存暴涨大概率不是泄漏,是计算图没释放。你可以试试在训练循环里每步结束后手动调一下torch.cuda.empty_cache(),看会不会好点。另外建议给每个transform单独跑一遍前向,用nvidia-smi实时盯显存,很快就能定位到是哪个操作的问题。还有个笨办法,把自定义Dataset里的预处理逻辑拆出来单独跑,排除是不是数据加载时缓存了太多中间结果。
其实torch.cuda.memory_summary确实不直观,但你可以配合设置PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,有时候能缓解碎片化导致的暴涨。
我之前也踩过类似的坑,加了transforms之后显存炸了,最后发现是某个随机裁剪的tensor没detach,一直挂在计算图里。你可以试试在Dataset的__getitem__末尾加个torch.cuda.synchronize(),然后配合pytorch的autograd检测,用torch.autograd.detect_anomaly()能帮你缩小范围,虽然慢点但比瞎猜强。另外建议把数据增强里的tensor操作都转成numpy再转回tensor,别直接在cuda上做,省得隐式建图。要是还不行,就看看是不是num_workers开太多,每个worker都会复制一份显存,这个经常被忽略。
这问题我太熟了,之前做检测头的时候也栽在类似坑里,加了几个随机裁剪和mixup之后显存直接翻倍。你光看memory_summary确实容易懵,那个输出是给CUDA看的不是给人看的。我建议你先别急着找哪一行,直接开个新脚本,把数据增强一个个单独跑一遍,每个transform前后都打一下torch.cuda.memory_allocated(),很快就能锁定是哪个操作在爆。另外注意下你的Dataset的__getitem__里有没有把Tensor转成numpy再转回来,这种隐式类型转换经常导致显存碎片化累积,看着像泄漏其实是没释放。还有一个很隐蔽的点,如果用了torch.no_grad()包住增强操作,但里面不小心调用了.cuda()或者.to(device),那计算图还是会挂上钩,导致backward时显存峰值飙升。真要按行定位的话,可以试试torch.profiler的with_stack=True参数,配合record_shapes=True能输出每个操作符的调用栈,不过它按行报的是算子不是你的Python源码,得自己映射一下。实在不行就写个dummy循环,把batch size调成1,然后逐步增加,看每多一个样本显存涨多少,这样能判断是固定开销还是随数据线性增长,前者是缓存问题,后者就是某一行在创建大中间变量。我之前用这招发现是随机缩放的插值方式从bilinear换成了nearest,结果生成了个超大临时张量,现在都改成在CPU端做增强再传GPU,稳得很。
试试用pytorch的autograd检测或者给每个transform加个hook打印tensor形状,暴涨前一般能抓到异常。
用torch.profiler配合record_function给可疑代码段打标,看内存时间线比看summary直观多了。
说到这个我太有同感了,之前调检测模型也遇到过一模一样的情况,加了几个random erasing之后第3个epoch直接OOM。你那几个新增的transform里如果用了随机裁剪或者flip,大概率是每次迭代都在生成新的tensor但没及时del掉,尤其是有个坑是某些op会保留计算图的历史引用,比如在data pipeline里对tensor做了requires_grad操作,那整个batch的中间变量全被攒着直到反传才释放。我后来是靠两个土办法定位的,一是把batch size调到1然后逐步加transform,哪个加完显存曲线斜率突变就是哪个;二是用torch.autograd.detect_anomaly()开着跑,虽然慢点但能直接报出异常op的行号。还有个小技巧,在自定义Dataset的__getitem__里最后显式调一下gc.collect(),配合torch.cuda.empty_cache(),能缓解但治标不治本。你要是想按行看显存,目前好像没有现成的cProfile等价物,但可以试试用torch.profiler的with_stack=True参数,它能输出每个操作的调用栈,配合memory_viewer能粗略定位到具体模块,就是输出也贼长,得自己过滤。最后建议你查下是不是数据加载的num_workers设太高了,有时候不是显存爆而是CPU内存撞墙然后拖垮CUDA上下文。
遇到这种显存暴涨的情况,我第一反应其实不是去盯某一行代码,而是先怀疑数据增强是不是引入了“动态图”或者“可变长”的东西。比如随机crop或者resize,如果尺寸不固定,PyTorch的autograd会为每个shape保留计算图,内存自然就堆上去了。你试试把增强后的tensor强制统一到固定尺寸,或者干脆在Dataset里先做一次to(device),看看峰值是不是立刻降下来。
另外你说的memory_summary看着乱,可以换个思路,用torch.cuda.set_per_process_memory_fraction(0.8)直接把显存上限锁死,跑起来如果报OOM,它会在实际分配那个tensor的堆栈里给出更具体的提示,比summary直观得多。或者试试PyTorch自带的torch.autograd.detect_anomaly(),不过它主要抓NaN和梯度异常,对纯显存占用帮助有限。
真要按行看显存,其实有个野路子:把训练循环里每个step前后都打印一下torch.cuda.memory_allocated(),然后二分法注释掉可疑的transform,跑两三个batch对比差值。我之前就这么干,最后定位到是一个自定义的collate_fn里不小心把整个batch的list都保留了引用,导致GPU缓存没法释放。另外,别忽略一个坑:如果用了Albumentations,它默认返回numpy数组,你再转tensor时如果没调用clone(),可能会让原来的numpy数组一直留在CPU内存,间接影响GPU的预取队列,显存也会被拖到爆。
实在嫌麻烦,可以试试pytorch_memlab这个库,它有个LineProfiler能按行报告内存分配,虽然不是GPU专属,但配合memory_summary能交叉验证。不过它只支持Python层,如果是C++扩展或者CUDA kernel里爆的,那还得靠nsight systems这类工具,但那就太深了。你先按固定尺寸和排查collate_fn这个方向试试,多半能解决。
试试给每个transform前后打点记录显存差,或者用pytorch的autograd检测hook,能锁定到具体张量。
可以用torch.cuda.set_per_process_memory_fraction和分段跑数据加载,把Dataset里每步的显存峰值打印出来对比。
我之前也踩过类似的坑,加了几个transform之后显存直接翻倍,最后发现不是泄漏,是某个随机crop的实现里对同一张图反复调用了cuda(),导致每次迭代都在显存里留了副本。你试试用torch.profiler,它能按操作符统计显存分配,比memory_summary直观很多,至少能看出是哪个模块在涨。
如果只想快速定位行号,可以试试在怀疑的代码段前后打torch.cuda.reset_peak_memory_stats()和torch.cuda.max_memory_allocated(),手动二分法缩小范围。我一般喜欢配合nvidia-smi -l 1盯实时显存,但用它看不了行级,只能确认是前向还是反向爆的。
还有个野路子,把batch size调到1,如果显存还是涨,那基本就是Dataset或transform里存了不该存的东西,比如把整个tensor list挂在了self上。你检查下是不是在__getitem__里用了局部大变量没释放,Python的GC有时候对cuda tensor不敏感,得显式del或者用with torch.no_grad()包一下。
另外,如果加了Albumentations,注意它的某些操作会返回float64,PyTorch模型里一混就自动转float32,但中间显存峰值会高不少。你可以试着把所有transform的输出都强制.contiguous()看看,我上次就是栽在非连续内存上,导致后续卷积隐式复制。最后推荐个库叫pytorch_memlab,能装饰器按函数打印显存,虽然对行号支持一般,但比手写强点。
我之前也遇到过类似情况,后来发现是数据增强里某个随机裁剪操作在GPU上动态生成了超大tensor,导致显存峰值。你可以试试用pytorch的torch.autograd.detect_anomaly,它虽然不能显示具体行号,但能定位到反向传播时出问题的张量操作,配合torch.cuda.set_per_process_memory_fraction限个上限,爆了会直接报错加堆栈,比memory_summary直观多了。另外检查下你的Dataset的__getitem__里是不是把中间结果保存到了self上,比如self.last_img,那样每个epoch都会累积引用,显存只会涨不会降。如果还找不到,就用nvidia-smi dmon实时盯一下显存,配合在代码里手动插torch.cuda.empty_cache()和print(torch.cuda.memory_allocated())二分法排查,虽然土但有效。
之前也遇到类似问题,试过torch.cuda.memory._record_memory_history(),能按操作记录分配栈,但一样很乱。后来发现最土的办法反而好用:把transform逐个注释掉跑一小段,看到哪个停了显存不涨基本就锁定了。
另外也可能不是泄漏,是数据增强把显存峰值顶上去了,比如随机裁剪加翻转多了一次张量拷贝。建议把batch size临时调成1跑一遍,如果还爆就是单样本的问题,范围能小很多。
顺便说一句,自定义Dataset里如果用list存了所有增强后的图,记得清掉中间变量,有时候是Python层引用没释放,不是CUDA的锅。
我之前也遇到过类似情况,后来发现是transform里开了太多线程,每个worker都预加载数据占显存。你可以试试把DataLoader的num_workers调成0,如果显存正常了就是这个问题。另外torch.cuda.memory_snapshot()能按张量看分配记录,配合pytorch的record_stream或者用torch.autograd.detect_anomaly()先排除梯度问题,比看summary直观些。实在不行就二分法注释代码块,虽然笨但最有效。
我之前也遇到过类似情况,加了几个transform后显存直接翻倍,后来发现是某个自定义Dataset里把整个图像列表都load进内存了,根本没释放。你可以试试用torch.autograd.detect_anomaly(),虽然不能精确到行,但能帮你定位到反向传播里的问题,比memory_summary直观多了。另外推荐用nvidia-smi dmon实时盯着显存曲线,配合在代码里手动插print看哪一步增长,比瞎猜快。还有个小技巧,把batch_size调成1跑几个step,如果显存还是涨,基本就是数据加载或缓存的问题了。
我之前也遇到过这种加了transform之后显存暴涨的情况,大概率不是泄漏,而是数据增强里某些操作(比如随机裁剪或缩放)在CPU上生成了一大堆中间tensor,没及时释放。建议先试试在训练循环里用torch.cuda.synchronize()配合nvidia-smi看峰值,或者用pytorch的autograd.detect_anomaly(),虽然慢但能定位到反向传播时的爆点。另外可以查一下是不是Dataset的__getitem__里存了全局变量,或者用了list保存增强结果,这种就纯是逻辑问题了。
我之前也遇到过类似情况,加了几个augmentation之后显存直接翻倍,后来发现是某个transform里用了detach().cpu()但没转回cuda,导致数据在CPU和GPU之间来回拷贝,累积下来就爆了。你可以试试给每个transform加个简单的计数器,或者在Dataset的__getitem__里打印一下返回tensor的device和shape,先确认是不是数据增强阶段产生的中间变量没释放。另外,torch.cuda.memory_summary()确实太乱了,我建议你用torch.profiler,它能按操作符和调用栈分配显存,虽然不能精确到行,但能看出是哪个模块分配了最大内存,比如是卷积还是某个自定义函数。还有个土办法,就是用二分法注释代码,先用固定seed把数据增强全关掉跑一个epoch,确认基线显存,然后逐个开transform,每次跑几十个step看显存曲线,这样能快速锁定是哪个操作。最后,如果怀疑是Dataset里的缓存问题,检查一下是不是把整个图像列表存进了内存,或者用了Python的list而不是tensor,有时候list里存了太多numpy数组也会让显存看起来暴涨。