ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch梯度累加实战:解决显存OOM并模拟大Batch训练

PyTorch梯度累加实战:解决显存OOM并模拟大Batch训练 在显存吃紧而 batch size 又不得不往大了撑的时候Gradient Accumulation梯度累加/梯度累积几乎是每个 PyTorch 玩家的必备技能。简单说它让你在相同显存下靠“分步计算、合并更新”的方式模拟出一个更大的 batch 在训练。这个技巧在目标检测、语义分割、大规模语言模型微调里特别常用只要是 batch size 一调大就 OOM 的场景你大概率都要回头找它。这篇内容就围绕 PyTorch 框架展开说清楚梯度累加的原理、标准写法、踩坑记录和调优技巧适合已经会跑通基础训练循环、想进一步压榨显存或复现大 batch 实验的读者。1. 为什么非要梯度累加显存墙和优化目标之间的死结1.1 一次前向传播到底发生了什么要理解梯度累加先得回到 PyTorch 的 autograd 机制。你调用loss.backward()的时候PyTorch 不会立刻把梯度清空而是把计算出来的梯度累加到每个叶子张量的.grad属性上。换句话说梯度是一个“累积器”不是“赋值器”。这个特性是梯度累加的底层基础很多人第一次听说梯度累加会以为需要手动改计算图其实完全不用——你只需要控制optimizer.step()和optimizer.zero_grad()的调用时机。举个例子假设你有一份 batch size 为 32 的数据显存只够跑 batch size 8。常规训练是8 个样本前向算 loss反向更新参数再取下一批 8 个样本。但梯度累加的做法是把数据分成 4 个 mini-batch每个 mini-batch 跑一次前向和反向但是不更新参数让梯度在.grad里自然累积跑完 4 个 mini-batch 后再统一调用optimizer.step()更新一次参数并清空梯度。这样一来参数更新时看到的梯度是 32 个样本梯度的平均等价于用 batch size 32 训练的效果。显存开销却只等于单次 8 个样本的显存开销这就是梯度累加最核心的价值。1.2 为什么显存会瞬间爆炸中间激活值才是大头很多人有个误区以为显存大头是模型参数其实训练阶段真正吃显存的是中间激活值activation。以 batch size 32 为例你一次前向会保留 32 个样本每一层的中间特征反向计算梯度还需要用到这些中间值所以训练态的显存是“模型参数 中间激活值 优化器状态”三层叠加其中激活值随 batch size 线性增长。batch size 一旦从 8 翻到 32激活值就翻了四倍显存不够就会直接 OOM。梯度累加的策略本质上是把“一次吃下大 batch 的激活值”拆成“多次吃下小 batch 的激活值”时间和显存互换。你付出的代价是训练时间变长因为同样的样本量你需要多做几次前向和反向调用Python 层的调用开销和 kernel 启动开销都会增加。1.3 梯度累加到底改变了什么数学逻辑从优化器角度看常规的随机梯度下降更新公式是[ \theta_{t1} \theta_t - \eta \cdot \frac{1}{B} \sum_{i1}^{B} abla L_i(\theta_t) ]其中 (B) 是 batch size(\eta) 是学习率。梯度累加做的是把 (B) 拆成 (B N \times M)其中 (N) 是 accumulation steps(M) 是单次实际送入网络的 micro-batch size。前 (N) 次反向都只把梯度加到.grad上第 (N) 次结束后再更新一次参数[ \theta_{t1} \theta_t - \eta \cdot \frac{1}{N} \sum_{j1}^{N} \left( \frac{1}{M} \sum_{i1}^{M} abla L_{j,i}(\theta_t) \right) ]从数学上可以证明当你不会在累加中途修改模型参数时累加得到的梯度均值就等于完整大 batch 的梯度均值。这也是为什么 accumulation steps 的选择直接影响训练效果——它决定的是“用多少个小 batch 的梯度合成一次有效更新”。1.4 术语辨析梯度累积、梯度累加、梯度累计很多中文资料里“梯度累积”“梯度累加”“梯度累计”混着用甚至 framework 的官方文档也不统一。严格来说Gradient Accumulation 翻译成“梯度累加”更准确因为它的核心动作是把多个backward()产生的梯度“累加”到.grad上“累积”更偏日常用语“累计”则是统计口径的词汇。不过你搜索的时候这三个词都是一回事不必纠结。另一个容易混的是 Gradient Accumulation 与 Gradient Checkpointing梯度检查点后者是降低中间激活值占用通过“前向时丢弃部分激活值、反向时重新计算”来省显存两者可以同时使用互不冲突下文会进一步说明。2. PyTorch 里梯度累加的标准写法和完整模板2.1 三个关键 API 的配合时序梯度累加的代码核心就三个 APIloss.backward()、optimizer.step()、optimizer.zero_grad()。理解它们的调用顺序就理解了梯度累加的全部。for i, (inputs, targets) in enumerate(train_loader): outputs model(inputs) loss criterion(outputs, targets) # 反向传播梯度累加到 .grad 中 loss.backward() # 每 accumulation_steps 次才更新一次参数 if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这段伪代码是梯度累加最精简的骨架。关键在于backward()每次都会把新梯度叠加到之前的梯度上所以不调用zero_grad()梯度就不会清零只有满足(i1) % accumulation_steps 0时我们才用当前累积起来的梯度做一次step()随后把梯度清空开始下一轮的累加。有个细节要注意如果训练的总步数不是 accumulation_steps 的整数倍末尾会剩下几个 batch 的梯度没有参与参数更新。最稳妥的做法是在一个 epoch 结束时检查一下是否有“残留梯度”如果是训练中途被中断或验证时务必先清理掉残留梯度避免干扰后续计算。2.2 一个可直接复用的最小模板我常用的模板比上面的伪代码稍微完善一些加了 loss 归一化和日志输出这里直接贴出来。import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision.models import resnet18 def train_one_epoch(model, loader, criterion, optimizer, scheduler, accumulation_steps4, devicecuda): model.train() optimizer.zero_grad() running_loss 0.0 for i, (inputs, targets) in enumerate(loader): inputs, targets inputs.to(device), targets.to(device) outputs model(inputs) loss criterion(outputs, targets) # 关键反向传播前先做 loss 归一化等价于对大 batch 的 loss 取平均 loss loss / accumulation_steps loss.backward() running_loss loss.item() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad() # 若使用带 step 的 scheduler通常是每一步有效更新后调度一次 if scheduler is not None: scheduler.step() # 处理最后残留的不足 accumulation_steps 的梯度 if (i 1) % accumulation_steps ! 0: optimizer.step() optimizer.zero_grad() return running_loss / len(loader)注意这里我特意把optimizer.zero_grad()放在循环外先调用一次保证从头开始是干净的。每次有效更新后也要清空否则下一个小批次的梯度会叠加上来导致更新数值莫名其妙变大。2.3 loss 归一化一个很多人搞错的关键点我见过不少人的梯度累加代码没有把每个 micro-batch 的 loss 除以 accumulation_steps。这样会导致什么假设 accumulation_steps 为 4最后一步更新时.grad里累积的梯度是 4 个 batch 的梯度之和相当于你用 4 倍大小的学习率更新参数训练初期经常直接发疯loss 乱跳或者后期模型反复震荡无法收敛。正确的归一化方式是在每次backward()之前对当前 micro-batch 的 loss 做除法loss loss / accumulation_steps loss.backward()这样 4 次累加后的梯度均值等于大 batch 的平均梯度。本质上完整大 batch 的 loss 是每个样本 loss 的均值你每个 micro-batch 的 loss 本身也是该 micro-batch 内样本的均值那么累加时如果不除以 N最终梯度就是所有样本梯度的均值再乘以 N并非真正的均值。还有一点如果你的 loss 是多个 loss 的加权和比如目标检测里的分类 loss 加回归 loss那么请把归一化放在加权求和之后或者对加权后的 loss 整体除以 accumulation_steps不要只对其中某一个 loss 做除法否则各 loss 之间的比例关系会被破坏。2.4 accumulation_steps 怎么算从显存预算倒推怎么确定 accumulation_steps原则很简单用实验测出你的显存可以承受的最大 batch size (M)假设你想模拟的完整 batch size 是 (B)那么[ accumulation_steps \lceil B / M \rceil ]比如你的显卡跑 batch size 16 刚好是极限分配给你的一张卡只有 24GB但你想模拟 batch size 64 的效果那 accumulation_steps 就是 4。这里我不建议直接把 M 拉到显存上限因为推进 tensor 拷贝、优化器状态、临时变量都会占额外显存留 10% 到 20% 的余量更稳。2.5 当 accumulation_steps 无法整除总样本数时总样本数不一定能被 batch size 整除最后一个 mini-batch 会变小这本身没问题。但如果len(train_loader)不是 accumulation_steps 的整数倍末尾会出现“到期清零”和“数据耗尽”之间的错位。两个选择一是像我上面模板那样在循环结束后主动把残留梯度做一次 step二是干脆设置 Drop Last强制丢弃最后不足一个完整 micro-batch 的数据。我个人更倾向后者因为训练循环语义更干净测试结果也更稳定。如果你用分布式采样器通常建议drop_lastTrue否则不同卡可能拿到不同长度的数据同步梯度时产生不必要的等待。3. 实操过程中的血泪经验这些坑我全踩过3.1 BatchNorm 与梯度累加的天然冲突BatchNormBN在训练时会维护一个 batch 内的均值和方差用来归一化当前数据。梯度累加把一个大 batch 拆成多个 micro-batch 后每个 micro-batch 的 BN 统计量是独立计算的这会导致模型看到的是“小 batch 的归一化统计量”而不是“大 batch 的归一化统计量”。batch size 越小BN 统计量的噪声越大比如 batch size 8 和 batch size 32 训练出来的 BN running_mean、running_var 差别肉眼可见。这个问题没有想象中容易绕开。如果你用的模型对 BN 敏感累积步数较大比如 accumulation_steps ≥ 8时模型可能比真正的大 batch 训练效果差。几个缓解办法用 SyncBatchNorm 替代普通 BN在分布式场景下多个卡的 micro-batch 会合并计算统计量效果更接近大 batch。如果单卡训练可以把 accumulation_steps 控制在 2 或 4不要贪太多。迁移到不使用 BN 的模型架构或者改用 GroupNorm、LayerNorm 这类与 batch 大小无关的归一化。我之前做语义分割时就踩过这个坑用较大的 accumulation_steps 在单卡上模拟大 batchBN 统计量一直抖后来换成 GroupNorm 才稳定下来。如果你的任务允许调整模型结构这是最省心的一条路。3.2 梯度裁剪必须放在累加完成之后梯度裁剪gradient clipping是为了防止梯度爆炸它应该作用在最终用于更新参数的那份完整梯度上。也就是说clip_grad_norm_或clip_grad_value_必须放在optimizer.step()之前并且要紧跟在满足 accumulation_steps 条件的代码块内部而不是在每个 micro-batch 的backward()后面都执行。错误示范for i, (inputs, targets) in enumerate(loader): loss compute_loss(inputs, targets) / accumulation_steps loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) # 错误 if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()如果每个 micro-batch 都做裁剪那么前几个 micro-batch 的梯度可能已经被缩到很小最后一个 micro-batch 的梯度却可能被原样保留累加结果是“截断过的求和无序混合”完全偏离真实大 batch 的梯度形态。正确位置for i, (inputs, targets) in enumerate(loader): loss compute_loss(inputs, targets) / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() optimizer.zero_grad()3.3 学习率调度器应该跟着有效更新步数走学习率调度器scheduler也有两种派系按 iteration 调度和按 epoch 调度。使用梯度累加后你要想清楚它到底以谁为单位。常见做法是把它当作“有效更新步effective step”的调度器也就是每执行一次optimizer.step()更新一次而不是每个 mini-batch 更新一次。比如CosineAnnealingLR设置T_max30如果你的 accumulation_steps4那相当于每 4 个 iteration 才走一个 epoch 的一小步等整个训练跑完学习率变化曲线和你设想的完全不同。更直观的表述是梯度累加让“一次参数更新”对应“accumulation_steps 个 mini-batch”所以 scheduler 的 step 频率要与参数更新频率保持一致否则学习率衰减速度会快几倍。3.4 与 AMP 混合精度一起用小心 grad scaler 的顺序PyTorch 的 AMPAutomatic Mixed Precision依赖GradScaler自动放大 loss避免 fp16 梯度下溢为 0。梯度累加和 AMP 一起使用时最常见的错误是把scaler.scale(loss)放在 loss 归一化之前。因为scaler.step(optimizer)内部会根据梯度是否出现 inf/nan 来决定是否跳过这步更新同时更新缩放因子。我的建议模板如下from torch.cuda.amp import autocast, GradScaler scaler GradScaler() optimizer.zero_grad() for i, (inputs, targets) in enumerate(loader): inputs, targets inputs.to(device), targets.to(device) with autocast(): outputs model(inputs) loss criterion(outputs, targets) / accumulation_steps scaler.scale(loss).backward() if (i 1) % accumulation_steps 0: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update() optimizer.zero_grad()注意scaler.unscale_(optimizer)必须在梯度裁剪前调用。如果你不调用unscale_裁剪拿到的是被放大的梯度裁剪阈值就失去意义了。3.5 分布式训练每个进程的梯度累加不能独立理解在 DataParallelDP或 DistributedDataParallelDDP下梯度累加的行为要额外小心。DDP 默认在每次backward()时通过 all-reduce 同步各卡梯度也就是说每个 micro-batch 反向时各卡之间会同步一次梯度。这和真正的大 batch 训练并不完全一致。真大 batch 是“每张卡算完自己的子 batch再同步合并”梯度累加是“每张卡算完 micro-batch 就同步一次反复同步多次”。如果你的模型使用 SyncBN这种频繁同步会让通信开销显著上升。另一个常见错误是在 DDP 里用梯度累加却忘了在最后一步更新前做一次额外的梯度同步。DDP 会自动处理好同步你只需要保证训练循环本身不破坏 data sampler 的对齐即可推荐把drop_lastTrue打开。3.6 验证和测试阶段不要梯度累加验证和测试阶段不需要反向传播也就没有梯度累加的问题。但验证模型之前一定要确认优化器状态是干净的。最好的做法是在验证循环开始前调用model.eval()并在torch.no_grad()上下文里执行前向。如果你是在训练中途插入验证那么验证前记得手动optimizer.zero_grad()一次把积累的残留梯度清掉省得后续模型状态可复现性变差。4. 常见问题排查与效果调优4.1 数值对不上如何快速验证梯度累加实现正确性很多人写完梯度累加后心里没底不确定自己的实现是否真的等价于大 batch。有一个可复现的快速验证方法固定随机种子分别用两种方式训练一小步对比模型参数的更新量。方式 A直接使用 batch size 32 的 batch 训练一次。方式 B使用 batch size 8 的 4 个 micro-batch 做梯度累加accumulation_steps4训练一次。如果实现正确两种方式的参数更新结果应该完全一致允许浮点误差误差一般在 1e-6 量级。注意BN 的情况特殊因为 micro-batch 的统计量不一致比较结果可能出现小幅偏差这是正常的如果用 LayerNorm 或无归一化层的小模型应当严格一致。import copy import torch from torch.nn.utils import parameters_to_vector def compare_gradients(): torch.manual_seed(0) model_a SimpleNet() model_b copy.deepcopy(model_a) # 方式A大batch直接更新 optimizer_a torch.optim.SGD(model_a.parameters(), lr0.01) loss_a criterion(model_a(big_batch_data), big_batch_label) optimizer_a.zero_grad() loss_a.backward() grad_a parameters_to_vector([p.grad for p in model_a.parameters()]) optimizer_a.step() # 方式B梯度累加 optimizer_b torch.optim.SGD(model_b.parameters(), lr0.01) optimizer_b.zero_grad() for i in range(accumulation_steps): loss_b criterion(model_b(small_batch_data[i]), small_batch_label[i]) / accumulation_steps loss_b.backward() grad_b parameters_to_vector([p.grad for p in model_b.parameters()]) optimizer_b.step() print((grad_a - grad_b).abs().max().item())只要这个值非常小你的梯度累加实现基本就是正确的。4.2 训练不稳定、loss 爆炸怎么办如果加了梯度累加后 loss 震荡加剧先检查你有没有做 loss 归一化。次数最多的问题就是这个。其次检查学习率和 warmup 设置。梯度累加相当于变相增加了有效 batch size而大 batch 训练通常需要更高学习率但学习率的增幅不是线性的。常见的经验法则是线性缩放规则batch size 翻倍学习率也翻倍但前提是已有 warmup 且训练足够长。不过这个规则在超大 batch 下并不总是成立所以我一般倾向于把学习率微调幅度控制在 0.5x 到 1.5x 之间配合 warmup 慢慢试探。4.3 加了梯度累加反而 OOM 了理论上梯度累加应该省显存为什么还会 OOM常见原因有几个一是你在每个 micro-batch 前向时保留了不必要的计算图引用比如把 loss 或 output 保存到了列表里导致多个 micro-batch 的中间激活值无法释放二是优化器状态或 AMP 的梯度缩放因子在一些情况下额外占显存三是检查代码里有没有在 micro-batch 上调用torch.cuda.empty_cache()——这个函数反而会拖慢训练而且不会减少已经被占用的显存因为你当前迭代还没结束。如果确认代码逻辑没问题还是 OOM可以把 micro-batch size 再调小一点或者同时开启 Gradient Checkpointing。要注意这两者并用的原理不同、占用也不同梯度检查点是牺牲时间换中间激活值空间梯度累加是把“大 batch 峰值”摊平成“小 batch 峰值”两者叠加时确实能处理非常大的有效 batch 需求。4.4 模型效果比大 batch 差BN 和噪声是主因前面提到 BN 统计量是最可能的原因。另一个原因是梯度累加带来的“伪等价”虽然梯度的期望相同但真实大 batch 对梯度的归一化是无偏的梯度累加时每个 micro-batch 的数据分布可能会有偏差尤其是 dataset 的类别分布不均匀。如果你的数据是顺序采样前几个 micro-batch 可能全部来自某个类别累积起来梯度就有偏。这时候你应该使用随机打乱的 DataLoader最好设shuffleTrue分布式场景用RandomSampler或DistributedSampler保证每个进程的样本尽量均匀。4.5 梯度累加的替代方案哪种才是你的最优解梯度累加不是唯一的大 batch 模拟方案。根据场景可以对比方案优点缺点适用场景梯度累加通用、代码简单、省显存明显训练耗时增加、BN 统计量有偏大多数单卡/多卡训练Gradient Checkpointing大幅降低激活值显存反向重新计算耗时增加约 20%-30%长序列、深层模型混合精度AMP显存减半、速度提升需要梯度过小风险控制和 scaler 调参有 NVIDIA GPU、模型大模型并行/张量并行支持超大模型工程复杂、通信成本高百亿参数以上模型重计算 梯度累加组合显存优化极限拉满训练耗时明显增加极端显存受限场景我个人建议单卡优先用 AMP 梯度累加两层叠加基本能解决大多数“显存不够又想跑大 batch”的问题。如果模型太大再考虑 Gradient Checkpointing。模型并行是真正的“硬核方案”工程复杂度高通常是大规模预训练才需要。4.6 训练时间变长怎样缓解梯度的“重复同步”开销梯度累加本质是用时间换显存所以训练时间变长是必然的但可以优化。一个经验是尽量提高单次 micro-batch 的吞吐而不是一味把 micro-batch 调小。micro-batch 如果太小比如 batch size 1 或 2GPU 的 kernel 启动开销占比大吞吐会很差。比如同样累加 8 步micro-batch 8 累加 8 步和 micro-batch 32 累加 2 步显存余量不同速度差异可能非常明显。你应该在显存允许范围内选一个尽量大的 micro-batch再调整 accumulation_steps 补齐有效 batch size这样吞吐最优。5. 进阶技巧梯度累加在生产环境中的应用5.1 与梯度检查点组合使用突破显存极限Gradient Checkpointing 的核心是把前向传播中保存的部分中间激活值丢弃反向时重新计算因此它和梯度累加的目标不同两者组合使用可以达到“显存优化叠加态”。在 NLP 大模型微调中我经常这么做model.gradient_checkpointing_enable() accumulation_steps 8 for i, batch in enumerate(loader): loss model(**batch).loss / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这种情况下当前向计算时 model 内部不会保存所有激活值反向时再重新算显存峰值进一步降低。代价是约 20%-30% 的额外耗时。如果你连这个都嫌慢就只能上模型并行或者换更大的显存了。5.2 动态调整 accumulation_steps针对不同数据分段有些场景需要模拟的是“变长大 batch”比如 NLP 里按序列长度动态 batch 训练或者不同数据集子集建议不同的有效 batch。理论上可以在每个 batch 前动态修改 accumulation_steps但 PyTorch 的optimizer.step()并不关心你是多少步累加只要计数条件满足就行。实现上给你一个计数器和条件判断动态调整是允许的。不过我个人不太建议频繁改变 accumulation_steps因为学习率调度器和 BN 统计量会因此更不稳定如果你必须要动态调整请同步调整 loss 归一化因子。5.3 从 Gradient Accumulation 到 Gradient Checkpointing 的取舍用一句话总结我的取舍经验如果只是 batch size 不够先上梯度累加如果模型本身很大、单样本激活值就很占显存先上梯度检查点如果又大又要大 batch那就两个一起上。前提是评估时间成本。我曾经在两个 16GB V100 上跑一个 7B 模型微调batch size 只能到 2用梯度累加 16 步模拟 32 的 batch训练速度慢了 4 倍左右当时还是能接受因为任务本身不赶时间。如果业务在线推理有延迟要求那就另当别论。5.4 和 EMA 或模型权重平均搭配时的注意点如果你用 Exponential Moving AverageEMA或 Stochastic Weight AveragingSWA这类基于权重平均的优化策略需要保证“一次 EMA 更新”对应“一次真实参数更新”。也就是说EMA 的num_updates计数应该放在optimizer.step()之后而不是每个 micro-batch 后都更新否则 EMA 会被过多的中间权重污染模型精度反而下降。5.5 日志里应该记录什么有效步数和平均 loss用梯度累加后日志系统也需要调整。常见错误是把每个 micro-batch 的 loss 都打到 TensorBoard导致曲线毛刺极多几乎无法观察收敛趋势。正确做法是以“有效更新步”为周期记录平均 loss。也就是每执行一次参数更新把过去 accumulation_steps 个 micro-batch 的平均 loss 打一个点。还可以额外记录以下几个指标effective_batch_size micro_batch_size * accumulation_steps * world_sizegrad_norm每次更新前的梯度范数用于判断是否梯度爆炸lr每个有效更新步的实时学习率这些信息能帮你快速判断训练状态远比打印每一步的 loss 有价值。6. 最后再分享一点我的个人体会做深度学习训练调优这些年我最大的感受是显存永远不够用计算资源永远紧张但很多工程问题不是靠蛮力换显卡解决的。梯度累加是 PyTorch 里少有的“改动极小、收益极大、但细节极多”的技巧。它的文档很短但真正稳定跑起来需要你对 autograd 机制、优化器行为、BN 特性和分布式通信有整体理解。我见过太多人栽在 loss 归一化、梯度裁剪位置、scheduler 步数这些不起眼的细节上希望这篇内容能帮你把坑提前避开。如果你要复现大规模论文的实验梯度累加几乎是绕不开的标配操作先在小任务上验证数值一致性再把服务稳定跑起来这套思路我在多个项目里反复验证过从来没让我失望过。
返回列表