全部版块 我的主页
› 论坛 › 数据科学与人工智能 › 人工智能
282 0
2026-05-12

如何让基于CNN(卷积神经网络)的模型更轻量化?直接使用该模型的小型版本不就行了吗?例如,对于ResNet(残差网络),如果ResNet-152感觉过于笨重,为什么不使用ResNet-101呢?或者对于DenseNet(密集连接网络),为什么不选择DenseNet-121而非DenseNet-169?——没错,这样做确实可行,但你必须为此牺牲一部分精度。基本上,如果你想要一个更轻量化的模型,就应该预期其精度也会随之下降。

但如果我告诉你,有一款模型比其基础版本更轻量化,却仍能在精度上与之抗衡呢?它就是CSPNet(Cross Stage Partial Network,跨阶段部分网络)。你会惊讶地发现,它能在有效降低计算复杂度的同时保持高精度——无需任何取舍!在本文中,我们将探讨CSPNet的架构,包括其工作原理以及如何从零开始实现它。

CSPNet的简要历史

CSPNet最早由王等人于2019年11月在题为《CSPNet:一种可增强CNN学习能力的新型骨干网络》的论文中提出[1]。CSPNet最初的提出是为了解决DenseNet的局限性。尽管DenseNet的计算量已经比ResNet小,但作者认为DenseNet本身的计算量仍然偏大。下面请看DenseNet的主要构建块(图1),以理解其中原因。

图1. DenseNet模型的主要构建块[2]。

在DenseNet的构建块——称为密集块(dense block)中,每个卷积层都会接收所有前序层的信息,这导致其存在大量冗余梯度信息,使得训练效率低下。我们可以将其类比为一名学生由5位不同的老师教授同一内容:这本身是有益的,因为学生可以从多个角度理解该主题。然而,在某些时候,这种方式会变得冗余,进而降低效率。对于DenseNet而言,我们可以将深层视为学生,将所有浅层的张量视为老师。在上述示例中,若假设H₄是我们的学生,那么x₀、x₁、x₂和x₃张量就相当于老师。不难想象,这位学生会被如此多的信息压得喘不过气来!

在深入探讨CSPNet之前,我其实有一篇专门介绍DenseNet的独立文章(参考文献[3])。如果你想全面了解该架构的工作原理,强烈建议你阅读这篇文章。

研究目标

CSPNet的目标是使网络具有更低的计算复杂度和更优的梯度组合。后者的原因在于,DenseNet中的大多数梯度信息都是相互重复的。需要注意的是,CSPNet并非一个独立的网络,而是一种应用于DenseNet的新范式。

现在,我们来看下面的图2,了解CSPNet是如何实现其目标的。你可以看到左侧的示意图:随着网络深度的增加,特征图的数量逐渐增多。如果你读过我之前关于DenseNet的文章,就会知道这本质上是我们可以通过增长率(growth rate)参数控制的——增长率即密集块内每个卷积层产生的特征图数量。事实上,这种特征图数量的增加正是作者所认为的计算瓶颈。

图2. 左:原始DenseNet构建块(与图1相同)。右:CSPNet版本的DenseNet构建块(称为CSPDenseNet)[1]。

通过应用跨阶段部分(Cross Stage Partial)机制,我们基本上可以降低DenseNet的计算量。观察右侧的示意图,我们可以看到从x₀延伸出一个额外的分支,直接连接到所谓的部分过渡层(Partial Transition Layer)。这种机制至少能带来两个优势,与我之前提到的目标一致:第一,由于密集块处理的特征图数量仅为原始数量的一半,我们可以节省大量计算资源;第二,由于增加了一条包含未处理特征图的路径,避免了冗余梯度信息,梯度信息变得更加多样化。简而言之,CSPNet的核心思想是消除DenseNet的计算冗余(通过跳跃路径),同时保留其特征复用特性(通过密集块)。

CSPNet详细架构

具体来说,原始特征图首先按通道维度分为两部分,每部分将通过不同路径进行处理。假设我们有64个输入通道,前32个特征图(第一部分)将跳过所有计算,而剩余的32个(第二部分)将由密集块处理。尽管拆分步骤非常简单,但合并步骤实际上并不简单。你可以在下面的图3中看到,我们有几种不同的合并机制。

图3. CSPNet中几种不同的特征组合方式[1]。

在称为“先融合”(fusion first,图c)的结构中,我们先将第一部分张量与经过密集块处理的第二部分张量进行拼接,然后再将它们传入过渡层。因此,选项(c)的实现非常直接,因为两个张量的空间维度完全相同,可以轻松进行拼接。

在我之前的文章[3]中提到,DenseNet的过渡层用于同时降低空间维度和通道数量。事实上,这一特性要求我们重新思考如何实现“后融合”(fusion last,图d)结构。本质上,这是因为过渡层会导致第二部分张量的空间维度小于第一部分张量。因此,从技术上讲,我们需要要么对第一部分分支应用步长为2的池化操作,要么在过渡层中省略下采样操作。通过这种方式,两个张量的空间维度将保持一致,从而可以进行拼接。

除了在特征组合之前或之后仅使用单个过渡层外,作者还提出了另一种方法,称为CSPDenseNet(图b)。我们可以将其视为(c)和(d)的组合,其中在张量拼接过程之前和之后各有一个过渡层。在这种情况下,第一个过渡层(位于第二部分分支中)将通过跨通道池化(即沿通道维度操作的池化层)进行通道缩减。同时,第二个过渡层将同时执行空间下采样和通道数量缩减。因此,在这种方法中,我们会进行两次通道缩减——至少根据我对论文中两个过渡层的理解是这样的,因为这些层内部的详细过程并未明确讨论。

实验结果

关于这些特征组合机制的实验结果,论文中解释道,后融合(d)优于先融合(c):前者可以显著降低计算复杂度,而精度仅出现非常轻微的下降。变体(c)虽然也能降低计算复杂度,但精度下降幅度较大。作者发现,变体(b)的结果比前两者更好。下面的图4展示了几组实验结果,显示了三种特征组合机制与基础模型的性能对比。然而,他们没有使用DenseNet,而是选择了PeleeNet来比较这些结构。

图4. 基础PeleeNet(对应图3中的(a))、CSPPeleeNet(b)、采用先融合方法的PeleeNet(c)和采用后融合方法的PeleeNet(d)的性能对比[1]。

从上图可以看出,CSP后融合(绿色)的性能确实优于CSP先融合(红色)。这是因为它的精度仅比基础模型下降了0.1%,而计算复杂度却降低了21%。与此同时,尽管CSP先融合成功将计算复杂度降低了26%,但精度下降幅度相当显著,比基础PeleeNet低1.5%。最令人印象深刻的是CSPPeleeNet变体(蓝色),即采用两个过渡层的结构。在这里我们可以清楚地看到,尽管计算复杂度降低了13%,但模型的精度实际上提高了0.2%——再次证明了“无取舍”的优势!

不仅如此,作者还尝试将CSPNet应用于其他骨干模型。下面的图5结果显示,CSPNet结构成功将DenseNet-201-Elastic和ResNeXt-50的计算复杂度分别降低了19%和22%。有趣的是,尽管模型复杂度降低,但ResNeXt模型的精度却有所提高,这与图4中CSPPeleeNet获得的结果一致。

图5. 实现CSPNet机制后,DenseNet-201-Elastic和ResNeXt-50的性能提升[1]。

CSPDenseNet的数学表达式

对于喜欢数学的读者,这里有一些你可能会感兴趣的符号。下面的图6和图7展示了DenseNet和CSPDenseNet块在前向传播阶段的数学表达式。

在DenseNet块中,x₁对应于第一个卷积层w₁基于输入张量x₀产生的张量。接下来,我们将原始张量x₀与x₁拼接,并将它们作为w₂层的输入(更准确地说,w实际上是卷积层的权重,而非卷积层本身)。随着网络深度的增加,我们不断产生更多的特征图,并将已有的特征图进行拼接。通过这种方式,我们基本上可以说,所有前序层的输出都成为当前层的输入。

图6. DenseNet块内前向传播的数学表示[1]。

CSPDenseNet的情况则有所不同。你可以在下面的符号中看到x₀’和x₀’’,这正是我们之前提到的第一部分和第二部分。x₀’’张量经过类似DenseNet块的处理,直到得到xₖ。接下来,这个密集块的输出被传递到第一个过渡层,记为wᵗ。得到的张量xᵗ随后与第一部分张量x₀’拼接,最终通过第二个过渡层wᵘ得到最终输出张量xᵘ。

图7. CSPDenseNet块内前向传播的数学表达式[1]。

CSPDenseNet实现

现在,让我们通过从零实现来更深入地了解CSPNet架构。尽管我们基本上可以将CSPNet结构应用于任何骨干网络,但在这里,我将把它应用于DenseNet模型,以与我之前展示的示意图和公式保持一致。下面的图8展示了完整的DenseNet架构。请记住,该架构中的每个密集块最初都遵循图3a中的DenseNet结构,而我们的目标是将所有这些密集块替换为图3b中所示的CSPDenseNet块。

图8. 完整的DenseNet架构[2]。

我们要做的第一件事是导入所需的模块并初始化可配置参数,如代码块1所示。GROWTH变量是增长率参数,表示密集块内每个瓶颈层(bottleneck)产生的特征图数量。接下来,CHANNEL_POOLING是我们用于调整第一个过渡层中跨通道池化机制行为的参数。这里我将该参数设置为0.8,意味着我们将把通道数量缩减到原始通道数的80%。COMPRESSION参数的作用与CHANNEL_POOLING变量类似,但它作用于第二个过渡层。最后,我们定义REPEATS列表,用于设置每个阶段的密集块内要初始化的瓶颈块数量。

# 代码块1

import torch
import torch.nn as nn

GROWTH          = 12
CHANNEL_POOLING = 0.8
COMPRESSION     = 0.5
REPEATS         = [6, 12, 24, 16]

瓶颈块实现

下面是要放置在密集块内的瓶颈块实现。这个Bottleneck类与我在DenseNet文章[3]中使用的完全相同。我直接从那里复制粘贴了代码,因为我们完全不需要修改这部分。只需记住,一个瓶颈块由一个1×1卷积层和一个3×3卷积层组成。

# 代码块2

class Bottleneck(nn.Module):
    def __init__(self, in_channels):
        super().__init__()
        
        self.relu = nn.ReLU()
        self.dropout = nn.Dropout(p=0.2)
        
        self.bn0   = nn.BatchNorm2d(num_features=in_channels)
        self.conv0 = nn.Conv2d(in_channels=in_channels, 
                               out_channels=GROWTH*4,          
                               kernel_size=1, 
                               padding=0, 
                               bias=False)
        
        self.bn1   = nn.BatchNorm2d(num_features=GROWTH*4)
        self.conv1 = nn.Conv2d(in_channels=GROWTH*4, 
                               out_channels=GROWTH,            
                               kernel_size=3, 
                               padding=1, 
                               bias=False)
    
    def forward(self, x):
        print(f'original\t: {x.size()}')
        
        out = self.dropout(self.conv0(self.relu(self.bn0(x))))
        print(f'after conv0\t: {out.size()}')
        
        out = self.dropout(self.conv1(self.relu(self.bn1(out))))
        print(f'after conv1\t: {out.size()}')
        
        concatenated = torch.cat((out, x), dim=1)              
        print(f'after concat\t: {concatenated.size()}')
        
        return concatenated

以下测试代码模拟了密集块内的第一个瓶颈块。请记住,架构中的第一个卷积层(使用7×7核的那个)产生64个特征图,但由于在CSPNet的情况下,我们只想处理其中的一半(第二部分张量),因此这里我们将使用一个具有32个特征图的张量进行测试。

# 代码块3

bottleneck = Bottleneck(in_channels=32)

x = torch.randn(1, 32, 56, 56)
x = bottleneck(x)

# 代码块3输出

original     : torch.Size([1, 32, 56, 56])
after conv0  : torch.Size([1, 48, 56, 56])
after conv1  : torch.Size([1, 12, 56, 56])
after concat : torch.Size([1, 44, 56, 56])

从上面的输出结果可以看到,过程结束时特征图的数量变为44,这个数字是通过输入通道数加上增长率得到的,即32 + 12 = 44。同样,如果你想更好地理解这个计算过程,可以查看我的DenseNet文章[3]。

密集块实现

现在,为了方便创建一系列瓶颈块,我们可以将其包装在下面代码块4的DenseBlock类中。之后,我们只需通过repeats参数指定要堆叠的瓶颈块数量即可。同样,这个类也是从我的DenseNet文章中复制粘贴的,因此我不再进一步解释。

# 代码块4

class DenseBlock(nn.Module):
    def __init__(self, in_channels, repeats):
        super().__init__()
        self.bottlenecks = nn.ModuleList()
        
        for i in range(repeats):
            current_in_channels = in_channels + i * GROWTH
            self.bottlenecks.append(Bottleneck(in_channels=current_in_channels))
        
    def forward(self, x):
        print(f'original\t\t\t: {x.size()}')
        
        for i, bottleneck in enumerate(self.bottlenecks):
            x = bottleneck(x)
            print(f'after bottleneck #{i}\t\t: {x.size()}')
            
        return x

为了检查我们的DenseBlock类是否正常工作,我们将使用下面的代码块5进行测试。这里我尝试模拟第一部分密集块处理的第二部分张量,该密集块包含一系列6个瓶颈块。

# 代码块5

dense_block = DenseBlock(in_channels=32, repeats=6)
x = torch.randn(1, 32, 56, 56)

x = dense_block(x)

下面是输出结果。在这里我们可以清楚地看到,每个瓶颈块都成功地将特征图数量增加了12。

# 代码块5输出

original             : torch.Size([1, 32, 56, 56])
after bottleneck #0  : torch.Size([1, 44, 56, 56])
after bottleneck #1  : torch.Size([1, 56, 56, 56])
after bottleneck #2  : torch.Size([1, 68, 56, 56])
after bottleneck #3  : torch.Size([1, 80, 56, 56])
after bottleneck #4  : torch.Size([1, 92, 56, 56])
after bottleneck #5  : torch.Size([1, 104, 56, 56])

第一个过渡层

记住,图3b中的CSPDenseNet变体使用了两个过渡层。在本节中,我们将讨论第一个过渡层,即用于处理第二部分分支中张量的过渡层。这里我们不会执行空间下采样,这就是为什么在下面的代码块6的__init__()方法中没有看到任何池化层的原因。相反,这里我们只执行跨通道池化,这可以看作是一种标准的池化操作,但沿通道维度进行。为了实现它,我们可以简单地使用一个1×1卷积(#(2))并指定我们想要的输出通道数量(#(1))。我们可以这样理解:在空间下采样过程中,我们基本上可以通过池化层或带步长的卷积层来实现,而在后者的情况下,它会从局部邻域中按特定权重聚合像素值。在跨通道池化的情况下,由于PyTorch没有专门的层用于此操作,我们可以简单地用逐点卷积层(pointwise convolution layer)代替,这样我们基本上可以沿通道维度聚合像素值。

# 代码块6

class FirstTransition(nn.Module):
    def __init__(self, in_channels, out_channels):
        super().__init__()
        
        self.bn   = nn.BatchNorm2d(num_features=in_channels)
        self.relu = nn.ReLU()
        self.conv = nn.Conv2d(in_channels=in_channels, 
                              out_channels=out_channels,   #(1)
                              kernel_size=1,               #(2)
                              padding=0,
                              bias=False)
        self.dropout = nn.Dropout(p=0.2)
     
    def forward(self, x):
        print(f'original\t\t: {x.size()}')
        
        out = self.dropout(self.conv(self.relu(self.bn(x))))
        print(f'after first_transition\t: {out.size()}')
        
        return out

代码块5输出的结果显示,第二部分张量经过密集块处理后,形状将为104×56×56。因此,在下面的测试代码中,我将使用这个张量形状来模拟该阶段的第一个过渡层。为了调整输出通道数量,我们可以简单地将输入通道数乘以我们之前初始化的CHANNEL_POOLING变量,如下面代码块7的第#(1)行所示。

# 代码块7

first_transition = FirstTransition(in_channels=104, 
                                   out_channels=int(104*CHANNEL_POOLING)) #(1)

x = torch.randn(1, 104, 56, 56)
x = first_transition(x)

运行上述代码后,我们可以看到特征图的数量从104缩减到83(原始数量的80%)。

# 代码块7输出

original          : torch.Size([1, 104, 56, 56])
after first_transition  : torch.Size([1, 83, 56, 56])

第二个过渡层

第二个过渡层的结构与第一个非常相似,不同之处在于这里我们还有一个步长为2的平均池化层,用于将空间维度减半(#(1))。

# 代码块8

class SecondTransition(nn.Module):
    def __init__(self, in_channels, out_channels):
        super().__init__()
        
        self.bn   = nn.BatchNorm2d(num_features=in_channels)
        self.relu = nn.ReLU()
        self.conv = nn.Conv2d(in_channels=in_channels, 
                              out_channels=out_channels, 
                              kernel_size=1, 
                              padding=0,
                              bias=False)
        self.dropout = nn.Dropout(p=0.2)
        self.pool = nn.AvgPool2d(kernel_size=2, stride=2)    #(1)
     
    def forward(self, x):
        print(f'original\t\t: {x.size()}')

        out = self.pool(self.dropout(self.conv(self.relu(self.bn(x)))))
        print(f'after second_transition\t: {out.size()}')
        
        return out

记住,进入第二个过渡层的张量是第一部分和第二部分张量的拼接结果。这就是为什么在下面的测试代码中,我将该层设置为接受32 + 83 = 115个特征图。与第一个过渡层类似,这里我们将这个特征图数量乘以COMPRESSION变量(#(1)),以进一步减少通道数量。

# 代码块9

second_transition = SecondTransition(in_channels=115, 
                                     out_channels=int(115*COMPRESSION))  #(1)

x = torch.randn(1, 115, 56, 56)
x = second_transition(x)

在下面的输出结果中,我们可以看到由于平均池化层,空间维度减半。同时,由于我们将COMPRESSION参数设置为0.5,特征图的数量也从115减少到57。

# 代码块9输出

original                : torch.Size([1, 115, 56, 56])
after second_transition : torch.Size([1, 57, 28, 28])

CSPDenseNet模型

所有组件准备就绪后,我们现在可以构建完整的CSPDenseNet架构,我将其分解在下面的代码块10a、10b和10c中。首先让我们关注代码块10a,在这里我根据图8中给出的结构初始化所有层。你可以在第#(1)行看到,我初始化了一个7×7卷积层,作为网络的输入层。该层之后是一个最大池化层(#(2))。这两个层都使用步长2,这意味着输入张量的空间维度将缩减到原始大小的四分之一。

# 代码块10a

class CSPDenseNet(nn.Module):
    def __init__(self):
        super().__init__()
        
        self.first_conv = nn.Conv2d(in_channels=3,         #(1)
                                    out_channels=64, 
                                    kernel_size=7,    
                                    stride=2,         
                                    padding=3,        
                                    bias=False)
        self.first_pool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)  #(2)
        channel_count = 64
        
        
        
        ##### Stage 0
        self.dense_block_0 = DenseBlock(in_channels=channel_count//2, 
                                        repeats=REPEATS[0])
        
        self.first_transition_0 = FirstTransition(in_channels=(channel_count//2)+(REPEATS[0]*GROWTH), 
                                                  out_channels=int(((channel_count//2)+(REPEATS[0]*GROWTH))*CHANNEL_POOLING))
        
        channel_count = (channel_count - (channel_count//2)) + int(((channel_count//2)+(REPEATS[0]*GROWTH))*CHANNEL_POOLING)
        
        self.second_transition_0 = SecondTransition(in_channels=channel_count, 
                                                  out_channels=int(channel_count*COMPRESSION))
        
        channel_count = int(channel_count*COMPRESSION)
        #####
        
        
        ##### Stage 1
        self.dense_block_1 = DenseBlock(in_channels=channel_count//2, 
                                        repeats=REPEATS[1])
        
        self.first_transition_1 = FirstTransition(in_channels=(channel_count//2)+(REPEATS[1]*GROWTH), 
                                                  out_channels=int(((channel_count//2)+(REPEATS[1]*GROWTH))*CHANNEL_POOLING))
        
        channel_count = (channel_count - (channel_count//2)) + int(((channel_count//2)+(REPEATS[1]*GROWTH))*CHANNEL_POOLING)
        
        self.second_transition_1 = SecondTransition(in_channels=channel_count, 
                                                  out_channels=int(channel_count*COMPRESSION))
        
        channel_count = int(channel_count*COMPRESSION)
        #####
        
        
        ##### Stage 2
        self.dense_block_2 = DenseBlock(in_channels=channel_count//2, 
                                        repeats=REPEATS[2])
        
        self.first_transition_2 = FirstTransition(in_channels=(channel_count//2)+(REPEATS[2]*GROWTH), 
                                                  out_channels=int(((channel_count//2)+(REPEATS[2]*GROWTH))*CHANNEL_POOLING))
        
        channel_count = (channel_count - (channel_count//2)) + int(((channel_count//2)+(REPEATS[2]*GROWTH))*CHANNEL_POOLING)
        
        self.second_transition_2 = SecondTransition(in_channels=channel_count, 
                                                  out_channels=int(channel_count*COMPRESSION))
        
        channel_count = int(channel_count*COMPRESSION)
        #####
        
        
        ##### Stage 3
        self.dense_block_3 = DenseBlock(in_channels=channel_count//2, 
                                        repeats=REPEATS[3])
        
        self.first_transition_3 = FirstTransition(in_channels=(channel_count//2)+(REPEATS[3]*GROWTH), 
                                                  out_channels=int(((channel_count//2)+(REPEATS[3]*GROWTH))*CHANNEL_POOLING))
        
        channel_count = (channel_count - (channel_count//2)) + int(((channel_count//2)+(REPEATS[3]*GROWTH))*CHANNEL_POOLING)
        #####
        
        
        self.avgpool = nn.AdaptiveAvgPool2d(output_size=(1,1))             #(3)
        self.fc = nn.Linear(in_features=channel_count, out_features=1000)  #(4)

仍然以上述代码块为例,我根据层所属的阶段对其进行分组。现在让我们关注我称为Stage 0的部分。在这里你可以看到,我们有一个密集块(dense_block_0)和第一个过渡层(first_transition_0)。这两个组件负责处理第二部分张量。接下来,我们初始化第二个过渡层(second_transition_0),用于处理第一部分和第二部分张量的拼接结果。由于通道数量会根据GROWTH、CHANNEL_POOLING、COMPRESSION和REPEATS变量动态变化,我们需要在每一步跟踪通道数量,以便模型能够根据这些变量进行自适应调整。我们对所有剩余阶段执行相同的操作,除了Stage 3——在这里我们不初始化第二个过渡层,因为此时我们不再进一步减少通道和空间维度。相反,我们将第一部分和第二部分张量的拼接结果直接传递到平均池化层(#(3))和分类层(#(4))。以上就是对代码块10a的讨论。

在进入forward()方法之前,我们还需要创建另一个函数:split_channels()。顾名思义,下面代码块10b中编写的这个函数用于将张量分为第一部分和第二部分。这里的if-else语句用于检查通道数量是奇数还是偶数。事实上,如果通道数量是偶数,我们可以直接将其分成两部分(#(4)),这非常简单。但如果通道数量是奇数,我们需要像第#(1)行和第#(2)行那样手动确定每部分的大小,然后再进行拆分(#(3))。

# 代码块10b

def split_channels(self, x):

        channel_count = x.size(1)

        if channel_count%2 != 0:
            split_size_2 = channel_count // 2            #(1)
            split_size_1 = channel_count - split_size_2  #(2)
            return torch.split(x, [split_size_1, split_size_2], dim=1)  #(3)

        else:
            return torch.split(x, channel_count // 2, dim=1)            #(4)

定义完__init__()和split_channel()方法后,我们现在可以在下面的代码块10c中实现forward()方法。一般来说,我们在这里所做的就是将张量按顺序向前传递。但现在让我们关注我称为Stage 0的部分。在这里你可以看到,张量经过first_pool层(#(1))后,我们使用之前声明的split_channels()函数将其分为两部分(#(2))。从这里开始,我们得到了part1和part2张量。我们将part1张量保持不变,直到该阶段结束。与此同时,对于part2张量,我们将用密集块(#(3))和第一个过渡层(#(4))对其进行处理。接下来,我们将得到的张量与part1张量拼接,创建跳跃连接(#(5))。然后,我们最终将其传递到第二个过渡层(#(6))。所有阶段都重复相同的步骤,直到我们最终到达输出层进行分类。只需记住,Stage 3有所不同,因为这里没有第二个过渡层。

# 代码块10c

def forward(self, x):
        print(f'original\t\t\t: {x.size()}')
        
        x = self.first_conv(x)
        print(f'after first_conv\t\t: {x.size()}')
        
        x = self.first_pool(x)      #(1)
        print(f'after first_pool\t\t: {x.size()}\n')
        
        
        
        ##### Stage 0
        part1, part2 = self.split_channels(x)    #(2)
        print(f'part1\t\t\t\t: {part1.size()}')
        print(f'part2\t\t\t\t: {part2.size()}')
        
        part2 = self.dense_block_0(part2)        #(3)
        print(f'part2 after dense block 0\t: {part2.size()}')
        
        part2 = self.first_transition_0(part2)   #(4)
        print(f'part2 after first trans 0\t: {part2.size()}')
        
        x = torch.cat((part1, part2), dim=1)     #(5)
        print(f'after concatenate\t\t: {x.size()}')
        
        x = self.second_transition_0(x)          #(6)
        print(f'after second transition 0\t: {x.size()}\n')
        
        
        
        ##### Stage 1
        part1, part2 = self.split_channels(x)
        print(f'part1\t\t\t\t: {part1.size()}')
        print(f'part2\t\t\t\t: {part2.size()}')
        
        part2 = self.dense_block_1(part2)
        print(f'part2 after dense block 1\t: {part2.size()}')
        
        part2 = self.first_transition_1(part2)
        print(f'part2 after first trans 1\t: {part2.size()}')
        
        x = torch.cat((part1, part2), dim=1)
        print(f'after concatenate\t\t: {x.size()}')
        
        x = self.second_transition_1(x)
        print(f'after second transition 1\t: {x.size()}\n')
        
        
        
        ##### Stage 2
        part1, part2 = self.split_channels(x)
        print(f'part1\t\t\t\t: {part1.size()}')
        print(f'part2\t\t\t\t: {part2.size()}')
        
        part2 = self.dense_block_2(part2)
        print(f'part2 after dense block 2\t: {part2.size()}')
        
        part2 = self.first_transition_2(part2)
        print(f'part2 after first trans 2\t: {part2.size()}')
        
        x = torch.cat((part1, part2), dim=1)
        print(f'after concatenate\t\t: {x.size()}')
        
        x = self.second_transition_2(x)
        print(f'after second transition 2\t: {x.size()}\n')
        
        
        
        ##### Stage 3
        part1, part2 = self.split_channels(x)
        print(f'part1\t\t\t\t: {part1.size()}')
        print(f'part2\t\t\t\t: {part2.size()}')
        
        part2 = self.dense_block_3(part2)
        print(f'part2 after dense block 2\t: {part2.size()}')
        
        part2 = self.first_transition_3(part2)
        print(f'part2 after first trans 2\t: {part2.size()}')
        
        x = torch.cat((part1, part2), dim=1)
        print(f'after concatenate\t\t: {x.size()}\n')
        
        
        
        x = self.avgpool(x)
        print(f'after avgpool\t\t\t: {x.size()}')
        
        x = torch.flatten(x, start_dim=1)
        print(f'after flatten\t\t\t: {x.size()}')
        
        x = self.fc(x)
        print(f'after fc\t\t\t: {x.size()}')
        
        return x

现在,让我们通过运行下面的代码块11来测试我们刚刚创建的CSPDenseNet类。这里我使用一个形状为3×224×224的虚拟张量,模拟一张224×224的RGB图像通过网络的过程。

# 代码块11

cspdensenet = CSPDenseNet()

x = torch.randn(1, 3, 224, 224)
x = cspdensenet(x)

下面是输出结果。在这里你可以看到,每次张量进入网络,我们的split_channels()方法都会正确地将其分为两部分(#(1–2))。然后,每个阶段内的瓶颈块也会正确地将第二部分张量的通道数量增加12,之后才传递到第一个过渡层。第一个过渡层本身成功地将通道数量减少了20%,如第#(3)行所示,模拟了跨通道池化机制。之后,得到的张量与第一部分的张量拼接(#(4)),并传递到第二个过渡层(#(5)),以进一步减少通道数量并将空间维度减半。所有阶段都重复相同的步骤,最终我们得到1000类的预测结果。

# 代码块11输出

original                  : torch.Size([1, 3, 224, 224])
after first_conv          : torch.Size([1, 64, 112, 112])
after first_pool          : torch.Size([1, 64, 56, 56])

part1                     : torch.Size([1, 32, 56, 56])    #(1)
part2                     : torch.Size([1, 32, 56, 56])    #(2)
after bottleneck #0       : torch.Size([1, 44, 56, 56])
after bottleneck #1       : torch.Size([1, 56, 56, 56])
after bottleneck #2       : torch.Size([1, 68, 56, 56])
after bottleneck #3       : torch.Size([1, 80, 56, 56])
after bottleneck #4       : torch.Size([1, 92, 56, 56])
after bottleneck #5       : torch.Size([1, 104, 56, 56])
part2 after dense block 0 : torch.Size([1, 104, 56, 56])
part2 after first trans 0 : torch.Size([1, 83, 56, 56])    #(3)
after concatenate         : torch.Size([1, 115, 56, 56])   #(4)
after second transition 0 : torch.Size([1, 57, 28, 28])    #(5)

part1                     : torch.Size([1, 29, 28, 28])
part2                     : torch.Size([1, 28, 28, 28])
after bottleneck #0       : torch.Size([1, 40, 28, 28])
after bottleneck #1       : torch.Size([1, 52, 28, 28])
after bottleneck #2       : torch.Size([1, 64, 28, 28])
after bottleneck #3       : torch.Size([1, 76, 28, 28])
after bottleneck #4       : torch.Size([1, 88, 28, 28])
after bottleneck #5       : torch.Size([1, 100, 28, 28])
after bottleneck #6       : torch.Size([1, 112, 28, 28])
after bottleneck #7       : torch.Size([1, 124, 28, 28])
after bottleneck #8       : torch.Size([1, 136, 28, 28])
after bottleneck #9       : torch.Size([1, 148, 28, 28])
after bottleneck #10      : torch.Size([1, 160, 28, 28])
after bottleneck #11      : torch.Size([1, 172, 28, 28])
part2 after dense block 1 : torch.Size([1, 172, 28, 28])
part2 after first trans 1 : torch.Size([1, 137, 28, 28])
after concatenate         : torch.Size([1, 166, 28, 28])
after second transition 1 : torch.Size([1, 83, 14, 14])

part1                     : torch.Size([1, 42, 14, 14])
part2                     : torch.Size([1, 41, 14, 14])
after bottleneck #0       : torch.Size([1, 53, 14, 14])
after bottleneck #1       : torch.Size([1, 65, 14, 14])
after bottleneck #2       : torch.Size([1, 77, 14, 14])
after bottleneck #3       : torch.Size([1, 89, 14, 14])
after bottleneck #4       : torch.Size([1, 101, 14, 14])
after bottleneck #5       : torch.Size([1, 113, 14, 14])
after bottleneck #6       : torch.Size([1, 125, 14, 14])
after bottleneck #7       : torch.Size([1, 137, 14, 14])
after bottleneck #8       : torch.Size([1, 149, 14, 14])
after bottleneck #9       : torch.Size([1, 161, 14, 14])
after bottleneck #10      : torch.Size([1, 173, 14, 14])
after bottleneck #11      : torch.Size([1, 185, 14, 14])
after bottleneck #12      : torch.Size([1, 197, 14, 14])
after bottleneck #13      : torch.Size([1, 209, 14, 14])
after bottleneck #14      : torch.Size([1, 221, 14, 14])
after bottleneck #15      : torch.Size([1, 233, 14, 14])
after bottleneck #16      : torch.Size([1, 245, 14, 14])
after bottleneck #17      : torch.Size([1, 257, 14, 14])
after bottleneck #18      : torch.Size([1, 269, 14, 14])
after bottleneck #19      : torch.Size([1, 281, 14, 14])
after bottleneck #20      : torch.Size([1, 293, 14, 14])
after bottleneck #21      : torch.Size([1, 305, 14, 14])
after bottleneck #22      : torch.Size([1, 317, 14, 14])
after bottleneck #23      : torch.Size([1, 329, 14, 14])
part2 after dense block 2 : torch.Size([1, 329, 14, 14])
part2 after first trans 2 : torch.Size([1, 263, 14, 14])
after concatenate         : torch.Size([1, 305, 14, 14])
after second transition 2 : torch.Size([1, 152, 7, 7])

part1                     : torch.Size([1, 76, 7, 7])
part2                     : torch.Size([1, 76, 7, 7])
after bottleneck #0       : torch.Size([1, 88, 7, 7])
after bottleneck #1       : torch.Size([1, 100, 7, 7])
after bottleneck #2       : torch.Size([1, 112, 7, 7])
after bottleneck #3       : torch.Size([1, 124, 7, 7])
after bottleneck #4       : torch.Size([1, 136, 7, 7])
after bottleneck #5       : torch.Size([1, 148, 7, 7])
after bottleneck #6       : torch.Size([1, 160, 7, 7])
after bottleneck #7       : torch.Size([1, 172, 7, 7])
after bottleneck #8       : torch.Size([1, 184, 7, 7])
after bottleneck #9       : torch.Size([1, 196, 7, 7])
after bottleneck #10      : torch.Size([1, 208, 7, 7])
after bottleneck #11      : torch.Size([1, 220, 7, 7])
after bottleneck #12      : torch.Size([1, 232, 7, 7])
after bottleneck #13      : torch.Size([1, 244, 7, 7])
after bottleneck #14      : torch.Size([1, 256, 7, 7])
after bottleneck #15      : torch.Size([1, 268, 7, 7])
part2 after dense block 2 : torch.Size([1, 268, 7, 7])
part2 after first trans 2 : torch.Size([1, 214, 7, 7])
after concatenate         : torch.Size([1, 290, 7, 7])

after avgpool             : torch.Size([1, 290, 1, 1])
after flatten             : torch.Size([1, 290])
after fc                  : torch.Size([1, 1000])

结语

就这样!我们已经成功学习了CSPNet,并在DenseNet骨干网络上实现了它。正如我之前提到的,我们实际上可以利用CSPNet的思想来改进任何其他骨干模型的性能,例如ResNet或ResNeXt。因此,我向你发起挑战:从零开始在这些模型上实现CSPNet。

说实话,我无法确认我的实现是100%正确的,因为该论文的官方GitHub仓库[4]没有提供PyTorch实现——但这至少是我从论文手稿中理解到的全部内容。如果你发现代码或我的解释中有任何错误,请告诉我。感谢阅读,我们下一篇文章再见!拜拜!

推荐学习书籍 《CDA一级教材》适合CDA一级考生备考,也适合业务及数据分析岗位的从业者提升自我。完整电子版已上线CDA网校,累计已有10万+在读~ !

免费加入阅读:https://edu.cda.cn/goods/show/3151?targetId=5147&preview=0

二维码

扫码加我 拉你入群

请注明:姓名-公司-职位

以便审核进群资格,未注明则拒绝

相关推荐
栏目导航
热门文章
推荐文章

说点什么

分享

扫码加好友,拉您进群
各岗位、行业、专业交流群