最近在调一个图像分割模型(DeepLabV3+),用RTX 3090训练,batch size设到4都直接OOM,但看别的项目同样的卡能跑到16。我检查了输入图像尺寸(512x512),也试了梯度累积,还是崩。
目前怀疑是不是自己写的DataLoader里做了什么骚操作,或者模型里某些层占用了大量中间变量?想问问大家一般怎么定位这种“显存泄漏”或者“缓存堆积”的问题?有没有工具或代码技巧能快速看到每层的显存占用?谢谢各位大佬。
PyTorch模型训练时显存爆炸,但batch size已经很小了,还有什么排查思路?
全部回复
共 145 条试试用torch.cuda.set_per_process_memory_fraction限制最大显存,再逐层hook打印输出尺寸,很快能定位到哪层在吃显存。
我最近也踩过类似的坑,建议你试试用torch.cuda.memory_summary()直接打印显存快照,能看清是哪个操作在爆。另外别忘了检查一下DataLoader里num_workers设了多少,有时候多进程会复制额外显存,还有模型里的BN层或者中间特征图保存太多也可能导致缓存堆积。
试试torch.cuda.memory_summary(),能直接看到每层的显存占用,另外检查下DataLoader里有没有重复缓存。
你这个情况我遇到过类似的,大概率不是batch size的问题,而是中间变量没释放。我建议你先用torch.cuda.max_memory_allocated()配合torch.cuda.reset_peak_memory_stats()在每个batch前后打印一下显存峰值,看看是不是逐次递增。如果每次迭代后显存不回落,那就是有变量没释放,常见坑包括DataLoader里用到了num_workers>0但没设置pin_memory=True,或者自定义的collate_fn里不小心把tensor挂到了全局变量上。
另外可以试试把模型切成几段,用torch.no_grad()包住前向传播的中间部分,或者用torch.utils.checkpoint.checkpoint来交换显存和计算。我上次调一个UNet++也是类似情况,最后发现是空洞卷积的中间特征图太大,而且PyTorch默认会保留所有中间变量用于反向传播。你可以用torch.profiler.profile来查看每层的memory_usage,或者直接用nvidia-smi配合gpustat实时监控,但更推荐用torchinfo或torchsummary直接把每层参数量和输出tensor尺寸打印出来。
还有,检查一下输入图像是不是真的被resize到了512x512,有时候DataLoader里忘了做变换,实际输入还是原图大小。如果以上都排查了还崩,可以试试torch.backends.cudnn.benchmark=False,虽然会慢一点,但有时候能避免缓存碎片导致的OOM。
这种情况我遇到过类似的,建议先试试torch.cuda.memory_summary()看完整显存快照,能直接定位到是哪一层爆的。另外可以检查下DataLoader的num_workers,设太高有时候会预加载过多数据到显存。还有DeepLabV3+的ASPP模块里空洞卷积的中间特征图其实挺吃显存的,可以尝试把输出通道降一点或者用checkpointing换空间。
这个我熟,之前用DeepLabV3+也踩过类似的坑。可以试试torch.cuda.memory_summary()看详细分配,或者用torch.autograd.set_detect_anomaly(True)配合CUDA_LAUNCH_BLOCKING=1跑一下,能定位到具体哪行爆的。另外检查下有没有把验证集的梯度也保留了,或者用了太大的num_workers导致内存溢出。
这个情况我太熟悉了,3090其实显存不算小,batch size 4都炸肯定有隐藏问题。除了DataLoader,建议你直接用torch.cuda.memory_summary()打印一下分配细节,能快速看到是哪个操作占了大头——有时候是损失函数或者后处理里悄悄创建了大tensor没释放。另外DeepLabV3+的ASPP模块里空洞卷积的中间特征图如果没做好梯度检查点,显存堆积会很夸张,可以试试用torch.utils.checkpoint把瓶颈层包起来。还有个小技巧:把输入图像换成随机噪声跑一遍前向,排除数据读取的锅,如果还炸就逐层hook看激活值形状。我之前遇到过是自定义的Dice Loss里把预测图和标签做了one-hot广播,那个临时tensor巨大,改写法就解决了。你可以在训练循环里加个torch.cuda.empty_cache()测一下是不是碎片问题,不过治标不治本。
试试用torch.cuda.memory_summary()看显存分配细节,能直接定位到哪个op在吃显存。另外检查下DataLoader里有没有把图像转成float64或者意外地存了多份副本,我之前就因为一个unsqueeze操作让中间变量翻倍了。还有DeepLabV3+的ASPP模块里空洞卷积的dilation rate大会产生大量临时张量,可以考虑用torch.no_grad()包住不参与反向传播的部分。
试试用torch.cuda.memory_summary()看分配细节,或者跑一小段关掉梯度检查每层的显存占用。
这种情况可以试试用torch.cuda.memory_summary()看下显存分配细节,或者把模型每层的input/output size打印出来排查。我之前遇到过类似问题,结果是ASPP模块里的空洞卷积扩张率太大导致中间特征图膨胀,调小一点就正常了。另外检查下DataLoader有没有在__getitem__里把图像多次复制或者缓存了不必要的变量,有时候问题出在预处理环节。
我最近也踩过类似的坑,建议先跑个简单的dummy input试试,如果还爆就是模型本身的问题,大概率是ASPP或者解码器里的空洞卷积和跳跃连接搞出来的中间变量。另外可以试试torch.cuda.memory_summary(),能直接看到每个操作占了多少显存,我上次就发现是某个上采样层没优化好。还有个偏方是关掉cudnn.benchmark,有时候自动调优会缓存一堆kernel导致显存慢慢涨上去。
我遇到过类似情况,后来发现是DataLoader里用了太多transform操作,尤其是随机裁剪和翻转带了额外的内存开销。建议你用torch.cuda.max_memory_allocated()看看峰值在哪,或者装个pytorch-memlab逐层打印显存占用,能快速定位到问题层。另外检查下模型里有没有用nn.BatchNorm但开了track_running_stats=False,那个也会让中间变量堆积。
我最近也遇到过类似问题,后来发现是DataLoader里用了太多的transform或者拼图操作,导致每个batch的中间变量没及时释放。建议先用torch.cuda.memory_summary()看下显存分配,或者用nvidia-smi配合py-spy实时追踪内存变化。另外检查下模型里有没有用了大kernel的ASPP或者多余的梯度缓存,有时候删几个hook就能省好多显存。
我之前也踩过类似的坑,3090跑DeepLabV3+按理说512分辨率不该这么惨,你试试把torch.no_grad()包在验证集外面,有时候是验证阶段反向传播的梯度没释放。另外检查一下backbone是不是用了预训练权重但没冻结BN,BatchNorm在训练模式下会缓存大量的running_mean和running_var,这个在显存里占得挺隐蔽的。
定位每层显存占用的话,推荐用torch.cuda.memory_summary()看峰值分配,或者干脆在forward里插几个torch.cuda.reset_peak_memory_stats()然后逐段打印当前占用,能明显看到是哪个模块暴涨。还有个更直接的骚操作:把batch size设成1跑一次,如果显存占用还是接近峰值,那基本就是模型结构或者DataLoader的问题,跟batch大小无关。
我怀疑你DataLoader里是不是用了pin_memory=True加上num_workers开太多,每个worker都会预加载一批图像到锁页内存,虽然不直接占显存,但会挤占CPU内存导致页面交换变慢。更常见的是你用了transforms里的随机缩放,如果尺寸不是固定512,实际输入可能变成1024甚至更大,那显存直接翻四倍。建议在dataset里打印一下每张图的tensor shape确认。
最后提个冷门的:检查一下有没有用torch.backends.cudnn.benchmark = True,这个在某些操作下会缓存多个算法的工作空间,比如空洞卷积,显存占个几百MB很正常,如果实在找不到原因就把它关掉试试。
试试torch.cuda.max_memory_allocated()打点,配合watch -n看显存曲线,多半是中间激活值没释放。
用pytorch的autograd记录图,排查下是不是有变量被意外retain了,或者试试torch.utils.bottleneck跑一遍。
试试关掉cudnn.benchmark和梯度检查点,另外用torch.cuda.memory_snapshot看下是不是某些算子缓存没释放。
我上次遇到类似问题,查出来是模型里一个自定义激活函数保留了整张特征图,换成inplace操作直接省了3个G。
我之前也遇到过类似的情况,batch size掉到2都爆,结果发现是输入尺寸没对齐,模型里的下采样倍数导致特征图在某个中间层突然翻倍了。你可以先用torch.cuda.set_per_process_memory_fraction限制一下显存,让程序在OOM前报错,然后配合torch.autograd.detect_anomaly去定位是哪一行loss反传时炸的。另外,检查一下是不是用了太多不必要的中间变量,比如在forward里保存了每个层的输出用于loss,这些都会堆积显存,能用del释放的就尽量释放。工具方面,torch.profiler挺好用的,能按时间线和内存分配排序,直接看到每个op的memory usage,比手动猜靠谱多了。还有个小坑,DataLoader里如果worker数开太多,每个worker会复制一份模型参数,也可能吃显存,虽然通常影响不大,但可以试试num_workers=0对比一下。最后,我怀疑你的梯度累积是不是真的生效了,有时候optimizer.zero_grad()的位置不对,或者累积步数没跑完就更新,会导致显存不降反升,这个可以打日志确认下。
之前也踩过类似的坑,3090跑DeepLabV3+按理说512输入完全没压力,先别急着怀疑数据加载,你试试把模型切到eval模式跑一个batch,如果显存还是爆那基本就是模型本身的问题。我遇到过最坑的是ASPP里那个并行的空洞卷积,每个分支都保留完整特征图,加上backbone的中间变量,峰值显存能翻好几倍。建议用torch.cuda.max_memory_allocated()在每一步之后打出来,对比前后差值就能看出来是哪个模块在疯狂吃显存,比手动猜靠谱得多。另外你的DataLoader里如果用了多进程并且开了pin_memory,有时候内存和显存会互相挤占,可以先把num_workers设成0试试。还有一个偏方是检查一下是不是loss或者metric里把整张图的feature detach之后又复制了一份,这种隐式操作特别容易忽略。如果实在找不到,可以试试torch.utils.checkpoint把某些重计算模块换成checkpoint,虽然慢一点但能省一大截峰值。最后提醒一下,3090的显存是24G,但如果你开着tensorboard或者别的进程也占了显存,那也会影响判断,最好拿nvidia-smi确认一下是不是真的就你一个进程在跑。
我之前也踩过类似的坑,建议先别急着怀疑DataLoader,直接上torch.cuda.set_per_process_memory_fraction配合nvidia-smi看实时曲线,或者用torch.profiler看每个op的显存分配。另外DeepLabV3+的ASPP和decoder部分中间特征图挺吃显存的,试试把输出stride调成16或者用混合精度训练,可能一下就能降下来。还有个隐藏点,检查下模型里有没有不小心把输入梯度设成True,那个会存一大堆中间变量。
试试torch.cuda.memory_summary(),能直接看到缓存和分配情况,大概率是中间变量忘detach了。
先用torch.autograd.detect_anomaly跑一遍,再设torch.cuda.empty_cache()看显存曲线,多半是backward时梯度图没释放。