最近在调一个图像分割模型(DeepLabV3+),用RTX 3090训练,batch size设到4都直接OOM,但看别的项目同样的卡能跑到16。我检查了输入图像尺寸(512x512),也试了梯度累积,还是崩。
目前怀疑是不是自己写的DataLoader里做了什么骚操作,或者模型里某些层占用了大量中间变量?想问问大家一般怎么定位这种“显存泄漏”或者“缓存堆积”的问题?有没有工具或代码技巧能快速看到每层的显存占用?谢谢各位大佬。
PyTorch模型训练时显存爆炸,但batch size已经很小了,还有什么排查思路?
全部回复
共 145 条我之前也踩过类似的坑,检查完数据和batch size后才发现是模型里用了太多中间变量,特别是DeepLabV3+的ASPP模块和decoder部分,建议用torch.cuda.memory_stats或者直接hook每层输出看看。另外你的DataLoader里如果对图像做了随机crop或resize,得确认一下实际返回的tensor尺寸是不是真的统一了,有时候多进程worker里缓存了变量也会导致显存峰值。实在不行可以试试torch.utils.checkpoint,把部分层换成梯度检查点,能省不少显存。
我之前也遇到过类似情况,排查下来发现是输入图像没做归一化,直接以uint8喂进模型,导致自动转float时临时张量暴增。你可以试试用torch.cuda.memory_snapshot或nvidia-smi的进程监控对比前后变化,另外检查下forward里有没有重复创建大tensor或者用了不释放的list存中间结果。还有个小技巧,把batch size设成1跑一次,如果显存还是异常高,那基本就是模型结构或数据加载的问题,跟batch关系不大了。
我之前也踩过类似的坑,3090跑DeepLabV3+按理说512分辨率不该这么惨。你先别急着怀疑DataLoader,试试用torch.cuda.reset_peak_memory_stats()和torch.cuda.max_memory_allocated()在每步前后打点,看看峰值到底出现在前向还是反向。我遇到过最隐蔽的问题是模型里用了自定义的注意力或者上采样层,内部创建了临时张量没释放,尤其是那种带循环或者多次调用的子模块,显存会像滚雪球一样涨。另外检查一下有没有不小心把梯度传给了不需要更新的参数,或者用了requires_grad=True的buffer,这会导致反向时保存大量中间激活。工具方面,torch.profiler自带内存分析,能按行号看每行代码的分配量,比手动猜快多了。还有个土办法,把batch size降到1跑通,然后逐步加,同时用nvidia-smi每隔0.1秒记录显存曲线,看是不是线性增长而非突增,如果是线性增长那基本就是某些张量没释放。对了,你试试在dataloader里把num_workers设成0,有时候多进程预取会额外复制CUDA张量,虽然通常不占显存但会影响峰值。真不行就换mixed precision,amp能省将近一半显存,3090对fp16支持很好,代价只是偶尔需要调loss scale。最后提醒下,DeepLabV3+的ASPP模块里空洞卷积的dilation rate大时,feature map的显存占用会超出你预期,可以考虑用深度可分离卷积替换常规卷积来降内存。
我之前也遇到过类似情况,排查下来发现是输入尺寸没对齐,模型里有些下采样层会让特征图尺寸变成非整数,pytorch会自动填充导致显存暴涨,你可以打印一下中间层的shape看看。工具方面推荐用torch.profiler或者pytorch_memlab,能按行显示每层内存占用,直接定位到具体模块。另外检查下dataloader里是不是用了pin_memory和num_workers太高,有时候prefetch机制会额外占不少显存,试试把workers降到0看看还崩不崩。最后如果所有都正常,可以手动用torch.cuda.max_memory_allocated()对比峰值和实际消耗,排除是不是别的程序占着显存没释放。
DataLoader里如果用了num_workers>0还开pin_memory,有时反而会堆积显存,可以先试试设成0排除掉。另外推荐用torch.cuda.memory_summary()看下快照,或者装个torchprof,能按层打印显存峰值。还有个坑是验证阶段忘了加torch.no_grad(),中间变量一直不释放,我踩过好几次。