尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyTorch CUDA out of memory排查与优化:从显存定位到分布式训练全指南

PyTorch CUDA out of memory排查与优化:从显存定位到分布式训练全指南 先说个现象训练还没跑两步控制台直接甩一行“CUDA out of memory”。这个提示我见得太多了尤其是刚把数据集加载好、模型刚搬上GPU正准备见证loss下降的时候突然就崩了。Pytorch的“内存不足”并不只发生在显存上CPU内存同样会爆只是报错形式不一样。很多人一遇到这个问题就盲目调小batch size结果训练效果变差问题还没解决。其实OOM有一套固定的排查思路和优化手段从定位到解决都能流程化。这篇文章我会结合自己在单卡、多卡训练和推理过程中踩过的坑把Pytorch运行时的内存不足问题拆开讲清楚。内容包括怎么从报错信息判断是显存还是内存不足、如何用工具定位是谁占用了显存、训练前模型和优化器层面怎么“减肥”、训练中数据加载和循环细节怎么“节流”、单卡实在放不下时怎么走多卡和分布式路线最后给一份常见问题速查表。适合刚接触Pytorch的初学者也适合已经被OOM折磨过几轮、想系统性解决问题的朋友。1. 先搞清楚内存到底满在哪里1.1 三类常见情况GPU显存、CPU内存、显存碎片Pytorch里的“内存不足”通常分三种情况第一种是GPU显存不足报错一般是“CUDA out of memory”。这是最常见的一种通常发生在把模型或中间张量放到GPU时显存不够了。第二种是CPU内存RAM不足报错一般是“MemoryError”或者直接进程被杀掉。这种情况在数据加载、CPU预处理、张量从GPU拷回内存时经常出现尤其是开了很多DataLoader的worker或者把数据集一次性全量加载到内存里。第三种是显存碎片化。进程总占用没到显存上限但内存被拆成很多小块无法分配一个大张量同样报“CUDA out of memory”。这个最迷惑人因为从nvidia-smi看显存明明还剩不少但Pytorch就是分配不出一个连续的大块。还有一种不是内存不足、但表现很接近的僵尸进程占着显存不放。训练结束后nvidia-smi里python进程还在显存被占了下一次跑就OOM。多数时候是因为代码里有后台线程、多进程没回收干净。1.2 从报错信息快速判断问题类型Pytorch的OOM报错里其实已经给出了关键信息。比如典型的一段RuntimeError: CUDA out of memory. Tried to allocate 2.00 MiB (GPU 0 has 23.65 GiB total capacity; 23.41 GiB already allocated; 0 bytes free; 23.55 GiB reserved in total by PyTorch)这段话很多人只看第一句就慌了实际上信息量很大Tried to allocate 2.00 MiB当前这一步想分配多少内存。这个值通常很小不是问题的根源只是压死骆驼的最后一根稻草。23.41 GiB already allocatedPytorch当前实际用于存放张量的内存。23.55 GiB reserved in total by PyTorchPytorch从CUDA驱动那里预留的总内存包括了已分配和未分配的缓存块。这个值约等于nvidia-smi里看到的进程显存占用。如果reserved远大于allocated说明存在大量“预留但没真正用上”的内存多半是缓存碎片问题了。如果两者都逼近显存上限那确实是张量占用太多需要从模型、batch size、激活值上下功夫。如果是CPU内存不足报错可能不会那么友好。Linux下经常是KilledWindows下可能是MemoryError。这时候优先检查DataLoader的num_workers、pin_memory以及是不是有人把整个数据集list加载进内存。1.3 先做一次显存体检两个基础命令加一个监控脚本遇到OOM我习惯先跑一遍“显存体检”确认当前状态再动代码。第一个命令是nvidia-smi直接看每块卡的显存总量、已用、空闲以及进程列表。如果显存被别人占了这里一眼就能看出来。Linux下可以用watch -n 1 nvidia-smi动态监控Windows下用nvidia-smi -l 1代替。第二个工具是Pytorch自带的torch.cuda.memory_summary()它会输出非常详细的分配信息包括当前分配、峰值分配、缓存块数量、碎片率等。在OOM发生前手动调用一次或者在OOM捕获异常里调用一次能拿到很多有价值的数据。再看GPU当前实时状态可以用import torch free, total torch.cuda.mem_get_info() used total - free print(ftotal: {total / 1024**3:.2f} GB) print(fused: {used / 1024**3:.2f} GB) print(ffree: {free / 1024**3:.2f} GB)峰值显存统计也非常有用torch.cuda.reset_peak_memory_stats() # 运行你的训练或推理代码 peak torch.cuda.max_memory_allocated() print(fpeak allocated: {peak / 1024**3:.2f} GB)峰值统计能帮你判断一个训练step到底峰值占用多少。很多人看nvidia-smi觉得占满了但实际张量只占了一部分剩下的都是缓存和碎片。这个区别对后续排查很关键。2. 定位OOM的实操方法从堆栈到变量的逐层排查2.1 用CUDA_LAUNCH_BLOCKING还原真实报错位置Pytorch的CUDA运算是异步的代码写在前报错未必当场出现而是等内核执行时才冒出来。这导致一个常见现象实际OOM发生在某个卷积层但堆栈却指向了后面的loss计算甚至优化器step。解决办法是设置环境变量CUDA_LAUNCH_BLOCKING1让每个CUDA操作同步执行这样报错堆栈能精确指向真正分配显存的那一行。export CUDA_LAUNCH_BLOCKING1在Python里也可以在代码开头设置import os os.environ[CUDA_LAUNCH_BLOCKING] 1代价是训练速度会变慢因为失去了异步计算的重叠能力。所以这只是排查手段定位到问题后记得关掉。还有一个环境变量PYTORCH_NO_CUDA_MEMORY_CACHING1它会让Pytorch禁用CUDA内存缓存每次分配都直接向驱动申请。这能非常直观地暴露内存碎片问题但也会显著拖慢训练。一般排查碎片问题时再用平时保持默认就好。2.2 用memory_summary找出元凶张量torch.cuda.memory_summary()输出的信息很详细主要看这几段Current usage当前已分配张量占用的显存。Peak usage历史峰值判断是不是瞬时冲高。Allocator state当前缓存块大小分布可以看是否存在大量碎片。如果Peak usage远超Current usage说明代码在某些时刻会创建非常大的中间张量但用完后释放了。这种情况可以盯住峰值出现的代码段多半是某个大矩阵运算或超大batch。如果Current usage本身就一直很高说明有张量被长期持有可能是模型参数、优化器状态或者某个没有释放的中间结果。在训练循环里加一个显存监控是很值得做的习惯import torch def print_memory_info(step): allocated torch.cuda.memory_allocated() / 1024**3 reserved torch.cuda.memory_reserved() / 1024**3 peak torch.cuda.max_memory_allocated() / 1024**3 print(fstep {step}: allocated{allocated:.2f}GB, reserved{reserved:.2f}GB, peak{peak:.2f}GB)每N个step打印一次能直观看到显存是稳定、逐步上涨还是突然飙升。逐步上涨基本就是泄漏了。2.3 排查隐藏的显存泄漏训练过程中的峰值监控显存泄漏最常见的情况是训练能跑但跑着跑着显存占用越来越高最终OOM。我遇到过的泄漏原因主要有三类第一类是tensor在循环中被list收集后没有释放。比如很多人习惯把每个batch的输出outputs.append(pred)最后统一cat。如果pred一直留在GPU上list里面累积的显存会越来越大。正确做法是边收集边转成普通Python标量或者用pred.cpu()转移到内存。第二类是自定义模型或训练循环中在不需要梯度的场景下构建了计算图。最常见的是验证阶段忘了torch.no_grad()导致每个batch都保留计算图验证跑完一轮显存就炸了。更隐蔽的是在torch.no_grad()作用域内调用了某些“不允许梯度”的操作但Pytorch为了一些算子依然会构建图注意代码缩进和上下文管理是否正确。第三类是检查点保存/加载的问题。训练到一半保存checkpoint加载时如果同时又保留了原来的模型和优化器两只完整的大模型同时存在显存直接翻倍。加载前先用del释放旧对象再torch.cuda.empty_cache()。定位泄漏峰值监控是最好用的工具分阶段看峰值变化如果每个epoch的峰值都在递增那肯定有东西没释放。3. 训练前的“减肥”方案模型与优化器层面的显存优化3.1 模型结构参数层面能做的那些事训练时显存占用的大头有三个部分模型参数、梯度、优化器状态以及前向传播时各层的激活值。很多人只盯着参数量其实激活值才是最容易爆的。一个粗略的显存构成公式是这样的模型参数参数量 × 每参数字节数。fp32是4字节fp16/bf16是2字节。梯度一般与模型参数同尺寸fp32下同样按4字节算。优化器状态Adam需要额外保存一阶和二阶动量通常是参数量的两倍空间。激活值依赖batch size、序列长度、特征图尺寸、模型深度。这部分最不可控也最容易被忽略。如果模型本身参数量太大最简单的办法是换更小的结构。对Transformer类模型减少隐藏层维度往往比减层数更立竿见影减少注意力头数或改用分组查询注意力GQA也能显著降低中间张量内存。对于自定义模型可以检查模型是否真的需要所有参数。比如embedding矩阵非常大时考虑让输入输出共享权重或者对embedding做低秩近似。这些改动虽然需要对模型结构有一定理解但确实是“省内存见效最快”的方式。如果实在不想动结构还有两个方案很常用一是混合精度训练二是梯度检查点。3.2 优化器状态Adam为什么吃显存换掉它能省多少很多人在报OOM时第一时间盯着模型却忽略了优化器状态才是吃显存的大户。以Adam为例它为每个参数额外保存一份一阶动量m和一份二阶动量v都是fp32。也就是说一个1B参数的模型光Adam的优化器状态就要占8GB。再加上模型参数4GB、梯度4GB和激活值显存自然起飞。如果换用SGD优化器状态几乎可以忽略不计但因为收敛效果通常不如Adam很多人不愿意换。更稳妥的替代方案是8bit优化器。用bitsandbytes库加载8bit版本的AdamW优化器状态直接砍到原来的四分之一左右。Adafactor。Pytorch官方实现里有内存占用比Adam低很多虽然收敛速度和最终效果与Adam有差异但在很多生成任务上差距不大。设置optimizer.zero_grad(set_to_noneTrue)。这不算换优化器但能让Pytorch把梯度清成None而不是0张量从而释放梯度内存。这个细节虽然小但在大模型训练时能省出一部分空间。另外Pytorch 2.x的AdamW支持fusedTrue选项在CUDA上会合并内核执行速度更快有时也能减少临时变量的峰值占用。如果GPU支持且Pytorch版本较新可以顺手打开。3.3 混合精度AMP与梯度检查点两个性价比最高的开关混合精度几乎是我排查OOM时第一个建议开启的功能。原理很简单前向和反向计算时把一部分张量以fp16存储和计算降低一半的内存占用同时保留fp32的“主权重”来维持训练的数值稳定性。这里说明一下AMP的实际显存收益是“接近减半但不到一半”因为优化器状态和主权重仍然是fp32。Pytorch的AMP代码很简洁from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for x, y in dataloader: optimizer.zero_grad() with autocast(): loss criterion(model(x), y) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()如果你是PyTorch 2.0以上推荐用新版接口torch.amp.autocast(cuda, dtypetorch.float16)写法更清晰。有个细节经常有人忽略保存checkpoint时要把scaler的状态也一起保存。否则训练到一半中断、恢复后梯度缩放系数对不上后续训练可能直接溢出变成NaN。梯度检查点Gradient Checkpointing是另一个开关思路是“拿时间换空间”。正常训练会保存每一层的前向激活值供反向传播使用开启检查点后不保存中间激活反向传播时再重新算一遍前向从而把激活值的内存从O(层数)降到一个极低水平。代价是训练时间增加20%-30%。使用方式from torch.utils.checkpoint import checkpoint def forward(self, x): return checkpoint(self.transformer_block, x, use_reentrantFalse)注意use_reentrantFalse这是PyTorch 2.x推荐的用法避开老版本里容易踩的requires_grad和版本检查的坑。这个方案对Transformer、BERT、GPT这类深层次模型效果极好但对浅层模型收益不大因为浅层模型激活值本来就少重算前向反而更不划算。4. 训练中的“节流”方案数据和循环细节优化4.1 batch size、梯度累积与学习率的关系调小batch size是最直接的控制显存手段但它会改变训练动态不能盲目调整。batch size减半后梯度估计的噪声变大模型收敛可能变慢严重时训练都不稳定。如果不想降低有效batch size可以梯度累积optimizer.zero_grad() accum_steps 4 for i, (x, y) in enumerate(dataloader): loss model(x, y) / accum_steps loss.backward() if (i 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()逻辑很简单连续多个小batch的梯度累加再统一更新参数。这样等效的batch size等于单个batch乘累积步数但显存只占单个batch的量。这里有几个坑务必注意累加时要把loss除以accum_steps否则等效学习率被放大了。必须在一个完整累积周期结束时调用optimizer.step()和optimizer.zero_grad()忘记清空梯度会导致参数更新路径完全错误而且极难发现。如果模型里有BatchNorm层梯度累积不能完全替代真实的大batch训练因为BN的均值和方差仍是按每个小batch计算的推荐换成LayerNorm或GroupNorm。学习率也需要配合调整。梯度累积增大有效batch后通常需要适当提高学习率但具体提高多少还是要看验证集的表现没有普适公式。4.2 千万别忽略的细节验证、推理与no_grad训练代码能跑验证阶段OOM的情况非常典型。很多人只在训练循环里用了model.train()验证时写了model.eval()但忘了包torch.no_grad()。model.eval()只是切换dropout和BN的状态并不会关闭autograd。验证和推理阶段正确写法model.eval() with torch.no_grad(): for x, y in val_loader: pred model(x) # 只保留结果数值不要保存GPU上的张量如果想更极致一点可以用torch.inference_mode()代替torch.no_grad()。inference_mode是no_grad的升级版会禁用更多与自动求导相关的机制推理速度更快内存分配也更少。还有一个容易踩的坑验证时把每个batch的输出都append到一个大list想最后一次性算指标。如果batch多、输出张量大这个list会把显存撑爆。正确做法是每个batch直接计算指标标量或者最多用pred.cpu()把张量搬到内存。4.3 DataLoadernum_workers、pin_memory、shuffle的隐藏开销OOM不一定是模型造成的也有可能是数据加载环节挤爆了内存。num_workers表示启动几个子进程加载数据。子进程数量过多时每个worker都会复制一部分数据集或预处理中间结果内存占用直接翻倍。尤其是在Windows上多进程用的是spawn方式开销比Linux的fork大很多这个问题更明显。具体建议如果数据集不大num_workers0反而最省内存虽然加载速度慢一点但不会出现多进程复制问题。如果数据集较大num_workers设为CPU核心数的一半左右即可不要贪多。prefetch_factor控制每个worker预取的batch数量默认2可以调为1来减少内存占用。pin_memoryTrue会为数据分配“页锁定内存”加快CPU到GPU的数据拷贝但它消耗的是CPU内存而不是显存。如果系统内存本来就紧张可以把它关掉。shuffleTrue本身不占太多内存但如果Dataset在__getitem__里做了大量图像解码、文本预处理每次读取都会产生临时对象多worker下内存压力也会变大。这种情况最好提前做一次预处理把结果缓存到磁盘或内存中。写Dataset时有个原则尽量做“懒加载”不要在初始化时把全部数据读进内存而是每次__getitem__只读取需要的那一份。如果你在一台内存不大的机器上训练这个原则能救你很多次。5. 当单卡真的放不下多卡与分布式方案5.1 DP和DDP看起来一样显存差很多单卡显存不够时最自然的想法是上多卡。但用多卡也有讲究。Pytorch里有两个现成的工具DataParallelDP和DistributedDataParallelDDP。网上很多老教程还在用DP因为写起来简单一行model nn.DataParallel(model)就完事。但DP在显存分配上有一个很明显的缺陷模型和初始参数被复制到每张卡但前向传播时每个batch会被切分发给各卡最后主卡通常是GPU 0要收集所有卡的梯度并计算loss主卡的显存压力会比其它卡高出一截。多卡训练时只有卡0爆显存多半就是用了DP。DDP虽然代码上复杂一些但每个进程独立负责一张卡通过ring all-reduce同步梯度显存分配更均衡训练速度也更快。一个最简单的DDP启动方式import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP dist.init_process_group(backendnccl) local_rank dist.get_rank() torch.cuda.set_device(local_rank) model DDP(model, device_ids[local_rank])配合torchrun启动torchrun --nproc_per_node4 train.py如果你的环境不支持torchrun也可以用torch.multiprocessing.spawn手动启进程。从训练效果和显存分布来看从DP迁移到DDP往往是解决“多卡训练某张卡OOM”的最有效手段。5.2 更进一步的ZeRO与FSDP低显存跑大模型的方向DDP虽然均衡了显存但每张卡仍然要完整保存一份模型参数、梯度和优化器状态。对于超大模型这依然不够。这时候可以考虑ZeRO或者FSDPFully Sharded Data Parallel。它们的核心思路是把参数、梯度、优化器状态分片到不同GPU上训练过程中需要哪部分再收集哪部分而不是每张卡都存一份完整副本。DeepSpeed的ZeRO有三个阶段Stage 1切分优化器状态。Stage 2切分优化器状态和梯度。Stage 3参数、梯度、优化器状态全部切分。Stage 2的Offload优化器到CPU可以让单卡显存大幅降低代价是速度明显变慢。如果你在单张卡上想跑一个比平时大一倍的模型可以从这个配置入手。Pytorch原生也有FSDP和ZeRO Stage 3思路类似from torch.distributed.fsdp import FullyShardedDataParallel as FSDP model FSDP(model)FSDP在Pytorch 2.x中已经比较成熟了支持CPU offload、混合精度、自动包装等对不想引入太多额外依赖的项目来说是很合适的选择。这类方案适合模型规模远超过单卡容量的场景。如果模型只是比显存大一点点优先用混合精度、梯度检查点等“轻量”手段因为ZeRO/FSDP的通信和调度开销都不小训练体验会明显变慢。5.3 显存碎片与CUDA缓存的高级控制显存碎片问题很隐蔽但实际中很常见。Pytorch在运行时会提前向CUDA驱动申请一大块显存然后自己管理分配和回收。训练过程中不断创建和销毁不同大小的中间张量会导致这块显存被切成很多不连续的小块。某个时刻想分配一个大张量时即使总剩余显存足够也因为没有连续块而报OOM。最简单的控制方式是设置环境变量export PYTORCH_CUDA_ALLOC_CONFmax_split_size_mb:128max_split_size_mb用来控制Pytorch缓存块的最大分割粒度。值越小越不容易产生大量小碎片但分配大块内存时可能更频繁地调用驱动。通常在32到512之间调整。如果你的模型存在明显的“大张量小张量混用”场景这个参数值得试。PyTorch 2.1及以上还有一个更省心的选项export PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True它允许Pytorch的缓存段动态扩展能大幅减少碎片化。在有连续显存压力、反复OOM且无法进一步降低batch size的场景下实测下来往往比手动调整max_split_size_mb更有效。缺点是会和某些自定义CUDA扩展或旧版本不兼容遇到奇怪的问题时可以关掉再试。torch.cuda.empty_cache()也能处理一部分碎片问题。它会清空Pytorch预留但未使用的缓存块把它返还给CUDA驱动。但很多人理解错了它并不会释放正在被张量使用的显存只是清掉缓存。所以最有效的用法是先del释放掉不用的张量再调用empty_cache()。频繁调用会拖慢速度最好在显存紧张的关键节点用比如每轮验证之后。如果是在共享服务器上还可以给进程设置显存上限torch.cuda.set_per_process_memory_fraction(0.9)这样进程最多使用90%的显存给别的进程留出一点空间也能在OOM之前更早暴露问题位置。配合CUDA_LAUNCH_BLOCKING1使用定位更准。6. 常见问题速查表与我的排查习惯6.1 典型现象速查表现象可能原因解决手段训练一开始就OOMbatch size过大或模型本身太大调小batch size开启AMP梯度检查点训练中期突然OOM峰值出现在某个大中间张量或数据中出现超长样本用CUDA_LAUNCH_BLOCKING定位检查输入shape是否有异常显存占用逐步升高直到崩溃显存泄漏多半有tensor被list持有或验证未关autograd检查循环中的列表缓存验证处加torch.no_grad()验证阶段OOMmodel.eval()但没写no_grad或保存了输出张量用inference_mode每个batch直接计算指标多卡训练只有卡0 OOM使用了DataParallel换成DistributedDataParallelWindows下数据加载时内存暴涨num_workers过多或pin_memory叠加内存不足调低num_workers关闭pin_memorynvidia-smi显示显存被占但进程已退出僵尸python进程未回收用tasklist或ps -aux找到PID并清理代码跑了很久没改过升级Pytorch后OOM分配策略或默认缓存行为变化检查PYTORCH_CUDA_ALLOC_CONF可尝试设置expandable_segments:True6.2 我的OOM排查主流程以及几个独门小技巧我处理OOM有一个固定顺序分享出来供参考第一步看nvidia-smi确认是否有别的进程占显存。如果是共享机器先看清楚自己还剩多少可用。第二步复现OOM时开启CUDA_LAUNCH_BLOCKING1拿到真实报错堆栈知道具体是哪一行爆的。第三步检查模型和batch size。如果一步就爆那就是模型或batch太大走AMP和梯度检查点如果是运行了很长一段时间才爆优先查泄漏和碎片。第四步加入峰值统计确定是瞬时峰值还是持续占用高。瞬时峰值可以针对大矩阵操作做特殊优化持续占用高需要检查长期持有的张量。第五步逐一尝试优化手段时每次只改一个变量不要同时开AMP又换优化器又改batch size否则出了问题都不知道是谁引起的。还有两个小技巧是我实际用下来觉得特别省心的第一个写一个自动探测batch size的脚本。遇到大模型时先用2的倍数从1开始尝试捕捉torch.cuda.OutOfMemoryError找到当前配置下能跑的最大batchimport torch def find_max_batch(model, sample_input): bs 1 while True: try: model(sample_input(bs)) bs * 2 except torch.cuda.OutOfMemoryError: torch.cuda.empty_cache() return bs // 2注意每次OOM后要清空缓存否则下一次尝试可能立刻失败。这个方法虽然暴力但能快速给batch size一个安全上限。第二个保存checkpoint时养成好习惯只保存state_dict不要保存整个model对象加载前先del旧模型和优化器再torch.cuda.empty_cache()。这个习惯能避免很多“跑着跑着显存越来越满”的问题。最后说一句OOM不可怕怕的是不判断原因就盲目调参。我现在的习惯是遇到OOM先冷静看日志确认是容量问题、碎片问题还是泄漏问题然后针对性处理。大部分时候AMP加梯度检查点就能解决80%的显存不够问题剩下的20%要么换优化器要么走多卡分布式。按这个思路排查基本不用再为OOM熬夜。
返回列表