别再只用SE模块了!手把手教你用PyTorch实现CBAM注意力(附完整代码)
·
超越SE模块:PyTorch实战CBAM注意力机制全解析
在计算机视觉领域,注意力机制已经成为提升卷积神经网络性能的关键组件。当大多数开发者还停留在使用SE(Squeeze-and-Excitation)模块时,CBAM(Convolutional Block Attention Module)已经展现出更强大的特征细化能力。本文将带你深入理解CBAM的工作原理,并手把手教你如何在PyTorch中实现这一先进注意力机制。
1. CBAM与SE模块的核心差异
SE模块通过全局平均池化获取通道注意力,而CBAM则从两个维度进行特征优化:
- 双注意力机制 :同时考虑通道和空间维度
- 特征聚合方式 :结合平均池化和最大池化的优势
- 计算效率 :保持轻量级设计,几乎不增加计算负担
下表对比了两种注意力模块的关键特性:
| 特性 | SE模块 | CBAM模块 |
|---|---|---|
| 注意力维度 | 仅通道 | 通道+空间 |
| 池化方式 | 平均池化 | 平均+最大池化 |
| 参数量 | 较少 | 略微增加 |
| 计算开销 | 低 | 中等 |
| 适用场景 | 分类任务 | 分类+检测+分割 |
提示:CBAM的空间注意力特别适合需要精确定位的任务,如目标检测和图像分割
2. CBAM模块的PyTorch实现
让我们从零开始构建CBAM模块。完整的实现包含通道注意力和空间注意力两个子模块。
2.1 通道注意力模块
import torch
import torch.nn as nn
class ChannelAttention(nn.Module):
def __init__(self, in_channels, reduction_ratio=16):
super(ChannelAttention, self).__init__()
self.avg_pool = nn.AdaptiveAvgPool2d(1)
self.max_pool = nn.AdaptiveMaxPool2d(1)
self.fc = nn.Sequential(
nn.Linear(in_channels, in_channels // reduction_ratio),
nn.ReLU(inplace=True),
nn.Linear(in_channels // reduction_ratio, in_channels)
)
self.sigmoid = nn.Sigmoid()
def forward(self, x):
avg_out = self.fc(self.avg_pool(x).view(x.size(0), -1))
max_out = self.fc(self.max_pool(x).view(x.size(0), -1))
out = avg_out + max_out
return self.sigmoid(out).unsqueeze(2).unsqueeze(3) * x
这段代码实现了CBAM的通道注意力部分,关键点包括:
- 同时使用平均池化和最大池化捕获不同统计特征
- 共享的全连接层减少参数数量
- Sigmoid激活生成0-1的注意力权重
2.2 空间注意力模块
class SpatialAttention(nn.Module):
def __init__(self, kernel_size=7):
super(SpatialAttention, self).__init__()
assert kernel_size % 2 == 1, "Kernel size must be odd"
self.conv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size//2)
self.sigmoid = nn.Sigmoid()
def forward(self, x):
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)
att = self.conv(concat)
return self.sigmoid(att) * x
空间注意力模块的特点:
- 沿通道维度进行平均和最大池化
- 使用卷积层融合空间信息
- 可调节的卷积核大小(通常设为7)
2.3 完整CBAM模块
将两个子模块组合起来,形成完整的CBAM:
class CBAM(nn.Module):
def __init__(self, in_channels, reduction_ratio=16, kernel_size=7):
super(CBAM, self).__init__()
self.channel_att = ChannelAttention(in_channels, reduction_ratio)
self.spatial_att = SpatialAttention(kernel_size)
def forward(self, x):
x = self.channel_att(x)
x = self.spatial_att(x)
return x
3. 将CBAM集成到常见网络架构
CBAM的灵活性使其可以轻松插入各种网络结构中。下面以ResNet为例,展示集成方法。
3.1 修改ResNet的基本块
class BasicBlockWithCBAM(nn.Module):
expansion = 1
def __init__(self, inplanes, planes, stride=1, downsample=None):
super(BasicBlockWithCBAM, self).__init__()
self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=3, stride=stride, padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(planes)
self.relu = nn.ReLU(inplace=True)
self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(planes)
self.downsample = downsample
self.stride = stride
self.cbam = CBAM(planes * self.expansion)
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
out = self.relu(out)
return out
3.2 在自定义网络中应用CBAM
对于自定义网络,CBAM可以灵活地插入到任何卷积层之后:
class CustomNet(nn.Module):
def __init__(self, num_classes=1000):
super(CustomNet, self).__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(kernel_size=3, stride=2, padding=1),
# 添加CBAM模块
CBAM(64),
nn.Conv2d(64, 128, kernel_size=3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
# 再次添加CBAM
CBAM(128),
nn.AdaptiveAvgPool2d((1, 1))
)
self.classifier = nn.Linear(128, num_classes)
def forward(self, x):
x = self.features(x)
x = x.view(x.size(0), -1)
x = self.classifier(x)
return x
4. 实战技巧与性能优化
在实际项目中应用CBAM时,以下几点经验值得注意:
- 位置选择 :CBAM通常放在残差连接之前,这样注意力可以同时作用于主路径和跳跃连接
- 缩减比例 :通道注意力的缩减比例(reduction_ratio)一般设为16,但对于小模型可以适当减小
- 初始化策略 :CBAM模块中的全连接层应采用适当的初始化,如Kaiming初始化
4.1 训练技巧
# 示例训练循环
model = ResNetWithCBAM(num_classes=1000)
optimizer = torch.optim.SGD(model.parameters(), lr=0.1, momentum=0.9, weight_decay=1e-4)
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)
criterion = nn.CrossEntropyLoss()
for epoch in range(100):
model.train()
for inputs, targets in train_loader:
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, targets)
loss.backward()
optimizer.step()
scheduler.step()
# 验证集评估
model.eval()
with torch.no_grad():
correct = 0
total = 0
for inputs, targets in val_loader:
outputs = model(inputs)
_, predicted = torch.max(outputs.data, 1)
total += targets.size(0)
correct += (predicted == targets).sum().item()
print(f'Epoch {epoch}, Accuracy: {100 * correct / total}%')
4.2 可视化注意力效果
理解CBAM如何影响特征图非常重要。以下代码展示了如何可视化注意力图:
import matplotlib.pyplot as plt
def visualize_attention(model, image):
# 前向传播获取中间特征
features = model.get_intermediate_features(image.unsqueeze(0))
# 可视化通道注意力
plt.figure(figsize=(12, 6))
for i in range(min(16, features.size(1))): # 显示前16个通道
plt.subplot(4, 4, i+1)
plt.imshow(features[0, i].detach().cpu().numpy(), cmap='viridis')
plt.axis('off')
plt.suptitle('Feature Maps with CBAM Attention')
plt.show()
在实际项目中,CBAM通常能带来1-2%的准确率提升,特别是在细粒度分类和目标检测任务中效果更为明显。相比SE模块,CBAM由于同时考虑了空间和通道信息,对于位置敏感的任务优势更加突出。
更多推荐


所有评论(0)