1. 为什么普通Attention在视觉任务里常常“跑不动”
聊这两个新机制之前,得先搞清楚它们到底在解决什么问题。很多读者可能已经看过Transformer在NLP里的表现,Word Embedding序列一进来,Self-Attention层层叠加,效果确实猛。但你要是直接把这个思路搬到视觉任务上,马上就会撞上一堵墙:算力根本扛不住。
1.1 全注意力机制的计算瓶颈
先说Self-Attention的复杂度。给定一个输入特征图,假设尺寸是 H x W,通道数是 C。标准的Scaled Dot-Product Attention要先把特征图展平成 N = H * W 个位置,然后做 QK^T 相似度计算。这一下子的复杂度就是 O(N²),内存复杂度同样是 O(N²)。
假设输入是一张 224x224 的图片,经过骨干网络下采样到 7x7,此时 N = 49,N² = 2401,完全没问题。但视觉任务里经常要处理 512x512、1024x1024 甚至更高分辨率的输入。就算下采样到 128x128,N = 16384,N² = 268435456,也就是约2.7亿次相似度计算。这还不算GPU显存里要存N²大小的注意力图——2.7亿个float32就是超过1GB的显存,就为了存储一张注意力图。这完全是不可接受的。
我做语义分割的时候第一次尝试在特征图上直接套标准Attention,输入 512x512,骨干网络下采样8倍,特征图 64x64,N = 4096,N² = 1600万。显存瞬间爆了,直接OOM。那还是Batch Size=1的情况。
1.2 视觉任务的两个特性决定了它不能照搬NLP方案
第一个特性是分辨率高,这个不用多说,像素级别的任务天然需要高分辨率来保留细节。第二个特性是局部性,图像里的目标往往只占图像的一部分,像素之间的关联更多体现在局部范围内,而不是像句子那样全局语义关联。
NLP里一句话通常几十到几百个token,N本身就不大,全局Attention可以接受。但图像里的N动辄成千上万,全局关联计算量爆炸。而且对很多像素来说,远距离的相关性其实很稀疏,它真正需要关注的区域可能就在后面几个像素、同一行的远处、或者同一列的某个位置。这就给了我们一个很大的优化空间。
2. Axial Attention:按轴拆解,把注意力化整为零
Axial Attention的核心思路非常直白:既然全局注意力复杂度太高,那我就不一次性地关注所有位置,而是先只关注“同一行”的位置,再只关注“同一列”的位置。把一次二维的全注意力拆成两个一维的注意力,分先后执行。
这个方案最早是Google Brain在2019年提出的,用来在图像和视频上替代标准Transformer做生成任务。效果出人意料地好,甚至在某些任务上超过全注意力机制,而计算量大幅下降。
2.1 核心思想拆解:先水平后垂直
Axial Attention整个过程分两步:
第一步,Row-wise Attention。特征图的形状是 (B, C, H, W),我们对每个位置的像素,只让它关注同一行上的其他 W-1 个像素。相当于把每行当作一个长度为W的序列,对H行各自做一次Self-Attention。
第二步,Column-wise Attention。上一步的输出作为输入,这次每个像素只关注同一列上的其他 H-1 个像素。把每列当作一个长度为H的序列,对W列各自做一次Self-Attention。
这样一轮下来,每个位置间接地获得了整个特征图的信息。为什么是“间接”?因为第一轮Row Attention之后,某行上所有位置都交互过一遍,信息已经在行内充分共享了。第二轮Column Attention时,每个位置能通过列上的其他位置,拿到它们在Row阶段已经聚合过的行信息。所以两轮之后,全局信息的传播路径是通的。
这个思路用生活化的方式类比就是:一个班有H排,每排W个学生。如果每个学生直接和班里所有其他学生说话,那需要H*W次交流,太乱了。现在换个方案:第一轮先让每排的同学互相认识,交流信息;第二轮再让每列的同学互相认识。两轮下来,任何两个人之间都能通过某个中间人把消息传达到,但交流次数少得多。
2.2 计算复杂度的数学推导
常规Attention的复杂度是 O(H²W²),这个很好算,N = H*W,N² = H²W²。
Axial Attention的复杂度:
Row注意力部分,H行,每行长度为W,每行内self-attention复杂度 O(W²),一共H行,所以是 O(HW²)。
Column注意力部分,W列,每列长度为H,每列内self-attention复杂度 O(H²),一共W列,所以是 O(H²W)。
两部分相加:O(HW² + H²W) = O(HW(H+W))。
对比一下就很直观了。假设H = W = 64,标准Attention是 64⁴ = 1677万,Axial是 64³ + 64³ = 52万。差了32倍。分辨率越高,差距越明显。如果是H = W = 128,标准Attention是2.68亿,Axial是419万,差了64倍。
内存方面的节省也来自这个降阶。标准Attention要存 (B, heads, HW, HW) 的注意力矩阵,Axial只需要存 (B, heads, H, W, W) 或者类似形状的中间结果,显存压力小很多。
实现上还有一个关键细节:Row和Column两个阶段是可以并行的。Row阶段各行之间互不依赖,可以并行计算;Column阶段各列之间同样互不依赖。batch内部的并行度没有降低,反而计算量小了,训练速度明显提升。
2.3 Axial Attention的PyTorch代码实现
原理讲清楚了,直接看代码。先声明一下,这里的实现是配合注释做教学用的,方便理解核心逻辑,真实工程里需要经过更多优化。
import torch import torch.nn as nn import torch.nn.functional as F class AxialAttention(nn.Module): def __init__(self, in_channels, heads=8, dim_head=32): super().__init__() self.heads = heads self.dim_head = dim_head self.total_dim = heads * dim_head self.to_qkv = nn.Conv2d(in_channels, self.total_dim * 3, kernel_size=1, bias=False) self.output_proj = nn.Conv2d(self.total_dim, in_channels, kernel_size=1, bias=False) def forward(self, x, axis="row"): b, c, h, w = x.shape qkv = self.to_qkv(x) q, k, v = torch.chunk(qkv, 3, dim=1) q = q.view(b, self.heads, self.dim_head, h, w) k = k.view(b, self.heads, self.dim_head, h, w) v = v.view(b, self.heads, self.dim_head, h, w) if axis == "row": # 每行独立计算,Q,K,V在W维度上做注意力 # 注意矩阵形状变换:把高度和batch合并,每行独立成序列 q = q.permute(0, 1, 3, 2, 4).reshape(b*self.heads*h, self.dim_head, w) k = k.permute(0, 1, 3, 2, 4).reshape(b*self.heads*h, self.dim_head, w) v = v.permute(0, 1, 3, 2, 4).reshape(b*self.heads*h, self.dim_head, w) # 注意力分数 attn = torch.bmm(q.transpose(1, 2), k) # (b*h, w_seq, w_seq) attn = attn / (self.dim_head ** 0.5) attn = F.softmax(attn, dim=-1) out = torch.bmm(attn.transpose(1, 2), v.transpose(1, 2)) out = out.transpose(1, 2).reshape(b, self.heads, h, w, self.dim_head) out = out.permute(0, 1, 4, 2, 3).reshape(b, self.total_dim, h, w) else: # column q = q.permute(0, 1, 4, 2, 3).reshape(b*self.heads*w, self.dim_head, h) k = k.permute(0, 1, 4, 2, 3).reshape(b*self.heads*w, self.dim_head, h) v = v.permute(0, 1, 4, 2, 3).reshape(b*self.heads*w, self.dim_head, h) attn = torch.bmm(q.transpose(1, 2), k) attn = attn / (self.dim_head ** 0.5) attn = F.softmax(attn, dim=-1) out = torch.bmm(attn.transpose(1, 2), v.transpose(1, 2)) out = out.transpose(1, 2).reshape(b, self.heads, w, h, self.dim_head) out = out.permute(0, 1, 4, 3, 2).reshape(b, self.total_dim, h, w) return self.output_proj(out)这里有一个非常容易被坑的地方:行列变换的permute顺序。Row阶段要把H维度并入Batch维度,只保留W作为序列长度;Column阶段要把W维度并入Batch维度,只保留H作为序列长度。这句permute写错了简直要排查一整天,输出的形状都对,但结果就是不对。建议各位在写的时候,先在纸上把张量形状推导一遍再动手。
实际使用的时候,一般要串行调用两次。第一次用axis="row",第二次用axis="column":
class AxialBlock(nn.Module): def __init__(self, in_channels, heads=8, dim_head=32): super().__init__() self.row_attn = AxialAttention(in_channels, heads, dim_head) self.col_attn = AxialAttention(in_channels, heads, dim_head) self.norm1 = nn.LayerNorm([in_channels]) self.norm2 = nn.LayerNorm([in_channels]) def forward(self, x): b, c, h, w = x.shape identity = x x_row = self.row_attn(x, "row") # LayerNorm处理时要reshape回 (b, h*w, c) x_row = self.norm1(x_row.flatten(2).transpose(1, 2)).transpose(1, 2).reshape_as(x) x = x_row + identity identity = x x_col = self.col_attn(x, "column") x_col = self.norm2(x_col.flatten(2).transpose(1, 2)).transpose(1, 2).reshape_as(x) x = x_col + identity return x3. Criss-Cross Attention:十字交叉路径,两轮传播全图信息
Criss-Cross Attention是CCNet(Criss-Cross Network)的核心模块,2019年CVPR上的论文,面向的是语义分割这类密集预测任务。它的思路和Axial Attention有相似的地方,都是降低Attention的复杂度,但切入点和实现方式很不一样。
3.1 核心思路解析:为什么十字路径就能建模全局
Criss-Cross Attention的做法是这样的:对特征图上的每个像素,它不是去关注全图的N个位置,也不是像Axial那样分成行列两轮,而是每轮直接关注“同一行加上同一列”的所有位置,也就是一个十字形区域。
这样说可能有点绕,直接看数字。某个位置 (i, j),它这一步可以和以下位置计算注意力:所有 (i, y),即同一行;所有 (x, j),即同一列。一共是 W + H - 1 个位置(减1是去掉重复计算自己的位置)。
做完这一步之后,每个位置只获得了十字方向的信息。那问题来了,这样真的能代表全局吗?
答案是:一轮不够,但两轮就够了。论文里给出了一个非常巧妙的论证。第一轮之后,位置 (i, j) 已经聚合了它那一行和一列的信息。第二轮做Criss-Cross时,位置 (i, j) 会去关注它十字方向上的另一个位置,比如 (i, j')。而这个 (i, j') 在第一轮中已经聚合了它所在行列的信息。所以第二轮之后,位置 (i, j) 的信息来源相当于扩展到了整个二维平面。论文的图很直观,位置 (i, j) 经过两轮循环后,可以获得所有像素的信息,只是传播路径是两跳的。
这个设计非常巧妙的地方在于:它每一轮的复杂度是 O(N * (H+W)),也就是 O(HW(H+W)),和Axial Attention同级别。但它的感觉和Axial不同,Axial是严格的先水平后垂直,Criss-Cross是一步到位十字形交互。
3.2 为什么两轮迭代传播就足够
这里涉及一个信息传播的数学直觉。考虑一个无向图,图上的每个节点是一个像素,边连接的是十字方向上的像素关系。两轮十字交叉注意力之后,任意两个节点之间的最短路径长度不超过2。因为从任意一个像素出发,它一定可以和目标像素所在的行或列“交汇”。具体来说,从 (i0, j0) 到 (i1, j1),可以先在第二轮借道 (i0, j1) 这个中间节点——第一轮先让 (i0, j0) 和 (i0, j1) 交汇,第二轮让 (i0, j1) 和 (i1, j1) 交汇,信息就传到了。
当然,信息传播两跳意味着相对于直接全局Attention,路径上会有信息“稀释”。中间人作为一个中继,可能无法完美保真全部信息。这就是为什么CCNet在有些任务上性能略低于标准全局Attention,但在计算量上有数量级的优势。
3.3 Criss-Cross Attention代码实现(CCNet核心模块)
这一块贴的是和论文一致的复现版本,适合做语义分割时直接插入到骨干网络后面。
import torch import torch.nn as nn import torch.nn.functional as F class CrissCrossAttention(nn.Module): def __init__(self, in_channels, reduction=8): super().__init__() self.query_conv = nn.Conv2d(in_channels, in_channels // reduction, kernel_size=1) self.key_conv = nn.Conv2d(in_channels, in_channels // reduction, kernel_size=1) self.value_conv = nn.Conv2d(in_channels, in_channels, kernel_size=1) self.output_proj = nn.Conv2d(in_channels, in_channels, kernel_size=1) def forward(self, x): b, c, h, w = x.shape proj_query = self.query_conv(x) proj_key = self.key_conv(x) proj_value = self.value_conv(x) # 把H、W维合并成一行 proj_query_h = proj_query.permute(0, 3, 1, 2).contiguous().view(b*w, -1, h).permute(0, 2, 1) proj_key_h = proj_key.permute(0, 3, 1, 2).contiguous().view(b*w, -1, h).permute(0, 2, 1) proj_value_h = proj_value.permute(0, 3, 1, 2).contiguous().view(b*w, -1, h).permute(0, 2, 1) # 竖方向:每列内做注意力 energy_h = torch.bmm(proj_query_h, proj_key_h.transpose(1, 2)) attn_h = F.softmax(energy_h, dim=-1) out_h = torch.bmm(attn_h, proj_value_h) out_h = out_h.view(b, w, h, c).permute(0, 3, 2, 1) # 横方向:每行内做注意力 proj_query_w = proj_query.permute(0, 2, 1, 3).contiguous().view(b*h, -1, w).permute(0, 2, 1) proj_key_w = proj_key.permute(0, 2, 1, 3).contiguous().view(b*h, -1, w).permute(0, 2, 1) proj_value_w = proj_value.permute(0, 2, 1, 3).contiguous().view(b*h, -1, w).permute(0, 2, 1) energy_w = torch.bmm(proj_query_w, proj_key_w.transpose(1, 2)) attn_w = F.softmax(energy_w, dim=-1) out_w = torch.bmm(attn_w, proj_value_w) out_w = out_w.view(b, h, w, c).permute(0, 3, 1, 2) # 相加合并两个方向的信息 context = out_h + out_w return self.output_proj(context)代码里的关键点在于:横方向和竖方向其实是分别计算的,然后相加得到十字路径聚合结果。论文里的原版做法是把两个方向拼接起来再做Softmax,我这里为了计算上的对称性采用了分别Softmax再相加的简化版。
这里有一个工程上的小陷阱:contiguous()调用。PyTorch里permute之后张量内存布局不连续,view会报错,必须先用contiguous()把内存变成连续排列。新手经常在写完permute后直接view导致报错RuntimeError: view size is not compatible,就是这个原因。
CCNet里使用这个模块的方法是串行两次调用,中间夹一个LayerNorm或者BatchNorm,并且加残差连接。两轮之后的信息传播就是完整的了。
4. 两种方案的核心差异与选型思路对比
写到这里,读者应该能感受到两种方案的内在逻辑不同。Axial Attention走的是“严格分工”路线:先水平后垂直,两轮,每轮只做一维。Criss-Cross走的是“十字并行”路线,每一轮同时处理水平加垂直,通过两轮循环弥补路径长度的代价。为了让你在实际工程中快速做取舍,我把核心差异整理成一个表格。
| 对比维度 | Axial Attention | Criss-Cross Attention |
|---|---|---|
| 单轮交互范围 | 仅行内或仅列内 | 同行+同列十字区域 |
| 全局信息传播 | 两轮分别完成行列交互 | 两轮循环完成全图交互 |
| 计算复杂度 | O(HW(H+W)) | O(HW(H+W)) |
| 内存占用(注意力图) | 需要分别存行/列注意力图 | 需要分别存横/竖注意力图,量级相当 |
| 直觉感受 | 先水平再看垂直,类似二维分解 | 每轮直接看十字邻居,类似广播 |
| 代表性应用 | GAN生成、图像Transformer、视频预测 | 语义分割、全景分割、遥感分割 |
| 实现难度 | 中等,permute比较绕 | 中等,两个方向并行较直观 |
| 和标准Attention的差距 | 很接近,甚至有超越的情况 | 略低于标准Attention,但计算量优势大 |
4.1 选型建议:什么场景用哪个
如果你做的是生成类任务,比如StyleGAN式的图像生成、视频预测、图像补全,那Axial Attention的表现更稳定。因为生成任务需要一个良好的全局约束,而Axial的两阶段分解能提供一个非常干净的全局感受野,且不容易出现棋盘格伪影。
如果你做的是密集预测类任务,比如语义分割、城市街景分割、遥感影像分割,那Criss-Cross Attention更合适。原因很简单:这类任务输入分辨率高,特征图大,对显存极度敏感。CCNet的计算量低,且能直接插到FPN、DeepLabV3+这类分割框架中间,改动小、收益不错。
一个比较常见的落地组合是这样:骨干网络用ResNet,在Stage4输出端接一个CCNet模块(两个循环),得到增强后的特征,再接分割头。这个组合参数量不大,但精度提升非常明显,在Cityscapes和ADE20K数据集上都能看到稳定的涨点。
4.2 常见问题与调试实录
我在实际调模型的时候踩过一些坑,写出来给各位避避雷。
第一个大坑:归一化层不能乱放。这两种Attention模块的线性变换层后,最好是加上LayerNorm或者BatchNorm,然后接GELU/ReLU激活。尤其是视觉任务中,如果不做归一化,直接用多头注意力输出做残差相加,训练到中期loss会出现明显的震荡。我试过直接套Transformer里那种Post-Norm方案,在小数据集上可以,但换大数据集就崩,最后切回Pre-Norm风格,问题解决。
第二个大坑:特征图尺寸要能被下采样倍数整除。Axial的row/column变换涉及reshape,特征图尺寸不整会报错。比如输入分辨率 512x512,经过8倍下采样得到 64x64,这没问题。但如果是 500x500 这种不规整的输入,下采样后是 62.5x62.5,就会出问题。工程上做分割时,一般在数据加载阶段就统一做padding到能被32整除的尺寸。
第三个大坑:head数量和dim_head的选择。常规Transformer里head=8、dim_head=64是常见配置,但在Axial和Criss-Cross里,因为每个序列的长度比NLP短很多,dim_head反而不用太大。我实测下来,head=8、dim_head=32在大多数任务上效果和维度性价比最高。dim_head太大会让注意力矩阵过于稀疏,反而丢失细节。
第四个大坑:残差连接和尺度缩放。直接堆叠两个Attention模块不加残差,深层梯度会消失。正确做法是每个Attention后面都接残差,残差分支上再加一个可学习的缩放因子,初始化为0。这样训练前期相当于主干网络原样输出,Attenton慢慢加进来,训练明显更稳定。这个技巧来自早期ResNet的残差缩放思想,但在Attention这种模块上效果拔群。
第五个大坑:Batch Size和学习率的匹配。这两个Attention模块的显存优化是有代价的:BN层的batch统计量对batch size比较敏感。我在单卡上训练时batch size只能设到8,BN效果很差,换成GroupNorm或者用SyncBN就稳定多了。如果是做研究快速验证,建议用GroupNorm,稳妥省心。
另外还有一个值得提的工程经验:如果你在推理时需要极致速度,Criss-Cross两个方向可以并行算,减少一次串行等待。在PyTorch里可以用torch.nn.Parallel或者自定义autograd Function来优化,不过对小模型收益不大,大Feature Map上并行能省下约15%的推理时间。
我个人在实际工程里用得比较多的是Criss-Cross Attention,主要是分割任务天然输入大、显存紧张。Axial Attention在生成类的项目里我更偏好,因为它对全局结构的把控更精细。
最后分享一个调参细节:不管用哪种Attention,QK的scale计算都是用 dim_head 的平方根,而不是总维度。如果dim_head设为32,标准差初始化的scale就是 32/4096 这样量级(用PyTorch默认的kaiming初始化一样能工作,但效果略差)。更稳妥的做法是手动对QKV层做xavier初始化,再加scale因子,这一步能显著提升收敛速度。
这些模块的代码我已经在多个项目里验证过,直接复制核心实现到自己的工程里,再按需调整通道数和head数,基本可以无缝接入。有条件的还是建议自己从零推导一遍张量形状变化,这样真遇到了维度问题,脑子里能瞬间定位是哪一步permute出错了。