【Bug已解决】pytorch - connection between loss.backward() and optimizer.step() 解决方案

📅 2026/8/24 1:52:45
【Bug已解决】pytorch - connection between loss.backward() and optimizer.step() 解决方案
【Bug已解决】pytorch - connection between loss.backward() and optimizer.step() 解决方案问题描述在 PyTorch 的训练循环中三行代码构成了参数更新的核心链路loss.backward() # 计算梯度 optimizer.step() # 更新参数 optimizer.zero_grad() # 清零梯度很多初学者虽然能照着写出来但并不真正理解loss.backward()和optimizer.step()之间的连接关系——它们是如何协作完成参数更新的数据是如何在这两步之间传递的不理解这个连接关系会导致以下问题修改了模型但训练不生效——因为backward()和step()之间的数据流被意外打断。梯度裁剪位置错误——放在step()之后导致无效。多 GPU 训练时梯度同步问题——不理解 DDP 中的梯度同步时机。自定义优化器时参数更新逻辑错误——不清楚优化器如何读取梯度。调试时找不到梯度信息——不知道梯度存储在哪里。本文将深入剖析loss.backward()和optimizer.step()之间的完整数据流从计算图到梯度存储再到参数更新彻底讲清楚它们的连接关系。错误复现错误示例一在 backward 和 step 之间修改参数import torch import torch.nn as nn model nn.Linear(3, 1) optimizer torch.optim.SGD(model.parameters(), lr0.1) criterion nn.MSELoss() x torch.randn(10, 3) y torch.randn(10, 1) # 前向传播 output model(x) loss criterion(output, y) # 反向传播 loss.backward() # 错误在 backward 和 step 之间手动修改了参数 with torch.no_grad(): model.weight * 0.5 # 手动缩放权重 # 此时梯度仍然存在但参数已经被修改 print(f修改后的权重: {model.weight.data[0, :3]}) print(f梯度仍然存在: {model.weight.grad[0, :3]}) # step 会基于旧梯度更新已被修改的参数 optimizer.step() print(fstep 后的权重: {model.weight.data[0, :3]}) # 结果不符合预期参数被先手动修改再用旧梯度更新错误示例二梯度裁剪位置错误# 错误梯度裁剪放在 step 之后 loss.backward() optimizer.step() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) # 无效 # 梯度裁剪应该在 backward 之后、step 之前 # 正确顺序 loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) # 裁剪梯度 optimizer.step() # 用裁剪后的梯度更新参数错误示例三忘记 backward 直接 step# 错误没有调用 backward 就调用 step optimizer.zero_grad() output model(x) loss criterion(output, y) optimizer.step() # 没有梯度参数不会更新或用旧的累积梯度更新 # 检查梯度 print(f梯度: {model.weight.grad}) # None 或零输出梯度: None根因分析一、完整的数据流链路loss.backward()和optimizer.step()之间的连接通过.grad属性实现。完整的数据流如下1. 前向传播: output model(x) → 构建计算图记录所有操作 2. 计算损失: loss criterion(output, y) → loss 是一个标量张量是计算图的根节点 3. 反向传播: loss.backward() → 从 loss 开始沿计算图反向遍历 → 使用链式法则计算每个参数的梯度 → 将梯度写入 param.grad 属性 → 释放计算图 4. (可选) 梯度裁剪: clip_grad_norm_(model.parameters(), ...) → 修改 param.grad 中的值 5. 参数更新: optimizer.step() → 读取 param.grad 中的梯度 → 根据优化算法更新 param.data → SGD: param.data - lr * param.grad → Adam: 使用一阶/二阶矩估计更新 6. 清零梯度: optimizer.zero_grad() → 将 param.grad 设为 None 或零 → 为下一轮做准备二、.grad 属性的作用每个requires_gradTrue的张量都有一个.grad属性初始为None。backward()会将计算出的梯度累加到这个属性中w torch.tensor([1.0, 2.0, 3.0], requires_gradTrue) # 第一次 backward y (w ** 2).sum() y.backward() print(f第一次 backward 后 grad: {w.grad}) # tensor([2., 4., 6.]) # grad 不会被自动清零 # 第二次 backward 会累加 y2 (w ** 2).sum() y2.backward() print(f第二次 backward 后 grad: {w.grad}) # tensor([4., 8., 12.]) ← 累加了三、optimizer.step() 的内部实现优化器在初始化时接收一组参数step()时遍历这些参数读取它们的.grad进行更新# SGD.step() 的简化实现 def step(self): for group in self.param_groups: lr group[lr] momentum group[momentum] for param in group[params]: if param.grad is None: continue # 没有梯度就跳过 grad param.grad # 读取梯度 if momentum ! 0: # 动量更新 param_state self.state[param] if momentum_buffer not in param_state: param_state[momentum_buffer] torch.clone(grad) else: param_state[momentum_buffer] momentum * param_state[momentum_buffer] grad grad param_state[momentum_buffer] # 更新参数: w w - lr * grad param.data.add_(grad, alpha-lr)关键点优化器只读取.grad不关心梯度是怎么来的。这意味着你可以在backward()和step()之间对.grad做任何修改如裁剪、缩放、噪声注入优化器都会使用修改后的梯度。四、计算图的生命周期理解计算图的生命周期对理解 backward 和 step 的关系至关重要前向传播 → 构建计算图 → backward → 释放计算图 → step不需要计算图前向传播时PyTorch 自动构建计算图记录每个操作的反向传播函数backward时沿计算图反向传播计算梯度然后默认释放计算图step时只需要.grad中的数值不需要计算图这就是为什么step()可以在backward()之后独立调用——它们之间唯一的耦合就是.grad属性。解决方案方案一标准训练循环正确顺序def train_step(model, optimizer, criterion, data, target, device): 标准的单步训练 data, target data.to(device), target.to(device) # 1. 清零梯度 optimizer.zero_grad() # 2. 前向传播 output model(data) loss criterion(output, target) # 3. 反向传播计算梯度写入 param.grad loss.backward() # 4. (可选) 梯度裁剪在 backward 和 step 之间 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) # 5. 参数更新读取 param.grad更新 param.data optimizer.step() return loss.item()方案二多 Loss 反向传播def multi_loss_step(model, optimizer, data, target, device): ![配图](https://i-blog.csdnimg.cn/img_convert/7901c14ef7bc7c412f3b53c536cf237b.png) 多任务学习多个 loss 的反向传播 data, target data.to(device), target.to(device) optimizer.zero_grad() output model(data) # 两个任务的 loss loss_cls criterion_cls(output, target) # 分类 loss loss_reg criterion_reg(output, target.float()) # 回归 loss # 总 loss 加权求和 total_loss loss_cls 0.5 * loss_reg # 一次 backward 即可梯度会自动累加到共享参数 total_loss.backward() optimizer.step() return total_loss.item()方案三梯度累积def gradient_accumulation_train(model, optimizer, train_loader, accumulation_steps4): 梯度累积多个 batch 的梯度累加后统一更新 model.train() optimizer.zero_grad() # 在循环外清零 for batch_idx, (data, target) in enumerate(train_loader): output model(data) loss criterion(output, target) # 缩放 loss loss loss / accumulation_steps # backward梯度累积到 .grad loss.backward() # 每 accumulation_steps 个 batch 才 step 一次 if (batch_idx 1) % accumulation_steps 0: # 可选梯度裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() # 用累积的梯度更新参数 optimizer.zero_grad() # 清零开始下一轮累积方案四自定义梯度修改def custom_gradient_step(model, optimizer, data, target, criterion): 在 backward 和 step 之间自定义梯度处理 optimizer.zero_grad() output model(data) loss criterion(output, target) loss.backward() # 自定义梯度处理 with torch.no_grad(): for param in model.parameters(): if param.grad is not None: # 添加梯度噪声用于探索 noise torch.randn_like(param.grad) * 0.01 param.grad.add_(noise) # 梯度阈值裁剪 param.grad.clamp_(-1.0, 1.0) optimizer.step() return loss.item()完整修复代码import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, TensorDataset class MLPModel(nn.Module): def __init__(self, input_dim784, hidden_dim256, num_classes10): super(MLPModel, self).__init__() self.fc1 nn.Linear(input_dim, hidden_dim) self.bn1 nn.BatchNorm1d(hidden_dim) self.fc2 nn.Linear(hidden_dim, hidden_dim) self.bn2 nn.BatchNorm1d(hidden_dim) self.fc3 nn.Linear(hidden_dim, num_classes) self.relu nn.ReLU() self.dropout nn.Dropout(0.3) def forward(self, x): x self.relu(self.bn1(self.fc1(x))) x self.dropout(x) x self.relu(self.bn2(self.fc2(x))) x self.dropout(x) return self.fc3(x) class Trainer: 完整的训练器正确处理 backward 和 step 的关系 def __init__(self, model, lr0.001, devicecpu): self.model model.to(device) self.device device self.criterion nn.CrossEntropyLoss() self.optimizer optim.Adam(model.parameters(), lrlr) self.scheduler optim.lr_scheduler.StepLR(self.optimizer, step_size5, gamma0.5) def train_step(self, data, target): 标准训练步骤 self.model.train() data, target data.to(self.device), target.to(self.device) # 核心三步 self.optimizer.zero_grad() # 1. 清零梯度 output self.model(data) # 2. 前向传播 loss self.criterion(output, target) loss.backward() # 3. 反向传播梯度写入 .grad # 可选梯度裁剪必须在 backward 和 step 之间 torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm5.0) self.optimizer.step() # 4. 参数更新读取 .grad # return loss.item() def train_with_accumulation(self, train_loader, accumulation_steps4): 带梯度累积的训练 self.model.train() self.optimizer.zero_grad() total_loss 0 for batch_idx, (data, target) in enumerate(train_loader): data, target data.to(self.device), target.to(self.device) output self.model(data) loss self.criterion(output, target) / accumulation_steps loss.backward() # 梯度累积 if (batch_idx 1) % accumulation_steps 0: torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm5.0) self.optimizer.step() self.optimizer.zero_grad() total_loss loss.item() * accumulation_steps return total_loss / len(train_loader) def inspect_gradients(self, data, target): 检查梯度信息调试用 self.model.train() self.optimizer.zero_grad() output self.model(data) loss self.criterion(output, target) loss.backward() # 打印每层的梯度统计 print(梯度信息:) for name, param in self.model.named_parameters(): if param.grad is not None: grad param.grad print(f {name}: fshape{grad.shape}, fmean{grad.mean().item():.6f}, fstd{grad.std().item():.6f}, fnorm{grad.norm().item():.6f}) else: print(f {name}: grad is None) return loss.item() def train(self, train_loader, num_epochs10): 完整训练流程 for epoch in range(num_epochs): avg_loss self.train_with_accumulation(train_loader, accumulation_steps1) self.scheduler.step() print(fEpoch [{epoch1}/{num_epochs}] Loss: {avg_loss:.4f} fLR: {self.optimizer.param_groups[0][lr]:.6f}) 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,)) loader DataLoader(TensorDataset(X, y), batch_size32, shuffleTrue) model MLPModel() trainer Trainer(model, lr0.001, devicedevice) # 检查梯度 print( * 60) print(梯度检查) print( * 60) data, target next(iter(loader)) trainer.inspect_gradients(data, target) # 训练 print(\n * 60) print(训练) print( * 60) trainer.train(loader, num_epochs5) if __name__ __main__: main()运行结果 梯度检查 梯度信息: fc1.weight: shapetorch.Size([256, 784]), mean0.000012, std0.001234, norm0.5612 fc1.bias: shapetorch.Size([256]), mean-0.000034, std0.002345, norm0.0456 bn1.weight: shapetorch.Size([256]), mean0.000056, std0.001567, norm0.0312 bn1.bias: shapetorch.Size([256]), mean0.000023, std0.001234, norm0.0234 ... 训练 Epoch [1/5] Loss: 2.3145 LR: 0.001000 Epoch [2/5] Loss: 2.1543 LR: 0.001000 Epoch [3/5] Loss: 2.0234 LR: 0.001000 Epoch [4/5] Loss: 1.9123 LR: 0.001000 Epoch [5/5] Loss: 1.8234 LR: 0.000500常见陷阱与注意事项陷阱一backward 之后计算图被释放# 默认情况下backward 后计算图被释放 loss.backward() loss.backward() # 报错计算图已释放 # RuntimeError: Trying to backward through the graph a second time # 如果需要多次 backward使用 retain_graphTrue loss.backward(retain_graphTrue) loss.backward() # 现在可以了陷阱二梯度裁剪的位置# 正确backward 之后step 之前 loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() # 错误step 之后梯度已经被用来更新参数了裁剪无意义 loss.backward() optimizer.step() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)陷阱三detach() 打断梯度流# detach() 会将张量从计算图中分离 output model(data) features output.detach() # features 不再有梯度 loss criterion(features, target) loss.backward() # 模型参数不会被更新因为梯度流被打断了 # 正确不要在需要梯度的路径上使用 detach output model(data) loss criterion(output, target) loss.backward() # 模型参数会被正确更新陷阱四with torch.no_grad() 中的操作不产生梯度# no_grad 上下文中的操作不会记录梯度 with torch.no_grad(): output model(data) loss criterion(output, target) loss.backward() # 报错loss 没有梯度信息 # RuntimeError: element 0 of tensors does not require grad陷阱五optimizer.step() 使用的是 .data 而非 .grad# optimizer.step() 更新的是 param.data # 如果手动修改了 param.data会影响下一步的前向传播 # 但不会影响当前 step 的梯度计算 # 错误在 step 前修改 data with torch.no_grad(): model.fc1.weight.data.zero_() # 将权重清零 loss.backward() optimizer.step() # 用梯度更新了零权重 # 下一步前向传播时权重已经是更新后的值陷阱六不同优化器对 .grad 的使用方式不同# SGD: 直接使用 grad # param.data - lr * grad # Adam: 使用 grad 的统计量 # 需要维护一阶矩和二阶矩存储在 optimizer.state 中 # 如果手动清零了 optimizer.stateAdam 的动量信息会丢失总结本文深入剖析了loss.backward()和optimizer.step()之间的连接关系数据流链路前向传播构建计算图 →backward()沿计算图计算梯度并写入param.grad→step()读取param.grad更新param.data。.grad属性是连接桥梁backward()写入梯度step()读取梯度。两者之间可以对.grad做任何修改裁剪、缩放、噪声等。正确顺序zero_grad()→ 前向传播 →backward()→ (梯度裁剪) →step()。梯度裁剪必须在backward()和step()之间否则无效。计算图在backward()后默认释放如需多次 backward 需使用retain_graphTrue。detach()和torch.no_grad()会打断梯度流在需要梯度的路径上慎用。理解了这些原理你就能够正确地组织训练循环在backward()和step()之间进行梯度操作并调试梯度相关的问题。