CBAM卷积注意力网络结构图解析与高效实现方案

1次阅读
没有评论

共计 1647 个字符,预计需要花费 5 分钟才能阅读完成。

image.webp

背景痛点:CNN 特征选择的局限性

传统 CNN 通过堆叠卷积层逐步提取特征,但存在两个明显问题:

  • 卷积核对所有通道和空间位置平等对待,无法自适应聚焦重要特征区域
  • 全局平均池化 (GAP) 操作会丢失空间细节信息,影响细粒度分类任务

这就像用相同力度扫描整张图片,无法对关键区域(如猫的耳朵、车轮纹理)进行针对性观察。

技术对比:主流注意力机制

SE 模块(2017)

  • 仅考虑通道维度注意力
  • 通过 GAP 压缩空间信息
# SE 模块结构示意
[GAP] -> [FC] -> [ReLU] -> [FC] -> [Sigmoid]

Non-local Networks(2018)

  • 计算所有像素点间的关联性
  • 计算复杂度高达 O(H^2W^2)

CBAM 创新点(2018)

  1. 通道 + 空间双注意力协同
  2. 采用 1D 卷积降低计算量
  3. 最大池化与平均池化特征融合

CBAM 卷积注意力网络结构图解析与高效实现方案

核心实现:PyTorch 代码详解

通道注意力模块

class ChannelAttention(nn.Module):
    def __init__(self, channels, reduction=16):
        super().__init__()
        # 共享 MLP 替代方案(用 1D 卷积实现)self.mlp = nn.Sequential(nn.Conv1d(channels, channels//reduction, 1),  # [B,C,1]
            nn.ReLU(),
            nn.Conv1d(channels//reduction, channels, 1)   # [B,C,1]
        )

    def forward(self, x):
        B, C, _, _ = x.shape
        # 双路池化特征 [B,C,1,1]
        avg_out = self.mlp(F.avg_pool2d(x, x.size()[2:]).view(B,C,1))
        max_out = self.mlp(F.max_pool2d(x, x.size()[2:]).view(B,C,1))
        # 特征融合 [B,C,1,1]
        scale = torch.sigmoid(avg_out + max_out).unsqueeze(-1)
        return x * scale

空间注意力模块

class SpatialAttention(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv = nn.Conv2d(2, 1, 7, padding=3)

    def forward(self, x):
        # 通道维度池化 [B,2,H,W]
        avg_out = torch.mean(x, dim=1, keepdim=True)
        max_out, _ = torch.max(x, dim=1, keepdim=True)
        concat = torch.cat([avg_out, max_out], dim=1)
        # 空间注意力权重 [B,1,H,W]
        scale = torch.sigmoid(self.conv(concat))
        return x * scale

完整 CBAM 集成

class CBAM(nn.Module):
    def __init__(self, channels):
        super().__init__()
        self.ca = ChannelAttention(channels)
        self.sa = SpatialAttention()

    def forward(self, x):
        x = self.ca(x)  # 通道注意力
        x = self.sa(x)  # 空间注意力
        return x

性能考量与实验数据

计算开销测试(ResNet50)

模型 FLOPs 参数量 Top-1 Acc
Baseline 4.1G 25.5M 76.2%
+SE 4.11G 28.1M 77.3%
+CBAM 4.13G 28.4M 78.1%

注意力热图可视化

避坑指南

  1. 输入归一化:推荐在 CBAM 前做 BatchNorm
  2. 小批量训练:当 batch_size<32 时,空间注意力改用 3 ×3 卷积
  3. 部署优化:将 sigmoid+ 乘法融合为 Scale 层

延伸思考:适配 Transformer

  1. 在 ViT 的 MLP 层后插入 CBAM
  2. 将空间注意力改为窗口划分
  3. 通道注意力作用于多头注意力维度

Colab 实践链接

正文完
 0
评论(没有评论)