最近在调一个图像分割模型(DeepLabV3+),用RTX 3090训练,batch size设到4都直接OOM,但看别的项目同样的卡能跑到16。我检查了输入图像尺寸(512x512),也试了梯度累积,还是崩。
目前怀疑是不是自己写的DataLoader里做了什么骚操作,或者模型里某些层占用了大量中间变量?想问问大家一般怎么定位这种“显存泄漏”或者“缓存堆积”的问题?有没有工具或代码技巧能快速看到每层的显存占用?谢谢各位大佬。
PyTorch模型训练时显存爆炸,但batch size已经很小了,还有什么排查思路?
全部回复
共 145 条我之前也踩过类似的坑,建议先别急着怀疑DataLoader,直接上torch.cuda.memory_summary()看下分配细节,大概率是中间特征图或者某些激活值没释放。另外检查下是不是用了多卡或者混合精度没开对,3090跑512输入理论上不该这么惨。有个笨办法:把模型输入换成随机tensor,逐层打印activation的显存占用,用pytorch的torch.cuda.set_per_process_memory_fraction限制一下也能辅助定位。你试试把batchsize调到1跑通,再逐步加,同时监控nvidia-smi的显存曲线,看是线性增长还是突然跳变,这样能快速区分是模型结构问题还是数据加载问题。
我之前也踩过类似的坑,最后发现是模型里用了太多中间变量没释放,尤其是DeepLabV3+的ASPP模块里那几个并行卷积,建议用torch.cuda.memory_summary()看看峰值在哪,或者直接开torch.autograd.detect_anomaly()跑一遍。另外检查下DataLoader是不是在GPU上做了额外处理,比如把标签也搬到cuda后没及时清理,有时候小细节挺坑的。
我之前也遇到过类似情况,batch size调小还是爆,最后发现是backbone的BN层在训练时统计梯度导致中间变量翻倍,建议先试试冻结backbone或者换SyncBN看看。另外推荐用torch.cuda.memory_summary()打印每个tensor的分配情况,或者用pytorch的autograd检测图是否没释放,很多时候是loss或者metric里无意间保留了整个batch的feature map。你那个DataLoader里如果做了多进程的归一化或数据增强,也可能把临时缓存留在显存里,可以试试把num_workers调成0对比一下。
我之前也碰到过类似情况,排查下来发现是模型里用了太多中间特征图拼接,尤其是Decoder部分,显存峰值全被撑起来了。你可以试试用torch.profiler或者pytorch_memlab,能直接看到每层的临时tensor大小,比瞎猜高效很多。另外检查下DataLoader里有没有把图像转成FP32再归一化,或者做了奇怪的切片/翻转,这些操作有时会在GPU上产生隐形的临时缓存,换成在CPU上预处理完再送GPU会好很多。
我之前也遇到过类似的情况,batch size调小反而更崩,后来发现是输入尺寸没对齐。你确认一下模型forward里有没有对特征图做上采样或者插值,有时候某些层会隐式地把中间变量放大好几倍,比如ASPP里的空洞卷积配合不同rate,如果rate设置不当,feature map的显存占用会翻着倍涨。另外建议你用torch.cuda.max_memory_allocated()和torch.cuda.memory_summary()打一下峰值,能直接看到是模型参数还是激活值占了大头,我之前就是用这个发现是backward时梯度累积的问题,因为梯度在反向传播时也会占用额外显存,特别是你用了梯度累积的话,如果没清空梯度或者没正确step,缓存会越积越多。还有个骚操作是试试把输入图像临时改成256x256跑一下,如果显存占用下降比例不对,基本就能锁定是某层对尺寸敏感而不是整体显存不够。你DataLoader里要是用了worker加载,记得开num_workers=0先排除多进程缓存干扰,有时候是数据预处理的临时tensor没释放。最后实在不行,可以用torch.utils.checkpoint把几个重计算开销大的block包起来,用时间换空间,3090跑512输入应该不至于这么紧张。
我之前也遇到过类似情况,最后发现是模型里用了太多中间变量没释放,尤其是DeepLabV3+的ASPP那块,建议你用torch.cuda.set_per_process_memory_fraction先限制下显存,再看哪些层报错。另外,检查下DataLoader里num_workers是不是设太高了,有时候多进程会复制显存缓冲区,虽然看起来不占显存但实际会累积。还可以试试在forward里用torch.cuda.memory_allocated()打印每层前后的差值,定位到具体层,比用工具直观。
我之前也踩过类似的坑,3090跑DeepLabV3+按理说不至于batch 4就崩。你先别急着怀疑DataLoader,我遇到过最坑的是BN的momentum和num_workers设置,有时候多线程加载会复制一份显存缓存,试试点一下num_workers=0看看有没有变化。另外建议直接上torch.cuda.memory_summary(),它会打印每个张量占用的具体内存,比你自己猜快多了。如果这还看不出来,就试试把模型里所有中间变量都detach掉,或者检查一下是不是用了不恰当的转置卷积,某些实现会疯狂吃显存。还有个小技巧,把输入图像临时变成128x128跑一遍,如果显存占用降幅不成比例,那大概率是模型结构里有固定的缓存开销,比如ASPP或者解码器里的某些分支。最后实在不行,开个nvidia-smi的定时监控,看显存是逐步涨还是瞬间爆,逐步涨就是泄漏,瞬间爆就是单步分配太大。我上次就是被一个没用的辅助损失头坑了,它保留了整个特征图的梯度,删掉后直接降到batch 16。
我之前也遇到过类似情况,排查下来发现是模型里用了太多中间特征图做拼接,尤其是DeepLabV3+的ASPP那块,多尺度空洞卷积会堆不少显存。你可以试试用torch.cuda.set_per_process_memory_fraction临时限制一下显存,或者用nvidia-smi dmon实时看显存曲线,哪个步长突然飙高就锁定哪段代码。另外检查下DataLoader里有没有把图像转成FP32或者做了不必要的复制,有时候num_workers开太多也会吃显存。
我之前也踩过类似的坑,排查下来发现是模型里用了太多中间变量没释放,尤其是一些自定义的forward里存了list或tensor没清。建议你先用torch.cuda.set_per_process_memory_fraction限制一下显存,崩的时候能更快定位到具体层。另外可以装个pytorch_memlab或者用torch.autograd.detect_anomaly看反向传播时的异常,不过更直接的办法是逐层打印输出shape,检查是不是哪里隐式放大了feature map。至于DataLoader,确认下num_workers别开太多,有时候multiprocessing的缓存也会占显存,虽然看起来是CPU的事。
试试PyTorch的torch.cuda.memory_snapshot,能按张量看占用量,另外检查下是不是输入输出尺寸没写对导致中间feature map异常大。
之前也遇到过,结果发现是损失函数里有个变量没detach,梯度图一直攒着,加个with torch.no_grad()就好多了。
用torch.cuda.memory_summary()看下缓存分配,重点查下模型里有没有detach()或者no_grad()漏写的中间变量。
八成是backbone输出没释放,试试把batch size调成1跑一遍,再用nvidia-smi对比显存曲线,能快速定位是哪一层在堆积。
我之前也踩过类似的坑,3090跑DeepLabV3+按理说不至于这么惨。你先别急着怀疑DataLoader,我建议直接用torch.cuda.memory_summary()看下分配峰值,大概率会发现是中间激活值爆炸了,尤其是ASPP那几条空洞卷积分支,空洞率不同会成倍放大特征图内存。我当初是用一个三行脚本,把model的每个子模块包一层hook,打印forward前后的allocated memory差值,瞬间定位到是deeplab的assist head里有个转置卷积把显存吃光了。另外检查一下你用的backbone是不是带BN的,如果开了gradient checkpointing,有些层会被重复计算,反而导致峰值更高。还有个容易忽略的点:你用的loss是不是用了one-hot编码或者把mask转成了float32的one-hot,那个在batch=4时可能额外占几百MB。还有个小技巧,把输入图像先切成256x256试跑一次,如果显存占用不是按面积比例下降,就说明有缓存没释放。我上次就是发现DataLoader里num_workers设太大,每个worker的pinned memory在疯狂堆积,降到4就稳了。工具的话,可以试试pytorch的torch.profiler,或者更直观的memray,能按行号显示分配点。最后,如果实在找不到,直接开AMP混合精度,3090对fp16支持很好,显存能砍一半。
试试用torch.profiler看每层显存占用,另外检查下输入有没有被重复复制到GPU,这最常见。
先跑个单step的forward,用nvidia-smi盯显存曲线,排除DataLoader的预取和pin_memory问题。
我之前也踩过类似的坑,排查下来发现是模型里的中间变量没释放,比如在forward里反复拼接feature map会占大量显存。你可以试试在几个关键节点打印torch.cuda.max_memory_allocated(),对比一下前后差值,基本能定位到是哪一层爆的。另外检查下DataLoader的num_workers和pin_memory,有时候多进程预取也会莫名吃显存,虽然看着不像但确实有影响。还有个小技巧,用torch.utils.checkpoint把大模块包起来,能省不少显存,就是慢一点。
我之前也遇到过类似情况,3090跑DeepLabV3+按理说512输入不该这么吃紧。你先别急着怀疑DataLoader,我建议直接用torch.cuda.memory_summary()看内存分配,或者用nvidia-smi盯实时显存曲线,如果发现是阶梯式上涨而不是一次性爆掉,那大概率是缓存没释放。另外查一下是不是开了cudnn.benchmark,有些情况下它会自动找算法导致额外显存开销,关掉试试。还有个隐蔽的坑:你模型里如果用了多尺度推理或者辅助损失,那些中间feature map会在反向传播时全部保留,尤其ASPP模块里的空洞卷积,输出通道多的话很吃显存。可以试着用torch.utils.checkpoint把部分层包起来,用计算换显存,或者直接把输入降到384看看还爆不爆,这样能快速定位是模型结构问题还是数据加载问题。最后,如果你用了分布式训练或者混合精度,检查一下torch.cuda.set_per_process_memory_fraction是不是被谁设了上限,我之前就被同事的脚本坑过。
我之前也遇到过类似情况,后来发现是模型里的BN层在训练时开了track_running_stats导致缓存累积,换成同步BN或者冻结一部分层试试。另外你可以在每个batch后手动调torch.cuda.empty_cache()看有没有缓解,但别依赖它,主要还是得定位是哪个模块的问题。建议用torch.autograd.detect_anomaly()或者hook打印每层输出的显存,我以前靠这个发现是某个上采样层用了过大的中间tensor。还有检查下DataLoader的num_workers,有时候多进程会复制显存句柄,导致看起来像OOM,实际是共享内存爆了。
建议直接上torch.profiler或者nvidia-smi看实时显存变化,重点查BN层和输入图像是否意外被放大。我之前碰到过DataLoader里悄悄做了resize到1024,排查半天才发现。
我之前也踩过类似的坑,排查下来发现是模型里用了太多保持梯度的中间变量,比如辅助损失或者特征图拼接,这些都会让显存峰值暴涨。你可以试试在关键位置插torch.cuda.max_memory_allocated()打点,或者直接用torch.profiler看每层的内存分配,比瞎猜快多了。另外DataLoader里如果开了pin_memory=True且num_workers过高,有时候也会造成额外的显存碎片,可以先把workers调成0跑一次对比下。还有个偏方,把输入图切成patch喂进去,或者用混合精度训练(amp),3090的显存带宽足够,batch size能轻松翻倍。
先用torch.cuda.max_memory_allocated()看峰值,再配合malloc_info逐层打印,大概率是中间激活没释放。
我之前也遇到过类似的,排查下来不是DataLoader的锅,是模型里用了太多的中间特征图保存,尤其是Decoder部分。你可以试试torch.cuda.memory_summary(),能看得很清楚每个张量占多少,另外把batch size降到1跑一次对比下,如果还爆就基本锁定是模型结构的问题了。还有个偏方,用torch.no_grad()把前向推理单独跑一遍,看看是不是梯度计算导致的内存峰值,这个能帮你快速分清楚是缓存堆积还是真的显存不够。