共计 2969 个字符,预计需要花费 8 分钟才能阅读完成。
CBAM(Convolutional Block Attention Module)是计算机视觉中一种轻量级的注意力模块,能显著提升模型对重要特征的敏感度,尤其在图像分类、目标检测等任务中表现出色。下面我们从原理到实践,一步步拆解 CBAM 的核心机制。

1. CBAM 双注意力机制图解
CBAM 包含两个串联的注意力模块:通道注意力和空间注意力。这种双注意力机制能让模型同时关注 ”what”(通道维度)和 ”where”(空间维度)的重要信息。
通道注意力模块
通道注意力通过全局平均池化和最大池化捕获通道间关系,其计算过程为:
M_c(F) = \sigma(MLP(AvgPool(F)) + MLP(MaxPool(F)))
F是输入特征图,形状为C×H×Wσ表示 Sigmoid 激活函数- 两个 MLP 共享权重,输出通道数为
C/r(r 为缩减比率)
空间注意力模块
空间注意力则关注特征图中的重要区域,计算公式为:
M_s(F) = \sigma(f^{7×7}([AvgPool(F); MaxPool(F)]))
f^{7×7}表示 7×7 卷积[;]表示沿通道维度的拼接
2. 与 SE-Net 的对比
相比 SE-Net 仅考虑通道注意力,CBAM 的主要优势在于:
- 增加了空间注意力,能定位重要区域
- 计算开销仅轻微增加(约 10% FLOPs)
- 在 ImageNet 上比 SE-Net 提升约 1.5% top- 1 准确率
3. PyTorch 实现详解
以下是完整的 CBAM 模块实现(带维度注释):
import torch
import torch.nn as nn
class CBAM(nn.Module):
def __init__(self, channels, reduction=16):
super().__init__()
# 通道注意力
self.avg_pool = nn.AdaptiveAvgPool2d(1)
self.max_pool = nn.AdaptiveMaxPool2d(1)
self.fc = nn.Sequential(nn.Linear(channels, channels // reduction),
nn.ReLU(),
nn.Linear(channels // reduction, channels)
)
# 空间注意力
self.conv = nn.Conv2d(2, 1, kernel_size=7, padding=3)
def forward(self, x):
b, c, _, _ = x.shape
# 通道注意力
avg_out = self.fc(self.avg_pool(x).view(b, c))
max_out = self.fc(self.max_pool(x).view(b, c))
channel = torch.sigmoid(avg_out + max_out).view(b, c, 1, 1)
x = x * channel
# 空间注意力
avg_out = torch.mean(x, dim=1, keepdim=True)
max_out, _ = torch.max(x, dim=1, keepdim=True)
spatial = torch.cat([avg_out, max_out], dim=1)
spatial = torch.sigmoid(self.conv(spatial))
return x * spatial
4. 嵌入 ResNet 示例
在 ResNet 中嵌入 CBAM 的典型方式(以 BasicBlock 为例):
class BasicBlock(nn.Module):
def __init__(self, inplanes, planes, stride=1):
super().__init__()
self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=3, stride=stride, padding=1)
self.bn1 = nn.BatchNorm2d(planes)
self.relu = nn.ReLU()
self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, padding=1)
self.bn2 = nn.BatchNorm2d(planes)
self.cbam = CBAM(planes) # 在残差连接前添加 CBAM
if stride != 1 or inplanes != planes:
self.downsample = nn.Sequential(nn.Conv2d(inplanes, planes, kernel_size=1, stride=stride),
nn.BatchNorm2d(planes)
)
else:
self.downsample = None
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.cbam(out) # 应用 CBAM
if self.downsample is not None:
identity = self.downsample(x)
out += identity
return self.relu(out)
5. 特征可视化方法
使用 matplotlib 可视化注意力权重:
import matplotlib.pyplot as plt
def visualize_attention(model, img_tensor):
# 获取中间层输出
activations = {}
def hook_fn(module, input, output):
activations['cbam'] = output[1] # 假设返回(特征图, 注意力权重)
handle = model.cbam.register_forward_hook(hook_fn)
with torch.no_grad():
_ = model(img_tensor.unsqueeze(0))
handle.remove()
# 绘制空间注意力热图
plt.imshow(activations['cbam'].cpu().squeeze(), cmap='jet')
plt.colorbar()
plt.title('Spatial Attention Map')
plt.show()
6. 实践调优建议
数据预处理
- 输入图像建议归一化到 [0,1] 或使用 ImageNet 的 mean/std
- 数据增强推荐:RandomResizedCrop + HorizontalFlip
超参数选择
- reduction 比率 r:通常选 8 -16(通道数少时选较小值)
- 放置位置:每个残差块后效果优于仅放在网络末端
计算开销
- 参数量增加:约 0.1%-0.5% (ResNet-50 增加约 2M 参数)
- FLOPs 增加:约 5 -15%(取决于网络深度)
7. 开放性问题思考
CBAM 的成功引发了一些有趣的扩展方向:
- 视频分析:如何将时空注意力结合?能否用 3D 卷积扩展空间注意力?
- 与 Transformer 融合:CBAM 的局部注意力能否与 Vision Transformer 的全局注意力互补?
- 轻量化设计:对于移动端设备,如何进一步压缩 CBAM 的计算量?
在实际项目中,我们发现 CBAM 对小目标检测和细粒度分类特别有效。建议初学者先从 CIFAR-10 等小数据集开始实验,逐步掌握注意力模块的调参技巧。
正文完
