(论文阅读)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. 所提出方法:
首先是模型框架。

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

特别的,cls token的混合是双向的,L to S以及S to L。具体地细节还需要看代码:为什么把cls token也放进去了,意义是什么?
总之,由于仅混合cls token 计算量要小一些。
3. 实验:
其实在今年这个结果不太重要了


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。变化维度也很清楚。没问题。
posted on 2025-04-30 11:28 Orange0005 阅读(224) 评论(0) 收藏 举报
浙公网安备 33010602011771号