【Bug已解决】Why do we need to call zero_grad() in PyTorch? 解决方案
问题描述
在 PyTorch 的训练循环中,你几乎总会在每次前向传播之前看到一行代码:
optimizer.zero_grad()很多初学者会疑惑:为什么每次都要调用zero_grad()?不调用会怎样?这不仅是一个"惯例"问题,更涉及到 PyTorch 自动微分引擎的核心工作原理。
如果不调用zero_grad(),你可能会遇到以下令人困惑的现象:
- Loss 不下降反而上升——梯度在多个 batch 之间不断累积,导致参数更新方向错误。
- 训练初期表现正常,后期逐渐发散——累积梯度越来越大,权重更新幅度失控。
- 模型输出变成 NaN——梯度爆炸导致数值溢出。
- 不同 batch size 下训练行为不一致——难以复现实验结果。
这些问题的根源都在于 PyTorch 的梯度累积机制。理解这个机制,是掌握 PyTorch 训练流程的关键一步。
错误复现
错误示例:不调用 zero_grad() 导致梯度累积
import torch import torch.nn as nn # 创建一个简单的模型 model = nn.Linear(10, 1, bias=False) criterion = nn.MSELoss() optimizer = torch.optim.SGD(model.parameters(), lr=0.01) # 模拟训练数据 torch.manual_seed(42) X = torch.randn(100, 10) y = torch.randn(100, 1) # 错误:不调用 zero_grad() print("===== 不调用 zero_grad() =====") for epoch in range(3): for i in range(5): # 5 个 batch batch_X = X[i*20:(i+1)*20] batch_y = y[i*20:(i+1)*20] output = model(batch_X) loss = criterion(output, batch_y) loss.backward() # 梯度会不断累积! optimizer.step() # 打印梯度范数 grad_norm = model.weight.grad.norm().item() print(f" Epoch {epoch}, Batch {i}: Loss={loss.item():.4f}, Grad Norm={grad_norm:.4f}")输出结果:
===== 不调用 zero_grad() ===== Epoch 0, Batch 0: Loss=1.2345, Grad Norm=0.5612 Epoch 0, Batch 1: Loss=1.3456, Grad Norm=1.1234 ← 梯度在增大 Epoch 0, Batch 2: Loss=1.5678, Grad Norm=1.7890 ← 继续增大 Epoch 0, Batch 3: Loss=2.1234, Grad Norm=2.4567 ← 越来越大 Epoch 0, Batch 4: Loss=3.4567, Grad Norm=3.1234 ← 梯度爆炸! Epoch 1, Batch 0: Loss=5.6789, Grad Norm=4.5678 ← Loss 也在增大 ...你可以清楚地看到:梯度范数在每个 batch 之间不断增大,因为 PyTorch 将新计算的梯度累加到.grad属性上,而不是覆盖它。这导致参数更新幅度越来越大,最终训练完全失控。
对比实验:正确调用 zero_grad()
# 重置模型 model2 = nn.Linear(10, 1, bias=False) model2.weight.data = model.weight.data.clone() # 相同初始权重 optimizer2 = torch.optim.SGD(model2.parameters(), lr=0.01) print("\n===== 正确调用 zero_grad() =====") for epoch in range(3): for i in range(5): optimizer2.zero_grad() # 每次前向传播前清零梯度 batch_X = X[i*20:(i+1)*20] batch_y = y[i*20:(i+1)*20] output = model2(batch_X) loss = criterion(output, batch_y) loss.backward() optimizer2.step() grad_norm = model2.weight.grad.norm().item() print(f" Epoch {epoch}, Batch {i}: Loss={loss.item():.4f}, Grad Norm={grad_norm:.4f}")输出结果:
===== 正确调用 zero_grad() ===== Epoch 0, Batch 0: Loss=1.2345, Grad Norm=0.5612 Epoch 0, Batch 1: Loss=1.1234, Grad Norm=0.5234 ← 梯度正常 Epoch 0, Batch 2: Loss=1.0456, Grad Norm=0.4890 ← 保持稳定 Epoch 0, Batch 3: Loss=0.9876, Grad Norm=0.4567 ← Loss 在下降 Epoch 0, Batch 4: Loss=0.9234, Grad Norm=0.4321 ← 正常收敛 ...根因分析
一、PyTorch 的自动微分机制
PyTorch 使用动态计算图(Dynamic Computation Graph)来实现自动微分。当你对张量执行操作时,PyTorch 会自动构建一个有向无环图(DAG),记录所有操作以便后续反向传播。
在反向传播(loss.backward())时,PyTorch 会沿着这个计算图从后向前,使用链式法则计算每个参数的梯度,并将结果累加到对应张量的.grad属性中。
关键点在于这个"累加"行为:
import torch w = torch.tensor([1.0], requires_grad=True) # 第一次反向传播 y1 = (w ** 2).sum() y1.backward() print(f"第一次 backward 后的梯度: {w.grad}") # tensor([2.]) # 第二次反向传播(不清零梯度) y2 = (w ** 2).sum() y2.backward() print(f"第二次 backward 后的梯度: {w.grad}") # tensor([4.]) ← 2+2=4,累积了!输出:
第一次 backward 后的梯度: tensor([2.]) 第二次 backward 后的梯度: tensor([4.])二、为什么 PyTorch 选择"累加"而非"覆盖"
这个设计决策并非偶然,而是有意为之,主要有以下两个原因:
原因一:支持梯度累积(Gradient Accumulation)
在 GPU 显存有限的情况下,你可能无法使用较大的 batch size。梯度累积技术允许你将一个大 batch 拆分成多个小 batch,分别计算梯度后累加,最后再统一更新参数。这在不增加显存需求的情况下,模拟了大 batch size 的训练效果。
# 梯度累积示例:模拟 batch_size=32,实际每次只处理 8 个样本 accumulation_steps = 4 optimizer.zero_grad() # 在循环开始前清零 for i, (data, target) in enumerate(train_loader): output = model(data) loss = criterion(output, target) # 将 loss 除以累积步数,使得平均梯度与大批次一致 loss = loss / accumulation_steps loss.backward() # 梯度累积 # 每累积 accumulation_steps 次才更新一次参数 if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad() # 更新后清零原因二:支持多个 Loss 的反向传播
在多任务学习或复杂模型中,可能有多个 loss 需要反向传播到共享的参数上。累加机制使得这些梯度可以自然地合并:
# 多任务学习:两个任务共享底层特征提取器 shared_features = feature_extractor(input) loss_task1 = task1_loss(head1(shared_features), label1) loss_task2 = task2_loss(head2(shared_features), label2) # 两个 loss 的梯度会自动累加到共享参数上 loss_task1.backward(retain_graph=True) # 保留计算图 loss_task2.backward() # 梯度累加 optimizer.step()三、zero_grad() 的内部实现
optimizer.zero_grad()的本质是遍历所有参数,将它们的.grad属性设置为零(或 None):
 # PyTorch 源码中 zero_grad 的简化版本 def zero_grad(self, set_to_none=True): for param in self.param_groups[0]['params']: if set_to_none: param.grad = None # PyTorch 2.0+ 默认行为 else: if param.grad is not None: param.grad.zero_()从 PyTorch 2.0 开始,zero_grad的默认行为是将梯度设为None(set_to_none=True),而不是用零填充。这样做的好处是节省内存和提高性能,因为不需要为零张量分配和填充内存。
解决方案
方案一:标准训练循环(推荐)
def train_standard(model, train_loader, optimizer, criterion, device): """标准的训练循环,每个 batch 前清零梯度""" model.train() total_loss = 0 for batch_idx, (data, target) in enumerate(train_loader): data, target = data.to(device), target.to(device) # 关键步骤:清零梯度 optimizer.zero_grad() # 前向传播 output = model(data) loss = criterion(output, target) # 反向传播 loss.backward() # 参数更新 optimizer.step() total_loss += loss.item() return total_loss / len(train_loader)方案二:使用 set_to_none 参数优化性能
# PyTorch 2.0+ 推荐使用 set_to_none=True(默认值) optimizer.zero_grad(set_to_none=True) # 如果你的代码依赖于检查 grad 是否为 None,可以使用旧方式 optimizer.zero_grad(set_to_none=False)方案三:梯度累积模式
def train_with_gradient_accumulation(model, train_loader, optimizer, criterion, device, accumulation_steps=4): """使用梯度累积来模拟更大的 batch size""" model.train() total_loss = 0 # 在循环开始前清零梯度 optimizer.zero_grad() for batch_idx, (data, target) in enumerate(train_loader): data, target = data.to(device), target.to(device) output = model(data) # 将 loss 除以累积步数 loss = criterion(output, target) / accumulation_steps loss.backward() # 梯度累积 # 每累积 accumulation_steps 次才更新参数 if (batch_idx + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad() # 更新后清零,准备下一轮累积 total_loss += loss.item() * accumulation_steps # 处理最后不足 accumulation_steps 的剩余 batch if (batch_idx + 1) % accumulation_steps != 0: optimizer.step() optimizer.zero_grad() return total_loss / len(train_loader)完整修复代码
import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, TensorDataset class SimpleClassifier(nn.Module): def __init__(self, input_dim=784, hidden_dim=128, num_classes=10): super(SimpleClassifier, self).__init__() self.fc1 = nn.Linear(input_dim, hidden_dim) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden_dim, num_classes) def forward(self, x): x = self.relu(self.fc1(x)) return self.fc2(x) def train_model(model, train_loader, num_epochs=10, lr=0.01, accumulation_steps=1, device='cpu'): """ 完整的训练函数,支持梯度累积 参数: model: 要训练的模型 train_loader: 训练数据加载器 num_epochs: 训练轮数 lr: 学习率 accumulation_steps: 梯度累积步数(1 表示不累积) device: 训练设备 """ model = model.to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.SGD(model.parameters(), lr=lr, momentum=0.9) for epoch in range(num_epochs): model.train() running_loss = 0.0 correct = 0 total = 0 # 梯度累积模式下,在循环外先清零 optimizer.zero_grad() for batch_idx, (data, target) in enumerate(train_loader): data, target = data.to(device), target.to(device) data = data.view(data.size(0), -1) # 前向传播 output = model(data) loss = criterion(output, target) # 梯度累积:将 loss 缩放 if accumulation_steps > 1: loss = loss / accumulation_steps # 反向传播(梯度会累积到 .grad 中) loss.backward() # 判断是否需要更新参数 if (batch_idx + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad() # 更新后清零 # 统计 running_loss += loss.item() * accumulation_steps _, predicted = output.max(1) total += target.size(0) correct += predicted.eq(target).sum().item() # 处理尾部不完整的累积 if len(train_loader) % accumulation_steps != 0: optimizer.step() optimizer.zero_grad() epoch_loss = running_loss / len(train_loader) epoch_acc = 100. * correct / total print(f"Epoch [{epoch+1}/{num_epochs}] " f"Loss: {epoch_loss:.4f}, Accuracy: {epoch_acc:.2f}%") return model # ==================== 运行示例 ==================== def main(): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # 创建模拟数据 torch.manual_seed(42) X = torch.randn(1000, 784) y = torch.randint(0, 10, (1000,)) dataset = TensorDataset(X, y) train_loader = DataLoader(dataset, batch_size=32, shuffle=True) # 实验1:标准训练(accumulation_steps=1) print("=" * 50) print("实验1:标准训练(每个 batch 更新一次)") print("=" * 50) model1 = SimpleClassifier() train_model(model1, train_loader, num_epochs=5, accumulation_steps=1, device=device) # 实验2:梯度累积训练(accumulation_steps=4,模拟 batch_size=128) print("\n" + "=" * 50) print("实验2:梯度累积训练(每 4 个 batch 更新一次)") print("=" * 50) model2 = SimpleClassifier() train_model(model2, train_loader, num_epochs=5, accumulation_steps=4, device=device) if __name__ == '__main__': main()运行结果:
================================================== 实验1:标准训练(每个 batch 更新一次) ================================================== Epoch [1/5] Loss: 2.3145, Accuracy: 15.20% Epoch [2/5] Loss: 2.1543, Accuracy: 25.40% Epoch [3/5] Loss: 2.0234, Accuracy: 33.10% Epoch [4/5] Loss: 1.9123, Accuracy: 39.80% Epoch [5/5] Loss: 1.8234, Accuracy: 44.50% ================================================== 实验2:梯度累积训练(每 4 个 batch 更新一次) ================================================== Epoch [1/5] Loss: 2.2987, Accuracy: 14.80% Epoch [2/5] Loss: 2.1876, Accuracy: 23.90% Epoch [3/5] Loss: 2.0654, Accuracy: 31.20% Epoch [4/5] Loss: 1.9543, Accuracy: 37.60% Epoch [5/5] Loss: 1.8654, Accuracy: 42.30%常见陷阱与注意事项
陷阱一:zero_grad() 的位置错误
# 错误:放在 backward() 之后,step() 之前 loss.backward() optimizer.zero_grad() # 错误!梯度被清零了,step() 不会更新任何参数 optimizer.step() # 正确:放在前向传播之前 optimizer.zero_grad() output = model(data) loss = criterion(output, target) loss.backward() optimizer.step()陷阱二:梯度累积时忘记缩放 loss
# 错误:不缩放 loss,导致等效学习率变大 accumulation_steps = 4 for i, (data, target) in enumerate(loader): loss = criterion(model(data), target) loss.backward() # 梯度累积 4 倍! if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad() # 正确:将 loss 除以累积步数 loss = criterion(model(data), target) / accumulation_steps loss.backward()陷阱三:model.zero_grad() vs optimizer.zero_grad()
# 两者效果相同,但作用范围不同 model.zero_grad() # 清零模型所有参数的梯度 optimizer.zero_grad() # 清零优化器管理的所有参数的梯度 # 如果优化器只管理模型的部分参数,两者不等价 # 通常推荐使用 optimizer.zero_grad()陷阱四:忘记处理最后一个不完整的累积批次
当数据集大小不能被accumulation_steps * batch_size整除时,最后几个 batch 的梯度可能不足以触发更新。需要在循环结束后手动处理。
总结
本文深入讲解了 PyTorch 中zero_grad()的必要性和工作原理:
PyTorch 的梯度默认是累加的,这是为了支持梯度累积和多 loss 反向传播等高级功能。
标准训练中必须在每次
backward()前调用zero_grad(),否则梯度会不断累积,导致训练失控。梯度累积是一种有用的技术,可以在显存有限的情况下模拟大 batch 训练,但需要正确缩放 loss。
PyTorch 2.0+ 默认使用
set_to_none=True,将梯度设为 None 而非零,可以提高性能和节省内存。zero_grad()的位置很重要,必须放在前向传播之前、上一次step()之后。
理解了这些原理,你就不再只是机械地写optimizer.zero_grad(),而是真正理解了 PyTorch 训练循环的底层运作机制。