☰
PyTorch CUDA out of memory排查与优化:从显存定位到分布式训练全指南
2026/10/1 1:02:07 网站建设 项目流程

先说个现象:训练还没跑两步,控制台直接甩一行“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 allocated:Pytorch当前实际用于存放张量的内存。
  • 23.55 GiB reserved in total by PyTorch:Pytorch从CUDA驱动那里预留的总内存,包括了已分配和未分配的缓存块。这个值约等于nvidia-smi里看到的进程显存占用。

如果reserved远大于allocated,说明存在大量“预留但没真正用上”的内存,多半是缓存碎片问题了。如果两者都逼近显存上限,那确实是张量占用太多,需要从模型、batch size、激活值上下功夫。

如果是CPU内存不足,报错可能不会那么友好。Linux下经常是Killed,Windows下可能是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(f"total: {total / 1024**3:.2f} GB") print(f"used: {used / 1024**3:.2f} GB") print(f"free: {free / 1024**3:.2f} GB")

峰值显存统计也非常有用:

torch.cuda.reset_peak_memory_stats() # 运行你的训练或推理代码 peak = torch.cuda.max_memory_allocated() print(f"peak allocated: {peak / 1024**3:.2f} GB")

峰值统计能帮你判断一个训练step到底峰值占用多少。很多人看nvidia-smi觉得占满了,但实际张量只占了一部分,剩下的都是缓存和碎片。这个区别对后续排查很关键。

2. 定位OOM的实操方法:从堆栈到变量的逐层排查

2.1 用CUDA_LAUNCH_BLOCKING还原真实报错位置

Pytorch的CUDA运算是异步的:代码写在前,报错未必当场出现,而是等内核执行时才冒出来。这导致一个常见现象:实际OOM发生在某个卷积层,但堆栈却指向了后面的loss计算甚至优化器step。

解决办法是设置环境变量CUDA_LAUNCH_BLOCKING=1,让每个CUDA操作同步执行,这样报错堆栈能精确指向真正分配显存的那一行。

export CUDA_LAUNCH_BLOCKING=1

在Python里也可以在代码开头设置:

import os os.environ["CUDA_LAUNCH_BLOCKING"] = "1"

代价是训练速度会变慢,因为失去了异步计算的重叠能力。所以这只是排查手段,定位到问题后记得关掉。

还有一个环境变量PYTORCH_NO_CUDA_MEMORY_CACHING=1,它会让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(f"step {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_none=True)。这不算换优化器,但能让Pytorch把梯度清成None而不是0张量,从而释放梯度内存。这个细节虽然小,但在大模型训练时能省出一部分空间。

另外,Pytorch 2.x的AdamW支持fused=True选项,在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', dtype=torch.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_reentrant=False)

注意use_reentrant=False,这是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 DataLoader:num_workers、pin_memory、shuffle的隐藏开销

OOM不一定是模型造成的,也有可能是数据加载环节挤爆了内存。

num_workers表示启动几个子进程加载数据。子进程数量过多时,每个worker都会复制一部分数据集或预处理中间结果,内存占用直接翻倍。尤其是在Windows上,多进程用的是spawn方式,开销比Linux的fork大很多,这个问题更明显。

具体建议:

  • 如果数据集不大,num_workers=0反而最省内存,虽然加载速度慢一点,但不会出现多进程复制问题。
  • 如果数据集较大,num_workers设为CPU核心数的一半左右即可,不要贪多。
  • prefetch_factor控制每个worker预取的batch数量,默认2,可以调为1来减少内存占用。
  • pin_memory=True会为数据分配“页锁定内存”,加快CPU到GPU的数据拷贝,但它消耗的是CPU内存而不是显存。如果系统内存本来就紧张,可以把它关掉。

shuffle=True本身不占太多内存,但如果Dataset在__getitem__里做了大量图像解码、文本预处理,每次读取都会产生临时对象,多worker下内存压力也会变大。这种情况最好提前做一次预处理,把结果缓存到磁盘或内存中。

写Dataset时有个原则:尽量做“懒加载”,不要在初始化时把全部数据读进内存,而是每次__getitem__只读取需要的那一份。如果你在一台内存不大的机器上训练,这个原则能救你很多次。

5. 当单卡真的放不下:多卡与分布式方案

5.1 DP和DDP:看起来一样,显存差很多

单卡显存不够时,最自然的想法是上多卡。但用多卡也有讲究。

Pytorch里有两个现成的工具:DataParallel(DP)和DistributedDataParallel(DDP)。网上很多老教程还在用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(backend="nccl") local_rank = dist.get_rank() torch.cuda.set_device(local_rank) model = DDP(model, device_ids=[local_rank])

配合torchrun启动:

torchrun --nproc_per_node=4 train.py

如果你的环境不支持torchrun,也可以用torch.multiprocessing.spawn手动启进程。从训练效果和显存分布来看,从DP迁移到DDP往往是解决“多卡训练某张卡OOM”的最有效手段。

5.2 更进一步的ZeRO与FSDP:低显存跑大模型的方向

DDP虽然均衡了显存,但每张卡仍然要完整保存一份模型参数、梯度和优化器状态。对于超大模型,这依然不够。

这时候可以考虑ZeRO或者FSDP(Fully 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_CONF=max_split_size_mb:128

max_split_size_mb用来控制Pytorch缓存块的最大分割粒度。值越小,越不容易产生大量小碎片,但分配大块内存时可能更频繁地调用驱动。通常在32到512之间调整。如果你的模型存在明显的“大张量+小张量混用”场景,这个参数值得试。

PyTorch 2.1及以上还有一个更省心的选项:

export PYTORCH_CUDA_ALLOC_CONF=expandable_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_BLOCKING=1使用,定位更准。

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换成DistributedDataParallel
Windows下数据加载时内存暴涨num_workers过多或pin_memory叠加内存不足调低num_workers,关闭pin_memory
nvidia-smi显示显存被占但进程已退出僵尸python进程未回收用tasklist或ps -aux找到PID并清理
代码跑了很久没改过,升级Pytorch后OOM分配策略或默认缓存行为变化检查PYTORCH_CUDA_ALLOC_CONF,可尝试设置expandable_segments:True

6.2 我的OOM排查主流程,以及几个独门小技巧

我处理OOM有一个固定顺序,分享出来供参考:

第一步,看nvidia-smi,确认是否有别的进程占显存。如果是共享机器,先看清楚自己还剩多少可用。

第二步,复现OOM时开启CUDA_LAUNCH_BLOCKING=1,拿到真实报错堆栈,知道具体是哪一行爆的。

第三步,检查模型和batch size。如果一步就爆,那就是模型或batch太大,走AMP和梯度检查点;如果是运行了很长一段时间才爆,优先查泄漏和碎片。

第四步,加入峰值统计,确定是瞬时峰值还是持续占用高。瞬时峰值可以针对大矩阵操作做特殊优化;持续占用高需要检查长期持有的张量。

第五步,逐一尝试优化手段时,每次只改一个变量,不要同时开AMP又换优化器又改batch size,否则出了问题都不知道是谁引起的。

还有两个小技巧,是我实际用下来觉得特别省心的:

第一个,写一个自动探测batch size的脚本。遇到大模型时,先用2的倍数从1开始尝试,捕捉torch.cuda.OutOfMemoryError,找到当前配置下能跑的最大batch:

import 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熬夜。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询