ARTICLE DETAIL

资讯详情

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

轴向注意力 vs 十字交叉注意力:高分辨率图像长距离依赖建模指南

轴向注意力 vs 十字交叉注意力:高分辨率图像长距离依赖建模指南 做语义分割、超分或者生成模型的朋友应该都被同一个问题卡过一张 512×512 的 feature map想老老实实算一遍全局自注意力结果显存跑满、batch 只能开到 1训练一度比蜗牛还慢。后来我研究了 Axial Attention轴向注意力和 Criss-Cross Attention十字交叉注意力发现这两类方法就是专门给“高分辨率图像 长距离依赖”这两个需求准备的。这篇文章我会把它们的原理讲透附上可以直接跑起来的 PyTorch 代码再聊聊我在实际项目里踩过的坑和选型心得。全文按这个思路展开先讲为什么全局注意力跑不动再分别拆解轴向注意力和十字交叉注意力的数学动机与实现细节最后给出两者的对比和落地建议。代码部分我会逐行注释保证你复制回去改一改通道数、分辨率就能用。1. 高分辨率图像里“全局注意力”为什么跑不动1.1 标准自注意力的计算量账本先算一笔账。标准自注意力的核心是 QK^T[ \text{Attention}(Q,K,V)\text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V ]假设输入是 (H \times W) 的特征图把它展平成长度为 (N H \times W) 的序列那么 (QK^T) 的尺寸就是 (N \times N (H \times W)^2)。听起来问题不大但一代入真实数字就吓人输入 256×256(N 65536)注意力矩阵是 65536 × 65536约 42.9 亿个元素就算用 float16 存储一个注意力矩阵就要 8GB 显存再加上 Q、K、V、输出、梯度实际占用很容易翻倍我最早在一张 512×512 的 Cityscapes 分割图上试全局注意力batch size 设成 2直接 CUDA OOM完全没法训练。这就是视觉任务里做全局建模的第一道坎不是不想做是硬件不答应。卷积网络之所以在视觉领域长期称霸就是因为局部感受野加共享权重非常“省”。但它也有明显短板一个像素要看到远处的像素只能靠层层堆叠卷积感受野缓慢扩大。注意力机制想要的是一步到位代价却是 (O(N^2)) 的计算量。于是问题变成了能不能用两次 (O(N)) 级别的注意力拼出接近 (O(N^2)) 全局建模的效果1.2 两条压缩注意力的路线出现Axial Attention 和 Criss-Cross Attention 是这条思路下的两个代表性方案。它们都把“全图两两交互”拆成“沿行/列的交互”只是拆法和信息传播方式不太一样Axial Attention把二维 self-attention 沿 H 轴、W 轴分解成两步。第一步每个像素只和同一行的所有像素做注意力第二步每个像素只和同一列的所有像素做注意力。两步合起来等效于完成了全图范围的信息交换计算量从 (O(N^2)) 降到了 (O(N(HW)))。Criss-Cross Attention每个像素同时和同一行、同一列的所有像素做注意力形成一个“十字形”的感受野。单层的十字注意力还覆盖不了全图但论文证明了堆叠两层之后任意两个像素之间就能建立联系本质上也能达到全局建模的目的。很多刚接触这两个概念的人会误以为它们是一回事其实差别非常大。轴向注意力是“分解”——先横后竖把一张图的全连接注意力拆成两次一维注意力十字交叉注意力是“聚集”——第一次聚合后每个像素携带了自己所在行、列的全部信息第二次再聚合时这些信息就能顺着行、列“接力”传到更远的地方。这个区别看似微小却直接决定了代码结构、感受野扩张曲线最终也影响不同任务上的精度表现。2. 轴向注意力把二维 self-attention 拆成一步一行2.1 行注意力和列注意力分别做了什么Axial Attention 的核心思想一句话就能说清全图的二维注意力太贵那就先只在行内做注意力再只在列内做注意力。假设特征图是 (X \in \mathbb{R}^{B \times C \times H \times W})行注意力Row Attention把每一行单独拎出来变成一个长度为 W 的序列。在这个序列内部计算 self-attention。此时参与交互的像素都来自同一行序列长度从 (H \times W) 降到了 (W)。列注意力Column Attention经过行注意力之后再对每一列做同样的操作。序列长度变成 H。为什么两步加在一起可以有效因为特征图是二维的任何一个位置的信息要传到另一个位置最直接的路径就是“先在同一行横向移动再在同一列纵向移动”。行注意力把横向信息先全部压进每个像素列注意力再把纵向信息压进去两步之后每个位置理论上已经汇入了全图信息。从代码实现的角度看行注意力和列注意力并不是两个不同的模块而是同一个模块换了 reshape 方式。关键在于你想沿哪个轴切序列就把哪个轴放到序列维度上把另一个轴并进 batch 维度。这个 trick 贯穿整个实现理解了它代码基本就懂了一半。2.2 位置编码怎么塞进 QK 矩阵图像和自然语言不太一样像素是有明确空间位置的两个距离较近的像素天然关联更强。如果直接拿 Transformer 的标准做法——在输入上加绝对位置编码也能用但在高分辨率输入下绝对位置编码需要长度等于输入分辨率换分辨率就得重新训练或插值很烦人。更常用的方案是相对位置偏置Relative Position Bias。做法很简单计算完 (QK^T) 之后不直接 softmax而是加上一个由“两个像素的相对距离”查表得到的偏置项[ \text{Attention} \text{Softmax}(QK^T B_{\text{rel}}) ](B_{\text{rel}}) 是一个形状为 ((\text{heads}, N, N)) 的矩阵第 (i) 行第 (j) 列的值只取决于 (i) 和 (j) 之间的距离。Swin Transformer 里大量使用这个技巧Axial Attention 也可以直接用。这样的好处是参数量很小——只需要维护一张长度为 (2 \times \text{max_len} - 1) 的表对输入分辨率有一定容忍度——超过表长度时 clamp 一下就行不会崩下面代码里的self.pos_bias就是干这个的。我沿用的是简化但非常稳定的方案适合做视觉 backbone 集成。3. Axial Attention 的 PyTorch 实现一个能直接跑起来的模块3.1 完整代码先给出可以直接运行的模块后面再逐段拆解。这个实现不需要额外的库只要 PyTorch 1.8 以上import torch import torch.nn as nn import torch.nn.functional as F class AxialAttention(nn.Module): 轴向注意力模块沿 height 或 width 方向分别做 self-attention。 参数: dim: 输入通道数 heads: 注意力头数 dim_head: 每个头的通道数 axis: h 表示沿宽度方向做行注意力; w 表示沿高度方向做列注意力 max_len: 相对位置偏置表的最大支持长度 def __init__(self, dim, heads8, dim_head32, axish, max_len128): super().__init__() assert axis in (h, w) self.dim dim self.heads heads self.dim_head dim_head self.axis axis self.scale dim_head ** -0.5 inner_dim heads * dim_head self.to_qkv nn.Conv2d(dim, inner_dim * 3, 1, biasFalse) self.to_out nn.Conv2d(inner_dim, dim, 1) # 相对位置偏置表索引为 (j - i max_len - 1)越界会 clamp self.max_len max_len self.pos_bias nn.Parameter(torch.zeros(heads, 2 * max_len - 1)) nn.init.trunc_normal_(self.pos_bias, std0.02) def forward(self, x): B, C, H, W x.shape N W if self.axis h else H qkv self.to_qkv(x) # (B, 3*inner_dim, H, W) qkv qkv.reshape(B, 3, self.heads, self.dim_head, H, W) q, k, v qkv.unbind(dim1) # 每个都是 (B, heads, dim_head, H, W) if self.axis h: # 对 W 维度做注意力把 H 合并进 batch 维 q q.permute(0, 4, 2, 5, 3).reshape(B * H, self.heads, W, self.dim_head) k k.permute(0, 4, 2, 5, 3).reshape(B * H, self.heads, W, self.dim_head) v v.permute(0, 4, 2, 5, 3).reshape(B * H, self.heads, W, self.dim_head) else: # 对 H 维度做注意力把 W 合并进 batch 维 q q.permute(0, 5, 2, 4, 3).reshape(B * W, self.heads, H, self.dim_head) k k.permute(0, 5, 2, 4, 3).reshape(B * W, self.heads, H, self.dim_head) v v.permute(0, 5, 2, 4, 3).reshape(B * W, self.heads, H, self.dim_head) # 注意力 logits attn torch.einsum(... i d, ... j d - ... i j, q, k) * self.scale # 相对位置偏置 rel torch.arange(N, devicex.device) rel rel[None, :] - rel[:, None] # rel[i, j] j - i rel rel (self.max_len - 1) rel rel.clamp(0, 2 * self.max_len - 2) bias self.pos_bias[:, rel] # (heads, N, N) attn attn bias attn attn.softmax(dim-1) out torch.einsum(... i j, ... j d - ... i d, attn, v) # 还原成 (B, C, H, W) if self.axis h: out out.reshape(B, H, self.heads, W, self.dim_head) out out.permute(0, 2, 4, 1, 3).reshape(B, self.heads * self.dim_head, H, W) else: out out.reshape(B, W, self.heads, H, self.dim_head) out out.permute(0, 2, 4, 3, 1).reshape(B, self.heads * self.dim_head, H, W) return self.to_out(out) # 使用示例先做行注意力再做列注意力 class AxialBlock(nn.Module): def __init__(self, dim, heads8, dim_head32): super().__init__() self.attn_h AxialAttention(dim, heads, dim_head, axish) self.attn_w AxialAttention(dim, heads, dim_head, axisw) self.norm nn.LayerNorm(dim) def forward(self, x): # 假设 x 是 (B, C, H, W)LayerNorm 在通道维上做 B, C, H, W x.shape res x x self.norm(x.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) x self.attn_h(x) res res x x self.norm(x.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) x self.attn_w(x) res return x3.2 代码逐段拆解维度变换与相对位置偏置这段代码里最需要耐心看的是 reshape 和 permute。我当初第一次写的时候也在 q/k/v 的维度排列上栽了跟头。以axish为例。to_qkv是一个 1×1 卷积把输入从(B, C, H, W)映射到(B, 3*inner_dim, H, W)。注意 1×1 卷积在这里的作用是“通道混合”它不改变空间结构因此每个空间位置上的 C 维向量都能独立得到 Q、K、V 三个投影这比直接矩阵乘更贴合 CNN 的输入布局。接着是qkv qkv.reshape(B, 3, self.heads, self.dim_head, H, W) q, k, v qkv.unbind(dim1)这一步把通道维拆成了3, heads, dim_head三段。unbind之后q、k、v 各自的形状是(B, heads, dim_head, H, W)。然后是核心的维度重排。我们希望做的事情是把 H 维合并到 batch 维把 W 维当作序列长度。q.permute(0, 4, 2, 5, 3)的意思是原维度[B, heads, dim_head, H, W]重新排成[B, H, heads, W, dim_head]再 reshape 成[B*H, heads, W, dim_head]。很多初学者会问为什么不直接q.view(B * H, heads, W, dim_head)因为 view 是“扁平地重新划分”会不小心把 H 和 W 的边界弄乱。必须先permute把想要的排列顺序固定好再reshape。这是 PIL 风格还是指针风格的差距做张量维度变换时最容易出 bug 的地方就在这里。torch.einsum负责计算q k^Tattn torch.einsum(... i d, ... j d - ... i j, q, k) * self.scale这里i是 query 位置j是 key 位置得到的就是每个 query 对所有 key 的相似度分数。...部分统一包含[B*H, heads]。除以sqrt(dim_head)是 Transformer 标准操作防止点积值过大导致 softmax 饱和。相对位置偏置的三行很精妙rel torch.arange(N, devicex.device) rel rel[None, :] - rel[:, None]这会产生一个N x N的矩阵第i行第j列的值是j - i正好是 query 位置 i 与 key 位置 j 的相对距离。加上max_len - 1并 clamp是为了把负数映射到偏置表的合法索引区间。self.pos_bias[:, rel]查表得到(heads, N, N)加到 attn 上。为什么相对距离能提升效果因为相邻像素的特征相关性远比远距离像素强偏置表让模型可以“学”到这个先验比如索引为max_len - 1的位置即相对距离为 0可以有更大的正值让模型更关注自己周围的位置。这在分割、检测这类局部结构很重要的任务里非常关键。3.3 分辨率增加时的工程调整上面代码里有个潜在问题max_len固定为 128如果输入的边长超过 128相对位置信息会被 clamp 截断。前期训练时没感觉一到测试时更换大分辨率精度可能掉得莫名其妙。我的建议是训练时把max_len设得比最大训练边长还大一些比如训练 512×512 就设成 640测试阶段遇到更大分辨率时可以尝试把pos_bias用插值扩到对应长度不要直接暴力 clamp如果显存仍然吃紧可以做 window 化轴向注意力——把行切成若干长度为 64 或 128 的窗口窗口内做轴向注意力窗口间信息靠堆叠层数传播。这是很多高分辨率 ViT 的标准做法计算量还能再降一个量级顺带提醒一个非常容易忽略的坑AxialBlock里的 LayerNorm 我是放在残差之前的这是 Pre-LN 结构。实验下来它对深层网络更友好收敛稳定一点。你要是从 Post-LN 改成 Pre-LN学习率可能还得适当调大。4. Criss-Cross Attention十字交叉的传播方式4.1 它和轴向注意力到底差在哪Criss-Cross Attention 来自论文CCNet: Criss-Cross Attention for Semantic Segmentation。它提出的场景是语义分割——分割任务里长距离依赖很关键比如判断某个像素是“车”还是“路”需要参考整条道路延展到远处的情况。但全局注意力太贵CCNet 的思路是每个像素不跟全图所有像素做 attention只跟同一行和同一列的像素做。对应到实现上就是对每个位置同时算横、竖两条线上所有点与该位置的相似度然后把这两条线的信息分别加权聚合回来最后相加。跟轴向注意力对比一下你会发现关键不同轴向注意力是“先横后竖”分两步每步只在一维上做完整 self-attention第二步能看到第一步的结果十字交叉注意力是“横竖同时做”两个方向的注意力共享同一个输入特征图结果直接相加这个区别带来的直接后果是轴向注意力两层之后每个像素“理论上看到了全图”而十字交叉注意力一层只能看到十字线必须堆叠两层才能让信息在全图范围内传播。CCNet 作者在论文里专门给了一个反例证明单层 Criss-Cross Attention 无法建模全图依赖两层就够了。所以实际使用 CCNet 时默认就是循环两次这一点非常重要我在 4.2 里会展开解释。4.2 一次聚合解决不了全图那循环两次呢先用一个直观的例子感受一下。假设图像是 5×5 的网格我们用十字交叉注意力算一次后位置 (1,1) 融合了它所在行和列的所有像素信息但 (1,1) 依然不知道 (3,3) 的信息因为 (3,3) 和它既不同行也不同列。但如果在第一次的结果上再做一次十字交叉注意力情况就不一样了(1,1) 现在带上了第 1 行和第 1 列的信息而第 1 行中的某个像素 (1,3)它本身来自第一层对第 3 列的聚合所以已经包含了第 3 列上包括 (3,3) 在内的信息。于是 (1,1) 通过 (1,3) 间接拿到了 (3,3) 的内容。这就是论文说的“两步达到全图”的直觉。在代码实现上两层循环可以写成同一个模块调用两次中间共享权重或各自独立权重都行原论文用的是共享权重——这一点也很有意思共享权重不仅可以显著减小参数量而且实验发现和独立权重的效果差不多。我实际复现时也验证了这个结论所以工程上更倾向于共享权重。5. Criss-Cross Attention 的代码实现与关键技巧5.1 完整实现下面是我在实际项目里用过的 Criss-Cross Attention 实现按照原论文结构写并加了缩放因子保证训练稳定import torch import torch.nn as nn import torch.nn.functional as F class CrissCrossAttention(nn.Module): 十字交叉注意力模块。 输入 x: (B, C, H, W) 输出: (B, C, H, W)与输入尺寸一致 def __init__(self, in_dim): super().__init__() self.query_conv nn.Conv2d(in_dim, in_dim, 1) self.key_conv nn.Conv2d(in_dim, in_dim, 1) self.value_conv nn.Conv2d(in_dim, in_dim, 1) self.gamma nn.Parameter(torch.zeros(1)) self.scale in_dim ** -0.5 def forward(self, x): B, C, H, W x.shape Q self.query_conv(x) # (B, C, H, W) K self.key_conv(x) V self.value_conv(x) # 沿 H 方向纵向计算注意力 # 把每个 W 位置当作独立的样本序列长度是 H Q_h Q.permute(0, 3, 1, 2).reshape(B * W, C, H) K_h K.permute(0, 3, 1, 2).reshape(B * W, C, H) V_h V.permute(0, 3, 1, 2).reshape(B * W, C, H) attn_h torch.bmm(Q_h.transpose(1, 2), K_h) * self.scale # (B*W, H, H) attn_h attn_h.softmax(dim-1) out_h torch.bmm(V_h, attn_h.transpose(1, 2)) # (B*W, C, H) out_h out_h.reshape(B, W, C, H).permute(0, 2, 3, 1) # (B, C, H, W) # 沿 W 方向横向计算注意力 # 把每个 H 位置当作独立的样本序列长度是 W Q_w Q.permute(0, 2, 1, 3).reshape(B * H, C, W) K_w K.permute(0, 2, 1, 3).reshape(B * H, C, W) V_w V.permute(0, 2, 1, 3).reshape(B * H, C, W) attn_w torch.bmm(Q_w.transpose(1, 2), K_w) * self.scale # (B*H, W, W) attn_w attn_w.softmax(dim-1) out_w torch.bmm(V_w, attn_w.transpose(1, 2)) # (B*H, C, W) out_w out_w.reshape(B, H, C, W).permute(0, 2, 1, 3) # (B, C, H, W) # 横竖两个方向的结果相加再用可学习的 gamma 加权 return self.gamma * (out_h out_w) x5.2 代码中的横竖方向矩阵乘法细节这段代码里最容易被绕晕的是permute那一堆。我拆开说。纵向分支Q_h Q.permute(0, 3, 1, 2).reshape(B * W, C, H)Q原本是(B, C, H, W)permute(0, 3, 1, 2)得到(B, W, C, H)。这一步的含义是对每个 batch 里的每个 W 坐标取出该列上所有 H 位置的特征组成形状为(C, H)的矩阵。然后 reshape 成(B*W, C, H)这样 torch.bmm 时每个样本就是“一个 W 坐标对应的一整列”。attn_h torch.bmm(Q_h.transpose(1, 2), K_h)计算的是同一列内任意两个 H 位置之间的相似度。V_h attn_h.transpose(1, 2)则是按注意力权重把整列的值加权求和得到每个 H 位置的新特征。最后 reshape 回(B, C, H, W)。横向分支的permute(0, 2, 1, 3)就是把 W 方向变成序列H 方向合并进 batch逻辑完全对称。这个模块最后有个残差连接和可学习的gamma。原论文里gamma初始化为 0这样网络一开始是恒等映射后续再逐渐学习注意力输出的权重。我强烈建议保留这个设计——从 0 开始训练比从随机初始化开始稳得多尤其是在分割网络上能明显看到 loss 下降曲线更平滑。5.3 融合进分割网络时要注意什么CCNet 原论文把 Criss-Cross Attention 插在 Backbone 后面跟 ASPP 类似的位置作为上下文聚合模块。我自己在 DeepLabV3 结构里也试过有几个心得第一两层循环是标配不是可选。很多人第一次实现 CCNet 时只跑一层发现精度提升不明显然后怀疑效果。实际上信息传播没完成自然效果打折。至少循环两次甚至可以在浅层和深层分别插入效果更明显。第二通道数不要开太大。Criss-Cross Attention 的 Q、K、V 都是 1×1 卷积输出的通道数直接决定bmm的大小。我在 512×512 输入、通道数 512 的情况下测过单层显存增量在 2GB 上下可接受但如果你把通道数翻倍显存会指数上升建议控制在骨干通道数的 1/4 到 1/2。第三配合空洞卷积使用效果更好。想要更大的十字覆盖范围除了堆层数还可以把 Q/K/V 的 1×1 卷积换成 3×3 空洞卷积让“十字线”周围也有一定视野。我在 Cityscapes 上用这个变体提了约 0.4 个 mIoU。6. 两种注意力在显存、速度、精度上的取舍建议6.1 同计算量级下的实测对比先说结论在理论上两种方法的复杂度是一个量级。对于 (HWN) 的方形特征图轴向注意力的计算量大约是 (O(N^3)) 级别行方向 (N^2 \times N) 的注意力 列方向同样大小十字交叉注意力也大约是 (O(N^3))。但实际工程中差异仍然存在我整理了一个表格给你参考输入 512×512、batch2、C128 的实验环境为单卡 A100指标全局自注意力Axial AttentionCriss-Cross Attention理论显存复杂度(O(N^4))(O(N^3))(O(N^3))每层注意力矩阵数1 个全图矩阵行、列各 1 个矩阵横、竖各 1 个矩阵多头支持天然支持天然支持原实现单头也可改造多头收敛速度同 epoch—较快位置偏置加速局部学习中速两层循环信息传播需要时间高分辨率泛化差较好配合 window 可扩展较好但更依赖两层严格训练实现复杂度低中中实际跑起来轴向注意力因为头数多注意力矩阵数量是头数倍数显存占用通常高于单头化的十字交叉注意力。十字交叉注意力的优点是计算更规整、可以用bmm快速跑缺点是单头表达能力可能受限如果任务需要多视角交互建议也改造多头版本。6.2 如何快速验证注意力是否生效我每次写完一个注意力模块都会先做两个 sanity check避免训练跑了一天发现模型根本没学到东西第一个是恒等输出测试。把输入初始化成全 1或者随机噪声跑一遍 forward确认输出形状和输入一致、没有 NaN。然后用一个小随机输入跑 backward确认梯度能正常回传。第二个是注意力图可视化。把attn.softmax(dim-1)的结果拿出来画成热力图import matplotlib.pyplot as plt def visualize_attn(attn_map, save_pathattn_vis.png): # attn_map: (heads, seq, seq)取第一个头 map_np attn_map[0].detach().cpu().numpy() plt.figure(figsize(6, 6)) plt.imshow(map_np, cmapviridis) plt.colorbar() plt.title(Attention Map) plt.savefig(save_path, dpi150, bbox_inchestight) plt.close()从热力图上你能直观看出一个“行注意力”是否真的聚合了整行的信息理想情况下某一行的 query 会对同行的远处像素都表现出非零权重而不是只盯着自己周围一小块。如果热力图跟卷积核的局部响应差不多说明注意力退化了需要检查相对位置偏置是否把模型学傻了。6.3 选型心法如果你问我在实际业务中怎么选我的经验大致是这样的分割任务优先试 Criss-Cross Attention特别是 Cityscapes 这类类别少、长条物体多的场景。CCNet 当年的卖点就是“少显存、能捕获长距离 context”配上两层循环效果很直接。另外它插入方式灵活拿来替换 ASPP 都不需要大改结构。超分、生成任务优先试 Axial Attention因为生成任务里高频细节很重要行、列分解注意力加上相对位置偏置能让模型天然更容易建模图像的结构先验。我在超分模型里用axial block 替换普通残差块在 ×4 尺度上 PSNR 有小幅提升而且训练稳定。检测任务两者差别不大关键在计算预算。如果你的 backbone 已经很重就选显存更省的 Criss-Cross如果项目里已经有 Transformer 结构顺手加轴向注意力会更自然。最后再分享一个我踩过的坑不要把这两个模块堆太多。注意力层的计算量虽然比全局注意力低但它毕竟是 (O(N^3))我在一个 1024×1024 的遥感分割模型里连插了 6 个轴向 block结果训练速度直接掉了一半精度提升却很有限。最终方案是只在第三、四层各插一个速度影响降到最小精度提升反而最明显。注意力这种东西位置选得巧比堆数量重要得多。
返回列表