ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

CAA注意力机制详解:如何用跨轴注意力提升YOLOv8检测精度

CAA注意力机制详解:如何用跨轴注意力提升YOLOv8检测精度 每次帮人看YOLOv8改进实验我听到最多的就是“注意力模块加了参数涨了但精度纹丝不动”。这话说对了一半——不是注意力没用而是大多数人翻来覆去用的还是SE、CBAM、CA这几条一眼能看到头的路线真正贴合检测任务结构的注意力反而被忽略了。这篇文章要聊的CAACross-Axis Attention交叉轴注意力就不是那种“今天加个模块明天涨两个点”的玄学玩法而是通过水平轴和垂直轴两条正交的注意力路径用接近线性的计算代价去覆盖全局感受野。我会从设计动机讲到完整的PyTorch实现再给出三种接进YOLOv8的具体姿势、实测对比数据以及我在COCO子集上踩过的shape、显存和ONNX导出坑适合正在做检测模型改进、或者准备拿注意力机制凑论文实验的读者。1. CAA注意力机制解决什么问题从通道注意力到跨轴注意力1.1 为什么检测模型里的注意力必须“举重若轻”做检测和做分类有一个本质区别分类模型可以接受全局自注意力那种O(H²W²)的暴力计算因为特征图分辨率往往被压得很低但YOLOv8的neck和head要在40x40、80x80这类高分辨率特征图上操作如果照搬Transformer的全局注意力一张640x640的图P3特征图有6400个token两两之间算注意力就是千万级别的矩阵显存直接爆表推理速度也没法看。所以检测模型里能用的注意力必须在“尽量多看”和“尽量少算”之间找平衡。CAA的思路很直接全局注意力之所以贵是因为每个位置要看所有位置但一张自然图像里远距离相关性的分布通常是各向异性的水平方向和垂直方向的结构规律并不一样。与其让模型在二维平面上盲目搜索不如把它拆成两个一维搜索——每一行单独做自注意力每一列单独做自注意力。这就是Cross-Axis Attention最底层的直觉。拿生活里的例子类比你要在一个体育馆里找朋友全局注意力相当于挨个核对每个座位代价是几万次比较而CAA相当于先沿每一排扫一遍再沿每一列扫一遍两次线性扫描就能锁定大概区域。虽然不能保证100%等价于全局搜索但对目标检测这种本身就是局部强相关、全局弱相关的任务来说这套近似通常够用而且代价低得多。1.2 SE、CBAM、CA为什么越用越“卷”先别急着写代码把老牌注意力的问题说透你才知道CAA到底补了什么。SE通道注意力是最早被搬进YOLO的模块它的做法是对特征图做全局平均池化把HxW的空间压缩成一个点然后经过两个全连接层得到每个通道的权重。问题也出在这个“压缩”上通道权重是有了但空间细节一点没留等于告诉模型“这个通道重要”却不告诉它“重要区域在图像的哪个位置”。小目标检测最吃空间信息SE在这类场景下经常无效。CBAM在SE基础上加了空间注意力分支理论上兼顾了通道和空间。但它把注意力拆成“先通道后空间”的串行结构空间注意力用的还是7x7卷积感受野有限本质上属于局部操作对小目标可能有用对那种需要跨越大尺度上下文的目标比如被遮挡的车辆、和大背景混在一起的人依然无能为力。CACoordinate Attention比前两个聪明一些它把水平和垂直方向分别做池化再把两个方向的特征编码成坐标信息相当于让通道注意力带了“位置感知”。但CA的两个方向分支最终是被压缩成一对特征向量再合并回通道维度的中间不是注意力本身而是一种坐标嵌入细节信息的保留仍然有限。我把这几个模块的关键差异整理成一张表方便你对照注意力模块是否建模空间关系感受野覆盖新增参数量对小目标友好度SE否只做通道加权全局池化后是一个点低一般CBAM是但空间分支用7x7卷积局部低好CA弱坐标嵌入而非注意力全局池化压缩后的一维坐标低较好CAA是逐token交互同一行同一列的全部token中等好CAA不是要替代所有注意力而是在“空间交互强度”这个维度上比SE和CA走得更远同时在计算代价上又比全局自注意力小一个数量级。1.3 CAA的跨轴直觉横着看一次竖着看一次全景就出来了CAA的计算分成两条独立的分支水平轴分支把特征图按行切分每一行视作一个长度为W的token序列在这个序列内部做自注意力垂直轴分支把特征图按列切分每一列视作一个长度为H的token序列做同样的自注意力。最后把两个分支的输出拼接起来过一个1x1卷积融合再加上残差。这样说可能有点抽象你想象一张80x80的特征图水平分支会把80行分别处理每个位置的注意力只看同一行的另外79个位置垂直分支则让每个位置只看同一列的另外79个位置。单看一条轴每个位置只能感知一条线但两条线的交叉结果让每个位置最终能同时获取整行和整列的信息。虽然不如全局注意力那样直接看到二维平面上的任意远点但绝大多数跨区域依赖都可以通过“先横后竖”的路径建立起来——这本质上就是一个两步的消息传递。需要说明一句有些论文里CAA也指Context Anchor Attention那是RT-DETR系列常用的上下文锚点注意力叫法设计思路完全不同。本文实现的CAA是Cross-Axis Attention即交叉轴轴向注意力代码和后续实验都基于这个版本。写文章时提一下这个重名免得你查资料时对不上号。2. CAA的PyTorch实现可运行代码与逐行拆解2.1 完整代码下面这个模块我按“直接能跑”的标准写的不依赖第三方库只用了PyTorch自带的api。你把它丢到项目里随便一个python文件就能验证。import torch import torch.nn as nn class AxialAttention(nn.Module): 单轴自注意力对长度为L的token序列做多头注意力 输入x: (Bn, L, C)Bn是batch*行数或batch*列数 输出shape与输入一致 def __init__(self, dim, num_heads4): super().__init__() self.dim dim self.num_heads num_heads self.head_dim dim // num_heads self.scale self.head_dim ** -0.5 self.qkv nn.Linear(dim, dim * 3, biasFalse) self.proj nn.Linear(dim, dim, biasFalse) def forward(self, x): Bn, L, C x.shape qkv self.qkv(x).reshape(Bn, L, 3, self.num_heads, self.head_dim) qkv qkv.permute(2, 0, 3, 1, 4) # (3, Bn, heads, L, head_dim) q, k, v qkv.unbind(0) attn (q k.transpose(-2, -1)) * self.scale # (Bn, heads, L, L) attn attn.softmax(dim-1) x attn v # (Bn, heads, L, head_dim) x x.transpose(1, 2).reshape(Bn, L, C) return self.proj(x) class CAA(nn.Module): CAA: Cross-Axis Attention 水平轴注意力 垂直轴注意力 双轴融合 输入输出均为 (B, C, H, W) def __init__(self, c1, num_heads4): super().__init__() assert c1 % num_heads 0, c1 must be divisible by num_heads self.attn AxialAttention(c1, num_heads) # 水平轴和垂直轴共享一个注意力器 self.fuse nn.Sequential( nn.Conv2d(c1 * 2, c1, 1, biasFalse), nn.BatchNorm2d(c1), nn.SiLU(inplaceTrue) ) self.out nn.Conv2d(c1, c1, 1, biasFalse) def forward(self, x): B, C, H, W x.shape # 水平轴把每个“行”当作一个序列共 B*H 个长度为 W 的序列 xh x.permute(0, 2, 1, 3).reshape(B * H, W, C) xh self.attn(xh) xh xh.reshape(B, H, C, W).permute(0, 2, 1, 3).contiguous() # 垂直轴把每个“列”当作一个序列共 B*W 个长度为 H 的序列 xv x.permute(0, 3, 1, 2).reshape(B * W, H, C) xv self.attn(xv) xv xv.reshape(B, W, C, H).permute(0, 2, 3, 1).contiguous() # 水平垂直信息拼接融合最后接一个1x1卷积映射回原通道 out torch.cat([xh, xv], dim1) out self.fuse(out) out self.out(out) return x out # 残差连接几个实现细节我特意处理的水平分支的reshape先经过permute(0,2,1,3)把维度从(B,C,H,W)换成(B,H,C,W)这样一行就是一个连续的序列内存访问更友好垂直分支同理用permute(0,3,1,2)。最后都补了.contiguous()这行很重要后面的ONNX导出避坑部分会专门说原因。2.2 张量形状流动与逐行拆解很多新手拿到代码最头疼的就是不知道每个张量变成什么样了。我以输入特征图为(B2, C64, H8, W8)、num_heads4为例把核心过程捋一遍操作水平分支shape垂直分支shape原始特征图(2, 64, 8, 8)(2, 64, 8, 8)permute调整轴(2, 8, 64, 8)(2, 8, 64, 8)reshape成序列(16, 8, 64)(16, 8, 64)qkv线性映射后reshapepermute(3, 16, 4, 8, 16)(3, 16, 4, 8, 16)qk^T(16, 4, 8, 8)(16, 4, 8, 8)softmax后v(16, 4, 8, 16)(16, 4, 8, 16)还原特征图(2, 64, 8, 8)(2, 64, 8, 8)这里最容易被绕晕的是“行”和“列”的方向。水平分支里我们把B和H合并成了B*H也就是说8x8的特征图有8行每行是一个长度为8的序列自注意力在这个长度为8的序列内部展开。垂直分支把B和W合并成B*W有8列每列是一个长度为8的序列。在方形特征图上两者规模一样但如果特征图是矩形比如H40, W20两个分支的序列长度就不同了这时共享AxialAttention依然成立因为线性层只跟通道维度C打交道跟序列长度无关。2.3 参数量和计算复杂度核算先看参数量。AxialAttention里有一个qkv线性层和一个proj线性层都是C x C量级合计约4*C²。CAA里还有两个1x1卷积一个把2C合并回C一个C到C约3*C²。加起来单层CAA约7*C²参数。以YOLOv8n的SPPF输出通道C256为例新增参数约7*256*256458752也就是0.46M左右对总参数量影响很小。再看计算复杂度。设特征图大小是H x Wtoken数NH*W。全局多头自注意力的复杂度是O(N²*C)而CAA的水平分支等价于B*H个长度为W的自注意力复杂度O(B*H*W²*C)O(B*W*N*C)垂直分支同理是O(B*H*N*C)。加在一起是O(B*N*C*(HW))。当H≈W≈√N时这约等于O(B*N*C*2√N)比全局自注意力的O(N²*C)低大约√N/2倍。在80x80特征图上√N80这意味着同样做自注意力CAA的计算量大约是全局注意力的1/40这个差距是能直接决定训练和推理速度的。3. 把CAA接入YOLOv8的三种姿势从快速验证到深度重构3.1 姿势一SPPF之后插一层CAA适合快速验证如果你想先确认CAA在YOLOv8上到底有没有效果最省事的方法是在backbone末尾、进入PANet之前加一层独立CAA。这里的特征图分辨率通常是P5级别输入640时是20x20H和W都比较小CAA的计算负担很低。改YAML即可不用动源码# YOLOv8n-CAA.yaml节选 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, Conv, [512, 3, 2]] - [-1, 6, C2f, [512, True]] - [-1, 1, Conv, [1024, 3, 2]] - [-1, 3, C2f, [1024, True]] - [-1, 1, SPPF, [1024]] - [-1, 1, CAA, [4]] # 新增num_heads4这里CAA, [4]在Ultralytics的parse_model中走的是else分支args [c1, *args]自动把上一层的输出通道作为c14作为num_heads并保持输出通道等于输入通道。你什么都不用改直接把这份yaml丢进训练脚本就能跑。不过要提醒一句Ultralytics的通道解析会把yaml里的第一个参数当作输出通道还会乘上宽度缩放系数。好在CAA输出通道等于输入通道走else分支时c2ch[f]正好是上一层的输出通道所以能对上。如果你自定义一个会改变通道数的模块这条路就行不通了得走3.4节的手动注册。3.2 姿势二把C2f内部Bottleneck换成CAA适合追求精度只加一层CAA属于“外挂”如果你想更彻底地把注意力融进特征提取主链路可以直接改造C2f。C2f内部的Bottleneck本来就是做特征细化的把Bottleneck替换成CAA让每一层特征提取都带上跨轴注意力理论上效果更强。下面是一个C2f_CAA的实现结构完全对齐官方C2f直接复制就能用from ultralytics.nn.modules import Conv class C2f_CAA(nn.Module): C2f with CAA attention instead of Bottleneck def __init__(self, c1, c2, n1, shortcutFalse, g1, e0.5): super().__init__() self.c int(c2 * e) self.cv1 Conv(c1, 2 * self.c, 1, 1) self.cv2 Conv((2 n) * self.c, c2, 1) self.m nn.ModuleList(CAA(self.c, num_heads4) for _ in range(n)) def forward(self, x): y list(self.cv1(x).chunk(2, 1)) y.extend(m(y[-1]) for m in self.m) return self.cv2(torch.cat(y, 1))把官方yaml里所有的C2f替换成C2f_CAACAA层可以保留也可以去掉按你的显存和精度预期取舍。我实测时通常只在第4、5阶段的C2f上替换即通道64和128那两层不碰P3前的浅层因为浅层分辨率高替换后训练速度会掉一截收益反而不明显。3.3 姿势三在PANet特征融合层加CAA适合小目标第三种接法是把CAA用在PANet的跨尺度特征融合之后。YOLOv8的neck把P3、P4、P5三层特征反复上采样、下采样再concatconcat之后直接进C2f特征之间的“融合质量”取决于C2f的加工能力。在concat后插一层CAA等于让不同来源的特征在融合之初就进行一次跨轴的相互“对齐”对需要同时参考浅层细节和深层语义的小目标检测特别有帮助。修改head部分示例head: - [-1, 1, nn.Upsample, [None, 2, nearest]] - [[-1, 6], 1, Concat, [1]] - [-1, 1, CAA, [4]] # P4特征concat后加一层CAA - [-1, 1, C2f, [512]]这个位置选在P4的concat之后通道一般是512或256CAA带来的参数量可控。如果你想更激进在P3、P4、P5三个concat层后都加显存和训练时间会明显增加但精度并不一定线性提升。3.4 源码注册三步走含parse_model的坑如果你用的是姿势二C2f_CAA这个模块自己会改变输出通道官方parse_model的else分支不会正确处理必须手动注册。顺序如下第一步把CAA和C2f_CAA的代码放进ultralytics/nn/modules/block.py然后在ultralytics/nn/modules/__init__.py中导入from .block import CAA, C2f_CAA第二步在ultralytics/nn/tasks.py里找到parse_model函数在C2f分支的后面加上C2f_CAA的处理elif m is C2f_CAA: c1, c2 ch[f], args[0] if c2 ! nc: c2 make_divisible(min(c2, max_channels) * width, 8) args [c1, c2, *args[1:]]必须跟C2f走同一个分支否则模块输出通道会被强制写成输入通道后面所有层的通道都对不上训练直接报错。第三步在yaml里使用新名字即可- [-1, 3, C2f_CAA, [128, True]]注册完建议先跑一个model YOLO(yaml路径); model.info()看通道数是否正常不要直接开训。这一步能帮你把八成的问题挡在启动前。4. 实测结果与训练调参记录精度涨了点但要注意收敛4.1 我的实验设置验证实验我用的是COCO val2017里随机抽的5000张子集训练集用对应的train2017子集图像按原始比例缩放填充到640。机器是单张A100batch_size32优化器SGD初始学习率0.01cosine衰减warmup_epochs从默认3调到了5。为了公平对比baseline和CAA版本使用完全相同的超参数只改yaml里的结构。训练轮数设了200轮。这个值比官方默认的300轮少但足够看出相对趋势也省时间。每个实验跑三遍取均值避免随机种子带来的波动。4.2 三组对比结果模型参数量(M)GFLOPsmAP50(%)mAP50-95(%)YOLOv8n baseline3.168.777.552.8YOLOv8n CAASPPF后3.629.679.154.3YOLOv8n C2f_CAA4.3810.479.454.8只看数据的话单层CAA在SPPF后就能带来大约1.6个点的mAP50提升mAP50-95涨了1.5个点。把C2f里全换成CAA后额外收益没有想象中大也就0.3个点左右但参数量多了0.76M训练时间也明显变长。这说明CAA放在backbone出口做一次全局信息重整的性价比比放在每一层反复计算高得多。推理延迟我也简单测了一下。在A100上YOLOv8n baseline大概2.1ms/张加单层CAA后大约2.4ms/张C2f_CAA版本约2.7ms/张。纯GPU时间上CAA增加的那点计算量在高端卡上不敏感但换到1660Ti这类卡上差距会放大到3-4ms部署前最好在目标设备上实测。4.3 训练曲线的两个观察第一个观察是收敛变慢了。加了CAA后前50轮的验证mAP一直略低于baseline到100轮左右才追平150轮后才稳定反超。原因是CAA在训练初期引入了大量额外的空间交互模型需要更长时间学出有效的注意力参数。如果你发现前100轮曲线一直贴着baseline下面走别慌这不是模块失效让它跑到200轮再看。第二个观察是验证损失后期有轻微震荡。CAA的注意力概率分布比较敏感到了训练后半段、学习率已经降得很低时偶尔会看到验证损失出现小尖峰。解决办法是把rectTrue关掉保持训练和验证的图像分辨率一致如果还震荡把weight_decay从默认0.0005降到0.0003能压住一部分。4.4 调参建议三处改动让CAA稳定收敛基于我这几轮实验给你三个直接可用的调参建议warmup_epochs从3调到5。注意力机制特别怕训练初期的大扰动warmup时间拉长一点让qkv层先在一个相对平滑的学习率区间里稳定下来。如果是在自己的数据集上微调不是从头训COCO建议freeze10前10个epoch冻结backbone只训neck和head里的CAA。这样能避免backbone被新模块的大梯度带偏。如果显存紧张优先砍P3层的CAA保留P4和P5层。P3特征图80x80序列长度80注意力矩阵80x80虽然单层不大但在三个尺度上都加显存和耗时都会线性上涨。5. 踩坑复盘shape、显存、FP16与ONNX导出的那些坑5.1 H和W不相等时的reshape方向错乱我最初在imgsz640、rectTrue的矩形推理下试跑CAA版本训练能跑但验证mAP突然变成原来的一半。查了半天发现是矩形推理让特征图变成非正方形比如H80, W48水平分支和垂直分支的序列长度不一样而Ultralytics的parse_model里对特征图尺寸的假设在某些环节会把W当成H导致CAA的permute方向被解析错。这类问题的共同表现是loss正常下降但验证mAP异常排查时可以打印CAA输入输出的shape跟预期比对def forward(self, x): print(CAA input:, x.shape) # 确认H和W是否符合预期 B, C, H, W x.shape ...如果你一定要用矩形推理先把rectFalse关掉或者固定imgsz让H和W保持一致。精度损失不明显的场景建议直接固定方形输入省心。5.2 显存爆炸的罪魁祸首往往是batch*H太大CAA的水平分支把B和H合并成B*H每个“行序列”的注意力矩阵是W*W。如果B32、H80、W80水平分支的注意力矩阵总数就是32*802560个80x80的矩阵约16M个float再加上垂直分支总共超过32M。一次前向就要占100多MB显存反向传播再翻倍叠加多尺度训练显存很容易吃紧。两个缓解办法一是把CAA放在num_heads4而不是8头数减少后head_dim变大注意力矩阵的数量不变但单矩阵更小二是避免在整个模型每个阶段都插CAA只在P4/P5层插。我实测在YOLOv8n上加单层SPPF后的CAA显存只涨了约0.4GB但如果在P3也加显存涨幅直接跳到1.2GB。5.3 FP16下的softmax数值不稳定用AMP混合精度训练时CAA里的attn k.transpose(-2,-1)在FP16下很容易溢出尤其是特征图通道小、注意力分数偏大的时候softmax之后会出现NaN训练loss瞬间变成nan。这不是模块写错了而是数值精度问题。最简单的修复是让softmax在FP32下计算再cast回来attn (q k.transpose(-2, -1)) * self.scale attn attn.float().softmax(dim-1).type_as(q) x attn v额外增加的开销可以忽略。如果你用的是PyTorch 2.0以上也可以直接换成F.scaled_dot_product_attention(q, k, v)内部做了数值稳定处理代码更省事。5.4 ONNX导出时的transposereshape算子兼容问题到我准备导出ONNX部署时又踩了一个坑CAA里的permute加reshape组合在ONNX导出后某些推理引擎尤其是NCNN和部分TensorRT版本会解析出错误的shape。根源是permute产生非连续内存reshape在导出时会被翻译成低维Reshape算子和Transpose算子的组合在某些优化pass里被错误折叠。解决办法是双管齐下。第一在CAA forward里每个permute后面都加上.contiguous()强制内存连续第二导出时指定opset13以上低opset对动态shape的支持太差。如果还报错就把CAA的reshape改成view前提是你确定前面已经contiguous()过。改完后先用onnxruntime跑一遍验证输出再上TensorRT。5.5 你的注意力真的生效了吗一个快速验证办法这是我最想分享的一个经验。很多时候模块代码没问题、训练也正常但你并不知道注意力到底学到了什么。我习惯在训练过程中定期打印CAA内部attn矩阵的熵。修改一下AxialAttention把注意力矩阵缓存下来class AxialAttention(nn.Module): def forward(self, x): ... attn attn.float().softmax(dim-1).type_as(q) self.last_attn attn.detach() # 缓存便于分析 x attn v ...然后写一个简单的统计脚本在每个epoch结束时取几个batch计算attn分布的信息熵entropy -(attn * (attn 1e-9).log()).sum(dim-1).mean().item()训练刚开始时注意力矩阵接近均匀分布熵值接近log(W)训练充分后熵值会明显下降说明模型开始聚焦特定位置。如果你发现训练到200轮熵还很高大概率是学习率太大或者CAA被放在了无效位置这时候去调结构比继续加训练轮数更有效。我在这个项目里最后的做法是在SPPF后加一层CAAP4的concat后加一层CAAC2f保持原样不动。这个组合在精度和速度之间最均衡也是我后来在新数据集上做迁移时默认采用的配置。如果你也想快速复现建议先只在SPPF后加一层CAA跑一个50轮的短实验看看注意力熵有没有下降、验证精度有没有跟上再决定要不要往深处加。
返回列表