如何在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
相关产品推荐
相关产品推荐

