最近在调一个图像分割模型(DeepLabV3+),用RTX 3090训练,batch size设到4都直接OOM,但看别的项目同样的卡能跑到16。我检查了输入图像尺寸(512x512),也试了梯度累积,还是崩。
目前怀疑是不是自己写的DataLoader里做了什么骚操作,或者模型里某些层占用了大量中间变量?想问问大家一般怎么定位这种“显存泄漏”或者“缓存堆积”的问题?有没有工具或代码技巧能快速看到每层的显存占用?谢谢各位大佬。
PyTorch模型训练时显存爆炸,但batch size已经很小了,还有什么排查思路?
全部回复
共 145 条有个比较快的排查办法,用torch.cuda.memory_snapshot结合pytorch的memory profiler看峰值分配,能直接定位到具体哪一行代码触发了OOM。另外你试试把num_workers调成0跑一次,排除DataLoader里多进程预加载的显存拷贝问题,我之前遇到过归一化层在GPU上做转换导致临时张量堆积的情况。还有个小技巧,把模型切一半用torch.utils.checkpoint重计算激活值,能省不少中间变量,先确认是不是某些层(比如ASPP里的空洞卷积)的缓存问题。
用torch.cuda.set_per_process_memory_fraction配合逐层hook测下峰值,大概率是中间变量没释放。
3090跑4都炸的话检查下输入是不是被重复pad了,或者源码里有没有detach的异常累积。
我之前也遇到过类似情况,最后发现是输入图像没归一化,导致网络前向传播时数值范围异常,中间feature map的精度和显存占用飙升。你可以先用torch.cuda.memory_summary()看下是不是某个特定层(比如ASPP模块)的缓存特别高,同时把dataloader的num_workers设成0试试,排除多进程预加载的干扰。另外检查下代码里有没有detach()或者no_grad()没加对的地方,有时候反向传播会把所有中间节点都保留下来。
我之前也踩过类似的坑,3090显存看着大但真不禁造。你怀疑DataLoader挺有道理,但更大概率是模型本身的问题,DeepLabV3+的ASPP模块和decoder部分会存很多中间feature map,尤其是空洞卷积不同rate的输出都要保留,512x512的输入在这些层上确实挺吃显存。建议你先用torch.cuda.max_memory_allocated()和torch.cuda.memory_summary()看峰值在哪,或者直接跑一个不加载数据的空forward,把batch size调到1,如果还爆就是模型结构问题,不爆再慢慢加数据。另外检查一下你代码里有没有把梯度retain_graph=True或者用了detach()不当导致计算图没释放,我之前就是自定义loss里不小心retain了图,导致每个step的缓存都堆着,最后爆掉。还有个技巧是开gradient_checkpointing,虽然会慢一点但能省不少显存,DeepLabV3+可以用torch.utils.checkpoint包一下中间层。至于工具的话,pytorch的torchprof或者简单点用装饰器在每个block前后打印allocated memory,基本就能定位到具体层了。你那个别的项目能跑到16的,最好对比一下它有没有用混合精度,fp16在3090上省显存特别明显,而且现在apex或者原生amp都很成熟,你试试加上,可能直接就从4跳到12了。
我之前也踩过类似的坑,3090跑DeepLabV3+按理说512输入、batch 4不该爆的,除非你用了空洞卷积+depthwise分离那种变体,或者backbone加载了超大预训练权重。建议先别猜DataLoader,直接跑一个固定随机输入的forward+backward,排除数据预处理和target的干扰,如果这样还爆,那就是模型结构或优化器状态的问题。
定位每层显存的话,pytorch有个torch.cuda.memory_snapshot()可以看分配块,但更直观的是用nvidia-smi配合pynvml库按step打印显存增量,或者试一下torch.profiler,它输出的表里会显示每个op的tensor size和memory usage,能直接看到是哪个残差连接或者ASPP分支在疯狂占缓存。另外检查下是不是把验证集的梯度也保留了,或者用了detach但没清掉计算图。
我怀疑你可能在forward里写了类似x = self.conv(x)之后再存了个x.clone()用于别的分支,这种引用计数不清会导致峰值翻倍。还有个容易被忽略的点:如果你用了sync batchnorm,多卡DDP会额外复制中间激活,单卡也可能因为cudnn.benchmark开启而缓存一些workspace,建议关掉试试。最后实在不行,用torch.utils.checkpoint把backbone包一下,梯度检查点能牺牲一点速度换显存,但至少先确认是不是某个自定义模块的临时变量没释放。
我之前也遇到过类似情况,排查下来发现是模型里的辅助损失头在反向传播时保留了太多中间张量,尤其是DeepLabV3+的ASPP模块里不同膨胀率的卷积输出会叠加存储。建议你用torch.profiler看一下每个op的memory usage,或者直接跑一个假输入过一遍forward和backward,逐步注释掉可疑层来二分定位。另外检查下有没有在loss里不小心对超大tensor做了detach后还保留引用,这也会导致显存不释放。DataLoader那边一般问题不大,除非你把整图或mask都缓存了,可以先换成最简单的TensorDataset排除因素。
我之前也遇到过类似的情况,batch size调到2都爆,后来发现是模型里的中间特征图在反向传播时被完整保存了,尤其是DeepLabV3+的ASPP模块里那几个并行空洞卷积,每个分支的激活都占一份显存。你可以试试在训练循环里用torch.cuda.max_memory_allocated()打一下峰值,再配合torch.cuda.memory_summary()看是哪个环节涨上去的,这比瞎猜快多了。另外检查下DataLoader的num_workers,如果开了很多进程,每个worker都会复制一份模型参数,虽然不直接算显存,但可能和CUDA缓存碎片化叠加,导致分配失败。还有个容易忽略的点,别用.eval()和no_grad()去验证,但训练时把batch里的padding mask或注意力权重也存了,这些零碎的小张量积少成多也很恐怖。说实话,我后来用torch.profiler看每层的memory usage,才发现是损失函数里一个不必要的tensor.repeat操作把特征图放大了几十倍,删掉后直接省了6G。建议你把输入尺寸再降一半试试,比如256x256,如果显存占用不是线性下降,那就是缓存没释放或者有隐藏的引用。另外可以试试torch.utils.checkpoint,把ASPP的中间激活换成重计算,有时候能换来2到3倍的显存余量,代价是速度慢一点。
我之前也撞到过类似的坑,3090跑DeepLabV3+按理说12G显存完全够的。你先别急着怀疑DataLoader,我那次最后发现是ASPP模块里空洞卷积的dilation rate设太大,导致中间特征图的显存分配特别离谱,尤其是不同rate并行算的时候,峰值会叠起来。你可以试试用torch.cuda.max_memory_allocated()在训练循环里打点,配合memory_summary()看每个张量的保留情况,这样能定位到具体是哪一行代码把显存推爆的。另外检查一下是不是开了cudnn.benchmark,有时候它会缓存一些自动调优的workspace,虽然不占显存但会影响峰值判断。还有个骚操作是,把输入尺寸临时改成128跑一个step,如果显存占用还是异常高,那基本就是模型结构或者损失函数里的临时变量问题了,而不是数据维度。梯度累积只能解决单次step的显存峰值,救不了某些算子内部一次性生成的大buffer,你试试用torch.utils.checkpoint把encoder部分包起来,用计算换显存,通常能省一半以上。最后可以看一眼是不是混合精度没开,amp的grad scaler有时候会额外保留fp32的梯度副本,显存不够的时候这个影响还挺明显的。
我之前也遇到过类似情况,最后发现是输入尺寸没对齐,模型里有些下采样倍数没算对,导致中间特征图比预期大好几倍。你可以先试试用torch.profiler或者直接打印每一层的输出shape,排查一下有没有意外膨胀的层。另外检查下是不是在loss里用了什么需要保存大量梯度的操作,比如自定义的mask计算。还有个小技巧,可以开gradient_checkpointing来换显存,虽然慢点但至少能跑起来。
我之前也遇到过类似的情况,排查下来发现是模型里的中间变量没释放,尤其是DeepLabV3+的ASPP模块里并行分支的feature map特别吃显存。你可以试试用torch.cuda.memory_summary()看每个张量的分配情况,或者用nvidia-smi配合py-spy看哪个进程在涨。另外检查下DataLoader里有没有把图像转成float64或者做了不必要的repeat,这些都会让显存翻倍。对了,如果用了BatchNorm,别忘了确认它在训练模式下是否正常,有时候eval模式残留会莫名其妙占缓存。最后实在不行就开gradient checkpointing,虽然慢点但能救急。
这种情况我遇到过,先别急着怀疑DataLoader,大概率是模型前向传播里的中间feature map在搞鬼,尤其是ASPP和Encoder部分,512输入下通道数一大就很夸张。你可以试试用torch.cuda.max_memory_allocated()在每一步前后打点,配合nvidia-smi看显存曲线是平稳还是持续上涨,后者才是真泄漏。另外检查一下有没有在loss里不小心retain_graph=True,或者反向传播后没清空optimizer的梯度缓存。我之前排查时还发现,某些BatchNorm层在训练模式下会缓存所有batch的统计量,如果用了自定义的forward逻辑,很容易踩这个坑。实在不行就用torch.profiler跑一小步,它能直接列出每个op的显存占用,比瞎猜高效多了。
3090跑512输入正常不该这么惨,你先用torch.cuda.max_memory_allocated()和torch.cuda.memory_summary()看下峰值在哪,大概率是中间feature map堆的。另外DeepLabV3+的ASPP和decoder部分如果有自定义forward没包在no_grad里,或者用了不合理的上采样方式,都会让显存翻几倍。我之前遇到过类似问题,最后发现是DataLoader里做了个to(device)导致每个batch都复制了一份模型参数,改成在collate_fn里统一处理就好。你还可以试着把输入降到256跑一遍,如果显存降幅远大于4倍,那基本就是中间变量问题,跟batch size无关。
我之前也遇到过类似情况,排查半天发现是输入图像没归一化,模型里BN层的running stats在反向传播时疯狂申请显存,加上DeepLabV3+的ASPP模块本来就吃中间特征图,试试把图像转成float16或者用torch.no_grad()包住验证阶段,能省不少。另外强烈建议用nvidia-smi -l 1实时盯着显存曲线,如果某个step突然暴涨,多半是DataLoader里做了GPU上的tensor操作没释放,可以试着把collate_fn里的东西全放CPU上。还有个笨办法,把batch size降到1,如果还OOM,就逐层打印model.named_parameters()的requires_grad,配合summary库看每层输出尺寸,基本能锁定是哪个层在偷吃显存。
我之前也遇到过一模一样的情况,3090按理说24G显存跑512的DeepLabV3+绝对够,batch 4都OOM肯定不是正常负载。我建议你先别急着怀疑DataLoader,因为数据加载的临时张量一般不会长期驻留显存,除非你用了pin_memory=True且num_workers开得特别高,但那种报错通常是CPU内存炸而不是显存。最可能的元凶其实是模型里的中间激活值,尤其是ASPP模块里那几个空洞卷积并行分支,它们会保留大量特征图直到反向传播结束,建议你用torch.cuda.set_per_process_memory_fraction先给进程设个上限,然后逐层打印参数的shape和requires_grad来对比。工具方面我强烈推荐你试一下torch.profiler,它带memory profiling功能,能直接按时间线看到每个操作分配了多少显存,比手动插hook省事得多。另外一个容易被忽略的点是,你检查一下模型里有没有用F.interpolate或者上采样层,有些实现会在前向时创建临时大张量,如果配合梯度检查点或者重计算反而更省。最后还有个歪招,你把batch size先设成1跑通,然后用一个循环手动累加梯度模拟batch 4,看显存峰值到底出现在哪个iteration,如果还是炸那基本就是模型结构问题,跟数据流无关。如果这些都排除了,建议你更新一下CUDA和cuDNN版本,3090是Ampere架构,老版本库有时会分配两倍显存做对齐。
我之前也遇到过类似情况,排查下来发现是模型里用了太多中间变量没释放,尤其是Decoder那块,建议先用torch.cuda.set_per_process_memory_fraction限制一下显存,看会不会提前报错,能快速判断是不是真超了。另外可以试试在dataloader里把pin_memory关掉,有时候这玩意儿在windows上会偷偷吃显存。逐层看占用的话,可以用torch.profiler或者简单点,在每个forward后面打印一下torch.cuda.memory_allocated(),对比一下就能看到哪层爆的。
我之前也踩过类似的坑,3090跑DeepLabV3+按理说512分辨率batch4不该炸的。你先别急着怀疑DataLoader,我遇到过最离谱的一次是模型里有个没用的辅助loss分支,它把中间层的超大特征图一直保着不释放,直接吃掉好几个G。建议你用torch.cuda.memory_summary()看下当前峰值在哪,或者干脆在关键层后面手动插torch.cuda.reset_peak_memory_stats()配合nvidia-smi dmon实时监看,这样能定位到具体是哪段代码突然涨显存。另外检查下是不是用了自动混合精度但忘了关梯度检查点,或者某些操作比如F.interpolate在反向时会产生额外缓存,这种小细节特别容易忽略。还有个土办法,把batch size降到1,如果还爆那基本就是模型结构问题,如果正常那就逐层加大batch找临界点,顺便看看是不是输入里混了不同尺寸的图没统一resize。最后提醒下,如果用了第三方预训练权重,有些实现会把BN设成可训练但实际又没跑满,也会导致显存异常。
试试torch.cuda.memory_summary(),能看每个缓存区占用,八成是中间变量没释放,检查下forward里有没有把大tensor存成list。
建议跑一下with torch.autograd.detect_anomaly(),能定位到具体哪行爆的,我之前遇到过就是自定义损失函数里多算了梯度。
我之前也碰到过类似情况,最后发现不是模型本身的问题,而是输入数据里有个意外的维度没被压缩。你可以先试试用torch.cuda.memory_summary()看看分配峰值在哪一步,如果显示是cudaMemcpy那大概率是数据传输的问题,否则就得看模型内部了。
另外我强烈建议你用torch.profiler,它自带内存分析功能,能直接看到每个op的临时buffer多大。我上次排查出是某个自定义的attention模块里用了大量的view和transpose,导致产生了巨大的中间张量,虽然最终输出很小但峰值特别吓人。你可以试试在forward里逐段打印tensor.shape和element_size,配合memory_allocated()函数,很快就能定位到是哪一层在爆。
还有个小技巧,如果你用了BatchNorm,试试在推理模式下跑一遍前向,看显存是否下降很多——这能区分是激活值缓存还是梯度问题。另外检查下有没有用detach()或者retain_graph=True的骚操作,我之前就是因为backward时retain_graph导致缓存层层累积,关掉就好了。
我遇到过类似的情况,最后发现是模型里用了太多中间变量没释放,尤其是那种在forward里反复做张量拼接的操作,显存峰值会特别高。你可以试试用torch.cuda.memory_summary()看内存分配情况,或者用nvidia-smi的监控配合逐步注释代码段来定位。另外检查下DataLoader里有没有把数据转到GPU后没及时del掉,有时候小batch也会因为缓存碎片导致OOM。
我之前也遇到过类似情况,最后发现是backbone的BN层在训练时统计梯度搞的鬼,试试用sync_bn或者冻结部分层看看。另外强烈建议用torch.profiler,它能直接打印每个op的显存占用,比肉眼猜快多了。你还可以检查一下输入图像有没有被意外转成连续内存,或者DataLoader里num_workers设太高导致缓存堆积,把workers降到2试试。