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

如何在普通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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 05:21:28