最近在调一个图像分割模型(DeepLabV3+),用RTX 3090训练,batch size设到4都直接OOM,但看别的项目同样的卡能跑到16。我检查了输入图像尺寸(512x512),也试了梯度累积,还是崩。
目前怀疑是不是自己写的DataLoader里做了什么骚操作,或者模型里某些层占用了大量中间变量?想问问大家一般怎么定位这种“显存泄漏”或者“缓存堆积”的问题?有没有工具或代码技巧能快速看到每层的显存占用?谢谢各位大佬。
PyTorch模型训练时显存爆炸,但batch size已经很小了,还有什么排查思路?
全部回复
共 145 条你这情况我上周刚踩过坑,最后发现是模型里有个上采样层用了F.interpolate,在计算图里把梯度张量全保留了。可以先试试with torch.no_grad()跑一遍forward,如果显存正常就说明是反向传播的问题。排查工具的话torch.cuda.memory_summary()挺好用的,能看到每个张量的分配情况,再配合nvidia-smi -l 1实时监控,基本能定位到具体层。
我之前也遇到过类似情况,最后发现是输入图像忘了归一化导致数值范围异常,中间特征图直接爆掉,你可以先检查下数据预处理。另外建议用torch.cuda.memory_summary()看下内存分配,或者用torch.profiler跑一下,能定位到具体哪一层峰值最高。还有个笨办法,把batch size设成1,如果还崩,基本就是模型或数据流的问题,跟显存容量无关了。
我之前也遇到过类似情况,最后发现是模型里的中间特征图没释放,尤其是ASPP那块并行分支叠加后特别吃显存。你可以试试用torch.cuda.memory_summary()看内存分配细节,或者把batch size调到1跑一次,如果还崩大概率是模型结构问题。另外检查下DataLoader里有没有对图像做不必要的复制或转置,有时候num_workers设太高也会导致缓存堆积,建议降到2试试。梯度累积虽然能缓解但治标不治本,重点还是得找到那个占显存的大头。
试试torch.cuda.memory_summary(),能看到每个tensor的分配情况,或者用pytorch的memory profiler hook逐层定位。
我之前也遇到过类似情况,排查下来多半不是DataLoader的锅,反而是模型前向里的中间变量在作祟,比如DeepLabV3+的ASPP模块里那几个空洞卷积,输出特征图叠加后特别吃显存。你可以试试用torch.cuda.memory_summary(device=None, abbreviated=False)看下内存分配,或者更直接点,在forward里对可疑层包个torch.profiler,能清楚看到每块分配的峰值。另外检查下是否不小心把梯度也存了,或者用了太大下采样率的特征图做上采样,有时候把输入缩到256跑一遍,如果显存骤降那基本就是输入和中间特征的问题了。
先关掉cudnn的benchmark试试,之前遇到过类似问题,是torch.backends.cudnn.benchmark设成True导致缓存爆的。
可以在模型forward里插几行torch.cuda.memory_summary()打印,重点看下空洞卷积和ASPP那块的显存峰值。
用torch.cuda.memory_summary()看下缓存分配,重点查下Decoder里有没有在循环里反复创建张量。
我之前遇到过类似问题,结果是backbone的BN层没设eval导致梯度图全存了。
我之前也踩过类似的坑,3090跑DeepLabV3+按说不该这么惨,batch4就爆基本不是显存容量问题,而是有东西在偷偷累积。你提到DataLoader,这个方向我觉得挺对的,先检查一下num_workers是不是开太多,有时候多进程加载会复制几份数据缓存,但更常见的是你transform里如果用了某些库(比如albumentations)的GPU版本,它会在显存里留中间结果。另外,模型本身也可能有隐患,DeepLabV3+的ASPP模块里空洞卷积如果dilation设得大,特征图padding会按比例膨胀,中间变量会异常大,你可以用torch.profiler或者简单的hook在forward里打印每层输出的shape和内存,定位到具体层再优化。还有个技巧是试一下torch.cuda.empty_cache()和设置torch.backends.cudnn.benchmark=False,虽然这俩治标不治本,但能帮你区分是缓存碎片还是真占用。如果用了混合精度,检查一下grad_scaler是不是没配合梯度累积正确更新,有时scale_factor会积累历史值。最后建议你跑一个纯forward不backward的测试,如果这都爆那就是模型结构问题,不爆就重点查反向传播里的梯度缓存,尤其是用了自定义loss或者多分支输出时,中间梯度会被保留到step结束。
之前也遇到过类似情况,最后发现是输入图像没做归一化导致数值范围太大,中间激活值爆炸式增长。建议先用torch.cuda.set_per_process_memory_fraction跑个小实验,配合nvidia-smi dmon实时看显存曲线,能快速区分是峰值还是持续累积。另外检查下模型里有没有用F.interpolate或者自定义forward里保存了多余tensor,DeepLabV3+的ASPP模块有时候会缓存多尺度特征,可以尝试torch.utils.checkpoint把部分层换掉。如果代码里用了accumulate_grad,记得确认optimizer.zero_grad是在每个mini-batch后调用的,不是每个accumulation step后。
我之前也遇到过类似情况,排查下来发现是输入到模型里的feature map在forward时被重复保存了,比如某些自定义模块里用了多次forward调用。你可以试试在推理模式下跑一遍(不更新梯度),如果显存占用明显下降,那就是中间变量缓存的问题。
另外建议用torch.cuda.memory_summary()看下具体是哪块内存峰值高,或者装个pytorch_memlab逐层打印,很快就能定位到异常的层。之前我把batch size调小但没改num_workers,DataLoader的pin_memory反而会占更多显存,关掉之后立竿见影。
对了,检查一下是不是用了梯度累积但没清空optimizer的state,有时候这会触发额外的显存开销。你可以先跑一个极小的batch(比如1)看是否还爆,如果还爆,基本上就是模型结构本身的问题了。
这种情况我遇到过,先别急着怀疑DataLoader,试试用torch.cuda.max_memory_allocated()打点,配合nvidia-smi看显存曲线,能大致判断是前向还是反向爆的。我上次查出来是模型里一个自定义的upsample操作保存了超大计算图,换成F.interpolate直接省了3G。另外检查下有没有在循环里把loss和输出append到list里,这种隐式引用会阻碍显存回收,有时候比模型本身还坑。
我之前也遇到过类似情况,batch size调小反而更崩,最后发现是backbone的BN层在作怪。你试试把模型切成几段,用torch.cuda.max_memory_allocated()打点,分别在forward和backward前后记录一下,能很快定位到是哪一层峰值暴涨。另外有个小技巧,把输入图像换成随机噪声,如果还OOM,那基本排除DataLoader的问题,多半是模型本身的中间激活太大,比如ASPP模块里的空洞卷积并行分支,每个分支的feature map都会驻留显存。你可以检查一下有没有在forward里保存了不必要的中间变量,比如为了可视化或loss计算把多个尺度的输出都return了,这些会阻止显存释放。还有个坑是有些库(比如某些版本的timm)会在forward里默认开启grad checkpointing的开关,反而增加显存占用。建议直接用torch.profiler看内存时间线,比手动加打印更直观。另外确认一下你的PyTorch版本和CUDA版本是不是匹配,有时候混合精度训练(AMP)在旧版本上会偷偷缓存动态loss scale的梯度,导致显存缓慢增长。最后实在不行,试试把模型换成torch.utils.checkpoint,虽然慢一点,但能把激活重计算,显存占用能降一个量级。
用torch.cuda.memory_summary()看下峰值在哪,多半是中间变量没释放或者backward时梯度叠加了。
试试把batch size调成1跑一遍,如果还爆就是模型本身的问题,查查ASPP那块的显存占用。
可能是backbone的pretrained权重里带了BN的running_mean/cache,试试冻结BN或者换SyncBN,之前遇到过类似情况。
用torch.cuda.memory_summary()看下分配在哪一层,或者跑个空tensor前向对比一下,多半是中间变量没释放。
之前也遇到过类似的情况,后来发现是模型里用了太多中间变量没释放,尤其是DeepLabV3+的ASPP那块,建议用torch.cuda.memory_summary()看下分配细节,或者把backward的retain_graph关掉试试。另外检查下DataLoader的num_workers是不是开太多,有时候多进程缓存也会吃显存,我之前把workers从8降到2就好了。还有个笨办法,把模型输入换成随机tensor一步步跑,排除数据增强或transform的锅。
试试torch.cuda.memory._record_memory_history,能逐行看分配,大概率是中间变量没释放,查下forward里的list或者detach。
用nvidia-smi看显存曲线,再配合pytorch的memory_summary(),重点查下BN层统计和backward时梯度缓存,你这配置不该爆。
试试点背靠背的torch.cuda.max_memory_allocated()打点,大概率是中间变量堆积,查下有没有在循环里保存loss或feature。
用nvidia-smi盯显存曲线,再配合torchsummary看每层输出,DataLoader一般不会背这锅,先看看是不是模型前向里用了太多临时tensor。
我之前也遇到过类似情况,最后发现是模型里用了太多中间feature map没及时释放,尤其是DeepLabV3+的ASPP模块里并行卷积很容易堆显存,你可以用torch.cuda.memory_summary()看看具体哪一层爆的,再配合nvidia-smi -l 1实时监控显存变化曲线。另外检查下DataLoader是不是用了pin_memory=True并且num_workers开太多,有时候数据预加载也会吃显存,改成False试试。还有个小技巧,把输入图像临时缩到128x128跑一遍,如果显存占用还是很高,基本就能确定是模型结构或缓存问题而不是数据维度问题。
试试torch.cuda.memory_summary()看缓存分配,重点查下forward里有没有保存大tensor的list,或者用torch.no_grad()跑一遍验证。
检查下有没有BN层的running stats在反向传播时被重复计算,或者模型里有没有非叶子节点导致autograd缓存堆积,用detach截断试试。
试试点一下torch.cuda.set_per_process_memory_fraction,先限制显存用量跑个前向,看是不是数据加载那步的pin_memory或者num_workers太多把显存挤爆了。我之前遇到过类似问题,最后发现是模型里用了大尺寸的中间feature map没及时释放,用torch.cuda.memory_summary()能看得很清楚。另外检查下有没有在循环里保留loss或输出的list,那个也会悄悄吃显存。