PyTorch反向传播实现及多输出模型损失解耦方法问询
PyTorch相关问题解答
一、PyTorch反向传播的基本实现步骤
反向传播的核心是计算损失对模型参数的梯度并更新参数,标准流程如下:
- 清空梯度:调用
optimizer.zero_grad()清除上一轮迭代的梯度累积,避免梯度叠加影响参数更新。 - 前向传播:将输入数据传入模型,得到预测输出结果。
- 计算损失:使用对应损失函数(如
CrossEntropyLoss、MSELoss)计算预测值与真实标签的损失值。 - 反向传播:调用
loss.backward(),PyTorch会自动遍历计算图,计算损失对所有可训练参数的梯度。 - 更新参数:调用
optimizer.step(),优化器根据计算出的梯度更新模型参数。
二、多输出模型实现损失互不影响的方案
你当前的代码中,共享层lin1的参数会同时被loss_a和loss_b的梯度更新,导致两个任务的训练互相干扰;而out_a仅受loss_a影响、out_b仅受loss_b影响。若要实现两个任务完全独立(包括底层特征提取环节),最直接的方式是拆分共享的特征层:
修改后的独立分支模型代码
import torch class NN(torch.nn.Module): def __init__(self): super(NN, self).__init__() # 为两个任务分别设置独立的底层特征层 self.lin1_a = torch.nn.Linear(784, 128) self.lin1_b = torch.nn.Linear(784, 128) self.out_a = torch.nn.Linear(128, 10) self.out_b = torch.nn.Linear(128, 6) def forward(self, x): # 任务A的独立前向流程 x_a = torch.nn.functional.relu(self.lin1_a(x)) out_a = self.out_a(x_a) # 任务B的独立前向流程 x_b = torch.nn.functional.relu(self.lin1_b(x)) out_b = self.out_b(x_b) return out_a, out_b model = NN() optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) criterion_a = torch.nn.CrossEntropyLoss() criterion_b = torch.nn.CrossEntropyLoss() # 训练循环中的反向传播逻辑 optimizer.zero_grad() y_pred_a, y_pred_b = model(x_train) loss_a = criterion_a(y_pred_a, y_train_a) loss_b = criterion_b(y_pred_b, y_train_b) # 分别反向传播,retain_graph=True保留计算图用于第二次反向传播 loss_a.backward(retain_graph=True) loss_b.backward() optimizer.step()
如果仅要求loss_a不影响out_b的参数、loss_b不影响out_a的参数,你当前的代码其实已经满足——因为loss_a对out_b参数的梯度为0,loss_b对out_a参数的梯度为0,相加后反向传播不会让一个任务的损失更新另一个任务的输出层参数。
若要连共享层也完全隔离,只能采用拆分特征层的方案,这是实现两个任务完全互不影响的最优解。
内容的提问来源于stack exchange,提问作者nike01705
相关产品推荐
相关产品推荐

