paper:Transformer in Transformer
official implementation:Efficient-AI-Backbones/tnt_pytorch at master · huawei-noah/Efficient-AI-Backbones · GitHub
third-party implementation:https://github.com/huggingface/pytorch-image-models/blob/main/timm/models/tnt.py
存在的问题
视觉 Transformer (ViT) 通过将输入图像分割成多个局部图像块(patches),然后计算这些块的表示及其相互关系,已经被用于图像识别任务。自然图像具有高度的复杂性,包含丰富的细节和颜色信息,传统 ViT 中的图像块划分粒度不足以有效挖掘不同尺度和位置的对象特征。
本文的创新点
文章指出,在局部图像块内部的注意力对于构建高性能的视觉 Transformer 同样至关重要。作者提出了一种新的架构,Transformer in Transformer (TNT),通过在局部图像块内进一步划分更小的图像块,并计算这些小块之间的注意力,以更精细地提取特征。
具体来说
- TNT架构:提出了一种新的 Transformer 结构,其中局部图像块被视为“视觉句子”,每个句子进一步被划分为更小的“视觉单词”。
- 双重注意力机制:在 TNT 中,引入了内部 Transformer 块来计算视觉词之间的特征和注意力,同时使用外部 Transformer 块来处理视觉句子。
- 计算效率:通过共享网络来独立计算每个视觉句子中视觉词的特征和注意力,使得参数和浮点运算(FLOPs)的增加可以忽略不计。
- 特征聚合:通过聚合视觉词和句子的特征来增强表示能力,从而在多个基准测试中取得了优于现有最先进视觉 Transformer 的结果。
- 性能提升:TNT模型在ImageNet等基准测试上表现显著提升,达到了81.5%的Top-1准确率,比具有相似计算成本的最先进视觉Transformer高出约1.7%。
方法介绍
给定一个2D图片,我们均匀地将其切分成 \(n\) 个图像块 \(\mathcal{X}=[X^1,X^2,...,X^n]\in \mathbb{R}^{n\times p\times p\times 3}\),其中 \((p,p)\) 是每个patch的分辨率。ViT直接使用transformer来处理patch序列但破坏了patch的内部结构。本文提出了 Transformer-iN-Transformer(TNT)架构来同时学习图像中的全局信息和局部信息。在TNT中,我们将patch作为表示图像的视觉句子“visual sentence”,每个patch被进一步切分成 \(m\) 个子图像块 sub-patch,即一个视觉句子是由视觉单词“visual word”的序列组成的

其中 \(x^{i,j}\in \mathbb{R}^{s\times s\times 3}\) 是第 \(i\) 个视觉句子的第 \(j\) 个视觉单词,\((s,s)\) 是sub-patch的大小,\(j=1,2,...,m\)。利用一个线性映射,我们将视觉单词转换成单词特征序列

其中 \(y^{i,j}\in \mathbb{R}^c\) 第 \(j\) 个单词embedding,\(c\) 是word embedding的维度,\(Vec(\cdot)\) 是向量化运算。
在TNT中,我们有两个数据流,一个处理句子,一个处理句子里的单词。对于单词embedding,我们用一个transformer block来探索视觉单词之间的关系

其中 \(l=1,2,...,L\) 是第几个block的索引,一共stack了 \(L\) 个block。第一个block \(Y_0^i\) 的输入就是式(5)中的 \(Y^i\)。像中所有的单词embedding在转换后为 \(\mathcal{Y}_l=[Y^1_l,Y^2_l,...,Y^n_l]\)。这可以看作一个inner transformer block,表示为 \(T_{in}\)。这个过程通过计算任意两个视觉单词之间的交互建立视觉单词之间的关系。
在句子层面,我们创建句子embedding memories来保存句子的表示:\(\mathcal{Z}_0=[Z_{class},Z^1_0,Z^2_0,...,Z^n_0]\in \mathbb{R}^{(n+1)\times d}\),其中 \(Z_{class}\) 是class token,并且它们都初始化为0。在网络每一层,单词embedding序列通过线性映射转换到句子embedding domain,并与句子embedding相加:

其中 \(Z^{i}_{l-1}\in \mathbb{R}^d\),FC用于匹配维度从而可以相加。通过相加,单词级别的特征加强了句子embedding的表示能力。我们用标准的transformer block来转换句子embedding:

这个outer block用来建模句子embedding之间的联系。
综上所述,TNT block的输入和输出包括visual word embeddings和sentence embeddings,如图1所示,因此TNT可以表述为:


Position emcoding. 空间信息是图像识别的一个重要因素。对于句子embedding和单词embedding我们都添加了相应的位置编码来保留空间信息,如图1所示。这里使用了标准的1D可学习的位置编码。每个句子都配有一个位置编码:

其中 \(E_{sentence}\in \mathbb{R}^{(n+1)\times d}\) 是句子位置编码。对于句子中的单词,每个单词的embedding都加上一个单词位置编码:

其中 \(E_{word}\in \mathbb{R}^{m\times c}\) 是单词位置编码,并在所有句子中共享。这样,句子位置编码保留了全局空间信息,而单词位置编码用来保存局部的相对位置。
实验结果
TNT不同变种的网络结构如表1所示

在ImageNet上的结果如表4所示,可以看到TNT的性能要优于ViT和DeiT。

代码解析
这里解析的是timm中实现的TNT,选用的模型是tnt_s_patch16_224,输入shape=(1, 3, 224, 224)。首先进行visual word embedding,代码如下。其中sentence patch的大小为16x16,word patch大小为4x4,self.proj是一个7x7-s4-p3-conv,word特征的dim=24,得到输出shape为(1, 24, 56, 56),这里用步长为4的卷积相当于提取出了word的特征。然后self.unfold以kernel_size=stride=4再从word特征通过滑动窗口聚合出每个4x4窗口的特征,nn.Unfold和卷积一样都是滑动窗口只不过卷积需要进行对应位置权重的相乘然后再相加而unfold直接把窗口内的值都取出来,具体用法见torch.nn.functional.unfold 用法解读-CSDN博客。前面说过句子patch大小为16x16,单词patch大小为4x4,因此一个句子就包含4x4个单词,self.proj提取出单词embedding后再通过unfold将每4x4个单词特征聚合到一起,相当于聚合了一个句子里的单词特征,得到输出shape为(1, 384, 196),其中384=4x4x24是一个句子内所有单词的embedding,196=(224/16)^2相当于图片中所有句子patch的数量。然后通过一些维度的变换并与单词级别的位置编码即pixel_pos相加最终的word embedding,shape为(196, 16, 24)。
class PixelEmbed(nn.Module):
""" Image to Pixel Embedding
"""
def __init__(self, img_size=224, patch_size=16, in_chans=3, in_dim=48, stride=4):
super().__init__()
img_size = to_2tuple(img_size)
patch_size = to_2tuple(patch_size) # (16,16)
# grid_size property necessary for resizing positional embedding
self.grid_size = (img_size[0] // patch_size[0], img_size[1] // patch_size[1]) # (224/16=14, 14)
num_patches = (self.grid_size[0]) * (self.grid_size[1]) # 196
self.img_size = img_size
self.num_patches = num_patches
self.in_dim = in_dim # 24
new_patch_size = [math.ceil(ps / stride) for ps in patch_size]
self.new_patch_size = new_patch_size # (4,4)
self.proj = nn.Conv2d(in_chans, self.in_dim, kernel_size=7, padding=3, stride=stride) # 3,24,4
self.unfold = nn.Unfold(kernel_size=new_patch_size, stride=new_patch_size) # https://blog.csdn.net/ooooocj/article/details/127029199
def forward(self, x, pixel_pos): # (1,3,224,224),(1,24,4,4)
B, C, H, W = x.shape
_assert(H == self.img_size[0],
f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]}).")
_assert(W == self.img_size[1],
f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]}).")
x = self.proj(x) # (1,24,56,56), 这里56x56是所有sub_patch的数量, 这一步相当于直接提取sub_patch的特征
x = self.unfold(x) # (1,384,196), 这一步对sub_patch再进行4x4的unfold相当于提取patch
x = x.transpose(1, 2).reshape(B * self.num_patches, self.in_dim, self.new_patch_size[0], self.new_patch_size[1]) # (1,196,384)->(196,24,4,4)
x = x + pixel_pos
x = x.reshape(B * self.num_patches, self.in_dim, -1).transpose(1, 2) # (196,16,24)
return x
然后第二步提取句子的embedding,这里代码中是对pixel_embed通过self.proj处理就得到了patch_embed,self.proj是一个nn.Linear线性层。我的理解是句子patch时由单词patch组成的,在ViT中单独提取句子patch无非就是卷积核和步长更大一些,所以这里对单词patch的特征进行learnable linear projection就得到了句子的embedding。然后加上class_token并与句子级别的位置编码相加就得到了最终的sentence embedding。
pixel_embed = self.pixel_embed(x, self.pixel_pos) # (196,16,24)
# pixel_embed中196表示16x16的patch的数量,(224/16)^2=196. 16表示每个patch内sub_patch的数量, (16/4)^2=16. 24则是每个sub_patch的特征dim
patch_embed = self.norm2_proj(self.proj(self.norm1_proj(pixel_embed.reshape(B, self.num_patches, -1)))) # (1,196,384) -> (1,196,384)
# 这里为什么pixel_embed用个线性层映射一下就变成patch_embed了?16x24就代表每个patch内的特征,因此proj一下 就变成patch_embed了
patch_embed = torch.cat((self.cls_token.expand(B, -1, -1), patch_embed), dim=1) # (1,197,384)
patch_embed = patch_embed + self.patch_pos
然后就是网络的主干部分,每一层输入单词特征和句子特征即pixel_embed和patch_embed,输出也包括pixel_embed和patch_embed。
for blk in self.blocks:
pixel_embed, patch_embed = blk(pixel_embed, patch_embed) # (196,16,24),(1,197,384)
每一层包含inner transformer block用来处理单词特征和outer transformer block用来处理句子特征,单词特征首先经过inner transformer block处理后与句子特征相加,然后再经过outer transformer block,通过将单词特征添加到句子特征中增强了句子特征的表示能力,具体如下
def forward(self, pixel_embed, patch_embed): # (196,16,24),(1,197,384)
# inner
pixel_embed = pixel_embed + self.drop_path(self.attn_in(self.norm_in(pixel_embed))) # (196,16,24)
pixel_embed = pixel_embed + self.drop_path(self.mlp_in(self.norm_mlp_in(pixel_embed))) # (196,16,24)
# outer
B, N, C = patch_embed.size()
patch_embed = torch.cat(
[patch_embed[:, 0:1], patch_embed[:, 1:] + self.proj(self.norm1_proj(pixel_embed).reshape(B, N - 1, -1))],
dim=1) # (1,197,384)
# torch.cat[(1,1,384), (1,196,384) + (1,196,384)]
patch_embed = patch_embed + self.drop_path(self.attn_out(self.norm_out(patch_embed))) # (1,197,384)
patch_embed = patch_embed + self.drop_path(self.mlp(self.norm_mlp(patch_embed))) # (1,197,384)
return pixel_embed, patch_embed