——只算“该算的地方“的稀疏注意力)
【医学图像分割模块】BiFormer / 双层路由注意力CVPR 2023—— 只算该算的地方的稀疏注意力医学图像分割里到处是高分辨率特征图。直接上全局注意力Self-Attention有两个老问题一是计算量随 token 数平方级爆炸二是大部分区域其实毫无关系比如肿瘤区和背景区却要两两算一遍——又慢又浪费。BiFormerCVPR 2023提出双层路由注意力Bi-Level Routing Attention, BRA先在区域region层面粗筛出每个区域最相关的少数几个区域再只在筛出来的区域里做细粒度 token-to-token 注意力。注意力变成内容感知的稀疏注意力——该算的才算计算量大幅下降而关键信息不丢。一、论文出处论文BiFormer: Vision Transformer with Bi-Level Routing Attention会议CVPR 2023论文链接https://arxiv.org/abs/2303.08810官方代码https://github.com/rayleizhu/BiFormer二、模块图截自论文原文 Figure 3图自 BiFormer 原文 Figure 3CVPR 2023。左半是整体金字塔架构Input → Stage 1~4每个 Stage 由Patch Embedding / Patch Merging 若干 BiFormer Block组成分辨率逐级减半、通道逐级翻倍右半虚线框展开BiFormer Block的内部DWConv 3×3 → LN → Bi-level Routing Attention → LN → MLP带残差。一句话主干是分层 ViT核心是里面的BRA 注意力——先在区域层路由再在少数相关区域里做注意力。三、核心思想与作用一句话注意力别全算——先在区域层面挑出相关的少部分区域只在里面做 token-to-token 注意力。拆解两步双层就在这区域级路由粗筛把特征图切成n×n个不重叠区域对每个区域算它和其他区域的相关度只保留 top-k 个最相关的区域其余丢弃。token 级注意力细算把上一步选中的 top-k 区域里的 key/value 聚到一起让本区域的每个 token 只对这一小撮 key/value 做注意力。为什么好计算量从全图两两算降到每个区域只和 top-k 个区域算复杂度显著下降内容感知路由是按内容动态决定的相关的地方才连比固定窗口/固定步长的稀疏注意力更灵活对高分辨率友好医学图像分辨率高、有效信息稀疏正好吃这个红利。在分割里的作用把编码器里的全局注意力换成 BRA在几乎不掉精度的情况下大幅省算力和显存尤其适合高分辨率 2D 医学图像、以及多尺度特征融合处。四、在 U-Net 里的插入位置BRA 是注意力模块用在需要全局建模的位置编码器深层 / 瓶颈层替换 Self-Attention 或 Transformer block做长距离依赖跳连skip connection融合处对多尺度特征做稀疏全局交互解码器上采样前在低分辨率、高通道处收益最大token 少、通道多。通道数越大、分辨率越高BRA 相对全注意力的优势越明显。五、复现代码PyTorch逐行中文注释下面是简化教学版保留了 BRA 的两个核心步骤区域路由 稀疏注意力代码清晰、可跑生产使用建议对照官方实现含位置编码、下采样等细节。https://github.com/rayleizhu/BiFormerimporttorchimporttorch.nnasnnclassBiLevelRoutingAttention(nn.Module):BiFormer 双层路由注意力简化教学版。 ① 区域级路由每个区域只保留最相关的 topk 个区域 ② 区域级 token-to-token 注意力只在选中的区域里算注意力。 def__init__(self,dim,num_heads8,n_win7,topk4):super().__init__()assertdim%num_heads0,dim 必须能被 num_heads 整除self.dimdim self.num_headsnum_heads self.n_winn_win# 把特征图切成 n_win × n_win 个区域self.topktopk# 每个区域只和 topk 个最相关区域做注意力self.head_dimdim//num_heads self.scaleself.head_dim**-0.5# 缩放因子防止点积过大self.qkvnn.Linear(dim,dim*3)# 一次投影出 q/k/vself.projnn.Linear(dim,dim)# 输出投影# 深度卷积做局部位置编码LePE补回被稀疏化丢失的局部信息self.lepenn.Conv2d(dim,dim,kernel_size5,padding2,groupsdim)defforward(self,x):B,N,Cx.shape HWint(N**0.5)# 假设方形特征图nself.n_win rs(H//n)*(W//n)# 每个区域内的 token 数# ① 生成 q/k/v[B, heads, N, head_dim]qkvself.qkv(x).reshape(B,N,3,self.num_heads,self.head_dim).permute(2,0,3,1,4)q,k,vqkv[0],qkv[1],qkv[2]# ② 按区域重塑[B, heads, R, rs, head_dim]R n*n 个区域defto_region(t):returnt.reshape(B,self.num_heads,n*n,rs,self.head_dim)q_r,k_r,v_rto_region(q),to_region(k),to_region(v)# ③ 区域代表向量 区域内 token 平均用来算区域间相关度q_repq_r.mean(dim3)# [B, heads, R, head_dim]k_repk_r.mean(dim3)# [B, heads, R, head_dim]aff(q_rep*self.scale) k_rep.transpose(-1,-2)# [B, heads, R, R] 区域相关度# ④ 双层路由的关键每个区域只留 topk 个最相关区域topkmin(self.topk,n*n)idxaff.topk(topk,dim-1).indices# [B, heads, R, topk]# ⑤ 按路由索引把选中的 k/v 区域聚起来[B, heads, R, topk, rs, head_dim]defgather_region(t):dt.shape[-1]t_expt.unsqueeze(3).expand(B,self.num_heads,n*n,topk,rs,d)indexidx[...,None,None].expand(B,self.num_heads,n*n,topk,rs,d)returntorch.gather(t_exp,2,index)k_ggather_region(k_r).reshape(B,self.num_heads,n*n,topk*rs,self.head_dim)v_ggather_region(v_r).reshape(B,self.num_heads,n*n,topk*rs,self.head_dim)# ⑥ 只在本区域 token × 选中区域 key之间做注意力attn(q_r*self.scale) k_g.transpose(-1,-2)# [B, heads, R, rs, topk*rs]attnattn.softmax(dim-1)outattn v_g# [B, heads, R, rs, head_dim]# ⑦ 还原回 token 序列[B, N, C]outout.reshape(B,self.num_heads,N,self.head_dim).permute(0,2,1,3).reshape(B,N,C)# ⑧ 加局部位置编码LePE 输出投影imgout.transpose(1,2).reshape(B,C,H,W)# [B, C, H, W]imgimgself.lepe(img)# 补局部位置信息outimg.flatten(2).transpose(1,2)# [B, N, C]outself.proj(out)returnout六、插入示例几行塞进你的网络# 例用 BRA 替换 U-Net 瓶颈层的全局注意力self.attnBiLevelRoutingAttention(dim256,num_heads8,n_win7,topk4)defforward(self,x):# x: [B, N, C]xxself.attn(x)# 残差接住替换原来的 SelfAttention / TransformerBlockreturnx七、实测经验与注意点粒度可调n_win区域划分和topk每区域连几个是两个核心旋钮——topk越小越省精度靠相关性保住一般n_win7/14、topk4是常用起点。复杂度相对全局注意力从随 token 数平方增长降为随区域数×topk 增长高分辨率下省得最多。踩坑输入必须是方形特征图HW否则要处理的区域切分和还原会变形原始实现里用H//n、W//n分别算dim % num_heads 0、H % n_win 0都要成立否则 reshape 报错稀疏化会丢局部信息务必保留 LePE 之类的局部位置编码教学版做了简化上生产请对照官方实现含 kv 下采样、位置编码细节。医学分割场景高分辨率切片、需要长距离依赖又要控显存时BRA 比全注意力更合适小数据集上建议配合预训练或强增广。八、完整工程 领取本文的完整可运行工程BRA 的.py、U-Net 插入 demo、n_win/topk可调配置、逐行中文注释我整理好了。领取方式关注公ZHU号MediVision回复「模块」——自动把《即插即用模块合集含 FasterNet / BiFormer / StarNet / RepViT / TransNeXt…统一接口一条 import 就能用》的入口发给你进群免费领。下一篇预告RepViTCVPR 2024——把 ViT 的高效设计搬回纯卷积用重参数化换一个又快又强的轻量主干。