(论文阅读)CrossViT: Cross-Attention Multi-Scale Vision Transformer for Image Classification

1. 论文

题目:CrossViT: Cross-Attention Multi-Scale Vision Transformer for Image Classification
代码: https://github.com/IBM/CrossViT
会议/期刊: ICCV 2021
摘要:

本文提出了一种双分支transformer(vit)用于学习多尺度特征。其中不同分支接受的patch token大小不同。接着通过交叉注意力混合cls token(一个分支)与所有patch toekn(另一个分支的)。并且两个分支大小也不同。


2. 所提出方法:

首先是模型框架。

image-20250430105744277

下面是本文的主要创新,不同分支之间的混合(Cross-Attention module)(本文还提出了三种直接简单地混合,详见文中):

image-20250430105936251

特别的,cls token的混合是双向的,L to S以及S to L。具体地细节还需要看代码:为什么把cls token也放进去了,意义是什么?

总之,由于仅混合cls token 计算量要小一些。


3. 实验:

其实在今年这个结果不太重要了

image-20250430110559739

image-20250430110634458


for i in range(self.num_branches):
    tmp = torch.cat((proj_cls_token[i], outs_b[(i + 1) % self.num_branches][:, 1:, ...]), dim=1)
    # 拼接了i branch的cls token与i+1 branch的patch token
    tmp = self.fusion[i](tmp)

# self.fusion 
x = x[:, 0:1, ...] + self.drop_path(self.attn(self.norm1(x)))
# 交叉注意力模块
class CrossAttention(nn.Module):
    def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0.):
        super().__init__()
        self.num_heads = num_heads
        head_dim = dim // num_heads
        # NOTE scale factor was wrong in my original version, can set manually to be compat with prev weights
        self.scale = qk_scale or head_dim ** -0.5

        self.wq = nn.Linear(dim, dim, bias=qkv_bias)
        self.wk = nn.Linear(dim, dim, bias=qkv_bias)
        self.wv = nn.Linear(dim, dim, bias=qkv_bias)
        self.attn_drop = nn.Dropout(attn_drop)
        self.proj = nn.Linear(dim, dim)
        self.proj_drop = nn.Dropout(proj_drop)

    def forward(self, x):

        B, N, C = x.shape
        q = self.wq(x[:, 0:1, ...]).reshape(B, 1, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3)  
        # B1C -> B1H(C/H) -> BH1(C/H)
        k = self.wk(x).reshape(B, N, self.num_heads, C//self.num_heads).permute(0, 2, 1, 3) 
        # BNC -> BNH(C/H) -> BHN(C/H)
        v = self.wv(x).reshape(B, N, self.num_heads, C//self.num_heads).permute(0, 2, 1, 3) 
        # BNC -> BNH(C/H) -> BHN(C/H)

        attn = (q @ k.transpose(-2, -1)) * self.scale  # BH1(C/H) @ BH(C/H)N -> BH1N
        attn = attn.softmax(dim=-1)
        attn = self.attn_drop(attn)

        x = (attn @ v).transpose(1, 2).reshape(B, 1, C)   
        # (BH1N @ BHN(C/H)) -> BH1(C/H) -> B1H(C/H) -> B1C
        x = self.proj(x)
        x = self.proj_drop(x)
        return x

KV中也包含了i branch的cls token。变化维度也很清楚。没问题。

代码来自:https://github.com/IBM/CrossViT

posted on 2025-04-30 11:28  Orange0005  阅读(224)  评论(0)    收藏  举报