论文信息

  • 标题:MobileViT: Light-weight, General-purpose, and Mobile-friendly Vision Transformer
  • 会议:ICLR 2022
  • 单位:Apple
  • 代码:https://github.com/apple/ml-cvnets
  • 论文:https://arxiv.org/pdf/2110.02178.pdf

引言:移动端视觉的"鱼与熊掌"难题

在2021年之前,移动端视觉任务几乎是轻量级CNN的天下。从MobileNet到EfficientNet,这些模型凭借精心设计的卷积结构和空间归纳偏置,能够在有限的计算资源下实现不错的性能。

通俗来说:CNN就像一个局部侦探,每次只看图像的一小块区域,然后逐层扩大视野。这种设计让它很擅长提取边缘、纹理等局部特征,而且计算效率很高。但它也有一个致命的缺点:只能进行局部建模,难以捕捉长距离的全局依赖关系。比如,要判断一张图里是猫还是狗,CNN需要堆叠很多层才能看到整个动物的轮廓。

与此同时,Transformer在NLP领域已经大杀四方。它的自注意力机制能够直接建模任意两个位置之间的关系,天生擅长全局信息处理。但当人们把Transformer搬到CV领域时,却遇到了一个大问题:ViT太笨重了。标准的ViT需要将图像拆成patch,然后对所有patch进行自注意力计算,复杂度是O(n2⋅d)O(n^2 \cdot d)O(n2d),其中nnn是patch的数量。对于一张224×224的图像,如果用16×16的patch,n=196n=196n=196,计算量还能接受;但如果要处理更高分辨率的图像,计算量会爆炸式增长,根本无法在手机上运行。

于是一个自然的问题出现了:能不能结合CNN和ViT的优势,打造一个既轻量又高效,还能在手机上流畅运行的视觉模型?

苹果公司在ICLR 2022上给出了答案:MobileViT。它提出了一个革命性的视角:将Transformer视为一种特殊的卷积操作。通过巧妙的设计,MobileViT既保留了CNN的空间归纳偏置和计算效率,又获得了ViT的全局建模能力。

实验结果令人震惊:在ImageNet-1k数据集上,MobileViT-S以约600万参数实现了78.4%的Top-1准确率,比同等参数的MobileNetv3高3.2%,比DeiT高6.2%。在MS-COCO目标检测任务上,MobileViT比MobileNetv3高5.7%,同时参数量还更小。


核心方法:Transformer即卷积

MobileViT的核心创新在于MobileViT Block,它用一种全新的方式将卷积和Transformer结合起来。不同于以往的"CNN+Transformer"简单拼接,MobileViT从计算范式层面重构了信息流动方式。

3.1 整体架构

MobileViT采用了经典的"stem-backbone-head"结构:

  • Stem:一个3×3卷积层,将输入图像从3通道映射到16通道
  • Backbone:由多个阶段组成,每个阶段交替使用MobileNetv2的倒残差块(MV2)和MobileViT Block。MV2块负责下采样和局部特征提取,MobileViT Block负责全局特征建模
  • Head:全局平均池化+全连接层,输出分类结果
    在这里插入图片描述
图片1:MobileViT整体架构图(出处:论文图1)

有趣的案例:为什么要交替使用MV2和MobileViT Block?

  • 如果全用MobileViT Block,计算量会太大,不适合移动端
  • 如果全用MV2,又无法捕捉全局依赖
  • 交替使用的设计是计算量和性能之间的完美平衡:在低分辨率的特征图上使用MobileViT Block,既能获得全局建模能力,又能控制计算量

3.2 MobileViT Block详解

MobileViT Block的核心思想是:保留卷积的展开和折叠操作,但将中间的局部矩阵乘法替换为Transformer的全局注意力机制

传统卷积可以分解为三个步骤:

  1. 展开(Unfold):将局部感受野内的像素展开为向量
  2. 局部处理(MatMul):与卷积核进行矩阵乘法
  3. 折叠(Fold):将结果重新组装为特征图

MobileViT Block将第二步的局部处理替换为全局的Transformer处理,这样每个像素都能看到其他所有像素,同时保留了卷积的空间归纳偏置。
在这里插入图片描述

图片2:MobileViT Block结构图(出处:论文图2)

一个标准的MobileViT Block包含三个步骤:

步骤1:局部特征提取

首先对输入张量X∈RH×W×CX \in \mathbb{R}^{H \times W \times C}XRH×W×C应用两个卷积层:

  • 一个3×3深度可分离卷积,提取局部空间信息
  • 一个1×1点卷积,将通道数从CCC投影到更高的维度ddd

得到局部特征图XL∈RH×W×dX_L \in \mathbb{R}^{H \times W \times d}XLRH×W×d

XL=Conv1×1(Conv3×3(X))X_L = \text{Conv}_{1 \times 1}(\text{Conv}_{3 \times 3}(X))XL=Conv1×1(Conv3×3(X))

通俗解释:这一步就像先让局部侦探把每个小区域的细节都调查清楚,为后续的全局分析做准备。

步骤2:全局特征建模

接下来是最关键的一步:将局部特征图转换为序列,用Transformer进行全局建模。

具体做法是:

  1. XLX_LXL拆分为NNN个非重叠的patch,每个patch的大小为h×wh \times wh×w(论文中使用2×2)
  2. 将每个patch展平为向量,得到序列P∈RN×(h⋅w⋅d)P \in \mathbb{R}^{N \times (h \cdot w \cdot d)}PRN×(hwd)
  3. 对这个序列应用LLL层Transformer编码器,建模跨patch的全局依赖
  4. 将输出序列重新折叠回特征图形状,得到全局特征图XG∈RH×W×dX_G \in \mathbb{R}^{H \times W \times d}XGRH×W×d

XG=Fold(Transformer(Unfold(XL)))X_G = \text{Fold}(\text{Transformer}(\text{Unfold}(X_L)))XG=Fold(Transformer(Unfold(XL)))

公式逐字母解释

  • XLX_LXL:局部特征图,形状为H×W×dH \times W \times dH×W×d
  • Unfold(⋅)\text{Unfold}(\cdot)Unfold():展开操作,将特征图拆分为patch序列
  • Transformer(⋅)\text{Transformer}(\cdot)Transformer():Transformer编码器,包含多头自注意力和MLP
  • Fold(⋅)\text{Fold}(\cdot)Fold():折叠操作,将序列重新组装为特征图
  • XGX_GXG:全局特征图,形状与XLX_LXL相同

有趣的案例:为什么用2×2的patch?

  • 如果用1×1的patch,序列长度会变成H×WH \times WH×W,计算量太大
  • 如果用4×4的patch,每个patch太大,会丢失太多细节信息
  • 2×2的patch是计算量和信息保留之间的完美平衡
步骤3:特征融合

最后,将全局特征图XGX_GXG通过一个1×1卷积降维到CCC通道,然后与原始输入XXX进行残差连接,得到最终的输出。

Y=X+Conv1×1(XG)Y = X + \text{Conv}_{1 \times 1}(X_G)Y=X+Conv1×1(XG)

通俗解释:这一步就像把全局指挥官的分析结果和局部侦探的调查结果结合起来,得到一个既包含细节又包含全局信息的综合判断。


实验结果:小模型的大能量

MobileViT在多个任务和数据集上进行了全面的测试,结果证明了它的优越性。

4.1 ImageNet-1k图像分类

表格1:ImageNet-1k分类结果对比(出处:论文表1)

模型 参数量(M) Top-1准确率(%)
MobileNetv1 4.2 70.6
MobileNetv2 3.5 72.0
MobileNetv3-Large 5.4 75.2
MNASNet 4.9 75.2
EfficientNet-B0 5.3 76.3
DeiT-Tiny 5.7 72.2
MobileViT-XXS 1.3 69.0
MobileViT-XS 2.3 74.7
MobileViT-S 5.6 78.4

分析

  • MobileViT-S以5.6M参数实现了78.4%的准确率,超过了所有同等参数的CNN和ViT模型
  • 比MobileNetv3-Large高3.2%,比EfficientNet-B0高2.1%
  • 比DeiT-Tiny高6.2%,证明了MobileViT的设计比纯ViT更适合移动端
    在这里插入图片描述
图片3:ImageNet准确率vs参数量对比(出处:论文图3)

4.2 MS-COCO目标检测

表格2:MS-COCO目标检测结果对比(出处:论文表2)

骨干网络 参数量(M) mAP@0.5
MobileNetv1 4.3 22.1
MobileNetv2 4.5 22.3
MobileNetv3 4.9 22.4
MNASNet 4.9 22.3
MobileViT-XS 2.7 24.2
MobileViT-S 5.9 27.7

分析

  • MobileViT-XS以2.7M参数实现了24.2的mAP@0.5,比MobileNetv3高1.8%,同时参数量还小45%
  • MobileViT-S以5.9M参数实现了27.7的mAP@0.5,比同等参数的CNN高5%以上
  • 证明了MobileViT作为通用骨干网络的强大能力,不仅适合分类,也适合检测等下游任务

4.3 语义分割

表格3:PASCAL VOC 2012语义分割结果对比(出处:论文表3)

骨干网络 参数量(M) mIOU
MobileNetv2 5.8 70.7
MobileViT-S 5.9 75.3

分析

  • MobileViT-S以几乎相同的参数量,比MobileNetv2高4.6%的mIOU
  • 证明了MobileViT的全局建模能力对语义分割任务特别有帮助,因为分割需要同时考虑局部细节和全局上下文

4.4 移动端推理速度

表格4:iPhone12上的推理延迟对比(出处:论文表4)

模型 参数量(M) 延迟(ms)
MobileNetv3-Large 5.4 8.2
MobileViT-XS 2.3 7.8
MobileViT-S 5.6 11.3

分析

  • MobileViT-XS的延迟比MobileNetv3-Large还低,同时准确率高2.5%
  • MobileViT-S的延迟略高,但准确率高3.2%,在很多应用场景下是可以接受的
  • 证明了MobileViT不仅在理论上高效,在实际移动端设备上也能流畅运行

深入理解:MobileViT到底在学什么

为了理解MobileViT的工作原理,论文进行了一系列可视化分析。

5.1 注意力可视化

在这里插入图片描述

图片5:MobileViT的注意力图(出处:论文图5)

分析

  • MobileViT的注意力几乎完全集中在物体的主体部分,比如猫的脸、狗的身体、鸟的翅膀
  • 这和人类的视觉注意力很像,说明MobileViT确实学会了识别图像中的语义相关区域
  • 不同于CNN的感受野是逐层扩大的,MobileViT在低层就能关注到全局信息

5.2 感受野分析

在这里插入图片描述

图片6:MobileViT的感受野(出处:论文图6)

分析

  • MobileViT的感受野比同等深度的CNN大得多
  • 在第3层,MobileViT的感受野就已经覆盖了整个图像,而ResNet-50需要到第10层以上才能达到
  • 这解释了为什么MobileViT能用更少的层达到更好的性能

核心代码实现

下面是PyTorch实现的MobileViT Block和完整的MobileViT模型,代码经过简化和注释,方便理解。

import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange, repeat

# ================================
# MobileNetv2 倒残差块 (MV2)
# ================================
class MV2Block(nn.Module):
    def __init__(self, in_channels, out_channels, stride=1, expansion_factor=4):
        super().__init__()
        self.stride = stride
        hidden_dim = int(in_channels * expansion_factor)
        
        # 主分支
        self.conv = nn.Sequential(
            # 1x1 升维
            nn.Conv2d(in_channels, hidden_dim, kernel_size=1, bias=False),
            nn.BatchNorm2d(hidden_dim),
            nn.SiLU(),
            # 3x3 深度可分离卷积
            nn.Conv2d(hidden_dim, hidden_dim, kernel_size=3, stride=stride, padding=1, groups=hidden_dim, bias=False),
            nn.BatchNorm2d(hidden_dim),
            nn.SiLU(),
            # 1x1 降维
            nn.Conv2d(hidden_dim, out_channels, kernel_size=1, bias=False),
            nn.BatchNorm2d(out_channels),
        )
        
        # 残差连接
        self.shortcut = nn.Sequential()
        if stride == 1 and in_channels != out_channels:
            self.shortcut = nn.Sequential(
                nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False),
                nn.BatchNorm2d(out_channels),
            )
    
    def forward(self, x):
        return self.conv(x) + self.shortcut(x) if self.stride == 1 else self.conv(x)

# ================================
# Transformer 编码器层
# ================================
class TransformerEncoder(nn.Module):
    def __init__(self, dim, heads, mlp_ratio=4.0, dropout=0.0):
        super().__init__()
        self.norm1 = nn.LayerNorm(dim)
        self.attn = nn.MultiheadAttention(dim, heads, dropout=dropout, batch_first=True)
        self.norm2 = nn.LayerNorm(dim)
        
        mlp_dim = int(dim * mlp_ratio)
        self.mlp = nn.Sequential(
            nn.Linear(dim, mlp_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(mlp_dim, dim),
            nn.Dropout(dropout)
        )
    
    def forward(self, x):
        # 多头自注意力 + 残差
        x_norm = self.norm1(x)
        attn_out, _ = self.attn(x_norm, x_norm, x_norm)
        x = x + attn_out
        
        # MLP + 残差
        x_norm = self.norm2(x)
        mlp_out = self.mlp(x_norm)
        x = x + mlp_out
        
        return x

# ================================
# MobileViT Block (核心创新)
# ================================
class MobileViTBlock(nn.Module):
    def __init__(self, in_channels, transformer_dim, num_heads, num_transformer_blocks=2, patch_size=2):
        super().__init__()
        self.patch_size = patch_size
        self.transformer_dim = transformer_dim
        
        # 步骤1:局部特征提取
        self.local_rep = nn.Sequential(
            nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1, groups=in_channels, bias=False),
            nn.BatchNorm2d(in_channels),
            nn.SiLU(),
            nn.Conv2d(in_channels, transformer_dim, kernel_size=1, bias=False),
            nn.BatchNorm2d(transformer_dim),
            nn.SiLU(),
        )
        
        # 步骤2:全局Transformer处理
        self.transformer = nn.Sequential(*[
            TransformerEncoder(transformer_dim, num_heads)
            for _ in range(num_transformer_blocks)
        ])
        
        # 步骤3:特征融合
        self.fusion = nn.Sequential(
            nn.Conv2d(transformer_dim, in_channels, kernel_size=1, bias=False),
            nn.BatchNorm2d(in_channels),
            nn.SiLU(),
        )
    
    def forward(self, x):
        B, C, H, W = x.shape
        p = self.patch_size
        
        # 确保H和W能被patch_size整除
        assert H % p == 0 and W % p == 0, f"Image size ({H}, {W}) must be divisible by patch size {p}"
        
        # 步骤1:局部特征提取
        x_local = self.local_rep(x)  # [B, d, H, W]
        
        # 步骤2:展开为patch序列
        x_patch = rearrange(x_local, 'b d (h p1) (w p2) -> b (h w) (p1 p2 d)', p1=p, p2=p)  # [B, N, p*p*d]
        
        # Transformer全局建模
        x_global = self.transformer(x_patch)  # [B, N, p*p*d]
        
        # 折叠回特征图形状
        x_global = rearrange(x_global, 'b (h w) (p1 p2 d) -> b d (h p1) (w p2)', h=H//p, w=W//p, p1=p, p2=p)  # [B, d, H, W]
        
        # 步骤3:特征融合 + 残差连接
        x_fusion = self.fusion(x_global)  # [B, C, H, W]
        out = x + x_fusion
        
        return out

# ================================
# 完整 MobileViT 模型
# ================================
class MobileViT(nn.Module):
    def __init__(self, img_size=256, in_channels=3, num_classes=1000, 
                 dims=[16, 32, 64, 96, 128, 160, 640],
                 depths=[2, 4, 3],
                 num_heads=[2, 4, 8],
                 patch_size=2):
        super().__init__()
        
        # Stem
        self.stem = nn.Sequential(
            nn.Conv2d(in_channels, dims[0], kernel_size=3, stride=2, padding=1, bias=False),
            nn.BatchNorm2d(dims[0]),
            nn.SiLU(),
        )
        
        # Stage 1: MV2 blocks
        self.stage1 = nn.Sequential(
            MV2Block(dims[0], dims[1], stride=1),
            MV2Block(dims[1], dims[2], stride=2),
        )
        
        # Stage 2: MV2 + MobileViT blocks
        self.stage2 = nn.Sequential(
            MV2Block(dims[2], dims[3], stride=2),
            *[MobileViTBlock(dims[3], dims[3], num_heads[0], patch_size=patch_size) for _ in range(depths[0])],
        )
        
        # Stage 3: MV2 + MobileViT blocks
        self.stage3 = nn.Sequential(
            MV2Block(dims[3], dims[4], stride=2),
            *[MobileViTBlock(dims[4], dims[4], num_heads[1], patch_size=patch_size) for _ in range(depths[1])],
        )
        
        # Stage 4: MV2 + MobileViT blocks
        self.stage4 = nn.Sequential(
            MV2Block(dims[4], dims[5], stride=2),
            *[MobileViTBlock(dims[5], dims[5], num_heads[2], patch_size=patch_size) for _ in range(depths[2])],
        )
        
        # Head
        self.head = nn.Sequential(
            nn.Conv2d(dims[5], dims[6], kernel_size=1, bias=False),
            nn.BatchNorm2d(dims[6]),
            nn.SiLU(),
            nn.AdaptiveAvgPool2d(1),
            nn.Flatten(),
            nn.Linear(dims[6], num_classes),
        )
        
        # 初始化权重
        self.apply(self._init_weights)
    
    def _init_weights(self, m):
        if isinstance(m, nn.Conv2d) or isinstance(m, nn.Linear):
            nn.init.trunc_normal_(m.weight, std=0.02)
            if m.bias is not None:
                nn.init.constant_(m.bias, 0)
        elif isinstance(m, nn.BatchNorm2d) or isinstance(m, nn.LayerNorm):
            nn.init.constant_(m.bias, 0)
            nn.init.constant_(m.weight, 1.0)
    
    def forward(self, x):
        x = self.stem(x)
        x = self.stage1(x)
        x = self.stage2(x)
        x = self.stage3(x)
        x = self.stage4(x)
        x = self.head(x)
        return x

# ================================
# 测试代码:创建 MobileViT-S 并验证
# ================================
if __name__ == '__main__':
    # 创建 MobileViT-S 模型
    model = MobileViT(
        img_size=256,
        dims=[16, 32, 64, 96, 128, 160, 640],
        depths=[2, 4, 3],
        num_heads=[2, 4, 8],
        num_classes=1000
    )
    
    # 生成随机输入
    x = torch.randn(2, 3, 256, 256)
    
    # 前向传播
    logits = model(x)
    
    print("输入图像 shape:", x.shape)
    print("输出 logits shape:", logits.shape)
    print("\n✅ MobileViT-S 完整运行成功!")
    print(f"模型参数量: {sum(p.numel() for p in model.parameters())/1e6:.2f}M")

结论与展望

MobileViT的提出是移动端视觉领域的一个里程碑。它证明了Transformer不仅能在云端大显身手,也能在移动端高效运行

MobileViT的核心贡献在于:

  1. 提出了"Transformer即卷积"的全新视角,从计算范式层面重构了卷积和Transformer的结合方式
  2. 设计了MobileViT Block,既保留了CNN的空间归纳偏置和计算效率,又获得了ViT的全局建模能力
  3. 打造了一系列轻量级模型,在多个任务上显著超越了同等参数的CNN和ViT模型

未来,MobileViT还有很多值得探索的方向:

  • 更高效的注意力机制,进一步降低计算量
  • 更好的自监督预训练方法,减少对标注数据的依赖
  • 针对移动端硬件的进一步优化,提高推理速度
  • 扩展到更多的视觉任务,如视频理解、图像生成等
Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐