逐层转发残差模块结果异常:PyTorch两种残差实现输出为何不同?
问题:PyTorch残差连接两种写法输出不一样的原因
我写了一段带残差连接的PyTorch代码:
all_module = [] for i in range(3): layer = nn.Sequential( nn.Conv1d(n_hidden_channels, n_hidden_channels), nn.LeakyReLU(), nn.Conv1d(n_hidden_channels, n_hidden_channels), nn.LeakyReLU() ) all_module.append(layer) module_list = nn.ModuleList(all_module) # 方法1 for layer in module_list: x = x + layer(x) print(x) # 方法2 for layer in module_list: y = torch.clone(x) for m in layer: y = m(y) x = x + y print(x)
运行后发现两种方法的输出结果不一样,这是为啥?
原因说明
两种方法输出不同根本不是实现逻辑的问题,而是你在同一段代码里顺序执行,导致方法2的输入已经被方法1修改了:
- 方法1跑完后,
x已经是经过3次残差叠加后的结果,不再是初始输入。 - 方法2接着用这个被改后的
x当输入,再跑3次残差,输出自然和方法1不一样。
如果要验证两种写法逻辑是否等价,得把它们分开跑:比如执行方法1前保存原始x,方法2用原始x重新初始化后再执行,或者把两种方法分成两段独立代码运行。
另外补充:两种写法的残差逻辑本身是完全一致的——方法1直接用Sequential的forward处理x,方法2手动遍历子模块处理克隆的y,只要输入x相同,输出结果肯定一模一样。
内容的提问来源于stack exchange,提问作者CA H
相关产品推荐
相关产品推荐

