PyTorch中zero_grad()、backward()和step()的协同工作机制解析
你的代码回顾
线性回归模型实现
class LinearRegressionModel2(nn.Module): def __init__(self): super().__init__() # 使用nn.Linear创建模型参数(线性变换层/全连接层) self.linear_layer = nn.Linear(in_features = 1, out_features = 1) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.linear_layer(x)
训练与测试循环
torch.manual_seed(42) model_1 = LinearRegressionModel2() # 定义损失函数 loss_fn = nn.L1Loss() # 等价于MAE # 定义优化器 optimizer = torch.optim.SGD(params = model_1.parameters(), lr = 0.01, ) epochs = 200 for epoch in range(epochs): model_1.train() # 1. 前向传播 y_pred = model_1(X_train) # 2. 计算损失 train_loss = loss_fn(y_pred, y_train) # 3. 重置梯度 optimizer.zero_grad() # 4. 反向传播 train_loss.backward() # 5. 更新参数 optimizer.step() ### 测试阶段 model_1.eval() with torch.inference_mode(): test_pred = model_1(X_test) test_loss = loss_fn(test_pred, y_test) # 打印日志 if epoch % 10 == 0: print(f"Epoch: {epoch} | Train Loss: {train_loss} | Test Loss: {test_loss}")
针对你的疑问的解答
1. backward()如何仅作用于一个数值型的损失值?
你看到的train_loss虽然是单个数值,但它本质是带计算图的Tensor——PyTorch在执行前向传播时,会自动记录所有运算的依赖关系(比如从输入X_train到y_pred,再到train_loss的每一步运算),每个Tensor都通过.grad_fn属性保存着生成它的运算逻辑。
当调用train_loss.backward()时,PyTorch会从这个损失值出发,沿着计算图反向遍历所有参与运算的可训练参数(也就是model_1里linear_layer的权重和偏置),根据链式法则计算每个参数对损失值的梯度,并把计算结果存储到对应参数的.grad属性中。
简单说:单个损失值是整个前向计算图的“终点”,反向传播从这里触发,就能回溯计算所有参数的梯度。
2. 执行backward()后,优化器是如何获取到梯度信息的?
优化器在初始化时已经和模型参数绑定了——你创建optimizer时传入了model_1.parameters(),这会把模型中所有可训练的nn.Parameter(本质是带requires_grad=True的Tensor)都注册到优化器里。
当backward()完成梯度计算后,每个参数的梯度都存在自身的.grad属性中,优化器调用step()时,会直接遍历自己注册的所有参数,读取它们的.grad值,然后按照对应的优化算法(比如SGD的参数 = 参数 - 学习率 * 梯度)更新参数。
不需要显式传递梯度,因为初始化时已经完成了参数的绑定,优化器知道该去哪些参数的.grad里取梯度。
3. zero_grad()、backward()、step()三者的协同逻辑
这三个步骤是PyTorch训练循环的核心流程,各司其职:
optimizer.zero_grad(): PyTorch默认会累加梯度(如果不清零,每次反向传播的梯度会叠加到上一次的结果上),这一步是把所有绑定参数的.grad属性重置为0,确保当前轮次的梯度计算不受上一轮的影响。train_loss.backward(): 触发反向传播,基于当前轮次的前向计算结果,计算每个参数对损失的梯度,并写入参数的.grad属性。optimizer.step(): 读取每个参数的.grad值,用优化算法更新参数的权重/偏置,完成一轮参数更新。
这三步必须按顺序执行,才能保证每一轮训练的梯度计算和参数更新都是独立且正确的。
内容的提问来源于stack exchange,提问作者seyit

