如何在普通MLP中添加跳跃连接(非ResNet实现)
为普通MLP添加跳跃连接的正确实现
问题分析
你用torch.nn.Sequential尝试添加跳跃连接的思路行不通,因为Sequential是线性执行的容器,只能按顺序传递数据,无法处理需要将不同层输出相加的分支逻辑。另外你代码里model.add_module(layer_1 + layer_2)的写法完全错误——add_module需要传入模块名称和模块实例,而且layer_1是add_module的返回值(实际为None),不能直接做加法操作。
正确实现方案
要在MLP里加跳跃连接,必须自定义nn.Module类,手动控制前向传播的数据流,实现分支相加的逻辑。以下是具体代码:
import torch import torch.nn as nn from torch.nn import ReLU class MLPWithSkip(nn.Module): def __init__(self, input_size=615, output_size=40): super().__init__() # 定义各层 self.layer0 = nn.Linear(input_size, 2048) self.act0 = ReLU() self.layer1 = nn.Linear(2048, 2048) self.act1 = ReLU() self.layer2 = nn.Linear(2048, 2048) self.act2 = ReLU() self.layer3 = nn.Linear(2048, output_size) def forward(self, x): # 第一层处理 x = self.act0(self.layer0(x)) # 保存跳跃连接的输入(这里是layer0+act0后的输出) skip_x = x # 中间两层处理 x = self.act1(self.layer1(x)) x = self.layer2(x) # 执行跳跃连接:将中间层输出与跳过的输入相加 x = x + skip_x # 后续层处理 x = self.act2(x) x = self.layer3(x) return x # 实例化模型 model = MLPWithSkip() # 测试模型输入输出 test_input = torch.randn(1, 615) output = model(test_input) print(output.shape) # 输出应为 torch.Size([1, 40])
灵活调整跳跃连接位置
如果需要调整跳跃连接的位置(比如从原始输入直接跳转到layer2之后),只需要在forward方法里保存对应位置的张量即可:
def forward(self, x): # 保存原始输入 original_x = x x = self.act0(self.layer0(x)) x = self.act1(self.layer1(x)) x = self.layer2(x) # 跳跃连接:原始输入与layer2输出相加(需保证维度匹配,这里加线性层做转换) x = x + self.skip_proj(original_x) # 需要在__init__里定义self.skip_proj = nn.Linear(615, 2048) x = self.act2(x) x = self.layer3(x) return x
注意:跳跃连接相加的两个张量必须维度完全一致,如果维度不同,需要添加一个线性层做维度转换,避免报错。
内容的提问来源于stack exchange,提问作者Internet Presence
相关产品推荐
相关产品推荐

