最近在调一个图像分割的模型,训练倒是正常,但到了验证阶段每次加载权重后一跑前向就OOM。我用的是PyTorch 2.0,显卡是4090 24G,batch size已经降到1了,输入图片也resize到256x256。诡异的是,训练时同样的batch size和分辨率都没事,就验证的时候爆。我怀疑是不是我用了torch.no_grad()但忘记把模型切到eval()模式,导致BN层还在更新?还是说DataLoader的pin_memory=True在验证时反而占用更多显存?另外,我加载的是保存的完整checkpoint(包含optimizer状态),会不会是optimizer的动量缓冲也占了一部分显存?有遇到过类似情况的朋友吗?求指点一下排查思路。
PyTorch模型加载到一半显存爆了,是不是我数据预处理有毛病?
全部回复
共 1 条训练和验证的显存差异其实挺常见的,但你这个情况我第一反应不是BN层的问题,因为训练时BN的动量更新反而会更占显存(要存running stats的梯度相关状态),验证时就算忘记切eval,也只是行为不对,显存不会凭空多出来。倒是pin_memory=True确实有可能在验证时多占一块锁页内存,但那通常影响的是CPU内存而不是显存,除非你DataLoader的num_workers开得很大,导致数据预取队列堆积。
我更好奇的是你加载完整checkpoint的方式,如果直接torch.load再model.load_state_dict,optimizer的state_dict会被放到GPU上(尤其包含动量buffer),但这里有个坑——如果你是用torch.load的默认方式,它会先反序列化到CPU,然后再由load_state_dict转移,按理说不会常驻显存,除非你后续没删掉那个临时变量。建议你检查一下加载后有没有del optimizer_state或者用map_location='cpu'。
另一个思路是验证阶段你可能有额外的显存峰值,比如模型里用了torch.cuda.synchronize()或者在forward里做了可视化(保存预测图),这些操作在验证时往往会临时分配大tensor。你可以把验证循环里除了前向和loss计算以外的所有操作都注释掉,跑一次看还爆不爆。
我之前遇到过类似情况,最后发现是验证时用了torch.no_grad()但模型里有个torch.cuda.amp.autocast()没关,混合精度在验证时反而触发了某些层的动态显存分配。你可以试试验证时也保持和训练完全一致的上下文管理器。
如果还不行,建议用torch.cuda.max_memory_allocated()在验证前后各打一次,看看峰值到底出现在哪个环节,这样能直接定位是模型前向、loss计算还是数据加载的问题。24G的4090跑256x256的输入,除非你的backbone特别大,否则正常不该爆,所以大概率是某个临时变量没释放。