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
相关产品推荐
相关产品推荐

