【Resnet】从零实现Resnet核心模块:BasicBlock与Bottleneck代码逐行解析
1. Resnet核心模块概述
第一次看到Resnet的网络结构图时,我完全被那些密密麻麻的连线搞晕了。直到亲手实现了BasicBlock和Bottleneck这两个核心模块,才真正理解了残差网络的精妙之处。简单来说,Resnet之所以能训练出上百层的深度网络而不出现梯度消失问题,全靠这两个模块中巧妙的"跳跃连接"设计。
想象一下你在学习新知识时,如果能把之前掌握的内容直接拿来用,而不是每次都从零开始,学习效率是不是会高很多?Resnet的残差模块就是基于类似的思路。每个模块不仅学习新的特征变换,还会保留原始输入的特征信息。这种设计让深层网络的训练变得可行,也是Resnet能在ImageNet竞赛中一战成名的关键。
BasicBlock和Bottleneck是Resnet中两种主要的残差模块,它们的区别就像普通公路和高速公路的区别。BasicBlock结构简单,适合浅层网络;Bottleneck通过1×1卷积压缩和恢复通道数,虽然复杂但计算量更小,适合构建深层网络。接下来我们就深入代码层面,看看它们的具体实现。
2. BasicBlock实现详解
2.1 BasicBlock结构解析
BasicBlock是Resnet中最基础的残差模块,主要用于Resnet18和Resnet34这类相对较浅的网络。它的结构就像一条直路加上一条捷径:主路进行两次3×3卷积变换,捷径则可能对输入做简单调整(当需要改变特征图尺寸或通道数时)。
先来看BasicBlock的初始化代码:
class BasicBlock(nn.Module):
expansion = 1
def __init__(self, inplanes, planes, stride=1, downsample=None, groups=1,
base_width=64, dilation=1, norm_layer=None):
super(BasicBlock, self).__init__()
if norm_layer is None:
norm_layer = nn.BatchNorm2d
if groups != 1 or base_width != 64:
raise ValueError('BasicBlock only supports groups=1 and base_width=64')
if dilation > 1:
raise NotImplementedError("Dilation > 1 not supported in BasicBlock")
self.conv1 = conv3x3(inplanes, planes, stride)
self.bn1 = norm_layer(planes)
self.relu = nn.ReLU(inplace=True)
self.conv2 = conv3x3(planes, planes)
self.bn2 = norm_layer(planes)
self.downsample = downsample
self.stride = stride
这段代码有几个关键点需要注意:
expansion=1表示这个模块不会扩展通道数,输出通道数与中间通道数相同- 只支持默认的groups和base_width设置,这是与Bottleneck的区别之一
- 创建了两个3×3卷积层,中间包含BN和ReLU
- downsample参数用于处理跳跃连接时需要调整输入的情况
2.2 BasicBlock前向传播解析
BasicBlock的前向传播过程清晰地展现了残差连接的核心思想:
def forward(self, x):
identity = x # 保存原始输入
out = self.conv1(x)
out = self.bn1(out)
out = self.relu(out)
out = self.conv2(out)
out = self.bn2(out)
if self.downsample is not None:
identity = self.downsample(x)
out += identity # 残差连接
out = self.relu(out)
return out
这里有个实际项目中的经验分享:我第一次实现时忘记在最后一个ReLU前做残差相加,结果网络根本无法训练。后来才明白,残差连接必须在激活函数之前完成,这样才能保证梯度可以畅通无阻地回传。
downsample的使用场景特别值得注意。当stride≠1或者输入输出通道数不一致时,就需要通过downsample来调整identity的尺寸和通道数。这个设计保证了无论主路如何变化,identity都能与之匹配相加。
3. Bottleneck实现详解
3.1 Bottleneck结构设计
Bottleneck模块是Resnet50及以上深度网络使用的残差模块,它的结构就像沙漏:先压缩通道数,再进行3×3卷积,最后恢复通道数。这种设计大幅减少了计算量,是构建深层网络的关键。
先看初始化代码:
class Bottleneck(nn.Module):
expansion = 4
def __init__(self, inplanes, planes, stride=1, downsample=None, groups=1,
base_width=64, dilation=1, norm_layer=None):
super(Bottleneck, self).__init__()
if norm_layer is None:
norm_layer = nn.BatchNorm2d
width = int(planes * (base_width / 64.)) * groups
self.conv1 = conv1x1(inplanes, width)
self.bn1 = norm_layer(width)
self.conv2 = conv3x3(width, width, stride, groups, dilation)
self.bn2 = norm_layer(width)
self.conv3 = conv1x1(width, planes * self.expansion)
self.bn3 = norm_layer(planes * self.expansion)
self.relu = nn.ReLU(inplace=True)
self.downsample = downsample
self.stride = stride
关键点解析:
expansion=4表示输出通道数是中间通道数的4倍- 使用1×1卷积先压缩通道数,减少3×3卷积的计算量
- 支持groups参数,这是与BasicBlock的重要区别
- 最后的1×1卷积将通道数扩展到planes*expansion
3.2 Bottleneck前向传播过程
Bottleneck的前向传播展示了更复杂的残差连接:
def forward(self, x):
identity = x
out = self.conv1(x)
out = self.bn1(out)
out = self.relu(out)
out = self.conv2(out)
out = self.bn2(out)
out = self.relu(out)
out = self.conv3(out)
out = self.bn3(out)
if self.downsample is not None:
identity = self.downsample(x)
out += identity
out = self.relu(out)
return out
在实际项目中,Bottleneck的通道数变化经常让人困惑。举个例子,假设输入通道是256,planes=64:
- 第一个1×1卷积将256压缩到64(当base_width=64时)
- 3×3卷积保持64通道
- 最后一个1×1卷积扩展到256通道(64×4)
这种"压缩-计算-扩展"的结构,使得3×3卷积可以在较小的通道数上进行,大幅减少了计算量。根据我的实测,在Resnet50中,Bottleneck的计算量只有相同通道数BasicBlock的约40%,这也是为什么深层网络都使用Bottleneck的原因。
4. 两种模块的对比与应用
4.1 结构与计算量对比
BasicBlock和Bottleneck虽然都实现了残差连接,但在结构和计算效率上有显著差异。我们可以用一个简单的例子来说明:
假设输入输出都是256通道,stride=1:
- BasicBlock:两个3×3卷积,参数量约为256×256×3×3×2=1,179,648
- Bottleneck:先64通道1×1卷积,然后64通道3×3卷积,最后256通道1×1卷积,参数量约为256×64×1×1 + 64×64×3×3 + 64×256×1×1=69,632
可以看到,Bottleneck的参数量只有BasicBlock的约6%,这种优势在深层网络中更加明显。
4.2 实际应用选择
在项目中如何选择使用哪种模块呢?根据我的经验:
-
对于Resnet18/34这类较浅的网络,使用BasicBlock更合适。因为网络不深,不需要过度压缩计算量,而且BasicBlock结构简单,训练更稳定。
-
对于Resnet50/101/152等深层网络,必须使用Bottleneck。我曾经尝试用BasicBlock构建Resnet50,结果不仅训练速度慢,而且准确率下降了约3%。
-
当需要自定义网络结构时,可以参考这样的原则:如果计算资源充足且网络不深(<30层),用BasicBlock;如果追求深度和效率,用Bottleneck。
在实现上,PyTorch官方代码通过_make_layer函数灵活地构建各层:
def _make_layer(self, block, planes, blocks, stride=1, dilate=False):
# ...
layers = []
layers.append(block(self.inplanes, planes, stride, downsample, self.groups,
self.base_width, previous_dilation, norm_layer))
self.inplanes = planes * block.expansion
for _ in range(1, blocks):
layers.append(block(self.inplanes, planes, groups=self.groups,
base_width=self.base_width, dilation=self.dilation,
norm_layer=norm_layer))
return nn.Sequential(*layers)
这个函数会根据传入的block类型(BasicBlock或Bottleneck)自动构建对应的残差层,非常灵活。在实际项目中,我经常复用这段代码来构建自定义深度的Resnet变体。
更多推荐


所有评论(0)