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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 01:06:01