PyTorch如何构建两套独立权重、双损失分别更新的神经网络?
实现方案
方案1:不拆分原有网络(推荐,改动最小)
不需要拆分网络,只要拆分参数管理、配合梯度冻结即可实现完全隔离的更新逻辑,和你原有前向逻辑完全兼容:
核心步骤
- 拆分两类参数,各自绑定独立优化器
- 训练时分两个阶段,分别更新对应权重,更新时冻结另一套权重的梯度计算
完整代码示例
import torch import torch.nn as nn import torch.optim as optim class Model(nn.Module): def __init__(self): super(Model, self).__init__() self.blue_1 = nn.Linear(2, 10) self.red_1 = nn.Linear(10, 5) self.blue_2 = nn.Linear(5, 4) self.red_2 = nn.Linear(4, 3) self.blue_3 = nn.Linear(3, 2) self.red_3 = nn.Linear(2, 1) def forward(self, x): x = torch.relu(self.blue_1(x)) x = self.red_1(x) x = self.blue_2(x) x = self.red_2(x) x = self.blue_3(x) x = self.red_3(x) return x net = Model() features = torch.rand((10,2)) # 10 inputs, each of 2D # 1. 拆分两类参数,分别绑定优化器 blue_params = [p for name, p in net.named_parameters() if 'blue' in name] red_params = [p for name, p in net.named_parameters() if 'red' in name] opt_blue = optim.Adam(blue_params, lr=1e-3) opt_red = optim.Adam(red_params, lr=1e-3) # 定义你自己的blue损失函数,这里举个示例 def loss_blue(output, target): return torch.mean(torch.square(output - target)) for epoch in range(3): # ---------------------- # 阶段1:更新red权重,冻结blue参数 # ---------------------- for p in blue_params: p.requires_grad = False for p in red_params: p.requires_grad = True pred = net(features) # 计算red损失(原代码里的损失) loss_red = torch.sum(torch.randint(0,10,(10,)) - pred) loss_red.backward() opt_red.step() opt_red.zero_grad() # ---------------------- # 阶段2:更新blue权重,冻结red参数 # ---------------------- for p in red_params: p.requires_grad = False for p in blue_params: p.requires_grad = True pred = net(features) # 计算blue损失,按你自己的需求替换 blue_target = torch.rand((10,1)) loss_b = loss_blue(pred, blue_target) loss_b.backward() opt_blue.step() opt_blue.zero_grad() print(f"Epoch {epoch}, loss_red: {loss_red.item()}, loss_blue: {loss_b.item()}")
方案2:拆分为两个独立网络
如果后续需要单独部署、复用两套权重,可以拆分实现,逻辑同样清晰:
拆分代码示例
class BlueNet(nn.Module): def __init__(self): super().__init__() self.blue_1 = nn.Linear(2, 10) self.blue_2 = nn.Linear(5, 4) self.blue_3 = nn.Linear(3, 2) class RedNet(nn.Module): def __init__(self): super().__init__() self.red_1 = nn.Linear(10, 5) self.red_2 = nn.Linear(4, 3) self.red_3 = nn.Linear(2, 1) # 初始化两个网络 blue_net = BlueNet() red_net = RedNet() # 分别定义优化器 opt_blue = optim.Adam(blue_net.parameters(), lr=1e-3) opt_red = optim.Adam(red_net.parameters(), lr=1e-3) # 前向逻辑完全对齐原实现 def full_forward(x): x = torch.relu(blue_net.blue_1(x)) x = red_net.red_1(x) x = blue_net.blue_2(x) x = red_net.red_2(x) x = blue_net.blue_3(x) x = red_net.red_3(x) return x
训练时的更新逻辑和方案1完全一致,更新其中一个网络时冻结另一个网络的requires_grad属性即可。
原理解释
PyTorch中参数的requires_grad属性设为False时,反向传播过程中不会计算该参数的梯度,梯度流到该位置就会自动截断,完全不会影响另一套权重的更新,两套权重的梯度计算和更新100%隔离。
内容的提问来源于stack exchange,提问作者Penguin
相关产品推荐
相关产品推荐

