You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中loss_fn与模型、优化器的关联及loss.backward()原理

关于PyTorch中loss.backward()的技术疑问解答

问题列表

  1. 代码中通过将model.parameters()传入optimizer,使优化器与模型建立关联,但loss_fn与模型、优化器无显式关联,那么loss.backward()具体是如何工作的?
  2. 若新增损失函数实例loss_fn_2 = torch.nn.MSELoss(reduction='sum'),执行loss_2 = loss_fn_2(y_pred, y)及loss_2.backward(),PyTorch如何识别loss_2与当前模型的关联?
  3. 若需构建两组独立的组件(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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.16 19:27:48