怎样实现梯度累积技术

GPU
小华
2026-08-17

梯度累积(Gradient Accumulation)是一种在显存有限的情况下,模拟大 batch size 训练的技术。其核心思想是:在多个小 batch 上分别计算梯度,但不立即更新参数,而是把梯度累加起来,累积到一定步数后再一次性更新模型参数并清空梯度

下面从原理、实现步骤、代码示例(PyTorch)以及注意事项几个方面说明。

一、梯度累积的原理

假设:

  • 目标 batch size = N
  • 显存只能支持 MM < N
  • 梯度累积步数 K = N / M

训练流程变为:

  1. 前向传播(小 batch)
  2. 反向传播(计算梯度)
  3. 不更新参数,只累积梯度
  4. 重复 K 次
  5. 用累积的梯度更新参数
  6. 清空梯度,进入下一轮

数学上等价于一次性用 N 个样本计算梯度。

二、实现步骤(通用)

  1. 设置 accumulation_steps
  2. 每个 step:
  • loss.backward()(梯度自动累加)
  1. accumulation_steps 次:
  • optimizer.step()
  • optimizer.zero_grad()

⚠️ 注意

  • loss 通常需要除以 accumulation_steps(或 batch size),避免梯度尺度不一致
  • zero_grad() 只在累积结束后调用

三、PyTorch 示例代码

1️⃣ 基本示例

model.train()
optimizer.zero_grad()
accumulation_steps = 4
for i, (inputs, targets) in enumerate(dataloader):
outputs = model(inputs)
loss = criterion(outputs, targets)
# 梯度累积时,loss 要归一化
loss = loss / accumulation_steps
loss.backward()
# 每 accumulation_steps 更新一次
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()

2️⃣ 完整训练循环示例

model.train()
optimizer.zero_grad()
accumulation_steps = 4
total_loss = 0.0
for epoch in range(num_epochs):
for i, (inputs, targets) in enumerate(dataloader):
outputs = model(inputs)
loss = criterion(outputs, targets) / accumulation_steps
loss.backward()
total_loss += loss.item()
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
print(f"Step {i+1}, Loss: {total_loss:.4f}")
total_loss = 0.0

四、与 Batch Size 的关系

实际 batch size累积步数等效 batch size
16464
8864

⚠️ 注意:

  • 学习率通常需要随等效 batch size 增大而调整
  • 大 batch 通常需要更大的 learning rate 或使用 warmup

五、常见注意事项

✅ 1. Batch Normalization

  • BN 的统计量是基于当前小 batch 的
  • 梯度累积 不会 增大 BN 的 batch size
  • 解决方案:
  • 使用 torch.nn.SyncBatchNorm
  • 或改用 LayerNorm / GroupNorm

✅ 2. 混合精度训练(AMP)

from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
optimizer.zero_grad()
accumulation_steps = 4
for i, (inputs, targets) in enumerate(dataloader):
with autocast():
outputs = model(inputs)
loss = criterion(outputs, targets) / accumulation_steps
scaler.scale(loss).backward()
if (i + 1) % accumulation_steps == 0:
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()

✅ 3. 梯度裁剪(Gradient Clipping)

torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

通常在 optimizer.step() 前执行。

六、什么时候使用梯度累积?

✅ 显存不足
✅ 想用大 batch size 提升稳定性
✅ 训练大模型(LLM、ViT、Diffusion)
❌ 不适合:

  • 对 BN 非常敏感的任务
  • 实时在线学习

如果你愿意,我可以:

  • ✅ 帮你把现有训练代码改成梯度累积版本
  • ✅ 结合 DDP / DeepSpeed / HuggingFace Trainer 讲实现
  • ✅ 讲 梯度累积与梯度同步 的区别

只要把你的代码或框架告诉我即可。

亿速云提供售前/售后服务

售前业务咨询

售后技术保障

400-100-2938

7*24小时售后电话

官方微信小程序