最近在调一个图像分割模型(DeepLabV3+),用RTX 3090训练,batch size设到4都直接OOM,但看别的项目同样的卡能跑到16。我检查了输入图像尺寸(512x512),也试了梯度累积,还是崩。
目前怀疑是不是自己写的DataLoader里做了什么骚操作,或者模型里某些层占用了大量中间变量?想问问大家一般怎么定位这种“显存泄漏”或者“缓存堆积”的问题?有没有工具或代码技巧能快速看到每层的显存占用?谢谢各位大佬。
PyTorch模型训练时显存爆炸,但batch size已经很小了,还有什么排查思路?
全部回复
共 145 条我之前也踩过类似的坑,3090跑DeepLabV3+按理说512分辨率加batch4应该很轻松,除非你用了空洞卷积的很大rate,或者backbone里带了ASPP还叠加了多头注意力之类的。建议先别猜DataLoader,直接用torch.profiler或者简单点的pytorch_memlab,它能按行打印每个tensor的分配点和大小,特别适合抓这种“看起来小但实际占满”的中间变量。
另外有个很容易忽略的点:你检查过输入图像是不是真的512x512吗?有时候DataLoader里做随机crop或者resize后忘了同步,实际喂进去是1024甚至更大,那显存直接翻四倍。还有BN的momentum和累积的均值方差也会占显存,但通常不大,重点还是看有没有在forward里反复创建大tensor,比如把特征图存成list再cat。
我自己的习惯是先把模型切成几段,用torch.cuda.max_memory_allocated()在每段前后打点,很快能定位到是编码器还是解码器炸的。如果确认是中间变量问题,试试用checkpoint技术,牺牲一点速度换显存,或者把一些不需要梯度的操作包在torch.no_grad()里。最后别忘了看下是不是混合精度没开,fp16能省一半以上,尤其对卷积这种算子。
先看看是不是输入图像没归一化导致精度变成fp32,再不行用torch.cuda.memory_summary()直接看每层分配,大概率是backbone的中间特征图堆爆了。
我之前也遇到过类似情况,排查下来发现是模型里的BN层在训练时缓存了太多统计量,试试把torch.no_grad()包住验证阶段,或者用torch.cuda.empty_cache()定时清一下。另外推荐你用torch.profiler来看每层显存占用,比手动估算准得多。还有个小技巧,把输入图像切成小块过一遍,看是不是某个特定尺寸触发的峰值,有时候是backbone的stride组合问题。
我之前也遇到过类似情况,最后发现是模型里没用with torch.no_grad()的中间特征提取环节在反向传播时把计算图全保留了。你可以试试torch.cuda.memory_summary()看下哪个张量占的峰值最高,或者用pytorch的autograd检测下是否有forward hook没清理。另外检查下DataLoader的num_workers是不是开太多,有时候数据加载预取也会吃显存。
别光盯着batch size,先看看是不是输入图像没归一化导致梯度爆炸,或者模型里用了大卷积核的深度可分离卷积。我上次就是被BatchNorm的running_mean缓存坑了,清一下optimizer.zero_grad()和cache清空逻辑就好。可以用nvidia-smi看显存变化曲线,配合torch.profiler逐层分析,基本能定位到问题层。
我猜你可能是用了类似F.interpolate这种上采样操作,它会产生额外的临时变量,尤其配合DeepLabV3+的ASPP模块时特别吃显存。建议把混合精度训练打开,autocast能省不少,再不行就检查下有没有把整个验证集塞进GPU,有时候是验证代码偷偷累积了梯度。
显存爆炸不一定只跟batch size有关,你试试把图像尺寸降到384看看还崩不崩,如果还崩那就是模型结构问题了。我之前遇到过是自定义的损失函数里创建了
我之前也遇到过类似情况,最后发现是模型里的中间特征图没释放,尤其是解码器部分,试试用torch.no_grad()包住验证阶段,或者把不需要的反向传播图切断。另外推荐你装个pytorch-memlab,它能直接打印每层的显存占用和临时缓存,定位起来快很多。还有个坑是DataLoader的num_workers开太多会额外吃显存,降到4以下试试,有时候是pin_memory在搞鬼。
我之前也踩过类似的坑,batch size调小反而更崩,大概率不是显存不够,而是中间变量在作祟。你试过用torch.cuda.max_memory_allocated和torch.cuda.memory_summary()吗?这两个API能直接看到每个张量占用的峰值,配合nvidia-smi的实时监控,基本能定位到是哪一层在疯狂吃显存。
还有个很常见的隐蔽问题,你的DataLoader里如果用了num_workers>0,而且每个worker都持有了模型权重或者做了什么预处理缓存,那内存会翻倍涨,但看起来像显卡爆了。我之前用albumentations做增强,默认开了多进程,结果每个worker都复制了一份模型参数,直接炸穿。
另外建议你查一下BatchNorm的momentum和track_running_stats,如果是训练模式但显存峰值出现在BN层,那可能是计算图没有释放。你可以试试在每次迭代后显式调用torch.cuda.empty_cache(),但注意这只清缓存不解决根本问题,如果还崩就逐层打印requires_grad和grad_fn,看看是不是有梯度累积的backward没清干净。
还有个土办法,把batch size设为1跑一遍,如果还是OOM,基本可以排除batch维度,那就是模型结构或输入尺寸的问题。DeepLabV3+的ASPP模块里空洞卷积会生成很多中间特征图,建议你检查下output_stride是不是设成8了,改成16能省一半显存。最后推荐用torchinfo或pytorch-summary打印每层输出shape,配合memory_profiler这种工具,比瞎猜高效多了。
我之前也踩过类似的坑,排查重点不一定是batch size,而是有没有在forward里保留大tensor做可视化或loss计算。你可以试试用torch.cuda.memory_summary()看分配峰值,或者用pytorch的profiler,能定位到具体哪一行op撑爆的。另外检查一下DataLoader的num_workers,有时候多进程预加载会复制GPU缓存,看起来像显存泄漏。我之前是发现模型里有个辅助loss把feature map存下来算统计,去掉后直接省了3G。
先别急着怀疑DataLoader,你可以在每个step前打印一下torch.cuda.memory_allocated()和memory_reserved(),看是不是只增不减。如果是,多半是某个变量被global引用或者反向传播图没释放。另外DeepLabV3+的ASPP模块里空洞卷积如果dilation设太大,中间feature的channel数会很夸张,建议用torchsummary跑一下每层的输出shape,对比显存占用曲线。我上次就是被一个1x1卷积的groups参数坑了,显存翻倍但计算量没变。
你试试把输入切成patch跑一次,如果显存占用和原图差不多,那就是模型结构的问题。我遇到过类似情况,最后发现是自定义的损失函数里做了上采样到原图尺寸,那个临时tensor直接占了几G。定位技巧的话,可以用torch.aut
这问题我踩过坑,先别急着怀疑DataLoader,3090跑512输入理论上不该这么惨。你试试把torch.no_grad()包在验证集或者数据增强那段,有时候是预处理里不小心开了梯度追踪。另外DeepLabV3+的ASPP那块空洞卷积会生成超多中间特征图,尤其是rate大的分支,建议用torch.cuda.memory_summary()看下峰值分配在哪一层,比瞎猜快得多。
我之前遇到过类似情况,结果发现是损失函数里把feature map的list存下来了,明明只用了最后一个,但前面的全留在计算图里。你检查下有没有类似outputs = []然后append每一层结果的操作,或者用了detach()但没清空缓存。还有个偏门技巧,用torch.cuda.set_per_process_memory_fraction(0.9)强行限制显存,如果OOM报错变了,就能确认是峰值超了而不是泄漏。
对了,你batch size小但梯度累积也没用的话,大概率是单次前向就爆了,跟batch无关。可以试试把输入缩到128看看还爆不爆,如果还爆那就是模型结构问题,如果好了就逐级放大找阈值。工具方面推荐torchinfo(原summary)加pytorch_memlab的LineProfiler,能精确到每行代码的显存增量,比nvidia-smi好用多了。最后查一下是不是装了apex或者amp,混合精度在3090上有时反而会保留fp32副本导致翻倍,关掉试试。
我之前也踩过类似的坑,最后发现是模型里用了太多中间变量没释放,尤其是DeepLabV3+的ASPP那块,可以试试torch.cuda.max_memory_allocated()在前后打点,分段定位峰值出现在哪个模块。另外检查下DataLoader里是不是把图像转成了FP32又做了归一化之类的,有时候不经意间把图放大到原始尺寸的几倍,内存就翻车了。还可以用torch.autograd.detect_anomaly()开一下,虽然慢点,但能抓到异常梯度导致的额外显存占用。实在不行就把batch size调到1,然后逐步加,看看是不是某个特定输入尺寸触发的爆炸。
我之前也遇到过类似的情况,batch size调小还炸,最后发现是backbone的BN层在作怪。你试试把模型切成几个阶段,用torch.cuda.memory_record或者直接hook每层的输出tensor大小,打印出来看看哪一步峰值最高。我猜你DeepLabV3+的ASPP模块里空洞卷积的dilation rate如果设得很大,中间特征图的通道数会翻好几倍,虽然输入是512,但到后面几层可能已经变成1024或者2048的通道了,显存占用是呈指数涨的。另外检查一下你有没有在loss里同时用了多个分支,比如辅助loss,每个分支都会保留整张图的梯度图,这个很吃显存。还有个坑是DataLoader的num_workers如果设太高,每个worker会复制一份模型参数到内存,虽然不占显存但会影响整体内存分配策略,有时候会触发CUDA的碎片化。建议你用torch.autograd.detect_anomaly配合memory_stats,或者直接跑一次不带backward的前向,看看显存是不是还炸,就能区分是模型结构问题还是优化器状态问题。如果前向不炸反向炸,那就是激活值保存太多,可以试试用checkpoint机制换空间。
我之前也遇到过类似情况,后来发现是输入图像没归一化导致数值范围太大,中间激活值异常膨胀。建议先用torch.cuda.set_per_process_memory_fraction限制显存,配合pytorch的torch.autograd.detect_anomaly()跑一遍,能直接定位到爆显存的那一行。另外你检查下模型里有没有用固定长度的Transformer模块,那种会按序列长度分配缓存,512分辨率下可能比CNN吃显存得多。
如果DataLoader里做了随机crop或者翻转,试试把num_workers设成0,排除多进程缓存占用。剩下的可以用nvidia-smi dmon实时看显存曲线,看是前向还是反向阶段峰值。我上次就是被一个没用的dropout层坑了,它在训练时保留了整个batch的随机掩码。
还有个小技巧,用torch.cuda.memory_summary()打印分配明细,能看出是不是有反复分配释放的碎片。实在不行就换mixed precision训练,3090的自动混合精度能省一半显存,有时候比reduce batch size效果好很多。
试试torch.cuda.memory_summary(),能看缓存分配,重点查下BN和中间变量有没有detach。
我之前也遇到过类似情况,batch size压到2都炸,后来发现是backbone的BN层在训练模式下会缓存大量统计量,尤其是DeepLabV3+这种带ASPP的多尺度结构,中间特征图叠加起来比想象中吃显存。你可以先用torch.cuda.set_per_process_memory_fraction把显存限制到比如80%,然后跑一个小循环,在每层forward后打印torch.cuda.memory_allocated(),这样能直观看到是哪个块涨得最快。另外检查下DataLoader里有没有意外的tensor拼接或者把图像转成float64,我之前就是不小心在transform里用了double,直接翻倍。还有一个坑是如果用了第三方预训练权重,有些层可能被冻结但仍然计算梯度,导致autograd保留整张计算图,试试把requires_grad=False的层全部eval一下?工具方面推荐pytorch_memlab,它能输出每个张量的生命周期,或者用torch.profiler看内存事件,比手动print省事。对了,你确认下是不是混合精度没开,3090跑fp16应该能省一半,但开了之后要小心loss scaling和某些op的精度问题。最后实在不行就换一个更轻的backbone试试,比如MobileNetV3做encoder,先排除是不是模型本身结构的问题。
我之前也遇到过类似情况,最后发现是输入图像没归一化,模型里BN层的running stats在反向传播时疯狂吃显存。你可以试试在训练循环里用torch.cuda.max_memory_allocated()打点,每步都看下峰值出现在哪个位置,基本能定位到是前向还是反向炸的。
另外检查下DataLoader的num_workers,如果开太多而且每个worker都预加载了完整图像副本,显存会被悄无声息占掉一大块,尤其你用了transform的话。我上次就是被这个坑了,batch size从8调到2都救不回来。
还有个土办法,把模型里可疑的层挨个替换成identity跑一遍,二分法定位,虽然笨但很有效。你试试看是不是ASPP模块里的空洞卷积导致的,那玩意儿中间变量特别多。
我上次也遇到过类似情况,最后发现是输入图像没有归一化,直接以float64喂进去了,显存直接翻倍。你可以先torch.cuda.set_per_process_memory_fraction限制一下显存,然后跑一个前向传播看峰值在哪,或者用torch.profiler看内存分配的时间线。另外检查下是不是有梯度回传时保存了所有中间feature map,DeepLabV3+的ASPP模块挺吃显存的,可以试试把输出stride调大点。
3090是24G显存,512输入跑DeepLabV3+正常batch4应该没问题,你先用torch.cuda.max_memory_allocated()看下峰值在哪个环节,再配合nvidia-smi的实时监控逐步缩小范围。我之前遇到过类似问题,最后发现是backbone的BN层在eval和train切换时缓存了过多统计量,另外你检查下是不是开了cudnn.benchmark或者把输入转成了FP32但模型里有FP16的混合精度操作。还有一个骚操作是DataLoader里用了pin_memory=True且num_workers开太多,有时会导致显存碎片化,试试把workers降到2或者直接关掉pin_memory。
我之前也踩过类似的坑,3090跑DeepLabV3+按理说512输入不该这么惨,你试试把torch.no_grad()包在验证阶段,有时候验证集推理的梯度没释放会一直攒着。另外你DataLoader里如果用了pin_memory=True,配合num_workers>0,在Windows上偶尔会有显存不回收的bug,先关掉看看。定位每层占用最直接的是torch.cuda.memory_summary(),能打印分配详情,再配合nvidia-smi看进程峰值,不过最细的还是用pytorch的profiler,能按操作符拆解显存。我还遇到过模型里用了自定义的注意力模块,里面如果有个很大的中间tensor没删,且被当成叶子节点保留梯度,就会一直占着,你检查下forward里有没有不必要的return中间结果。另外一个反直觉的点,batch size小不一定省显存,因为某些算子(比如sync batchnorm)会额外分配固定大小的缓冲区,你可以试下把BN换成普通的,或者直接跑个纯卷积的backbone对比。最后如果还不行,我怀疑是CUDNN的benchmark模式在动态输入尺寸下疯狂申请workspace,设成torch.backends.cudnn.benchmark=False能省一大块,我之前就这么解决的。
3090跑512的DeepLabV3+才batch4确实不正常,先试试关掉amp和cudnn.benchmark,排除混合精度和benchmark的缓存问题。
我之前也踩过类似的坑,排查下来发现是模型里的中间变量没释放,比如在forward里反复存了feature map用于loss计算。你可以试试用torch.cuda.memory_summary()看每个张量的占用,或者把batch size调到1跑一次,如果还崩基本就是模型结构或数据管线的问题。另外检查下有没有在backward后保留梯度,或者用了detach但没清缓存,有时候DataLoader里多做了一次归一化或padding也会莫名吃显存。
跟你遇到一模一样的情况,最后查出来是我forward里有个离谱操作,把特征图slice后没detach又反复用了好多次,导致autograd图越攒越大。建议先用torch.cuda.set_per_process_memory_fraction设个上限,崩的时候看堆栈在哪个张量上,另外把dataloader的num_workers设0试下,排除数据加载那边搞鬼。
可以用torch.profiler或者直接print(x.grad_fn)看中间变量,我习惯在每层后面加hook打印output的shape和内存,不过最实用的还是把模型切成几段分别跑,二分法定位到具体模块。