PyTorch中loss_fn与模型、优化器的关联及loss.backward()原理
关于PyTorch中
loss.backward()的技术疑问解答 问题列表
- 代码中通过将
model.parameters()传入optimizer,使优化器与模型建立关联,但loss_fn与模型、优化器无显式关联,那么loss.backward()具体是如何工作的? - 若新增损失函数实例
loss_fn_2 = torch.nn.MSELoss(reduction='sum'),执行loss_2 = loss_fn_2(y_pred, y)及loss_2.backward(),PyTorch如何识别loss_2与当前模型的关联? - 若需构建两组独立的组件(
model_a、loss_fn_a、optimizer_a与model_b、loss_fn_b、optimizer_b),该如何实现它们的互相隔离?
问题解答
1. loss.backward()的工作机制
PyTorch的自动梯度依赖计算图实现:
- 执行模型前向传播(
y_pred = model(xx))时,会自动构建计算图,图中包含输入xx到输出y_pred的所有运算节点,以及模型中所有requires_grad=True的可训练参数(比如Linear层的weight和bias)。 - 损失值
loss由loss_fn(y_pred, y)计算得到,是计算图的最终节点。 - 调用
loss.backward()时,PyTorch从loss节点出发,沿计算图反向遍历,通过链式法则计算每个可训练参数的梯度,并将梯度存储在参数的.grad属性中。 - 优化器初始化时已绑定模型参数,后续
optimizer.step()会读取这些.grad值更新对应参数。
整个过程无需loss_fn与模型/优化器显式关联,计算图已通过前向传播将损失、模型参数、输入输出串联。
2. 新损失函数与模型的关联逻辑
loss_2与模型的关联依然由计算图维持:
y_pred是模型前向传播的输出,属于模型计算图的一部分。loss_2 = loss_fn_2(y_pred, y)是基于y_pred的运算,因此loss_2会成为原计算图的新终点节点。- 调用
loss_2.backward()时,PyTorch会沿计算图反向追溯,找到所有参与y_pred计算的可训练参数(即当前模型的参数),计算这些参数相对于loss_2的梯度。 - 损失函数是无状态工具类,只要损失值由模型输出计算而来,就会自动关联到模型的计算图。
3. 实现两组独立组件的隔离
要让两组组件完全隔离,核心是保证它们的计算图、参数更新流程完全独立,具体步骤:
- 定义独立模型:创建两个参数空间完全独立的模型实例:
model_a = torch.nn.Sequential(torch.nn.Linear(3, 1), torch.nn.Flatten(0, 1)) model_b = torch.nn.Sequential(torch.nn.Linear(3, 1), torch.nn.Flatten(0, 1)) - 绑定独立优化器:优化器初始化时分别传入对应模型的参数,确保各自只管理自身模型的参数:
optimizer_a = torch.optim.RMSprop(model_a.parameters(), lr=1e-3) optimizer_b = torch.optim.RMSprop(model_b.parameters(), lr=1e-3) - 独立执行训练流程:对两组组件分别执行完整训练步骤,流程互不干扰:
# 训练model_a y_pred_a = model_a(xx_a) loss_a = loss_fn_a(y_pred_a, y_a) optimizer_a.zero_grad() # 仅清空model_a参数的梯度 loss_a.backward() # 仅计算model_a参数的梯度 optimizer_a.step() # 仅更新model_a的参数 # 训练model_b y_pred_b = model_b(xx_b) loss_b = loss_fn_b(y_pred_b, y_b) optimizer_b.zero_grad() # 仅清空model_b参数的梯度 loss_b.backward() # 仅计算model_b参数的梯度 optimizer_b.step() # 仅更新model_b的参数
通过以上方式,两组组件的计算图、梯度、参数更新完全隔离,不会互相影响。
附参考代码
import torch import math # Create Tensors to hold input and outputs. x = torch.linspace(-math.pi, math.pi, 2000) y = torch.sin(x) # Prepare the input tensor (x, x^2, x^3). p = torch.tensor([1, 2, 3]) xx = x.unsqueeze(-1).pow(p) # Use the nn package to define our model and loss function. model = torch.nn.Sequential( torch.nn.Linear(3, 1), torch.nn.Flatten(0, 1) ) loss_fn = torch.nn.MSELoss(reduction='sum') # Use the optim package to define an Optimizer that will update the weights of # the model for us. Here we will use RMSprop; the optim package contains many other # optimization algorithms. The first argument to the RMSprop constructor tells the # optimizer which Tensors it should update. learning_rate = 1e-3 optimizer = torch.optim.RMSprop(model.parameters(), lr=learning_rate) for t in range(2000): # Forward pass: compute predicted y by passing x to the model. y_pred = model(xx) # Compute and print loss. loss = loss_fn(y_pred, y) if t % 100 == 99: print(t, loss.item()) # Before the backward pass, use the optimizer object to zero all of the # gradients for the variables it will update (which are the learnable # weights of the model). This is because by default, gradients are # accumulated in buffers( i.e, not overwritten) whenever .backward() # is called. Checkout docs of torch.autograd.backward for more details. optimizer.zero_grad() # Backward pass: compute gradient of the loss with respect to model # parameters loss.backward() # Calling the step function on an Optimizer makes an update to its # parameters optimizer.step() linear_layer = model[0] print(f'Result: y = {linear_layer.bias.item()} + {linear_layer.weight[:, 0].item()} x + {linear_layer.weight[:, 1].item()} x^2 + {linear_layer.weight[:, 2].item()} x^3')
内容的提问来源于stack exchange,提问作者A.A
相关产品推荐
相关产品推荐

