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

如何在PyTorch的Branch类中添加Skip Connection且无需新建类

给PyTorch的Branch类添加跳连接(无需新建类)

需求说明

需要在现有PyTorch的Branch类中加入跳连接(Skip Connection),且不新建额外类。原代码如下:

class Branch(nn.Module):
    def __init__(self, channels, strides):
        super(Branch, self).__init__()
        
        self.conv_layer = nn.Sequential()
        for i in range(1, len(channels)):
            self.conv_layer.add_module(f'conv_{i}', nn.Conv2d(in_channels=channels[i-1], out_channels=channels[i], kernel_size=3, stride=strides[i-1], padding=(0, 1))),
            self.conv_layer.add_module(f'bn_{i}', nn.BatchNorm2d(channels[i])),
            self.conv_layer.add_module(f'relu_{i}', nn.ReLU())
            self.conv_layer.add_module(f'conv2_{i}', nn.Conv2d(in_channels=channels[i], out_channels=channels[i], kernel_size=3, padding='same')),
            self.conv_layer.add_module(f'bn2_{i}', nn.BatchNorm2d(channels[i])),
            self.conv_layer.add_module(f'relu2_{i}', nn.ReLU())

        self.projector = nn.LazyLinear(256)
        
    def forward(self, x):
        x = self.conv_layer(x)
        x = x.view(x.size(0), x.size(1) * x.size(2), x.size(3)).permute(0, 2, 1)
        x = self.projector(x)
        return x

解决方案

原代码使用单一Sequential串联所有层,无法直接实现跳连接(需要保留每组卷积的输入并与输出相加)。以下是修改后的代码,直接在原类内实现跳连接:

import torch.nn as nn

class Branch(nn.Module):
    def __init__(self, channels, strides):
        super(Branch, self).__init__()
        self.res_blocks = nn.ModuleList()
        
        for i in range(1, len(channels)):
            in_ch = channels[i-1]
            out_ch = channels[i]
            stride = strides[i-1]
            
            # 定义每组残差块的卷积层(不含最后ReLU,留到跳连接相加后)
            block = nn.Sequential(
                nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=stride, padding=(0, 1)),
                nn.BatchNorm2d(out_ch),
                nn.ReLU(),
                nn.Conv2d(out_ch, out_ch, kernel_size=3, padding='same'),
                nn.BatchNorm2d(out_ch)
            )
            self.res_blocks.append(block)
            
            # 当输入输出通道不同/步长不为1时,添加1x1卷积做维度匹配;否则用恒等映射
            if in_ch != out_ch or stride != 1:
                self.add_module(f'proj_{i}', nn.Conv2d(in_ch, out_ch, kernel_size=1, stride=stride))
            else:
                setattr(self, f'proj_{i}', nn.Identity())
        
        self.projector = nn.LazyLinear(256)
        
    def forward(self, x):
        # 遍历每个残差块,实现跳连接
        for i, block in enumerate(self.res_blocks, 1):
            # 保存当前输入作为残差
            residual = x
            # 用投影层调整残差维度(如果需要)
            residual = getattr(self, f'proj_{i}')(residual)
            
            # 经过残差块处理
            x = block(x)
            # 跳连接相加
            x = x + residual
            # 激活函数
            x = nn.ReLU()(x)
        
        # 原有的展平与投影逻辑
        x = x.view(x.size(0), x.size(1) * x.size(2), x.size(3)).permute(0, 2, 1)
        x = self.projector(x)
        return x

关键改动说明

  • 将原单一Sequential拆分为ModuleList存储每组残差块,方便单独处理每组的输入输出
  • 为每组残差块添加维度匹配层:当通道数变化或步长不为1时,用1x1卷积调整残差维度,保证可以与卷积输出相加;否则用恒等映射
  • 在forward中手动实现跳连接逻辑:保存输入→处理→相加→激活,完全符合标准残差结构
  • 全程未新建额外类,仅对原Branch类内部结构和逻辑修改

内容的提问来源于stack exchange,提问作者Qusai Triboush

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 18:15:19