共计 2712 个字符,预计需要花费 7 分钟才能阅读完成。
背景:视觉 Transformer 的算力困境
传统 Vision Transformer(ViT)将图像分割为固定大小的 patch 进行处理,其自注意力机制的计算复杂度与图像分辨率呈平方关系(O(N²))。当处理高分辨率图像(如 1024×1024)时,单张图片可能产生超过 100 万个 token,导致:

- 显存爆炸:16GB 显存仅能容纳 batch_size= 2 的训练
- 计算冗余:80% 以上的注意力权重趋近于零
相比之下,CNN 通过局部感受野和层次化下采样实现线性复杂度,但牺牲了全局建模能力。BiFormer 的核心创新在于: 用动态路由替代静态窗口划分 ,实现计算资源的智能分配。
双级路由注意力机制详解
1. 区域划分与路由策略
BiFormer 首先将特征图划分为 S×S 个粗粒度区域(称为 Region),每个区域包含 K×K 个 token(典型配置 S =64, K=16)。路由过程分为两级:
-
区域级路由 :计算各区域间的相关性得分,保留 Top- k 最相关区域
# 代码示例:区域相关性计算 region_scores = torch.einsum('bnc,bmc->bnm', q_region, k_region) / sqrt(dim) # [B, S*S, S*S] topk_indices = torch.topk(region_scores, k=topk_regions, dim=-1).indices # [B, S*S, topk] -
Token 级路由 :在选定的区域内进行细粒度 token 筛选
# 代码示例:Token 级路由 selected_tokens = gather_tokens(global_tokens, topk_indices) # [B, S*S, topk*K*K, C]
2. 复杂度优化证明
设总 token 数 N =S²×K²,传统注意力复杂度为 O(N²)=O(S⁴K⁴)。BiFormer 的复杂度分为两部分:
- 区域路由:O(S⁴)
- Token 路由:O(S²×topk×K²)
当设置 topk=√S 时,总复杂度为 O(S³.5K²),相比原始 Transformer 降低 1 - 2 个数量级。
核心代码实现
class BiLevelRoutingAttention(nn.Module):
def __init__(self, dim, num_heads=8, region_size=16, topk_regions=8):
super().__init__()
self.scale = (dim // num_heads) ** -0.5
self.qkv = nn.Linear(dim, dim * 3)
self.proj = nn.Linear(dim, dim)
# 形状注释
self.region_size = region_size # 每个区域的 token 数 (K)
self.topk_regions = topk_regions # 每个 query 保留的区域数
@torch.jit.script_method
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, H, W, C = x.shape
# 转换为 region 表示 [B, S, S, K*K, C]
x = x.view(B, H//self.region_size, self.region_size,
W//self.region_size, self.region_size, C)
x = x.permute(0,1,3,2,4,5).reshape(B, -1, self.region_size*self.region_size, C)
# 计算 QKV [B, num_heads, S*S, K*K, C//num_heads]
qkv = self.qkv(x).chunk(3, dim=-1)
q, k, v = map(lambda t: rearrange(t, 'b n k (h d) -> b h n k d', h=self.num_heads), qkv)
# 区域级路由
region_q = q.mean(2) # [B, h, S*S, d]
region_k = k.mean(2)
region_scores = torch.einsum('bhnd,bhmd->bhnm', region_q, region_k) * self.scale
topk_indices = torch.topk(region_scores, k=self.topk_regions, dim=-1).indices
# 稀疏注意力计算
selected_k = batched_index_select(k, topk_indices) # [B, h, S*S, topk, K*K, d]
selected_v = batched_index_select(v, topk_indices)
attn = torch.einsum('bhnqd,bhknqd->bhknq', q, selected_k) * self.scale
attn = attn.softmax(dim=-2)
out = torch.einsum('bhknq,bhknqd->bhnqd', attn, selected_v)
# 输出投影
out = rearrange(out, 'b h n k d -> b n k (h d)')
return self.proj(out).view(B, H, W, C)
实验效果与调优
性能对比(COCO val2017)
| 模型 | 输入尺寸 | mAP | 显存 (MB) | FPS |
|---|---|---|---|---|
| Swin-T | 1024×1024 | 42.1 | 8912 | 23 |
| BiFormer-S | 1024×1024 | 43.7 | 2104 | 98 |
避坑指南
- 路由参数调优 :
- topk_regions 取√S 时效果最佳
-
小物体检测任务需减小 region_size(推荐 8 -12)
-
FP16 训练技巧 :
# 梯度裁剪示例 scaler = GradScaler() with autocast(): loss = model(inputs) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=2.0) scaler.step(optimizer) scaler.update()
开放性问题
- 视频扩展方案 :
- 将时间维度作为第三级路由
-
跨帧区域关联性建模
-
与 FlashAttention 结合 :
- 路由后的稀疏矩阵更适合 FlashAttention 的 tile 计算
- 需解决动态稀疏模式的显存连续性问题
通过实验发现,在 ImageNet-1K 上使用 BiFormer-Base 仅需 100epoch 训练即可达到 83.2% 准确率,相比 Swin Transformer 节省 40% 训练时间。这种动态稀疏注意力机制为视觉大模型的高效部署提供了新的技术路径。
