什么是梯度累积技术详解

GPU
小华
2026-09-16

梯度累积(Gradient Accumulation)技术详解

一、什么是梯度累积

梯度累积是一种在深度学习训练过程中,在不增加显存占用的前提下,模拟更大 batch size 训练效果的技术。
其核心思想是:

将多个小 batch 的前向和反向传播分开进行,把计算得到的梯度累加(accumulate),直到达到目标 batch size 后,再统一更新一次模型参数

二、为什么需要梯度累积

1. 显存限制

  • GPU 显存有限
  • 大 batch size 会导致:
  • 激活值(activations)占用显存暴增
  • 直接 OOM(Out of Memory)

2. 训练稳定性

  • 大 batch 通常:
  • 梯度更稳定
  • 收敛更平滑
  • 但硬件不允许直接用大 batch

3. 解决思路

用时间换空间

  • 小 batch 计算
  • 多次累积梯度
  • 等效于大 batch

三、数学原理

假设目标 batch size 为 B,实际可用 batch size 为 b,累积步数为:
[
K = frac{B}{b}
]

普通训练(无累积)

每次:
[
theta = theta - eta cdot g
]

梯度累积训练

对于第 k 个 micro-batch:
[
g_{acc} = g_{acc} + g_k
]
k == K 时:
[
theta = theta - eta cdot frac{g_{acc}}{K}
]

注意:是否除以 K 取决于框架实现(有些自动平均,有些需手动处理)

四、训练流程图

┌────────────┐
│ micro-batch 1 │ → forward → backward → 累加梯度
├────────────┤
│ micro-batch 2 │ → forward → backward → 累加梯度
├────────────┤
│     ...      │
├────────────┤
│ micro-batch K │ → forward → backward → 累加梯度
│              │ → optimizer.step()
│              │ → optimizer.zero_grad()
└────────────┘

五、代码示例

PyTorch 示例

accumulation_steps = 4
optimizer.zero_grad()
for i, (x, y) in enumerate(dataloader):
pred = model(x)
loss = criterion(pred, y)
loss = loss / accumulation_steps
loss.backward()
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()

关键点说明

  • loss / accumulation_steps:保证梯度尺度正确
  • optimizer.step() 只在累积完成后调用
  • zero_grad() 防止梯度泄漏

六、梯度累积 vs 大 Batch

对比项大 Batch梯度累积
显存占用
计算效率稍低
数值结果完全一致近似一致
实现复杂度
在理想情况下,二者训练结果几乎等价

七、常见注意事项

1. Batch Normalization

  • BN 依赖当前 batch 统计量
  • 梯度累积时:
  • BN 仍按 micro-batch 计算
  • 可能与真实大 batch 有偏差

✅ 解决方案:

  • 使用 SyncBN
  • 或训练后期再开大 batch

2. 学习率

  • 等效 batch 变大
  • 通常可 线性放大学习率

3. 随机性

  • Dropout
  • 数据顺序
  • 仍按 micro-batch 执行

八、典型应用场景

  • ✅ 大模型训练(LLM)
  • ✅ 显存受限的 GPU
  • ✅ 多卡通信成本高的场景
  • ✅ 科研实验复现大 batch 设置

九、总结

梯度累积的本质:

用多次小batch的反向传播,等价替代一次大batch的反向传播
它是:
  • 显存友好
  • 工程实用
  • 大模型时代必备技术

如果你愿意,我也可以:

  • 画一张更直观的示意图
  • 对比 TensorFlow / PyTorch 实现差异
  • 讲梯度累积 + 混合精度的组合用法
亿速云提供售前/售后服务

售前业务咨询

售后技术保障

400-100-2938

7*24小时售后电话

官方微信小程序