1. 初识PSPNet:语义分割的利器

第一次接触PSPNet是在处理一个医学影像分割项目时,当时需要从CT扫描图中精确分割出肺部组织。传统方法效果总是不尽如人意,直到尝试了PSPNet,效果提升明显。PSPNet(Pyramid Scene Parsing Network)是2017年CVPR提出的语义分割网络,它的核心创新在于引入了金字塔池化模块(PSP模块),能够有效捕获多尺度上下文信息。

PSPNet特别适合处理复杂场景下的语义分割任务。比如在自动驾驶中,需要同时识别远处的交通标志和近处的行人;在遥感图像分析中,要区分不同大小的建筑物和植被。这些场景都要求模型具备捕捉不同尺度信息的能力,而这正是PSPNet的强项。

与FCN、U-Net等经典分割网络相比,PSPNet在以下几个关键点上做了改进:

  • 多尺度特征融合:通过金字塔池化模块,同时考虑全局和局部特征
  • 上下文信息增强:不同大小的池化核帮助模型理解场景的全局结构
  • 细节保持能力:在特征提取阶段保留更多空间信息

2. 轻量化PSPNet设计:MobileNetV2主干网络

2.1 为什么选择MobileNetV2

在实际项目中,我们经常需要在资源受限的环境(如移动设备或边缘计算设备)部署模型。原版PSPNet使用ResNet作为主干网络,虽然性能优秀,但计算量和参数量都较大。经过多次实验对比,我发现MobileNetV2是一个理想的轻量化替代方案。

MobileNetV2的核心是倒残差结构(Inverted Residuals)和线性瓶颈层(Linear Bottlenecks)。与常规残差块先压缩再扩张不同,倒残差结构先扩张通道数再进行深度可分离卷积,最后压缩回原通道数。这种设计在保持模型容量的同时大幅减少了计算量。

class InvertedResidual(nn.Module):
    def __init__(self, inp, oup, stride, expand_ratio):
        super(InvertedResidual, self).__init__()
        self.stride = stride
        assert stride in [1, 2]
        
        hidden_dim = round(inp * expand_ratio)
        self.use_res_connect = self.stride == 1 and inp == oup
        
        layers = []
        if expand_ratio != 1:
            layers.append(nn.Conv2d(inp, hidden_dim, 1, 1, 0, bias=False))
            layers.append(nn.BatchNorm2d(hidden_dim))
            layers.append(nn.ReLU6(inplace=True))
            
        layers.extend([
            nn.Conv2d(hidden_dim, hidden_dim, 3, stride, 1, 
                     groups=hidden_dim, bias=False),
            nn.BatchNorm2d(hidden_dim),
            nn.ReLU6(inplace=True),
            nn.Conv2d(hidden_dim, oup, 1, 1, 0, bias=False),
            nn.BatchNorm2d(oup),
        ])
        
        self.conv = nn.Sequential(*layers)

2.2 主干网络改造技巧

将PSPNet的主干网络替换为MobileNetV2时,有几个关键点需要注意:

  1. 下采样策略调整:原版MobileNetV2有5次下采样,但语义分割通常只需要3-4次。我通常保留前4个下采样层,这样能在保持足够空间分辨率的同时提取深层特征。

  2. 特征通道匹配:MobileNetV2最后的特征通道数(320)与ResNet(2048)不同,需要相应调整PSP模块的输入通道数。

  3. 预训练权重使用:加载在ImageNet上预训练的MobileNetV2权重可以显著提升模型性能。在PyTorch中可以直接使用torchvision提供的预训练模型:

from torchvision.models import mobilenet_v2

backbone = mobilenet_v2(pretrained=True).features

3. PSP模块的优化实现

3.1 标准PSP模块解析

PSP模块的核心思想是通过不同尺度的池化操作捕获多尺度上下文信息。标准的PSP模块包含四个并行分支:

  1. 1x1的全局平均池化(捕获全局上下文)
  2. 2x2的自适应平均池化
  3. 3x3的自适应平均池化
  4. 6x6的自适应平均池化

这些不同尺度的特征图经过卷积处理后,会被上采样到原始尺寸并拼接在一起。这种设计让模型同时"看到"局部细节和全局场景。

class PSPModule(nn.Module):
    def __init__(self, in_channels, pool_sizes=[1,2,3,6]):
        super(PSPModule, self).__init__()
        out_channels = in_channels // len(pool_sizes)
        
        self.stages = nn.ModuleList([
            self._make_stage(in_channels, out_channels, size) 
            for size in pool_sizes
        ])
        
        self.bottleneck = nn.Sequential(
            nn.Conv2d(in_channels + out_channels*len(pool_sizes), out_channels, 
                     kernel_size=3, padding=1, bias=False),
            nn.BatchNorm2d(out_channels),
            nn.ReLU(inplace=True),
            nn.Dropout2d(0.1)
        )
    
    def _make_stage(self, in_channels, out_channels, size):
        prior = nn.AdaptiveAvgPool2d(output_size=(size, size))
        conv = nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)
        return nn.Sequential(prior, conv)
    
    def forward(self, x):
        input_size = x.size()[2:]
        pyramids = [x]
        pyramids.extend([
            F.interpolate(stage(x), size=input_size, mode='bilinear', align_corners=True)
            for stage in self.stages
        ])
        output = self.bottleneck(torch.cat(pyramids, dim=1))
        return output

3.2 轻量化改进技巧

在实际部署中,我发现可以通过以下方法进一步优化PSP模块:

  1. 减少池化分支:对于小尺寸输入(如256x256),可以去掉6x6池化分支,只保留1x1、2x2和3x3分支。

  2. 通道压缩:在拼接各分支特征前,先用1x1卷积压缩通道数,减少后续计算量。

  3. 深度可分离卷积:将PSP模块中的常规卷积替换为深度可分离卷积,可以大幅减少参数量。

class LightPSPModule(nn.Module):
    def __init__(self, in_channels, pool_sizes=[1,2,3]):
        super().__init__()
        out_channels = in_channels // 4
        
        self.stages = nn.ModuleList([
            nn.Sequential(
                nn.AdaptiveAvgPool2d(size),
                nn.Conv2d(in_channels, out_channels, 1, bias=False),
                nn.BatchNorm2d(out_channels),
                nn.ReLU6(inplace=True)
            )
            for size in pool_sizes
        ])
        
        self.global_pool = nn.Sequential(
            nn.AdaptiveAvgPool2d(1),
            nn.Conv2d(in_channels, out_channels, 1, bias=False),
            nn.BatchNorm2d(out_channels),
            nn.ReLU6(inplace=True)
        )
        
        self.conv = nn.Sequential(
            nn.Conv2d(in_channels + out_channels*(len(pool_sizes)+1), 
                     out_channels, 1, bias=False),
            nn.BatchNorm2d(out_channels),
            nn.ReLU6(inplace=True)
        )

4. 训练技巧与优化策略

4.1 损失函数选择

语义分割常用的损失函数是交叉熵损失,但对于类别不平衡的数据集(如医学图像中病灶区域通常很小),单纯的交叉熵效果不佳。经过多次实验,我发现结合Dice Loss和交叉熵的复合损失函数效果最好。

class MixedLoss(nn.Module):
    def __init__(self, alpha=0.5):
        super().__init__()
        self.alpha = alpha
        self.ce_loss = nn.CrossEntropyLoss()
        
    def dice_loss(self, inputs, targets):
        smooth = 1.0
        inputs = torch.softmax(inputs, dim=1)
        
        iflat = inputs.flatten(2)
        tflat = targets.flatten(2)
        
        intersection = (iflat * tflat).sum(2)
        dice = (2. * intersection + smooth) / (iflat.sum(2) + tflat.sum(2) + smooth)
        return 1 - dice.mean()
    
    def forward(self, inputs, targets):
        ce = self.ce_loss(inputs, targets)
        dice = self.dice_loss(inputs, targets)
        return self.alpha * ce + (1 - self.alpha) * dice

4.2 数据增强策略

有效的数

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐