BiFormer入门指南:理解Bi-Level Routing Attention在Vision Transformer中的应用

1次阅读
没有评论

共计 2284 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

背景与痛点

Vision Transformer(ViT)通过将图像分割为 patch 序列并应用自注意力机制,在多个视觉任务上取得了显著成果。然而,标准的全局自注意力计算复杂度为 O(N²),其中 N 是 patch 数量。对于高分辨率图像(如 224×224 分为 196 个 16×16 patch),这会导致巨大的计算和内存开销,限制了 ViT 在资源受限设备上的应用。

BiFormer 入门指南:理解 Bi-Level Routing Attention 在 Vision Transformer 中的应用

技术对比

与标准 ViT 和 Swin Transformer 相比,BiFormer 通过双级路由注意力机制实现了显著改进:

  • 计算复杂度:标准 ViT 为 O(N²),Swin 为 O(N),BiFormer 为 O(N√N)
  • 内存占用:在 ImageNet-1K 上,BiFormer 比 ViT 减少约 40% 内存使用
  • 准确率:在相似计算量下,BiFormer Top- 1 准确率比 Swin 高出 1 -2%

核心原理

Bi-Level Routing Attention(BRA)机制包含三个关键步骤:

  1. 区域划分:将输入特征图划分为 S×S 个区域(Region),每个区域包含 K×K 个 token

  2. 路由选择

  3. 第一级:选择最相关的 M 个区域(M≪S²)
  4. 第二级:在每个选中区域内选择 N 个最相关 token(N≪K²)
  5. 路由函数:$R = softmax(QW_r)$,其中 W_r 是可学习路由权重

  6. 注意力计算

  7. 对选中的 token 应用标准自注意力
  8. 未选中 token 通过最近邻插值获得注意力结果

代码实现

以下是 PyTorch 实现的 BRA 核心模块:

import torch
import torch.nn as nn
import torch.nn.functional as F

class BiLevelRoutingAttention(nn.Module):
    def __init__(self, dim, num_heads=8, region_size=7, topk_region=4, topk_token=16):
        super().__init__()
        self.dim = dim
        self.num_heads = num_heads
        self.region_size = region_size
        self.topk_region = topk_region
        self.topk_token = topk_token

        # 投影层
        self.qkv = nn.Linear(dim, dim * 3)
        self.proj = nn.Linear(dim, dim)

        # 路由权重
        self.routing_weight = nn.Parameter(torch.randn(dim, num_heads))

    def forward(self, x):
        B, H, W, C = x.shape
        # 划分区域
        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**2, C)

        # 计算查询向量和路由分数
        qkv = self.qkv(x).reshape(B, -1, 3, self.num_heads, C//self.num_heads)
        q, k, v = qkv.unbind(2)

        # 双级路由选择
        routing_score = torch.einsum('bnhd, dh->bnh', q.mean(dim=2), self.routing_weight)
        region_score = routing_score.mean(dim=-1)  # 区域级分数
        token_score = routing_score  # token 级分数

        # 选择 topk 区域和 token
        _, region_indices = region_score.topk(self.topk_region, dim=1)
        _, token_indices = token_score.topk(self.topk_token, dim=2)

        # 注意力计算(省略完整实现)# ...

        return x

性能分析

下表展示了不同输入尺寸下的性能对比(基于 BiFormer-Tiny 模型):

输入尺寸 FLOPs (G) 内存 (MB) 延时 (ms)
224×224 1.2 450 15.2
384×384 3.6 1200 42.8
512×512 6.4 2100 78.5

适用场景建议:
– 移动设备:224×224 输入
– 服务器端:可考虑 384×384 或更大输入

避坑指南

  1. 路由参数调优
  2. topk_regiontopk_token 需要平衡计算量和模型性能
  3. 建议初始值:topk_region=4,topk_token=16

  4. 混合精度训练

  5. 路由计算部分建议保持 FP32 精度
  6. 可使用 torch.cuda.amp 自动管理

  7. 学习率设置

  8. 路由权重学习率应为其他参数的 1 /5-1/10
  9. 建议使用分层学习率策略

实践建议

在自定义数据集上微调 BiFormer 的步骤:

  1. 数据准备:
  2. 确保图像尺寸能被 region_size 整除
  3. 推荐使用 224×224 输入

  4. 模型初始化:

  5. 从官方预训练模型开始
  6. 冻结除分类头外的所有层训练几个 epoch

  7. 微调策略:

  8. 初始学习率:1e-4
  9. 使用余弦退火调度
  10. batch size 至少 32 以保证路由稳定性

延伸思考

Bi-Level Routing 机制可以扩展到其他视觉任务:

  1. 目标检测:
  2. 对 ROI 区域应用路由选择
  3. 减少检测头计算量

  4. 语义分割:

  5. 在解码器阶段使用路由注意力
  6. 重点关注边界区域

  7. 视频理解:

  8. 时空双重路由
  9. 选择关键帧和关键区域

通过本文的介绍,希望读者能够理解 BiFormer 的核心思想并成功应用于自己的项目中。双级路由注意力提供了一种有效的计算 - 精度平衡方案,特别适合资源受限的应用场景。

正文完
 0
评论(没有评论)