【Bug已解决】Understanding accumulated gradients in PyTorch 解决方案
2026/9/16 5:12:02 网站建设 项目流程

【Bug已解决】Understanding accumulated gradients in PyTorch 解决方案

本文深入解析 PyTorch 中梯度累积(Accumulated Gradients)的原理、常见误区与正确实现方式,帮助你在显存受限场景下模拟更大 batch size 进行训练。

问题描述

在深度学习训练中,batch size 是影响模型收敛的关键超参数。受限于 GPU 显存,很多时候无法使用足够大的 batch size。梯度累积(Gradient Accumulation)通过在多个 mini-batch 上前向传播并累积梯度,累积到一定步数后统一执行一次反向更新,从而在不增加显存占用的前提下模拟更大的 batch size。

但在 PyTorch 中实现梯度累积时,开发者经常遇到以下问题:

  1. 梯度没有被正确清零optimizer.zero_grad()调用时机不对,导致梯度不断累加,训练发散。
  2. loss 缩放问题:累积多个 step 的 loss 后没有正确缩放,导致有效学习率偏大或偏小。
  3. 与 AMP 配合使用时的 GradScaler 问题scaler.scale(loss).backward()scaler.step()调用时机混乱。
  4. DDP 场景下的梯度同步no_sync()上下文管理器使用不当,导致梯度在累积期间被意外同步。

错误复现

以下是一个典型的错误实现:

import torch import torch.nn as nn from torch.utils.data import DataLoader, TensorDataset model = nn.Sequential( nn.Linear(784, 256), nn.ReLU(), nn.Linear(256, 10) ).cuda() data = torch.randn(1000, 784).cuda() labels = torch.randint(0, 10, (1000,)).cuda() dataset = TensorDataset(data, labels) batch_size = 16 accumulation_steps = 4 dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) # 错误实现:每个 step 都调用 zero_grad 和 step,梯度累积完全失效 for epoch in range(5): for i, (inputs, targets) in enumerate(dataloader): outputs = model(inputs) loss = criterion(outputs, targets) loss.backward() optimizer.step() optimizer.zero_grad() # 每步都清零,累积无效

更危险的错误是完全忘记zero_grad()

# 错误:没有 zero_grad,梯度无限累积 for epoch in range(5): for i, (inputs, targets) in enumerate(dataloader): loss = criterion(model(inputs), targets) loss.backward() optimizer.step() # 没有 zero_grad()!梯度会无限累积,loss 迅速变为 NaN

根因分析

1. PyTorch 梯度累积的底层机制

PyTorch 的 Autograd 在调用loss.backward()时,会将梯度累加到各参数的.grad属性中,而非覆盖。这是设计决策,正是为了支持梯度累积:

w = torch.tensor([1.0], requires_grad=True) (w ** 2).sum().backward() print(w.grad) # tensor([2.]) (w ** 2).sum().backward() print(w.grad) # tensor([4.]) -- 累加了!

因此梯度累积的核心逻辑是:累积阶段多次backward()step()zero_grad();更新阶段step()zero_grad()

2. Loss 缩放的必要性

累积accumulation_steps步后更新时,等效 loss 应是各 mini-batch loss 的平均值。不缩放会导致等效学习率偏大accumulation_steps倍:

# 正确:loss = loss / accumulation_steps,使累积梯度 = mean(grad_i)

3. AMP 场景的复杂性

GradScalerscaler.step()时检查梯度是否溢出。非更新步调用scaler.step()会导致 scaler 错误跳过更新或调整缩放因子。

4. DDP 梯度同步

DDP 中每次backward()都触发 all-reduce 同步。累积阶段应使用model.no_sync()跳过同步,只在最后一步正常同步。

解决方案

方案一:基础梯度累积(单 GPU)

import torch import torch.nn as nn from torch.utils.data import DataLoader, TensorDataset class SimpleModel(nn.Module): def __init__(self, input_dim=784, hidden_dim=256, num_classes=10): super().__init__() self.fc1 = nn.Linear(input_dim, hidden_dim) self.relu = nn.ReLU() self.dropout = nn.Dropout(0.2) self.fc2 = nn.Linear(hidden_dim, num_classes) def forward(self, x): x = self.relu(self.fc1(x)) x = self.dropout(x) return self.fc2(x) def train_with_gradient_accumulation(): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = SimpleModel().to(device) data = torch.randn(2000, 784) labels = torch.randint(0, 10, (2000,)) dataset = TensorDataset(data, labels) batch_size = 16 accumulation_steps = 4 dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10) epochs = 10 model.train() for epoch in range(epochs): total_loss = 0.0 num_updates = 0 optimizer.zero_grad() for i, (inputs, targets) in enumerate(dataloader): inputs, targets = inputs.to(device), targets.to(device) outputs = model(inputs) loss = criterion(outputs, targets) scaled_loss = loss / accumulation_steps # 关键:缩放 loss scaled_loss.backward() total_loss += loss.item() if (i + 1) % accumulation_steps == 0: torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() optimizer.zero_grad() num_updates += 1 # 处理最后一个不完整组 if (len(dataloader) % accumulation_steps) != 0: torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() optimizer.zero_grad() num_updates += 1 scheduler.step() print(f"Epoch {epoch+1}/{epochs}, Avg Loss: {total_loss/len(dataloader):.4f}, Updates: {num_updates}") return model if __name__ == '__main__': train_with_gradient_accumulation()

方案二:封装通用训练器

import torch import torch.nn as nn from torch.utils.data import DataLoader from typing import Optional class GradientAccumulationTrainer: def __init__(self, model, optimizer, criterion, accumulation_steps=1, max_grad_norm=1.0, use_amp=False, scheduler=None, device='cuda'): self.model = model.to(device) self.optimizer = optimizer self.criterion = criterion self.accumulation_steps = accumulation_steps self.max_grad_norm = max_grad_norm self.use_amp = use_amp self.scheduler = scheduler self.device = device self.scaler = torch.cuda.amp.GradScaler(enabled=use_amp) self.global_step = 0 self.accumulation_count = 0 def train_step(self, inputs, targets, step_in_epoch, total_steps): inputs, targets = inputs.to(self.device), targets.to(self.device) self.model.train() with torch.cuda.amp.autocast(enabled=self.use_amp): outputs = self.model(inputs) loss = self.criterion(outputs, targets) scaled_loss = loss / self.accumulation_steps if self.use_amp: self.scaler.scale(scaled_loss).backward() else: scaled_loss.backward() self.accumulation_count += 1 should_update = (self.accumulation_count >= self.accumulation_steps) or (step_in_epoch == total_steps - 1) if should_update: if self.max_grad_norm is not None: if self.use_amp: self.scaler.unscale_(self.optimizer) torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.max_grad_norm) if self.use_amp: self.scaler.step(self.optimizer) self.scaler.update() else: self.optimizer.step() self.optimizer.zero_grad() self.global_step += 1 self.accumulation_count = 0 return loss.item() def train(self, dataloader, epochs): for epoch in range(1, epochs + 1): total_loss = 0.0 total_steps = len(dataloader) ![配图](https://i-blog.csdnimg.cn/img_convert/809cc0eba38616a9986eb376e055893e.png) for step, (inputs, targets) in enumerate(dataloader): loss = self.train_step(inputs, targets, step, total_steps) total_loss += loss if step % 100 == 0: lr = self.optimizer.param_groups[0]['lr'] print(f"Epoch {epoch} Step {step}/{total_steps} | Loss: {loss:.4f} | LR: {lr:.2e}") if self.scheduler: self.scheduler.step() print(f"Epoch {epoch} | Avg Loss: {total_loss/total_steps:.4f}") print(f"训练完成!总更新步数: {self.global_step}")

方案三:DDP 分布式场景

import torch import torch.nn as nn import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader, DistributedSampler, TensorDataset import os def setup_ddp(rank, world_size): os.environ['MASTER_ADDR'] = 'localhost' os.environ['MASTER_PORT'] = '12355' dist.init_process_group('nccl', rank=rank, world_size=world_size) def train_ddp_with_accumulation(rank, world_size): setup_ddp(rank, world_size) device = torch.device(f'cuda:{rank}') torch.cuda.set_device(device) model = nn.Sequential(nn.Linear(784, 256), nn.ReLU(), nn.Linear(256, 10)).to(device) model = DDP(model, device_ids=[rank]) data = torch.randn(2000, 784) labels = torch.randint(0, 10, (2000,)) dataset = TensorDataset(data, labels) sampler = DistributedSampler(dataset, num_replicas=world_size, rank=rank, shuffle=True) dataloader = DataLoader(dataset, batch_size=16, sampler=sampler) accumulation_steps = 4 criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3) for epoch in range(5): sampler.set_epoch(epoch) model.train() optimizer.zero_grad() for i, (inputs, targets) in enumerate(dataloader): inputs, targets = inputs.to(device), targets.to(device) is_last = (i + 1) % accumulation_steps == 0 or i == len(dataloader) - 1 # 关键:非最后一步用 no_sync 避免梯度同步 context = model.no_sync() if not is_last else torch.enable_grad() with context: outputs = model(inputs) loss = criterion(outputs, targets) / accumulation_steps loss.backward() if is_last: torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() optimizer.zero_grad() if rank == 0 and i % 50 == 0: print(f"Epoch {epoch}, Step {i}, Loss: {loss.item()*accumulation_steps:.4f}") dist.destroy_process_group()

完整修复代码

以下是一个生产级别的完整实现,集成梯度累积、AMP、梯度裁剪、学习率调度和检查点保存:

"""完整的梯度累积训练脚本""" import torch import torch.nn as nn from torch.utils.data import DataLoader, TensorDataset, random_split import math, os, time, json from typing import Optional, Dict class Config: data_size = 5000 input_dim = 784 num_classes = 10 hidden_dim = 512 dropout = 0.2 batch_size = 16 accumulation_steps = 8 # 等效 batch size = 128 epochs = 20 learning_rate = 1e-3 weight_decay = 0.01 max_grad_norm = 1.0 use_amp = True warmup_steps = 100 min_lr = 1e-6 checkpoint_dir = './checkpoints' save_every_n_steps = 200 log_every_n_steps = 50 class ClassifierModel(nn.Module): def __init__(self, config): super().__init__() self.net = nn.Sequential( nn.Linear(config.input_dim, config.hidden_dim), nn.BatchNorm1d(config.hidden_dim), nn.ReLU(), nn.Dropout(config.dropout), nn.Linear(config.hidden_dim, config.hidden_dim // 2), nn.BatchNorm1d(config.hidden_dim // 2), nn.ReLU(), nn.Dropout(config.dropout), nn.Linear(config.hidden_dim // 2, config.num_classes) ) def forward(self, x): return self.net(x) class WarmupCosineScheduler: def __init__(self, optimizer, warmup_steps, total_steps, min_lr, max_lr): self.optimizer = optimizer self.warmup_steps = warmup_steps self.total_steps = total_steps self.min_lr = min_lr self.max_lr = max_lr self.current_step = 0 def step(self): self.current_step += 1 if self.current_step < self.warmup_steps: lr = self.max_lr * (self.current_step / self.warmup_steps) else: progress = min(1.0, (self.current_step - self.warmup_steps) / max(1, self.total_steps - self.warmup_steps)) lr = self.min_lr + 0.5 * (self.max_lr - self.min_lr) * (1 + math.cos(math.pi * progress)) for pg in self.optimizer.param_groups: pg['lr'] = lr return lr class GradAccumTrainer: def __init__(self, config): self.config = config self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') self.model = ClassifierModel(config).to(self.device) self.optimizer = torch.optim.AdamW(self.model.parameters(), lr=config.learning_rate, weight_decay=config.weight_decay) self.criterion = nn.CrossEntropyLoss() self.scaler = torch.cuda.amp.GradScaler(enabled=config.use_amp and self.device.type == 'cuda') total_steps = (config.data_size // config.batch_size) * config.epochs // config.accumulation_steps self.scheduler = WarmupCosineScheduler(self.optimizer, config.warmup_steps, total_steps, config.min_lr, config.learning_rate) self.global_step = 0 self.best_loss = float('inf') self.history = [] os.makedirs(config.checkpoint_dir, exist_ok=True) def _save_checkpoint(self, path, extra=None): ckpt = {'model': self.model.state_dict(), 'optimizer': self.optimizer.state_dict(), 'scaler': self.scaler.state_dict(), 'global_step': self.global_step, 'best_loss': self.best_loss} if extra: ckpt.update(extra) torch.save(ckpt, path) def train(self, train_loader, val_loader=None, resume=None): if resume and os.path.exists(resume): ckpt = torch.load(resume, map_location=self.device) self.model.load_state_dict(ckpt['model']) self.optimizer.load_state_dict(ckpt['optimizer']) self.scaler.load_state_dict(ckpt['scaler']) self.global_step = ckpt['global_step'] print(f"恢复训练: global_step={self.global_step}") acc_steps = self.config.accumulation_steps print(f"开始训练 | 设备: {self.device} | 等效batch: {self.config.batch_size * acc_steps}") for epoch in range(1, self.config.epochs + 1): self.model.train() epoch_loss = 0.0 num_batches = 0 self.optimizer.zero_grad() acc_count = 0 for step, (inputs, targets) in enumerate(train_loader): inputs = inputs.to(self.device, non_blocking=True) targets = targets.to(self.device, non_blocking=True) with torch.cuda.amp.autocast(enabled=self.config.use_amp and self.device.type == 'cuda'): outputs = self.model(inputs) loss = self.criterion(outputs, targets) scaled_loss = loss / acc_steps self.scaler.scale(scaled_loss).backward() epoch_loss += loss.item() num_batches += 1 acc_count += 1 should_update = (acc_count >= acc_steps) or (step == len(train_loader) - 1) if should_update: if self.config.max_grad_norm is not None: self.scaler.unscale_(self.optimizer) torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.config.max_grad_norm) self.scaler.step(self.optimizer) self.scaler.update() self.optimizer.zero_grad() self.scheduler.step() acc_count = 0 self.global_step += 1 if self.global_step % self.config.log_every_n_steps == 0: lr = self.optimizer.param_groups[0]['lr'] print(f"Epoch {epoch} | Step {self.global_step} | Loss: {epoch_loss/num_batches:.4f} | LR: {lr:.2e}") avg_loss = epoch_loss / num_batches val_loss, val_acc = (None, None) if val_loader: val_loss, val_acc = self.validate(val_loader) self.history.append({'epoch': epoch, 'train_loss': avg_loss, 'val_loss': val_loss, 'val_acc': val_acc}) print(f"Epoch {epoch}/{self.config.epochs} | Train: {avg_loss:.4f}", end='') if val_loss: print(f" | Val: {val_loss:.4f} Acc: {val_acc:.2%}", end='') print() current = val_loss if val_loss else avg_loss if current < self.best_loss: self.best_loss = current self._save_checkpoint(os.path.join(self.config.checkpoint_dir, 'best.pt'), {'epoch': epoch}) self._save_checkpoint(os.path.join(self.config.checkpoint_dir, 'final.pt')) with open(os.path.join(self.config.checkpoint_dir, 'history.json'), 'w') as f: json.dump(self.history, f, indent=2) print(f"训练完成!总步数: {self.global_step} | 最佳: {self.best_loss:.4f}") @torch.no_grad() def validate(self, loader): self.model.eval() total_loss, correct, total = 0.0, 0, 0 for inputs, targets in loader: inputs, targets = inputs.to(self.device), targets.to(self.device) with torch.cuda.amp.autocast(enabled=self.config.use_amp): outputs = self.model(inputs) loss = self.criterion(outputs, targets) total_loss += loss.item() correct += outputs.max(1)[1].eq(targets).sum().item() total += targets.size(0) return total_loss / len(loader), correct / total def main(): config = Config() data = torch.randn(config.data_size, config.input_dim) labels = torch.randint(0, config.num_classes, (config.data_size,)) dataset = TensorDataset(data, labels) train_size = int(0.8 * len(dataset)) train_ds, val_ds = random_split(dataset, [train_size, len(dataset) - train_size]) train_loader = DataLoader(train_ds, batch_size=config.batch_size, shuffle=True) val_loader = DataLoader(val_ds, batch_size=config.batch_size, shuffle=False) trainer = GradAccumTrainer(config) trainer.train(train_loader, val_loader) if __name__ == '__main__': main()

常见陷阱与注意事项

1. zero_grad 调用时机

最常见错误是每个 mini-batch 后都调用zero_grad(),使梯度累积失效。正确做法是只在step()之后调用:

# 正确 if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad() # 只在更新后清零

2. BatchNorm 与梯度累积的冲突

BatchNorm 在每个 mini-batch 上计算统计量,梯度累积无法真正模拟大 batch 对 BN 的影响。解决方案:使用GroupNorm/LayerNorm替代,或使用SyncBatchNorm(DDP 场景),或确保 mini-batch 足够大。

3. Dropout 的随机性

Dropout 在每个 mini-batch 独立采样 mask,梯度累积不会改变这一行为,等效效果可能与真正大 batch 略有不同。

4. 学习率调整

使用梯度累积时,学习率应按等效 batch size 设置。经验法则:等效 batch size 翻倍,学习率可增加约 1.4 倍。

5. AMP GradScaler 注意事项

# 正确的 AMP + 梯度累积 scaler.scale(loss / accumulation_steps).backward() # 累积阶段 if should_update: scaler.unscale_(optimizer) # 先 unscale 再裁剪 clip_grad_norm_(model.parameters(), max_norm) scaler.step(optimizer) scaler.update() # 更新缩放因子 optimizer.zero_grad()

scaler.update()应在每次scaler.step()后调用。

6. 最后一个不完整 batch

len(dataloader) % accumulation_steps != 0时,直接用不足的梯度更新(推荐),或设drop_last=True丢弃。

7. DDP no_sync 正确使用

is_last = (i + 1) % accumulation_steps == 0 with model.no_sync() if not is_last else torch.enable_grad(): loss.backward() if is_last: optimizer.step() optimizer.zero_grad()

no_sync()只影响反向传播梯度同步,不影响前向传播和 BatchNorm 统计。

8. 梯度裁剪时机

裁剪应在所有梯度累积完成后、step()之前执行。AMP 下需先unscale_()再裁剪。

9. 检查点完整性

保存检查点时除模型和优化器状态,还应保存scaler状态和global_step

checkpoint = { 'model': model.state_dict(), 'optimizer': optimizer.state_dict(), 'scaler': scaler.state_dict(), 'global_step': global_step, }

10. 验证梯度累积正确性

可通过比较梯度范数验证:

# 真实大 batch 的梯度范数 vs 梯度累积的梯度范数,应非常接近 grad_norm = torch.norm(torch.cat([p.grad.flatten() for p in model.parameters() if p.grad is not None]))

总结

梯度累积是 PyTorch 中模拟大 batch size 训练的重要技术,其核心原理是利用 Autograd 的梯度累加特性。正确实现需要注意三个关键点:按累积步数缩放 loss在正确的时机调用step()zero_grad()与 AMP/DDP 正确配合

本文从底层机制出发,详细分析了梯度累积的原理和常见陷阱,提供了从基础实现到生产级训练器的完整代码。关键要点:loss.backward()会累加梯度而非覆盖;optimizer.zero_grad()只在step()后调用;AMP 下用scaler.unscale_()后再裁剪梯度;DDP 下用no_sync()避免累积期间的冗余同步。掌握这些细节,就能在显存受限的场景下灵活运用梯度累积策略,有效提升训练效果。

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

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

立即咨询