"""CBAM (Convolutional Block Attention Module) — 残差初始化版本 通道注意力:学习哪些特征通道更重要(类似 SE/ECA) 空间注意力:学习哪些空间位置更重要(对遮挡、密集场景有帮助) 残差初始化:alpha 初始=0 → 输出恒等,不破坏预训练特征。 """ import torch import torch.nn as nn class ChannelAttention(nn.Module): def __init__(self, channels, reduction=16): super().__init__() mid = max(channels // reduction, 8) self.avg_pool = nn.AdaptiveAvgPool2d(1) self.max_pool = nn.AdaptiveMaxPool2d(1) self.fc = nn.Sequential( nn.Conv2d(channels, mid, 1, bias=False), nn.ReLU(inplace=True), nn.Conv2d(mid, channels, 1, bias=False), ) self.sigmoid = nn.Sigmoid() def forward(self, x): return self.sigmoid(self.fc(self.avg_pool(x)) + self.fc(self.max_pool(x))) class SpatialAttention(nn.Module): def __init__(self, kernel_size=7): super().__init__() self.conv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size // 2, bias=False) 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) return self.sigmoid(self.conv(torch.cat([avg_out, max_out], dim=1))) class CBAM(nn.Module): def __init__(self, c1=None, c2=None, reduction=16, kernel_size=7): super().__init__() self.reduction = reduction self.kernel_size = kernel_size self.alpha = nn.Parameter(torch.zeros(1)) self.ca = None self.sa = None if c1 is not None: self._init_modules(c1) def _init_modules(self, channels): self.ca = ChannelAttention(channels, self.reduction) self.sa = SpatialAttention(self.kernel_size) def forward(self, x): if self.ca is None: self._init_modules(x.shape[1]) self.ca = self.ca.to(x.device) self.sa = self.sa.to(x.device) attn = self.ca(x) * self.sa(x) return x * (1.0 + self.alpha * (attn - 1.0)) def register_cbam(): import ultralytics.nn.tasks as tasks if "CBAM" not in vars(tasks): tasks.CBAM = CBAM