PyTorch自定义多输入Block在nn.Sequential传reverse参数报错如何解决
错误原因
nn.Sequential 的默认forward方法仅支持传递位置输入,且内部逻辑是将前一个模块的返回值直接作为后一个模块的唯一输入,无法透传reverse、logdet这类额外的关键字参数,也不支持处理多返回值的传递,所以你传入的reverse参数会被nn.Sequential的forward方法拦截,找不到对应参数就报错。
解决方案
你可以自己实现一个适配多输入多输出的堆叠容器,替代原生的nn.Sequential即可,具体实现如下:
步骤1:先修正Block的小bug
你给出的Block定义中__init__的父类调用有误,先修正:
class Block(nn.Module): def __init__(self, num_channels): # 原写法super(InvConv, self).__init__()错误,父类应为Block super(Block, self).__init__() self.num_channels = num_channels # Initialize with a random orthogonal matrix w_init = np.random.randn(num_channels, num_channels) w_init = np.linalg.qr(w_init)[0].astype(np.float32) self.weight = nn.Parameter(torch.from_numpy(w_init)) def forward(self, x, logdet, reverse=False): ldj = torch.slogdet(self.weight)[1] * x.size(2) * x.size(3) if reverse: weight = torch.inverse(self.weight.double()).float() logdet = logdet - ldj else: weight = self.weight logdet = logdet + ldj weight = weight.view(self.num_channels, self.num_channels, 1, 1) z = F.conv2d(x, weight) return z, logdet
步骤2:实现自定义堆叠容器
class BlockSequential(nn.Module): def __init__(self, blocks): super().__init__() # 用ModuleList存储模块,保证参数能被PyTorch正确注册 self.blocks = nn.ModuleList(blocks) def forward(self, x, logdet, reverse=False): # 反向推理时自动倒序遍历Block,符合可逆网络常用逻辑 iter_blocks = reversed(self.blocks) if reverse else self.blocks for block in iter_blocks: x, logdet = block(x, logdet, reverse=reverse) return x, logdet
步骤3:使用自定义容器堆叠Block
完全匹配你要求的循环生成+调用方式如下:
# 循环生成10个Block完成堆叠 features = [] for i in range(10): features.append(Block(num_channels=48)) self.features = BlockSequential(features) # 直接调用即可 x = torch.Tensor(np.random.rand(2, 48, 8, 8)) z, final_logdet = self.features(x, logdet=0, reverse=False)
内容的提问来源于stack exchange,提问作者Nikoo_Ebrahimi
相关产品推荐
相关产品推荐

