import torch
import torch.nn as nn

class AtrousModule(nn.Module):
    def __init__(self, in_ch, out_ch, kernel_size, padding, dilation, is_bn=True):
        super(AtrousModule, self).__init__()
        self.atrous_conv = nn.Conv2d(in_ch, out_ch, kernel_size=kernel_size,
                                    stride=1, padding=padding, dilation=dilation, bias=False)
        self.is_bn = is_bn
        if is_bn:
            self.bn = nn.BatchNorm2d(out_ch)
        self.relu = nn.ReLU(inplace=True)
    
    def forward(self, x):
        x = self.atrous_conv(x)
        if self.is_bn:
            x = self.bn(x)
        x = self.relu(x)
        return x

class AvgPoolModule(nn.Module):
    def __init__(self, in_ch, out_ch, is_bn=True):
        super(AvgPoolModule, self).__init__()
        self.avg_pool = nn.AdaptiveAvgPool2d((1, 1))
        self.conv_pool = nn.Conv2d(in_ch, out_ch, kernel_size=1, padding=0)
        self.is_bn = is_bn
        if is_bn:
            self.bn_pool = nn.BatchNorm2d(out_ch)
        self.relu = nn.ReLU(inplace=True)
    
    def forward(self, x):
        x = self.avg_pool(x)
        x = self.conv_pool(x)
        if self.is_bn:
            x = self.bn_pool(x)
        x = self.relu(x)
        return x


class ASPP(nn.Module):
    def __init__(self, in_channels, out_channels, is_bn=True):
        super(ASPP, self).__init__()
        self.is_bn = is_bn
        # ASPP共包含5个并行的block : 1x1 conv; 
        self.block_1x1 = nn.Conv2d(in_channels, out_channels, kernel_size=1, padding=0)
        if is_bn:
            self.bn1 = nn.BatchNorm2d(out_channels)
        self.relu = nn.ReLU(inplace=True)
        # 3个膨胀率不同的3x3 conv;
        self.atrous_r6 = AtrousModule(in_channels, out_channels, kernel_size=3, padding=6, dilation=6, is_bn=is_bn)
        self.atrous_r12 = AtrousModule(in_channels, out_channels, kernel_size=3, padding=12, dilation=12, is_bn=is_bn)
        self.atrous_r18 = AtrousModule(in_channels, out_channels, kernel_size=3, padding=18, dilation=18, is_bn=is_bn)
        # 全局平均池化
        self.block_avg = AvgPoolModule(in_channels, out_channels, is_bn=is_bn)
        # 最后的输出层将所有5个并行block进行融合
        self.out_conv = nn.Conv2d(out_channels * 5, out_channels, kernel_size=1)
        if is_bn:
            self.out_bn = nn.BatchNorm2d(out_channels)
        
    def forward(self, x):
        # 1x1 conv
        if self.is_bn:
            out_1x1 = self.relu(self.bn1(self.block_1x1(x)))
        else:
            out_1x1 = self.relu(self.block_1x1(x))
        # dilation rate r6 r12 r18
        out_r6 = self.atrous_r6(x)
        out_r12 = self.atrous_r12(x)
        out_r18 = self.atrous_r18(x)
        # avg pool + resize
        out_avg = self.block_avg(x)
        out_avg = nn.functional.interpolate(out_avg, size=out_1x1.size()[2:], 
                                            mode='bilinear', align_corners=True)
        out = torch.cat([out_1x1, out_r6, out_r12, out_r18, out_avg], dim=1)
        if self.is_bn:
            out = self.relu(self.out_bn(self.out_conv(out)))
        else:
            out = self.relu(self.out_conv(out))

        # 打印内部变量信息
        print("[ASPP] out_1x1 size : ", out_1x1.size())
        print("[ASPP] out_r6  size : ", out_r6.size())
        print("[ASPP] out_r12 size : ", out_r12.size())
        print("[ASPP] out_r18 size : ", out_r18.size())
        print("[ASPP] out_avg size : ", out_avg.size())

        return out


x_in = torch.rand((4, 16, 64, 64))
aspp_block = ASPP(in_channels=16, out_channels=32, is_bn=True)
x_out = aspp_block(x_in)
print("input size : ", x_in.size())
print("output size : ", x_out.size())

代码使用PyTorch实现了Atrous Spatial Pyramid Pooling (ASPP) 模块,这是一种用于语义分割任务的多尺度特征提取结构,常用于DeepLab系列模型中。ASPP通过并行应用不同膨胀率的空洞卷积(Atrous Convolution)和全局池化,来捕捉图像中的多尺度上下文信息,同时保持特征图的分辨率。该实现支持可选的Batch Normalization (BN),并在forward中打印内部变量尺寸以便调试。

代码整体结构清晰:首先定义辅助模块(AtrousModule和AvgPoolModule),然后构建ASPP类,最后进行测试。假设输入为[4, 16, 64, 64](batch=4,通道=16,高度/宽度=64),out_channels=32,is_bn=True。基于计算,特征图尺寸变化如下(所有尺寸均为[batch, channels, height, width]):

  • out_1x1: [4, 32, 64, 64]
  • out_r6: [4, 32, 64, 64]
  • out_r12: [4, 32, 64, 64]
  • out_r18: [4, 32, 64, 64]
  • out_avg: [4, 32, 64, 64](经过interpolate上采样)
  • out(最终输出): [4, 32, 64, 64]

这些尺寸通过卷积公式计算得出:卷积输出尺寸 Hout=Hin+2×padding−kernel_size−(dilation−1)×(kernel_size−1)stride+1H_{out} = \frac{H_{in} + 2 \times padding - kernel\_size - (dilation - 1) \times (kernel\_size - 1)}{stride} + 1Hout=strideHin+2×paddingkernel_size(dilation1)×(kernel_size1)+1(stride=1)。实际运行时,代码中的print语句会输出这些尺寸。

下面将按照代码的逻辑结构,从导入库开始,一步一步逐段分析。每段解释包括代码的目的、关键概念、实现细节,以及为什么这样设计。

1. 导入库
import torch
import torch.nn as nn
  • 目的:导入PyTorch核心库,用于定义神经网络模型和张量操作。
  • 详细解释
    • import torch:PyTorch主库,提供张量(如torch.rand)和自动微分功能。
    • import torch.nn as nn:神经网络模块,提供Conv2d、BatchNorm2d、ReLU、AdaptiveAvgPool2d等预定义层。
  • 为什么需要这些库:ASPP是一个深度学习模块,需要PyTorch的模块化构建方式来定义卷积、池化和激活层。无其他外部库依赖,便于移植。
2. 空洞卷积模块:AtrousModule
class AtrousModule(nn.Module):
    def __init__(self, in_ch, out_ch, kernel_size, padding, dilation, is_bn=True):
        super(AtrousModule, self).__init__()
        self.atrous_conv = nn.Conv2d(in_ch, out_ch, kernel_size=kernel_size,
                                    stride=1, padding=padding, dilation=dilation, bias=False)
        self.is_bn = is_bn
        if is_bn:
            self.bn = nn.BatchNorm2d(out_ch)
        self.relu = nn.ReLU(inplace=True)
    
    def forward(self, x):
        x = self.atrous_conv(x)
        if self.is_bn:
            x = self.bn(x)
        x = self.relu(x)
        return x
  • 目的:定义一个可配置的空洞卷积单元,包括空洞Conv、可选BN和ReLU激活,用于ASPP中的多率分支。
  • 详细解释
    • __init__:初始化层。nn.Conv2d(in_channels, out_channels, kernel_size, stride=1, padding, dilation, bias=False):2D空洞卷积,bias=False避免偏移(常由BN处理)。dilation是膨胀率(dilation rate),控制采样间隔。padding根据dilation调整(e.g., padding=6 for dilation=6,确保尺寸不变)。is_bn控制添加nn.BatchNorm2d(批标准化,稳定训练)。nn.ReLU(inplace=True):就地激活,节省内存。
    • forward:前向传播。顺序:Conv → [BN] → ReLU。
  • 数学概念:空洞卷积输出尺寸:Hout=Hin+2×padding−dilation×(kernelsize−1)−1stride+1H_{out} = \frac{H_{in} + 2 \times padding - dilation \times (kernel_size - 1) - 1}{stride} + 1Hout=strideHin+2×paddingdilation×(kernelsize1)1+1。bias=False因为BN会添加可学习偏置。
  • 为什么这样设计:模块化,便于在ASPP中复用不同dilation的分支。is_bn可选,支持实验(如无BN的轻量模型)。这实现了ASPP的多尺度捕捉:不同dilation对应不同感受野(e.g., dilation=6的3x3核感受野=13x13)。
3. 全局平均池化模块:AvgPoolModule
class AvgPoolModule(nn.Module):
    def __init__(self, in_ch, out_ch, is_bn=True):
        super(AvgPoolModule, self).__init__()
        self.avg_pool = nn.AdaptiveAvgPool2d((1, 1))
        self.conv_pool = nn.Conv2d(in_ch, out_ch, kernel_size=1, padding=0)
        self.is_bn = is_bn
        if is_bn:
            self.bn_pool = nn.BatchNorm2d(out_ch)
        self.relu = nn.ReLU(inplace=True)
    
    def forward(self, x):
        x = self.avg_pool(x)
        x = self.conv_pool(x)
        if self.is_bn:
            x = self.bn_pool(x)
        x = self.relu(x)
        return x
  • 目的:定义全局平均池化单元,用于ASPP的图像级特征提取,包括AdaptiveAvgPool、1x1 Conv、可选BN和ReLU。
  • 详细解释
    • __init__nn.AdaptiveAvgPool2d((1, 1)):自适应平均池化,将任意尺寸特征图池化为1x1(全局平均)。然后nn.Conv2d(in_ch, out_ch, kernel_size=1, padding=0):1x1卷积调整通道。is_bn添加BN。
    • forward:AvgPool → Conv → [BN] → ReLU。
  • 数学概念:全局平均池化:xout[c]=1H×W∑i,jx[i,j,c]x_{out}[c] = \frac{1}{H \times W} \sum_{i,j} x[i,j,c]xout[c]=H×W1i,jx[i,j,c],捕捉整个特征图的统计信息。
  • 为什么这样设计:在ASPP中,这分支提供全局上下文(image-level features),补充空洞卷积的局部/中等尺度。1x1 Conv确保与其他分支通道一致,便于后续cat融合。
4. ASPP模块:ASPP
class ASPP(nn.Module):
    def __init__(self, in_channels, out_channels, is_bn=True):
        super(ASPP, self).__init__()
        self.is_bn = is_bn
        # ASPP共包含5个并行的block : 1x1 conv; 
        self.block_1x1 = nn.Conv2d(in_channels, out_channels, kernel_size=1, padding=0)
        if is_bn:
            self.bn1 = nn.BatchNorm2d(out_channels)
        self.relu = nn.ReLU(inplace=True)
        # 3个膨胀率不同的3x3 conv;
        self.atrous_r6 = AtrousModule(in_channels, out_channels, kernel_size=3, padding=6, dilation=6, is_bn=is_bn)
        self.atrous_r12 = AtrousModule(in_channels, out_channels, kernel_size=3, padding=12, dilation=12, is_bn=is_bn)
        self.atrous_r18 = AtrousModule(in_channels, out_channels, kernel_size=3, padding=18, dilation=18, is_bn=is_bn)
        # 全局平均池化
        self.block_avg = AvgPoolModule(in_channels, out_channels, is_bn=is_bn)
        # 最后的输出层将所有5个并行block进行融合
        self.out_conv = nn.Conv2d(out_channels * 5, out_channels, kernel_size=1)
        if is_bn:
            self.out_bn = nn.BatchNorm2d(out_channels)
        
    def forward(self, x):
        # 1x1 conv
        if self.is_bn:
            out_1x1 = self.relu(self.bn1(self.block_1x1(x)))
        else:
            out_1x1 = self.relu(self.block_1x1(x))
        # dilation rate r6 r12 r18
        out_r6 = self.atrous_r6(x)
        out_r12 = self.atrous_r12(x)
        out_r18 = self.atrous_r18(x)
        # avg pool + resize
        out_avg = self.block_avg(x)
        out_avg = nn.functional.interpolate(out_avg, size=out_1x1.size()[2:], 
                                            mode='bilinear', align_corners=True)
        out = torch.cat([out_1x1, out_r6, out_r12, out_r18, out_avg], dim=1)
        if self.is_bn:
            out = self.relu(self.out_bn(self.out_conv(out)))
        else:
            out = self.relu(self.out_conv(out))

        # 打印内部变量信息
        print("[ASPP] out_1x1 size : ", out_1x1.size())
        print("[ASPP] out_r6  size : ", out_r6.size())
        print("[ASPP] out_r12 size : ", out_r12.size())
        print("[ASPP] out_r18 size : ", out_r18.size())
        print("[ASPP] out_avg size : ", out_avg.size())

        return out
  • 目的:定义ASPP模块,实现多尺度特征提取,包括5个并行分支和融合层。
  • 详细解释
    • __init__:is_bn全局控制BN。1x1分支:nn.Conv2d(in_channels, out_channels, kernel_size=1, padding=0),捕捉局部特征。三个AtrousModule:dilation=6/12/18,padding相应调整(padding=dilation,确保尺寸不变)。全局池化:AvgPoolModule。融合:nn.Conv2d(out_channels * 5, out_channels, kernel_size=1),1x1卷积降维。
    • forward:计算每个分支(1x1、r6/r12/r18、avg)。avg通过nn.functional.interpolate上采样到与其他分支相同尺寸(bilinear模式,align_corners=True确保边界对齐)。torch.cat沿dim=1拼接(通道维)。然后Conv → [BN] → ReLU。print语句调试内部尺寸。
  • 数学概念:拼接后通道= out_channels * 5,1x1 Conv融合:线性变换多尺度特征。
  • 为什么这样设计:5分支捕捉不同尺度(局部到全局),膨胀率6/12/18是经验值(DeepLab标准)。interpolate确保尺寸一致。print便于验证每个分支输出尺寸不变。
5. 测试部分
x_in = torch.rand((4, 16, 64, 64))
aspp_block = ASPP(in_channels=16, out_channels=32, is_bn=True)
x_out = aspp_block(x_in)
print("input size : ", x_in.size())
print("output size : ", x_out.size())
  • 目的:创建随机输入,实例化ASPP,运行前向传播,并打印输入/输出尺寸。
  • 详细解释torch.rand(shape)生成[0,1)均匀分布随机张量,作为模拟输入。aspp_block实例化(in=16, out=32, BN=True)。x_out = aspp_block(x_in)调用forward。print验证输入[4,16,64,64],输出[4,32,64,64](尺寸不变,通道变化)。
  • 为什么:调试模型,确保端到端工作。随机输入模拟骨干网络输出。
总体总结
  • 代码流程:辅助模块 → ASPP组装 → 测试验证。
  • 适用场景:语义分割的特征提取层,如DeepLabv3中置于骨干末端。捕捉多尺度上下文,提升mIoU。
  • 潜在改进:添加Dropout防过拟合;用更高级骨干(如ResNet);训练时结合CrossEntropyLoss。
  • 运行输出示例(基于计算):
    [ASPP] out_1x1 size : torch.Size([4, 32, 64, 64])
    [ASPP] out_r6 size : torch.Size([4, 32, 64, 64])
    [ASPP] out_r12 size : torch.Size([4, 32, 64, 64])
    [ASPP] out_r18 size : torch.Size([4, 32, 64, 64])
    [ASPP] out_avg size : torch.Size([4, 32, 64, 64])
    input size : torch.Size([4, 16, 64, 64])
    output size : torch.Size([4, 32, 64, 64])
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐