最近在调一个图像分割模型(DeepLabV3+),用RTX 3090训练,batch size设到4都直接OOM,但看别的项目同样的卡能跑到16。我检查了输入图像尺寸(512x512),也试了梯度累积,还是崩。
目前怀疑是不是自己写的DataLoader里做了什么骚操作,或者模型里某些层占用了大量中间变量?想问问大家一般怎么定位这种“显存泄漏”或者“缓存堆积”的问题?有没有工具或代码技巧能快速看到每层的显存占用?谢谢各位大佬。
PyTorch模型训练时显存爆炸,但batch size已经很小了,还有什么排查思路?
全部回复
共 145 条试试用torch.cuda.memory_summary()看分配细节,或者hook一下中间变量,可能某些层保留了太多计算图。
我个人也觉得这问题挺常见的,很多时候不是batch size的锅,而是模型里某些算子或者中间变量把显存给吃了。比如DeepLabV3+里的ASPP模块,尤其是空洞卷积和上采样层,如果框架没做inplace优化,很容易产生大量临时张量。你可以试试用torch.cuda.max_memory_allocated()和torch.cuda.memory_summary()来抓一下具体哪个操作后显存飙升,这个比肉眼猜靠谱多了。
另外建议检查下DataLoader里的num_workers,如果设得太大,Python的多进程会额外复制一些缓存,特别是图像增强用了albumentations这类库时,容易把pin_memory的缓冲池撑爆。我自己遇到过类似情况,明明batch size只有2,但DataLoader prefetch factor设高了直接爆显存。
还有个小技巧:把输入图像缩到256x256先跑一遍,如果还爆,基本就是模型结构或损失函数里有显存泄漏;如果正常,那八成是输入尺寸或数据预处理堆了太多中间变量。梯度累积本身不减少峰值显存,只能平滑占用,所以排查时先关掉它。
工具方面,PyTorch的torchinfo(原torchsummary)可以看每层参数量和输出尺寸,但显存占用建议用torch.cuda.set_per_process_memory_fraction做边界测试,或者直接用nvidia-smi加watch实时盯。如果还找不到,可以考虑把模型里可疑的模块(比如ASPP)单独拎出来跑一次前向,看是不是某个子模块有bug。
我遇到过类似的情况,分享一个排查思路:试试用torch.cuda.max_memory_allocated()在每步训练后打峰值显存,或者用torch.autograd.profiler测一下各层的内存消耗,DeeplabV3+的ASPP模块里空洞卷积对中间变量很敏感。另外建议检查DataLoader的num_workers,设太高会导致数据预加载时显存堆积,还有输入图像有没有偷偷存了梯度或没用detach的变量。
这问题我也踩过坑,3090跑512x512按理说不至于batch size 4就炸,除非你模型里塞了什么奇怪的模块。建议先试试用torch.cuda.memory_summary()看下显存分配,能直接看到是参数、梯度还是中间变量占大头。如果中间变量爆了,多半是DataLoader里数据预处理时没释放,或者模型前向传播时某些层(比如ASPP里的空洞卷积)产生了大量临时特征图——可以用torch.no_grad()逐层跑一次,对比激活前后的显存变化。另外,检查下有没有不小心把验证集的梯度也保留了,或者多进程DataLoader的worker数设太高导致缓存堆积,我上次就是worker设了8又没加prefetch_factor限制,直接多吃了2G。还有个骚操作:把输入切成patch跑一次,如果显存正常下降,那大概率是模型结构层面的问题,可以考虑用checkpointing或梯度检查点来换空间。
试试用torch.cuda.memory_summary(),能直接看到每层的显存占用,另外检查下有没有梯度回传时保留中间变量的操作。
同用3090搞分割的路过,你这个情况我太熟了。除了图像尺寸,建议你检查一下输入有没有做奇怪的transform,比如ToTensor之后又手动加了一些归一化或者padding,那中间变量会多出好几倍。我之前就是DataLoader里顺手做了个随机crop但忘了释放旧tensor,结果显存直接翻倍。
还有一个杀手锏是torch.cuda.memory_summary(),跑完一个batch之后调用它能列出每个CUDA分配点的内存占用,特别适合抓那些意料之外的中间变量。另外你可以试试用torch.no_grad()包装一下验证阶段,有时候验证集也开梯度就会疯狂累积计算图。
对了,DeepLabV3+的ASPP模块里空洞卷积的dilation rate如果设得太大,虽然参数没涨但特征图的内存占用会非线性增加。你可以把backbone换成ResNet50试试看是不是解码器的问题,或者干脆用torch.jit.script把模型trace一下再跑,有时候动态图会留下一些临时buffer。梯度累积只是缓解,根源还得看哪个层在吃显存。
用torch.cuda.max_memory_allocated()分段打印一下,大概率是中间特征图缓存没释放,或者验证集也在算梯度。
确实,这种问题挺折磨人的,batch size 4都爆显存肯定不正常。我建议先排除下DataLoader的坑,比如是不是在__getitem__里把图像或者mask转成了float32然后又没释放,或者用了太大的数据增强缓存(比如albumentations的返回值没及时清理)。你可以试试直接跑一个空循环,每次只加载一个batch但不做前向传播,看显存会不会持续上涨,这样就能快速定位是不是数据加载的问题。
另外模型本身也可能有隐藏的显存大户,尤其是DeepLabV3+里的ASPP模块,空洞卷积的中间特征图如果没被及时释放,叠加起来会很恐怖。推荐用torch.cuda.memory_summary()打印详细的分配日志,或者用torch.cuda.memory_allocated()在训练循环里每步输出对比,看哪一步突然暴涨。也可以试试把模型里的BatchNorm换成SyncBN或者GroupNorm,有时候BN的统计变量在分布式或大batch设置下会偷偷多占。
还有个骚操作就是手动用torch.no_grad()包装一下推理时不需要梯度的层,比如backbone的前几层,或者把一些中间变量del掉再主动调torch.cuda.empty_cache(),虽然治标不治本但能帮你验证是不是缓存堆积。如果还是搞不定,建议用nvidia-smi -l 1实时监控显存变化,同时对比别人的开源代码,看看是不是你的输入尺寸虽然写了512x512,但实际在预处理阶段被resize成更大尺寸了(比如某些增强库默认会padding到512x512但实际中间过程用了更大图)。
我之前也遇到过类似的情况,后来发现是DataLoader里用了太多的pin_memory,加上num_workers开太大,导致内存碎片化间接影响了显存分配。建议你先用torch.cuda.memory_summary()打印一下当前显存快照,看看到底是哪个操作在爆,或者试试在forward里加个torch.cuda.empty_cache()观察变化。另外注意一下DeepLabV3+的ASPP模块里空洞卷积的中间变量是否被保留了下来,有时候用with torch.no_grad()包一下推理部分能省不少。
试试用torch.cuda.max_memory_allocated和summary工具看每层显存,我上次就是被中间特征图缓存坑的。
我遇到过类似问题,排查时发现是DataLoader的num_workers开太多导致显存被预分配占满,降到4就正常了。另外可以试试torch.cuda.memory_summary(),这工具能直接打印每层的显存占用,定位到具体是哪个op在吃显存。还有看看有没有不小心把验证集的梯度也保留了,关掉torch.no_grad能省不少。
我遇到过类似的情况,建议你用torch.cuda.memory_summary()直接看显存分配,能清楚看到是模型参数还是中间变量占了大头。另外检查下DataLoader里有没有把整张图或者不必要的缓存留在计算图里,比如意外的detach或clone操作。还有DeepLabV3+的ASPP模块和骨干网络的后几层其实挺吃内存的,可以试试用checkpointing换空间,或者调低输出stride。
试试用torch.cuda.memory_summary()看显存分配,再排查下DataLoader里有没有多余的变量没释放。
我之前也遇到过类似情况,后来发现是模型里的ASPP模块里空洞卷积的并行分支把中间变量堆太多了。你可以试试用torch.cuda.memory_summary()打印显存快照,或者开一下torch.backends.cudnn.deterministic看看是不是缓存没释放。另外检查下DataLoader的num_workers和pin_memory,有时候多进程预加载也会偷偷占显存。
试试torch.cuda.memory_summary(),能直接看到每层的缓存占用,另外检查下DataLoader里有没有重复加载图像没释放。
这种问题我踩过好多次坑,建议先用torch.cuda.memory_summary()看下峰值显存在哪一步爆的,顺便检查下DataLoader里是不是有num_workers开太多或者pin_memory=True导致缓存没释放。另外DeepLabV3+的ASPP模块里空洞卷积对显存消耗挺大的,可以试试把输出stride调成16或者用torch.no_grad()包住不需要梯度的部分。
我遇到过类似的情况,后来发现是DataLoader里的num_workers设得太高,或者pin_memory=True导致缓存堆积,可以先调成0试试。另外推荐用torch.cuda.memory_summary()看显存分配,或者装个pytorch_memlab逐层跟踪,我上次就是靠它发现ASPP模块里有个膨胀卷积的中间变量没释放。顺便问下,你的backbone是不是用了预训练权重?有时候加载方式不对也会多占显存。
试试用torch.cuda.memory_summary()看具体哪块儿吃显存,另外检查下DataLoader里num_workers是不是开太多导致缓存堆积。
检查一下你的DataLoader里是不是开了过多的num_workers或者pin_memory,有时候这些设置会额外缓存不少显存。另外可以试试torch.cuda.memory_summary(),它能打印出每个CUDA tensor的分配情况,排查哪层占得多。我之前遇到过相似问题,最后发现是损失函数里有个不必要的detach()操作导致了内存堆积。
试试用torch.cuda.memory_summary()看下分配细节,能抓到是不是哪层没释放。另外检查下DataLoader的num_workers,有时候开太多worker会在显存里缓存一堆数据。我上次就是被一个没用with torch.no_grad()的验证步骤坑了,前向传播的中间变量一直攒着。