最近在调一个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%,这正常吗?
全部回复
共 55 条3090跑256的patch按理说8的batch不该爆,你试试把dataloader的num_workers调低或者pin_memory关掉,有时候数据加载的缓存也会占显存。混合精度反而更吃显存可能是grad scalar或者某些op在fp16下没优化好,建议用torch.cuda.memory_summary()看下具体分配,碎片化确实存在但通常不会这么早炸。我怀疑跟你的loss计算或者中间变量没释放有关,检查下有没有把整个tensor存进list之类的操作。实在不行就换梯度累积,先用小batch跑通再说。
这情况我碰到过,nvidia-smi看到的占用和PyTorch实际分配经常对不上,因为缓存块是复用的,但碎片化严重时新申请大块内存就会OOM。混合精度反而更吃显存很可能是autocast只改了前向,但loss缩放和梯度回传时某些操作还是FP32,显存峰值反而上去了。建议先试下torch.cuda.memory_summary()看下分配明细,另外可以开环境变量PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,能缓解碎片化问题。还有个小技巧,把batch size减半跑一个epoch看显存曲线,能判断是不是数据加载或缓存累积导致的。
大概率是pytorch缓存碎片化加峰值波动,nvidia-smi看的是当前占用不是峰值,试试monkey patch或者pytorch_memlab抓峰值分配。
跑20轮才炸估计是某些tensor没释放,autocast反而吃显存可能是中间激活没对齐,开个gradient checkpointing试试。
这题我太熟了,3090跑UNet按理说256的patch加batch 8应该很轻松,你这种情况大概率不是显存真的不够,而是碎片化加缓存膨胀叠加的结果。nvidia-smi看到的是驱动层占用,PyTorch的缓存池可能已经把显存分成很多小块,新tensor申请不到连续空间就会报OOM,哪怕总量看着还有余量。empty_cache只释放空闲块,治标不治本,关键要减少动态shape或者频繁创建临时tensor。
混合精度反而更吃显存这个现象我遇到过,autocast如果配合GradScaler,有时候会把master weight和梯度都保留在fp32,加上临时变量,反而比纯fp32还多。你可以试试显式把模型参数和优化器状态都转到fp16,或者用torch.cuda.memory_summary()看详细分配,那个比nvidia-smi准得多。
另外我怀疑你loss正常下降但显存慢慢涨,可能是DataLoader的num_workers在后台预取数据,或者某个op在accumulate梯度时保留了中间变量。建议先跑一个epoch固定随机种子,用memray或者pytorch的memory profiler逐层打印,基本能定位到是哪个模块在漏。实在不行就换梯度累积,把batch拆成4个小的,显存压力会小很多。
你这现象太典型了,nvidia-smi看的60%是当前实际占用的显存,但PyTorch的缓存分配器会提前把显存划走,报错时其实是你申请的新tensor已经超过了缓存池的上限,empty_cache只是释放了未用的缓存块,但已经分配出去的碎片不会还回去。混合精度理论上该省显存,但autocast如果配合GradScaler,在反向传播时会把loss缩放后的梯度也存成fp32,加上UNet的跳跃连接和中间激活值,某些层反而会多存一份数据,感觉更吃显存很正常。
我建议你查一下是不是有某个batch的数据形状异常,比如最后一个batch不足8张,导致PyTorch动态图反复重新分配内存,碎片化累积到20个epoch才爆发。你可以用torch.cuda.memory_summary()看详细分配记录,或者开一下PYTORCH_NO_CUDA_CACHE环境变量试试。另外试试把batch size降到4,如果显存占用立刻降下来,那就是缓存碎片问题,可以改用torch.utils.checkpoint来切分激活值,代价是慢一点但显存能压住。
我之前跑3D分割也遇到过类似情况,最后发现是DataLoader的num_workers太多,每个worker都复制了一份CUDA上下文,显存悄悄被吃掉了。你可以把worker设为0跑一个epoch对比下,如果显存曲线稳定了,那就是这个原因。可视化工具的话,nvidia的Nsight Systems能看分配时间线,但配置麻烦,先用memory_summary和分段训练定位吧。
这情况我遇到过,大概率不是真的显存不够,而是PyTorch的缓存分配器把显存占着不还,加上训练过程中张量碎片化越来越严重,到第20个epoch才触发OOM。你可以试试torch.cuda.set_per_process_memory_fraction或者用pytorch的memory_stats接口查一下缓存和实际使用的差值。混合精度理论上该省显存,但如果你loss scaling或者梯度累积没调好,反而可能临时多出几个大tensor,建议把autocast范围缩小到forward pass试试。可视化的话可以用nvidia-smi的循环采样或者pytorch的memory_snapshot,能直接看每个张量的存活时间。
大概率是PyTorch的缓存分配器在搞鬼,它会把显存块留着复用,nvidia-smi看到的是实际占用,但分配器可能已经预占了更多。你可以试试PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb=128或者garbage_collection_threshold,这俩参数对碎片化挺有效。混合精度理论上该省显存,但autocast配AMP时如果梯度没同步,反而可能让缓存更碎。另外,你试试把batch size降到4跑到崩溃点,如果显存还是慢慢涨,那就基本锁定是缓存碎片而不是数据量问题了。可视化的话,torch.cuda.memory_snapshot()配合pytorch_memlab能看分配历史,比盯着nvidia-smi靠谱。
这情况我碰到过,大概率就是PyTorch的缓存分配器在搞鬼,显存碎片化导致明明总量够用但就是分配不出连续块,尤其UNet这种多尺度特征图特别容易触发。混合精度按理说更省显存,但你得确认下是不是loss scaling或者某些算子回退到FP32了,可以看看autocast的日志。建议试试PyTorch的memory_stats接口,或者用nvidia-smi加个--query-gpu=memory.used,memory.total循环看,能明显看到缓存峰值比实际占用高得多。还有个土办法,把batch size降到4跑几个epoch看看显存增长曲线,要是还涨那就是代码里有张量没释放,查查dataloader的num_workers是不是开太多。
3090跑256的UNet不该爆,试试把batch降到4看loss曲线稳不稳,大概率是缓存碎片加验证集没清梯度。
这情况太典型了,我赌五毛钱就是PyTorch的缓存分配器在作妖。nvidia-smi看到的是整个显存的使用,但PyTorch自己有个内存池,训练时申请的显存释放后不会立刻还给驱动,而是留着复用,所以占用率看着不高但实际上池子已经撑爆了。你那个慢慢涨上去的现象,大概率是某个中间变量或者梯度累积导致的碎片化,empty_cache只是清空未使用的缓存块,对已经分配但没释放的张量没辙。
混合精度反而更吃显存的话,检查下是不是autocast没包住整个forward+loss,导致某些算子强制回落到fp32,或者你用了gradscaler但没配合使用,这时候master weight和fp32梯度副本反而多占一份空间。建议你装个pytorch的memory profiler,或者用torch.cuda.memory_summary()看下哪一行分配的峰值,这比nvidia-smi直观多了。
另外20个epoch才爆,很可能是某个数据集样本尺寸不统一,到那个batch刚好触发了峰值。试试把batch size降到4或者6,同时把pin_memory关掉,有时候它预加载的锁页内存也会悄悄吃掉不少显存。如果还不行,就手动把输入resize成固定尺寸,或者用累积梯度模拟大batch,这样至少能稳定跑完。
nvidia-smi显示的占用率确实容易骗人,它统计的是整个显存池的使用,而PyTorch的缓存分配器会把释放的块留着复用,所以监控看着不高但实际可用块已经被切碎了。你可以试试看torch.cuda.memory_summary(),能清楚看到缓存块和碎片情况。混合精度理论上该省显存,但如果loss scaling或者梯度scale没调好,反而可能触发额外显存分配,建议先关掉autocast跑个基准对比一下。另外20个epoch才爆,也可能是某个中间变量在特定输入下尺寸变大,比如UNet的skip connection在batch norm统计更新时临时保存了完整图,可以查一下是不是有显存峰值。
大概率是碎片化,3090的缓存块一旦错位就收不回来,建议把dataloader的pin_memory关了试试。
大概率是PyTorch的缓存分配器没释放碎片,试试设PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,我换完就好了。
显存涨到后面才爆大概率是碎片化加缓存累积,nvidia-smi看的60%是当前实际占用,但CUDA context预留的缓存块可能已经碎得没法分配连续内存了。empty_cache只清空未使用的缓存,对碎片没帮助,可以试试把dataloader的num_workers调低或者用pin_memory=False看有没有变化。混合精度理论上该省显存,但autocast对某些op(比如卷积)的临时张量反而可能保留fp32,建议用torch.profiler看一下具体是哪个模块在涨。我之前用pytorch_memlab的Snapshots定位过类似问题,你也可以试试。
大概率是pytorch的缓存碎片化了,试试max_split_size_mb参数或者换dataloader的pin_memory看看。
大概率是PyTorch缓存碎片化,试试设置PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb=128,或者用torch.cuda.memory_summary()看下分配细节。
nvidia-smi看的是峰值不是实时分配,PyTorch缓存池占着不还,试试max_split_size_mb参数。
大概率是碎片化+缓存没复用,试试给dataloader加个num_workers和pin_memory,或者把batch先调小再慢慢加回去。
混合精度吃显存可能是loss scaling的临时张量没释放,建议用torch.profiler看下峰值在哪一步。
我之前跑3D分割也遇到过一模一样的情况,nvidia-smi那个占用率其实不准,它显示的是整个GPU的分配情况,PyTorch缓存池里的显存不一定都被算进去。你把PYTORCH_CUDA_ALLOC_CONF设成max_split_size_mb=128试试,或者用torch.cuda.memory_summary()看下详细分配,大概率是碎片问题。混合精度按理说应该省显存,但如果loss scale或者梯度累积没调好,反而会多占,可以试试关掉autocast对比一下。
另外一个排查思路是看是不是某个特定操作(比如attention或者上采样)在长时间运行后触发了cudnn的benchmark重新搜索,导致临时显存暴涨。我建议你盯一下每个epoch的峰值显存,用torch.cuda.max_memory_allocated()打出来,如果每个epoch都在涨那可能是数据泄露或者缓存没释放干净,如果只是最后突然爆那基本就是碎片了。
这情况我也踩过坑,大概率就是PyTorch的缓存分配器没把显存还给驱动,nvidia-smi看到的是驱动占用的总数,跟进程实际可用的块不是一回事。混合精度理论上该省显存,但autocast对某些算子(比如BatchNorm)反而会保留fp32的额外buffer,加上你显存缓慢上涨更像是有个Tensor在某个分支没被释放,长周期累积导致的。建议你跑的时候用pytorch的memory_summary()看下分代分配,或者试试torch.cuda.memory_snapshot()转成trace文件用chrome://tracing打开,能定位到具体是哪行代码在申请内存。另外检查下dataloader是不是有num_workers>0导致显存拷贝泄漏,我上次就是被这个坑了。