最近在调一个UNet做医学图像分割,输入patch是256x256,batch size设的8,显卡是3090(24G)。刚开始训练loss下降挺正常,结果跑到第20个epoch左右,突然报CUDA out of memory。但我用nvidia-smi看显存占用才60%左右,而且显存是慢慢涨上去的,不是一下子爆的。我怀疑是不是PyTorch的缓存机制在搞鬼,还是有内存碎片化的问题?试过torch.cuda.empty_cache()也没啥用。另外,我用了混合精度(autocast),但感觉反而更吃显存了?有没有大佬遇到过类似情况,求指点排查方向,或者有没有什么工具能可视化显存分配?谢谢了。
PyTorch训练到一半显存爆掉,但看占用率才60%,这正常吗?
全部回复
共 72 条这情况太典型了,不是碎片化就是缓存没释放。你试试把dataloader的num_workers调到0,顺便检查下有没有变量不小心被梯度带住了,比如loss里用了不该用的中间量。混合精度没降显存大概率是batch里有一两个样本特别大,导致autocast的dynamic loss scaling在反复调整,可以关掉gradscaler看下峰值。工具的话推荐pytorch的torch.cuda.memory._dump_snapshot,或者直接用nvidia-smi的--query-gpu=memory.used,memory.total --loop=1盯着实时变化,比看占用率直观多了。
大概率是缓存碎片化,试试pytorch的分配器日志或者用nvidia-smi跟踪峰值显存,别急着上empty_cache。
这情况我碰到过,大概率就是PyTorch的缓存分配器没把显存还给驱动,加上混合精度下autocast偶尔会额外保留fp32的梯度副本,显存占用就慢慢爬上去了。empty_cache只是清空未使用的缓存块,对碎片化其实帮助有限。建议你试试用torch.cuda.memory_summary()看下具体是哪些张量在占内存,或者开一下PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,这个对碎片化挺有效的。另外3090跑256patch按理说很宽裕,你也可以查下是不是数据加载时把整个数据集都怼到显存里了。
这情况太典型了,3090跑UNet按理说24G绰绰有余,八成是PyTorch的缓存分配器在搞鬼,显存碎片化确实会这样,empty_cache只是清空未使用的缓存块,没法整理碎片。你可以试试在dataloader里加个pin_memory=False,或者把batch size降到4看还涨不涨,能快速定位问题。混合精度理论上该省显存,但如果loss缩放或者梯度检查点设置不对,反而可能触发额外内存分配,建议先关掉autocast对比一下。工具的话,torch.cuda.memory_summary()能看详细分配,或者用nvidia的Nsight Systems抓一下内存曲线,比nvidia-smi直观多了。
3090跑256的patch才8的batch按理说很宽裕啊,你这个涨到20个epoch才爆有点意思,我怀疑是PyTorch的缓存块在反复分配和释放中产生了碎片化,尤其你开了autocast,梯度缩放和master weight会额外占一份fp32的副本,显存峰值反而可能比纯fp32更难看。empty_cache只清空未使用的缓存池,对已经分配给tensor的碎片没辙,你可以试试看把dataloader的pin_memory关掉,或者用torch.cuda.memory_summary()看下实际分配明细,那个比nvidia-smi准多了。另外你loss是不是每次step都在涨?有些BN层在混合精度下统计量会漂移,导致显存需求慢慢爬升,可以检查下是不是有变量被意外拉进了计算图。我之前调分割模型也遇到过类似,最后是换成梯度累积每4步更新一次,batch降到4,反而稳定了,你可以先验证下是不是缓存问题,跑个10个epoch用nvidia-smi --query-gpu=memory.used,memory.total --format=csv -l 1盯一下曲线,如果还是平缓上涨再考虑代码层面。顺便问下你用没用什么第三方库比如segmentation-models-pytorch,有些封装会额外缓存中间特征图,这个特别坑。
3090跑256的patch还爆显存确实不太正常,不过你观察到的“占用60%但OOM”很可能是PyTorch的缓存分配器把显存块预留了但没完全用满,nvidia-smi看到的是物理占用,跟进程内部分配不是一回事。混合精度理论上该省显存,但如果loss scaling或者梯度相关buffer没处理好,反而可能多占。建议你试试把batch size降到4跑几个epoch对比一下,如果峰值显存没有线性下降,基本就是碎片化或缓存问题。工具方面可以看下torch.cuda.memory_summary(),或者用pytorch的memory profiler,能打印每个张量的分配情况。
显存慢慢涨然后突然爆,大概率不是碎片化,而是某个地方持有了计算图没释放,比如你把loss或者中间变量存到list里做日志了。混合精度下autocast区域里的bn或者某些op可能会偷偷转fp32,显存反而比纯fp32高,这个坑我也踩过。建议用torch.cuda.memory_summary()看下allocated和reserved的差距,再用py3nvml或者wandb的system面板盯一下epoch间的增量。另外检查下dataloader的num_workers和pin_memory,有时候是host侧累积拖累了device。
显存慢慢涨多半是没detach或loss累积了,用torch.cuda.memory_summary看下分配更准。
显存慢慢涨多半是碎片化,试试torch.cuda.memory_summary看下分配情况,另外amp有时反而会多占显存。
显存慢慢涨大概率是内存泄漏,检查下有没有没detach的张量累积在list里,混合精度有时确实会多占。
显存慢慢涨上去然后突然OOM,这基本就是典型的内存泄漏了。建议重点查一下validation或者日志那块有没有把tensor一直挂在graph上,比如累加loss的时候用了total_loss += loss而不是total_loss += loss.item()。另外混合精度本身不会更吃显存,但如果scale没处理好导致某些中间激活被保留,反而更糟。可视化可以试试torch.cuda.memory_summary(),能看到allocated和reserved到底差多少。
这个现象其实挺常见的,nvidia-smi显示的占用率确实会骗人,它反映的是当前时刻的分配情况,但PyTorch的缓存分配器会预留一大块显存池,碎片化的时候你看着还有空间,实际连续的大块已经没了。你这种慢慢涨上去然后突然爆的情况,八成是某个地方持有了计算图没释放,比如在训练循环里不小心把loss或者中间变量存进了list,或者验证阶段没加torch.no_grad()。混合精度那个感觉更吃显存,有可能是autocast配合某些op时反而产生了额外的fp32副本,特别是UNet里那些skip connection的concat操作。建议你用torch.cuda.memory_summary()打一下,能看到allocated和reserved的差距,再用torch.cuda.memory_snapshot()配合可视化工具看碎片分布。另外检查一下dataloader的num_workers和pin_memory,有时候worker里的张量没及时释放也会累积。可以试试把batch size降到4跑一遍,如果还是同样epoch爆,那基本就是代码里有泄漏而不是显存不够。