最近在调一个图像分割模型(DeepLabV3+),用RTX 3090训练,batch size设到4都直接OOM,但看别的项目同样的卡能跑到16。我检查了输入图像尺寸(512x512),也试了梯度累积,还是崩。
目前怀疑是不是自己写的DataLoader里做了什么骚操作,或者模型里某些层占用了大量中间变量?想问问大家一般怎么定位这种“显存泄漏”或者“缓存堆积”的问题?有没有工具或代码技巧能快速看到每层的显存占用?谢谢各位大佬。
PyTorch模型训练时显存爆炸,但batch size已经很小了,还有什么排查思路?
全部回复
共 20 条看到这个问题挺有共鸣的,我之前调一个3DUNet的时候也被显存折磨过。你可以试试用torch.cuda.memory_summary()直接看当前分配情况,或者装个pytorch_memlab的hooks,能逐层打印显存占用,排查中间变量非常直观。另外我怀疑是不是在DataLoader里用了太多worker或者开了pin_memory,有时候这玩意儿会把缓存占满,尤其你图像预处理如果用了随机变换之类的,可能会在CPU端积攒很多浮点数据没及时释放。还有一个坑是梯度累积,如果没配合no_grad或者detach,计算图会一直挂着,每个batch的中间变量都不会释放,建议检查一下loss.backward()之后的optimizer.zero_grad()是不是写在了正确位置。另外DeepLabV3+的ASPP模块里空洞卷积的dilation rate如果设得太大,特征图尺寸虽然不变但内存碎片会爆炸,可以试试把输出stride调成16或者16+8的混合模式。你可以先用一个小数据跑一个简化的forward,逐层print tensor的shape,很快就能锁定是哪一层在吃显存。
可以用torch.cuda.memory_summary()看一下内存分配细节,能直观看到哪块在爆。另外DeepLabV3+的ASPP模块里空洞卷积的中间特征图挺吃显存的,试试把输出stride调成16或者8看看?还有检查下DataLoader里有没有把图像转成float64或者做了什么不必要的repeat操作,这些细节很容易忽略。
试试把梯度检查点打开,或者用torch.cuda.memory_summary()看下哪个层的缓存没释放。
我最近也踩过类似的坑,你可以在训练脚本里加个torch.cuda.memory_summary()看看峰值在哪段代码,或者用nvidia-smi盯着显存变化。另外DeepLabV3+的ASPP模块里空洞卷积的中间变量特别吃显存,试一下把输出stride调大或者减少空洞率,可能比折腾DataLoader见效快。
这种情况我也遇到过,检查下DataLoader里有没有把图像转成float64或者做了什么奇怪的数据增强,显存会翻倍。另外可以用torch.cuda.memory_summary()看具体分配,或者试试torch.no_grad()跑一次前向,排除梯度存储的影响。我上次发现是ASPP模块里的空洞卷积缓存太多,换用torch.cuda.empty_cache()手动清理也没用,最后把inplace=True打开才缓解。
这个情况我遇到过类似的,排查下来往往是中间特征图或者辅助损失那块没释放。你可以试试用torch.cuda.memory_summary()打印一下内存快照,看看是不是某些变量没及时detach。另外检查一下DataLoader里有没有用pin_memory=True但num_workers开太多,有时候worker进程堆积也会吃掉额外显存。另外我习惯用nvidia-smi配合gpustat实时看显存波动,比单纯看batch size更直观。
我之前也遇到过类似问题,后来发现是DataLoader里把图像转成float16时没处理好,导致中间变量缓存爆炸。你可以试试用torch.cuda.memory_summary()看具体哪一层在吃显存,或者把模型切成几段跑一次前向,观察峰值内存变化。另外检查下是不是开了梯度检查点或者用了同步BN,有时候这些默认设置会偷偷多占资源。
建议先试试torch.cuda.set_per_process_memory_fraction限制一下显存占比,看是直接崩还是慢慢涨,能区分显存泄漏还是单次分配过大。排查中间变量可以用torch.cuda.memory_summary()或者pytorch的memory_profiler,能看到每个操作分配的显存。另外检查下DataLoader里有没有对tensor做不必要的clone或者detach操作,尤其是自定义collate_fn容易产生缓存堆积。如果用了同步BN或者梯度检查点,也可能增加临时变量开销。
试试用torch.cuda.memory_summary()看显存分配,或者hook一下中间层的梯度保留情况,很多时候是backward时的中间变量没释放。
这个坑我也踩过,除了检查DataLoader里有没有把不需要的变量挂在self上,建议用torch.cuda.memory_summary()看下具体哪块显存爆的,能精确到每层的缓存。另外DeepLabV3+的ASPP模块里空洞卷积如果dilation rate设太大,中间特征图会特别占显存,可以试试把输出stride改成16或者用checkpoint技术换时间省空间。
试试torch.cuda.memory_summary(),能直接看到每层的显存占用,或者把BatchNorm换成SyncBN看看有没有改善。
这问题我太有同感了,之前调一个三阶段检测器也是batch size调到2还爆,后来排查了半天发现是DataLoader里有个to(device)写在了循环外面,导致每次迭代都往显存里塞一份新的缓存,你可以先检查下是不是dataloader的worker数开太高或者pin_memory=True和某些操作冲突了。另外推荐用torch.cuda.memory_summary()打印完整显存快照,或者装个pytorch_memlab库,能画出每层tensor的分配路径,特别适合找那种梯度更新后没释放的中间变量。你既然怀疑模型层的问题,可以试试把输入切成小图跑一次前向,然后逐层打印tensor的shape和requires_grad状态,有时候是某些自定义模块里的detach()用错了导致计算图不释放。还有个小技巧,把optimizer.zero_grad()改成set_to_none=True模式,能省掉显存里的梯度缓存。如果还不行,建议检查下torch.backends.cudnn.benchmark设置,某些场景下它会导致cudnn缓存膨胀到离谱的程度。
我遇到过类似问题,后来发现是DataLoader的num_workers设太高导致显存碎片化,降到4就解决了。你试试在训练循环里用torch.cuda.reset_peak_memory_stats()配合torch.cuda.max_memory_allocated()看峰值,能定位到具体哪一步爆的。另外检查下模型里有没有用nn.BatchNorm之外的自定义层,有些操作会缓存大量中间结果,用torch.no_grad()或del临时变量手动释放一下。
建议装个torchinfo或pytorch_memlab,跑模型前先hook一下每层的forward就能看到显存分配,我上次也是在DeepLabV3+的ASPP模块里发现有个膨胀卷积的中间变量没释放。另外检查下DataLoader是不是开了pin_memory或者num_workers太大,有时候多进程预加载也会把显存堆满,尤其你用3090这种卡,小batch下显存碎片也可能是个问题。
我最近也踩过类似的坑,建议你试试用torch.cuda.memory_summary()看每步显存峰值,能直接定位到哪层爆的。另外检查下DataLoader里是不是开了pin_memory=True但num_workers设太高,这个组合有时候会导致缓存堆积。还有确认下模型里有没有用nn.BatchNorm但没设track_running_stats=False,某些情况下会额外占显存。
我最近也踩过类似的坑,排查下来发现是模型里用了大量的中间特征保存,比如ASPP模块里并行卷积的显存占用比想象中大得多。你可以试试torch.cuda.memory_summary()直接看每块的分配情况,或者用torchinfo打印各层参数量和激活尺寸,很可能有些层输出通道数设得太大了。另外检查下DataLoader里有没有把多余的变量挂到GPU上,比如标签one-hot编码时不小心留在显存里。
我遇到过类似的情况,排查后发现是DataLoader里用了太多自定义的数据增强操作,每个batch都会创建大量中间变量导致显存溢出。建议你用torch.cuda.memory_summary()看总显存分配,或者加个hook监控每层的输出尺寸,有时候是某些层的中间激活值没被释放。另外检查下是不是开了梯度检查点或者cudnn.benchmark,这两个在某些场景下也会爆显存。
试试用torch.cuda.memory_summary()看每步显存变化,重点查下DataLoader里有没有存了额外的计算图。
这个太典型了,我之前用DeepLabV3+也踩过类似的坑。你怀疑DataLoader和中间变量是对的,但我觉得更大概率是ASPP模块或者backbone里的空洞卷积导致的显存碎片化。有个小技巧:用torch.cuda.memory_summary()可以打出完整的显存分配情况,配合nvidia-smi -l 1实时监控,重点看reserved memory和allocated memory的差值。另外你可以试试把模型切成几段,用torch.no_grad()单跑每段看显存峰值,我上次就是这么发现是ASPP里并行的空洞卷积叠加后中间变量没释放。还有一点——检查下你的输入图像是不是真的被预处理成512x512了,有时候DataLoader里忘了做resize,实际读进去的图是原始尺寸,那就直接炸了。梯度累积虽然能缓解,但本质问题还在,建议先把batch size设成1,然后逐层打印每层的output和gradient的显存占用,用torch.cuda.max_memory_allocated()能看到峰值具体在哪一步爆的。如果这些都没问题,那就看看是不是用了什么第三方backbone的预训练权重,有些权重会自带额外的buffer变量。
我个人经验是,DeepLabV3+的ASPP模块和backbone里的空洞卷积确实容易吃显存,尤其是你用了ResNet101的话,中间特征图叠加起来很夸张。建议先用torch.cuda.memory_summary()打印一下内存快照,看看是不是某个特定层在反向传播时缓存了过多梯度。另外,你用的输入是512x512,但有些实现会在Encoder内部做多尺度特征融合,比如把不同空洞率的特征图cat在一起,这个操作如果不及时释放中间变量,显存会瞬间飙高。可以试试在DataLoader里把pin_memory设为False,或者检查一下num_workers是不是设得太大,有时候多进程预加载会导致显存碎片化。还有一个野路子:把torch.backends.cudnn.benchmark设为False,虽然会慢一点,但有时能缓解显存分配不连续的问题。你还可以用torch.utils.checkpoint(梯度检查点)对某些层进行激活重计算,牺牲一点时间换空间。对了,你的模型是用的预训练权重吗?如果是在ImageNet上训的,可以试试冻结backbone的前几层,只训练decoder部分,这样显存占用能降不少。