yolo核心组件4:CA 注意力机制 (Coordinate Attention)

CA 注意力机制 (Coordinate Attention)

!abstract 论文信息

  • 论文标题: Coordinate Attention for Efficient Mobile Network Design
  • 作者: Qibin Hou, Daquan Zhou, Jiashi Feng
  • 发表: CVPR 2021
  • 论文地址: https://arxiv.org/abs/2103.02907
  • 核心贡献: 提出坐标注意力机制,将位置信息嵌入通道注意力中,使移动端网络能够关注大区域且保留精确的位置信息

一、核心思想

CA (Coordinate Attention) 的核心思想是将位置信息编码到通道注意力中。与 SE 只考虑通道关系、CBAM 使用局部卷积捕获空间信息不同,CA 通过两个方向(水平和垂直)的全局池化,分别沿宽度和高度方向聚合特征,从而保留了精确的位置信息。

CA 的关键特点:

  1. 位置感知:不同于全局池化丢失位置信息,CA 在水平和垂直方向分别编码,保留了空间位置
  2. 长程依赖:通过沿一个方向的全局池化,可以捕获该方向的长程依赖关系
  3. 双分支结构:水平和垂直两个方向并行处理,生成两个注意力图
  4. 移动端友好:设计轻量,适合部署在移动端设备上

二、模块结构

2.1 整体流程

复制代码
输入特征图 X [C×H×W]
        │
  ┌─────┴─────┐
  │           │
水平方向      垂直方向
全局池化      全局池化
(H方向)      (W方向)
  │           │
[C×H×1]    [C×1×W]
  │           │
  Concat     (维度变换)
  │
[C×(H+W)×1]
  │
Conv1×1 + BN + h_swish
  │
  ┌─────┴─────┐
  │           │
水平注意力    垂直注意力
  │           │
Sigmoid     Sigmoid
  │           │
Xh × X      Xw × X
  │           │
  输出特征图

2.2 Coordinate Attention 详细结构

复制代码
输入 X [C×H×W]
      │
  ┌───┴───┐
  │       │
  │   Coordinate
  │   Information
  │   Embedding (CIE)
  │       │
  │   ┌───┴───┐
  │   │       │
  │ AvgPool  AvgPool
  │ 沿W方向  沿H方向
  │   │       │
  │ [C×H×1] [C×1×W]
  │   │       │
  │   └───┬───┘
  │       │
  │   Concat + Transform
  │   Conv1×1 + BN + h_swish
  │       │
  │   Split
  │   ┌───┴───┐
  │   │       │
  │ Conv1×1  Conv1×1
  │ +Sigmoid +Sigmoid
  │   │       │
  │   └───┬───┘
  │       │
  └───┬───┘
      │
  输出 X' [C×H×W]

三、数学公式

3.1 坐标信息嵌入 (Coordinate Information Embedding)

水平方向(沿宽度 W 聚合):

zch(h)=1W∑0≤i<Wxc(h,i)z_c^h(h) = \frac{1}{W} \sum_{0 \leq i < W} x_c(h, i)zch(h)=W10≤i<W∑xc(h,i)

输出 zh∈RC×H×1z^h \in \mathbb{R}^{C \times H \times 1}zh∈RC×H×1,每个通道保留了高度方向的位置信息。

垂直方向(沿高度 H 聚合):

zcw(w)=1H∑0≤j<Hxc(j,w)z_c^w(w) = \frac{1}{H} \sum_{0 \leq j < H} x_c(j, w)zcw(w)=H10≤j<H∑xc(j,w)

输出 zw∈RC×1×Wz^w \in \mathbb{R}^{C \times 1 \times W}zw∈RC×1×W,每个通道保留了宽度方向的位置信息。

3.2 坐标注意力生成

将两个方向的特征拼接后通过共享变换:

f=δ(BN(Conv1×1(zh;zw)))f = \delta\Big(BN\big(Conv_{1\times 1}(z\^h; z\^w)\big)\Big)f=δ(BN(Conv1×1(zh;zw)))

其中:

  • ⋅;⋅\\cdot;\\cdot⋅;⋅ 表示空间维度的拼接,f∈RC/r×(H+W)f \in \mathbb{R}^{C/r \times (H+W)}f∈RC/r×(H+W)
  • δ\deltaδ 为 h-swish 激活函数
  • rrr 为缩减率(默认16)

然后将 fff 沿空间维度分割:

fh∈RC/r×H,fw∈RC/r×Wf^h \in \mathbb{R}^{C/r \times H}, \quad f^w \in \mathbb{R}^{C/r \times W}fh∈RC/r×H,fw∈RC/r×W

3.3 注意力权重生成

gh=σ(Conv1×1(fh))∈RC×H×1g^h = \sigma\Big(Conv_{1\times 1}(f^h)\Big) \in \mathbb{R}^{C \times H \times 1}gh=σ(Conv1×1(fh))∈RC×H×1

gw=σ(Conv1×1(fw))∈RC×1×Wg^w = \sigma\Big(Conv_{1\times 1}(f^w)\Big) \in \mathbb{R}^{C \times 1 \times W}gw=σ(Conv1×1(fw))∈RC×1×W

3.4 最终输出

yc(i,j)=xc(i,j)×gch(i)×gcw(j)y_c(i, j) = x_c(i, j) \times g_c^h(i) \times g_c^w(j)yc(i,j)=xc(i,j)×gch(i)×gcw(j)


四、代码实现

python 复制代码
import torch
import torch.nn as nn

class h_sigmoid(nn.Module):
    """Hard Sigmoid 激活函数"""
    def __init__(self, inplace=True):
        super().__init__()
        self.relu = nn.ReLU6(inplace=inplace)

    def forward(self, x):
        return self.relu(x + 3) / 6

class h_swish(nn.Module):
    """Hard Swish 激活函数"""
    def __init__(self, inplace=True):
        super().__init__()
        self.hsigmoid = h_sigmoid(inplace=inplace)

    def forward(self, x):
        return x * self.hsigmoid(x)

class CoordAtt(nn.Module):
    """Coordinate Attention 模块"""
    def __init__(self, channels, reduction=32):
        super().__init__()
        mid_channels = max(8, channels // reduction)

        # 共享变换层
        self.pool_h = nn.AdaptiveAvgPool2d((None, 1))  # 沿W方向池化
        self.pool_w = nn.AdaptiveAvgPool2d((1, None))  # 沿H方向池化

        self.conv1 = nn.Conv2d(channels, mid_channels, kernel_size=1,
                               stride=1, padding=0, bias=False)
        self.bn1 = nn.BatchNorm2d(mid_channels)
        self.act = h_swish()

        # 分支卷积
        self.conv_h = nn.Conv2d(mid_channels, channels, kernel_size=1,
                                stride=1, padding=0, bias=False)
        self.conv_w = nn.Conv2d(mid_channels, channels, kernel_size=1,
                                stride=1, padding=0, bias=False)

    def forward(self, x):
        b, c, h, w = x.size()

        # 沿两个方向进行全局池化
        x_h = self.pool_h(x)  # [B, C, H, 1]
        x_w = self.pool_w(x).permute(0, 1, 3, 2)  # [B, C, W, 1]

        # 拼接并变换
        y = torch.cat([x_h, x_w], dim=2)  # [B, C, H+W, 1]
        y = self.conv1(y)
        y = self.bn1(y)
        y = self.act(y)

        # 分割
        x_h, x_w = torch.split(y, [h, w], dim=2)
        x_w = x_w.permute(0, 1, 3, 2)  # [B, C/r, 1, W]

        # 生成注意力权重
        att_h = self.conv_h(x_h).sigmoid()  # [B, C, H, 1]
        att_w = self.conv_w(x_w).sigmoid()  # [B, C, 1, W]

        # 应用注意力
        return x * att_h * att_w

4.2 YOLO 集成版本

python 复制代码
class C2f_CA(nn.Module):
    """在C2f模块后添加CA注意力"""
    def __init__(self, c1, c2, n=1, shortcut=False, e=0.5):
        super().__init__()
        self.c2f = C2f(c1, c2, n, shortcut, e)
        self.ca = CoordAtt(c2)

    def forward(self, x):
        x = self.c2f(x)
        x = self.ca(x)
        return x

五、在YOLO中的应用

5.1 适用场景

CA 注意力在以下场景中特别有效:

场景 原因
小目标检测 位置信息对小目标定位至关重要
密集目标场景 精确的空间注意力帮助区分相邻目标
自动驾驶 车辆、行人等目标需要精确定位
遥感图像 目标位置信息对检测很重要

5.2 YOLO 配置示例

yaml 复制代码
# 在YOLOv8的Backbone中添加CA
backbone:
  - [-1, 1, Conv, [64, 3, 2]]
  - [-1, 1, Conv, [128, 3, 2]]
  - [-1, 3, C2f, [128, True]]
  - [-1, 1, Conv, [256, 3, 2]]
  - [-1, 6, C2f, [256, True]]
  - [-1, 1, CoordAtt, [256]]       # 在P3后添加CA
  - [-1, 1, Conv, [512, 3, 2]]
  - [-1, 6, C2f, [512, True]]
  - [-1, 1, Conv, [1024, 3, 2]]
  - [-1, 3, C2f, [1024, True]]
  - [-1, 1, CoordAtt, [1024]]      # 在P5后添加CA

5.3 注意力机制对比

注意力 位置信息 参数量 适用场景
SE 无(全局池化) 中等 通用分类
CBAM 局部(7×7卷积) 中等 通用检测
ECA 无(全局池化) 极低 轻量化模型
CA 精确(行列编码) 位置敏感任务

六、优缺点

优点

  1. 精确的位置编码:通过水平和垂直方向的独立池化,保留了精确的空间位置信息
  2. 长程依赖:沿某一方向的全局聚合可以捕获该方向的长程依赖关系
  3. 轻量化设计:参数量远少于 CBAM 和 SE,适合移动端部署
  4. 双方向互补:水平和垂直方向的注意力互补,增强了空间感知能力
  5. 即插即用:可灵活嵌入各种网络架构

缺点

  1. 二维设计限制:CA 的设计基于二维特征图,扩展到三维场景需要额外设计
  2. 计算开销:相比 ECA 稍高,但在可接受范围内
  3. 池化压缩:沿某方向的全局池化仍会损失该方向上的局部细节
  4. 对通道数敏感:缩减率的选择会影响性能,需要适当调参

参考

  1. Hou, Q., Zhou, D., & Feng, J. (2021). Coordinate Attention for Efficient Mobile Network Design. CVPR 2021.
  2. https://arxiv.org/abs/2103.02907
  3. https://github.com/houqb/CoordAttention
相关推荐
m沐沐2 小时前
【计算机视觉】人脸识别三大经典算法:LBPH、Eigenfaces、FisherFaces 原理与实战
图像处理·人工智能·深度学习·opencv·算法·机器学习·计算机视觉
今天AI了吗2 小时前
Spring AI 框架实战:Java 后端集成大模型的架构设计与工程落地
java·人工智能·python·spring·机器学习
酉鬼女又兒3 小时前
[特殊字符]零基础入门AI:归纳演绎、假设空间、归纳偏好、NFL、过拟合与欠拟合、模型评估选择、超参数、性能度量、混淆矩阵、P-R曲线和F1
人工智能·windows·python·深度学习·安全·机器学习·矩阵
道影子3 小时前
阴阳引力多维基元:超越二进制的计算革命
人工智能·深度学习·微服务·云原生·架构
学习中.........3 小时前
从 Karpathy 手写 BPE 到 CS336 `train_bpe` 作业实现
人工智能·机器学习·语言模型·自然语言处理
手写码匠3 小时前
Android 17 灵魂拷问深度解析:隐私、大屏、AI 端侧全面适配实战
人工智能·深度学习·算法·aigc
JackHCC4 小时前
Meta-WHALE:逐层融合统一推荐模型
人工智能·机器学习
若殇丶球4 小时前
二、PointNet算法
深度学习
道影子5 小时前
《道德经》031兵者不祥,胜以丧礼处之
人工智能·深度学习·算法