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

Python中如何高效将一个DNN类的输出传递给另一个DNN类(不使用全局变量)

PyTorch双DNN参数传递解决方案

你当前使用全局变量传参的方式隐患极大,多batch训练、多卡并行时很容易出现值覆盖、梯度计算异常,下面是两种高效无额外开销的实现方案:

方案1:直接修改Net2的forward传参(最灵活,推荐优先使用)

直接把第一个网络输出的x作为入参传给Net2的forward方法,调用时先实例化两个网络,跑第一个网络拿到x后再传入第二个网络即可。

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        # 你的网络结构定义
        ...
    def forward(self, x):
        # 你的前向计算逻辑
        ...
        return x, z


class Net2(nn.Module):
    def __init__(self):
        super().__init__()
        # 你的网络结构定义
        ...
    def forward(self, v, x):  # 新增x作为入参
        # 注意torch.cat第一个参数是张量列表,dim指定拼接维度,根据你的需求调整
        y = torch.cat([v, x], dim=1)
        # 后续计算逻辑
        ...
        return y

# 调用示例
net = Net()
net2 = Net2()
input_x = ... # 第一个网络的输入
v = ... # 第二个网络的输入v
out1_x, out1_z = net(input_x)
out2 = net2(v, out1_x)

方案2:封装联合模型(适合两个网络串联复用场景)

如果两个网络需要固定串联使用,可以封装为一个统一的Module,内部管理两个子网实例,对外只需要暴露必要的输入参数即可:

class CombinedNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = Net()
        self.net2 = Net2()
    
    def forward(self, net_input, v):
        x, z = self.net(net_input)
        y = self.net2(v, x)
        # 可以根据需要返回需要的输出,比如同时返回z和y
        return y, z

# 调用示例
combined_net = CombinedNet()
input_x = ...
v = ...
out_y, out_z = combined_net(input_x, v)
  • 以上两种方案都没有任何额外计算、内存开销,梯度可以正常反向传播,支持两个网络的联合训练或单独训练
  • torch.cat的dim参数需要根据你的张量形状调整,例如2D张量(样本数+特征数)在特征维度拼接就设置dim=1

内容的提问来源于stack exchange,提问作者Mohamed Nabih

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 08:36:01