PyTorch推理、反向传播、参数更新阶段GPU显存分配原理及OOM排查
PyTorch GPU显存分配机制与OOM问题解答
PyTorch的GPU显存占用主要分为4类:模型参数、前向传播中间激活、参数梯度、优化器状态,同时PyTorch采用惰性显存分配+显存缓存复用机制:不会将不再使用的显存立刻归还给CUDA驱动,而是会预留为内部缓存复用,因此nvidia-smi显示的占用是PyTorch预留的总显存,不是当前实际激活使用的显存。
接下来结合你的日志逐一解答疑问:
1. 前向传播output = model(input, comb)新增约3G显存的原因
你日志中check6显存为3847MiB,check7执行完前向传播后为6725MiB,差值约2.9G。
前向传播过程中,PyTorch会自动保存计算图上所有中间层的激活值,这些值是反向传播计算梯度的必要输入,其总大小和batch size、模型每层通道数、特征图尺寸正相关。你本次前向产生的所有中间激活总大小约3G,因此显存出现对应涨幅。
2. 反向传播loss.backward()新增约3G显存的原因
你日志中check9显存为6725MiB,check10执行完反向传播后为9761MiB,差值约3G。
反向传播的核心是基于前向保存的中间激活,计算每个可训练参数的梯度,梯度的总大小和模型可训练参数的总大小基本一致:你模型参数本身占3.7G,刨除不需要计算梯度的层、显存对齐开销后,梯度的实际占用约3G,和涨幅匹配。同时反向传播过程中会逐步释放前向保存的中间激活,因此你看到的净涨幅就是梯度的存储开销。
3. 优化器更新optimizer.step()新增约6.3G显存的原因
你日志中check10显存为9761MiB,check11执行完step后为16053MiB,差值约6.3G。
你使用的是Adam优化器,它会为每个可训练参数维护2个状态变量:一阶矩估计(动量)、二阶矩估计,这两个变量的大小和参数本身完全相同,也就是说总优化器状态占用等于2倍的可训练参数总大小。你模型参数3.7G,对应2倍就是7.4G,刨除不需要优化的参数、显存碎片化开销后,实际申请的显存约6.3G,和涨幅匹配。另外Adam的状态变量是第一次调用step()时才会惰性初始化分配显存,因此你之前的检查点看不到这部分占用,第一次执行step就会一次性申请,直接将显存打满。
可选优化方案
- 降低batch size,可直接减少前向中间激活的显存占用
- 启用PyTorch混合精度训练
torch.cuda.amp,可将显存占用降低约50% - 无需参与梯度计算的张量及时调用
.detach()从计算图剥离,避免无效存储 - 若业务允许可改用SGD优化器:仅需维护1个动量状态(不加动量则无额外状态),优化器显存占用比Adam低50%以上
内容的提问来源于stack exchange,提问作者Ambrose
相关产品推荐
相关产品推荐

