最近在调一个图像分割模型(DeepLabV3+),用RTX 3090训练,batch size设到4都直接OOM,但看别的项目同样的卡能跑到16。我检查了输入图像尺寸(512x512),也试了梯度累积,还是崩。
目前怀疑是不是自己写的DataLoader里做了什么骚操作,或者模型里某些层占用了大量中间变量?想问问大家一般怎么定位这种“显存泄漏”或者“缓存堆积”的问题?有没有工具或代码技巧能快速看到每层的显存占用?谢谢各位大佬。
PyTorch模型训练时显存爆炸,但batch size已经很小了,还有什么排查思路?
全部回复
共 145 条我之前也踩过类似的坑,而且也是DeepLabV3+,最后发现问题不在输入尺寸,而是ASPP模块里的空洞卷积并行分支,每个分支的中间特征图都会保留梯度,显存是成倍往上叠的。你可以试一下用torch.utils.checkpoint把几个重计算点包起来,能省不少,但会牺牲一点速度。排查工具的话,torch.profiler或者直接给model挂一个forward hook去打印每层输出的shape和占用,比nvidia-smi直观得多。另外建议检查一下DataLoader里的collate_fn,我之前因为把图片转成list再拼接,导致额外复制了一份tensor到显存,后来改成torch.stack就解决了。还有个容易忽略的点,优化器如果是Adam,它的状态缓存也会吃显存,可以试试换SGD或者LAMB对比一下,排除是模型还是优化器的问题。最后实在不行,试试把backbone换成MobileNetV3之类的轻量版,或者干脆用混合精度训练,3090的bf16支持很好,通常能压一半以上。
之前遇到过类似情况,最后发现是backbone的pretrained权重没冻结,BN层在训练时统计量更新导致中间激活翻倍,你试试把BN换成GN或者冻结前几层看看。另外排查显存别光看batch size,用torch.cuda.max_memory_allocated()分段打点,配合nvidia-smi看显存曲线,能明显看出是前向还是反向爆的。如果DataLoader里有多进程num_workers太高,也可能因为共享内存堆积拖累显存,先降到2试试。
试试关掉gradient checkpointing再开,有些模型和显存优化器会互相冲突,我之前就这样找到的锅。
我之前也遇到过类似情况,batch size调小还爆显存多半不是数据加载的问题,而是模型中间激活值太大。你可以试试用torch.cuda.memory_summary()看下内存分配,或者装个pytorch_memlab,能逐层打印显存占用,定位到具体是哪层爆的。另外检查下是不是开了gradient checkpointing但没生效,或者某些层用了大卷积核导致的临时张量堆积。还有个野路子:把输入图像先降到256试试,排除是输入尺寸的问题,然后再二分排查模型结构。
我猜你DataLoader里可能做了些奇怪的transform,比如在GPU上做数据增强,或者把整个batch的tensor都留在显存里没释放。建议先用torch.profiler看看内存峰值出现在哪个阶段,顺便查下是不是用了类内缓存没清,比如自定义的池化层或者注意力机制里存了太多中间结果。另外3090显存40G,512的图跑DeepLabV3+按理说4的batch不该爆,说不定是某个算子的内部实现有bug,换个backbone试试?
巧了,我之前调一个UNet也这样,最后发现是BN层的running stats在反向传播时被重复计算,导致梯度图越攒越大。你可以试试把batch size改成1,如果还OOM那基本就是模型结构问题了,用torchsummary看下每层的输出shape
我之前也踩过类似的坑,排查思路其实可以分两步走:先别急着改代码,用torch.cuda.max_memory_allocated(0)配合那段逐步跑模型各模块的代码,把每个block的显存峰值打出来,这样能直接看出是不是某个模块(比如ASPP或者decoder)在反向传播时保留了太多中间变量。另外,你提到DataLoader,建议检查下num_workers是不是设太高了,有时候多进程会额外复制显存缓存,还有pin_memory=True在3090上偶尔会有奇怪的内存堆积,可以试着关掉对比下。我上次就是靠torchinfo加这个逐步排查法,发现是自定义的辅助损失函数里多算了一次全尺寸特征图,删掉就省了快3G。
用torch.cuda.memory_summary()看看是不是数据加载时把图堆在显存里了,顺便查下backbone的BN层有没有开eval。
试试nvidia-smi看显存是不是被别的进程占了,我之前遇到过缓存被锁住不释放的情况。
我之前也遇到过类似情况,batch size调小反而更崩,后来发现是模型里用了太多中间变量没释放,尤其是DeepLabV3+的ASPP那块,可以试试torch.cuda.max_memory_allocated()分段打印,或者直接hook一下每层的输出尺寸,看看是不是哪层突然爆了。另外DataLoader里如果做了多进程预处理,偶尔会有缓存没清干净的问题,可以先把num_workers设0跑一次排除掉这个因素。建议你从简单的输入直接前向一次,逐层打印显存,比瞎猜快很多。
用torch.cuda.memory_summary()看下缓存分配,多半是中间变量没释放,查查backbone的梯度图。
试试关掉算梯度那层的保留图,或者把图片尺寸缩到256跑一次对比,基本就能定位是模型还是数据的问题。
试试关掉gradient checkpointing再看,或者用torch.cuda.memory_summary()看下哪层峰值最大。
用torch.profiler跑一下,能直接看到每个op的显存占用,比瞎猜快多了。
试试用torch.cuda.set_per_process_memory_fraction限制显存,再配合nvidia-smi看峰值,大概率是中间变量没释放。
我之前也遇到过类似情况,最后发现是模型里用了太多中间变量没释放,尤其是DeepLabV3+的ASPP那块,建议你用torch.cuda.memory_summary()看下每个张量的分配情况,或者用nvidia-smi -l 1实时盯着。另外检查下DataLoader里有没有在GPU上做数据增强,有些操作会偷偷缓存大量临时结果。我之前就是因为在collate_fn里不小心把整张图都搬到了cuda上,导致显存直接翻倍。
我之前也遇到过类似情况,最后发现是模型里用了太多中间变量没释放,尤其是Decoder那块。你可以试试在训练循环里加torch.cuda.empty_cache()看有没有缓解,但更重要的是用torchsummary或者直接hook每层输出,打个log看哪一层峰值最高。另外检查下DataLoader有没有把不需要的tensor也搬到GPU上,collate_fn里有时候会莫名存了整张图。你这3090跑512输入按理说很宽裕,八成是某个操作把batch维度复制了,比如one-hot或者索引展开之类的。
我之前也遇到过类似情况,排查下来发现是输入尺寸没对齐,模型下采样后特征图维度异常导致中间变量爆炸,你可以先打印每层的输出shape看看。另外推荐用torch.cuda.memory_snapshot()结合pytorch的memory profiler,能直接看到每个tensor的分配情况,比瞎猜快很多。还有个小技巧,把gradient_checkpointing打开试试,虽然会慢点但能显著省显存。你检查过BN层的momentum或者模型里有没有用float64吗?有时候这些细节也会让显存占用翻倍。
试试关掉梯度和输入张量的requires_grad,或者用torch.cuda.memory_summary()看下是不是某些层的缓存没释放。
可以用torch.profiler或者hooks打印每层输出尺寸,多半是中间特征图太大,检查下ASPP里的空洞卷积叠加。
我之前也遇到过类似情况,最后发现是backbone的BN层在训练时统计梯度导致中间变量没释放,你试试把model的train模式改成eval跑一下前向,如果显存降了基本就是这块。排查的话可以用pytorch的torch.cuda.set_per_process_memory_fraction配合nvidia-smi看实时占用,或者直接把batch size设成1看单样本占用,再逐步加层定位。还有个土办法,在dataloader里把pin_memory关掉,有时候多进程预取也会堆显存。
先用torch.cuda.set_per_process_memory_fraction限制显存跑一波,再配合nvidia-smi看趋势,基本能定位是缓存堆积还是单步峰值。
另外查下DataLoader里有没有对tensor做detach或clone,或者模型里用了大kernel的dilation,3090的24G不该这么脆。
先用torch.cuda.memory_summary()看下峰值在哪,大概率是中间变量没释放,或者输入尺寸没你想象的那么小。
我之前也踩过类似的坑,排查方向其实可以先不用猜DataLoader,直接给模型喂一个固定的假tensor跑一次forward和backward,看显存曲线就能排除数据侧的问题。另外你试过把输入尺寸降到256看看吗?如果显存立刻降下来,那大概率是某些层(比如ASPP或者decoder)对分辨率特别敏感。工具的话,pytorch自带torch.cuda.memory_snapshot()可以看分配块,或者用nvidia-smi dmon实时盯显存变化,比手动打印好使。还有个冷门技巧,把batch size设成1跑一遍,如果还是OOM,那就是模型本身有静态显存占用异常,跟batch没关系了。
试试在训练循环里加torch.cuda.reset_peak_memory_stats()然后打印各层输出,多半是中间变量没释放。还有检查一下BN层或自定义forward里有没有不小心保存了大tensor。
用torch.cuda.max_memory_allocated和summary看下每层缓存,多半是中间变量没释放,查下backbone的梯度。
可以试试把batch size设成1跑一遍,再用nvidia-smi监控显存曲线,能快速区分是数据加载还是模型本身的问题。