Fast Vision Mamba: Pooling Spatial Dimensions for Accelerated Processing

TL;DR

FastVim以交替空间平均池化将SSM并行步数减半,2048²图像推理提速72.5%。

cs.CV 🟡 进阶级 2025-02-02 25 次浏览
Saarthak Kapse Robin Betz Srinivasan Sivanandan
FastVim Vision Mamba 状态空间模型 空间池化 高分辨率视觉

核心发现

方法论

FastVim基于Vision Mamba(Vim),在1D卷积后将二维token网格沿行或列做无参数平均池化,把h×w序列压缩为h或w个token。选择性扫描仍使用Mamba的输入依赖参数B、C、Δ,并在扫描后repeat恢复网格;网络块之间交替转置与池化方向,使跨行、跨列信息逐步传播。

关键结果

  • 在ImageNet-1K监督训练中,FastVim-T/S/B分别取得75.4%、81.1%、82.6% Top-1;对应Vim为76.1%、80.5%、经LayerNorm稳定化后的82.6%。FastVim-S在保持精度的同时,FLOPs为4.43G,低于Vim-S的5.9G。
  • 在H100、batch size 128、float32测试中,2048×2048图像的SSM层获得324%相对加速,整体推理最高提速72.5%。FastVim-T在224和2048分辨率分别减少35%和38.5% FLOPs;1024分辨率以上速度超过ViT。
  • MAE预训练的FastMaskVim在ImageNet-1K上达到Base 83.0%、Large 84.9%、Huge 86.1%,448×448达到86.7%。JUMP-CP上FastChannelVim为73.6%(patch/8),高于ChannelViT的74.8%则未超过;ChannelVim为83.0%,FastChannelVim为83.1%。

研究意义

论文针对高分辨率视觉中token数量随面积增长、Mamba扫描仍形成瓶颈的问题,展示了稀疏而有规律的上下文交互可以维持性能。它为病理、显微镜和卫星图像提供了比全token扫描更高吞吐、较低显存的路线,也说明Mamba与Transformer在上下文压缩上的适配性不同。其价值不仅是单一模型加速,还在于池化模块可嵌入VMamba、MambaVision等架构。

技术贡献

核心工程贡献是把SSM扫描长度从L=h²降为h,使并行扫描由log(h²)=2log(h)变为log(h)。算法1明确了reshape、pool、输入依赖投影、离散化、SSM scan、repeat和skip connection流程。FastMaskVim通过稀疏索引和按固定宽度归一化处理遮罩网格;FastChannelVim则兼容每通道token化、排序HCS及不同扫描顺序。

新颖性

与ViT中直接减少token或稀疏注意力不同,FastVim不改变块间token数量,也不引入可学习池化参数,而是在Mamba上下文扫描前压缩、扫描后复制。交替行列池化是关键:单一方向会永久削弱对应方向交互,交替设计使多层网络恢复二维信息传播。

局限性

  • repeat会把同一池化表示广播给多个位置,短期内损失细粒度空间差异;方法依赖多层交替传播,浅层网络或极端纹理任务可能受损。
  • 论文主要报告Vim及其扩展,未系统证明在所有VMamba、MambaVision或非规则大规模数据上的收益;池化和repeat虽几乎不增加FLOPs,却有实际内存访问开销。

未来方向

未来可研究可学习或内容自适应池化、局部残差细节通道,以及更适合非方形和稀疏网格的扫描。还需在更多高分辨率分割、检测、病理MIL和多通道遥感数据上验证,并系统比较不同Mamba骨干、混合精度和硬件实现。

AI 总览摘要

视觉模型正面对一个越来越现实的难题:图像分辨率升高时,token数量按面积增长。Vision Transformer的自注意力具有二次复杂度;Vision Mamba(Vim)借助选择性状态空间模型和并行扫描,将交互成本降为线性,但扫描长度仍等于全部token数。论文作者指出,2048×2048图像会令这一瓶颈重新显现。

FastVim的做法十分简洁:在1D卷积后,把二维token网格沿列或行做平均池化,再用Mamba选择性扫描压缩后的序列,最后repeat回原网格。每个块转置网格并交替池化方向,因此信息不会永远局限在单一方向。这样,正方形网格的扫描长度从h²变为h,并行步骤从2log₂h降至log₂h。

结果显示,FastVim在ImageNet-1K上达到Tiny 75.4%、Small 81.1%、Base 82.6% Top-1,基本保持Vim精度;2048²图像上整体推理最高提速72.5%,SSM部分相对加速324%。MAE预训练的FastMaskVim在448分辨率达到86.7%,JUMP-CP的FastChannelVim达到73.6%。研究表明,简单的结构化信息压缩即可显著缓解高分辨率Mamba瓶颈,但其细节损失、硬件开销及跨架构泛化仍需继续研究。

深度分析

研究背景

S4通过结构化矩阵加速状态空间序列建模,Mamba进一步让B、C和Δ依赖输入,实现选择性扫描。Vim将Mamba扩展到图像,使用前向与后向SSM及1D卷积;相比ViT,避免注意力二次token交互。但图像token数L=H×W/P²,分辨率增加会使SSM扫描并行深度达到log(L),高分辨率吞吐仍受限。

核心问题

论文要解决的不是参数量,而是上下文扫描的实际步数。对h×h token网格,Vim需处理h²个位置,并行深度为log(h²)=2log(h)。如果只减少token,可能破坏空间关系;若仅依靠并行scan,长序列的同步、显存和内存访问仍成为瓶颈。

核心创新

第一,FastVim在扫描前沿一个空间维度做mean pooling,把h×w压缩为h或w。第二,块间交替转置和池化方向,使行、列信息都能传播。第三,FastMaskVim用稀疏索引适配MAE、DINOv2和病理非规则网格。第四,FastChannelVim适配ChannelViT式每通道token化,并处理扫描顺序与hierarchical channel sampling。

方法详解

  • �� 输入x∈R^(B×L×D),reshape为(B,h,w,D)。
  • �� 经过Norm、扩展层和Conv1D后,沿列或行平均池化,得到(B,h,1,D)。
  • �� 由Linear_N生成输入依赖B、C,由softplus(Parameter+sΔ(x))生成Δ;再按零阶保持离散化A、B。
  • �� 对压缩序列执行Forward/Backward Mamba selective scan。
  • �� 将输出沿池化维repeat,接skip connection和LayerNorm;下一块转置网格,交替方向。
  • �� FastMaskVim按非掩码token索引转置并按固定w归一化;FastChannelVim在空间或通道优先序列上扫描。

实验设计

ImageNet-1K含1.28M训练图像和50K验证图像;监督训练300 epochs、AdamW、batch 1024、初始学习率1e-3、5 epoch warmup、weight decay 0.05。MAE预训练1600 epochs、mask ratio 0.75。效率测试在H100、batch 128、float32上比较FastVim、Vim和ViT。JUMP-CP使用BR00116991板,127K/45K/45K训练、验证、测试图像及8通道数据,评估160类扰动预测。

结果分析

FastVim-T/S/B的Top-1为75.4/81.1/82.6%,FLOPs为1.17/4.43/17.23G。FastVim-T相比Vim-T在224分辨率少35% FLOPs,在2048少38.5%。2048分辨率时SSM层加速324%,整体加速72.5%。FastMaskVim-Huge在448图像达到86.7%。JUMP-CP patch/8上ChannelVim与FastChannelVim分别83.0%和83.1%,明显高于ChannelViT的74.8%。

应用场景

直接应用包括高分辨率分类、目标检测、语义或实例分割、MAE/DINOv2表征学习、病理切片MIL、显微镜细胞扰动预测及卫星多通道成像。使用者需要规则或可索引的token网格、支持Mamba scan的实现和足够显存;非规则网格应采用FastMaskVim,多通道数据可采用FastChannelVim。

局限与展望

池化后repeat会抹平同一行或列内的差异,性能依赖跨块交替传播;单块或浅层模型可能难以恢复局部细节。池化并不显著降低所有模块的FLOPs,MLP和gating仍按原token数运行,且repeat带来内存访问开销。论文还显示该策略在ViT上失败,说明其收益依赖Mamba的递归上下文机制;未来需扩大跨架构、任务和硬件验证。

通俗解读 非专业人士也能看懂

把一张超高清图片想成一座巨大的仓库,每个小格子都是一个观察员。原来的Vim让所有观察员排成长队,依次把消息传给下一个人;虽然可以用很多人同时推进,但队伍太长时仍然慢。FastVim先让同一列或同一行的观察员开一个短会,把许多相似消息合成一份,再让短队伍快速传递,最后把结果发回原来的格子。

如果每次都只开列会,横向信息会缺失;所以下一层改开行会,再下一层开列会。几轮之后,每个格子都能间接听到各个方向的消息。这个办法不增加新的可训练零件,像是重新安排会议,而不是雇更多人。

代价是短会会丢掉一些细节:同一行中两个完全不同的小物体可能被合成一个平均印象。因此它特别适合需要整体背景和长距离联系的任务,高度依赖细粒度边界的任务仍需额外细节通道。实验中,2048×2048图片推理最多快72.5%,说明合理压缩能让超高清处理更实用。

简单解释 像给14岁少年讲一样

想象你在玩一款超大地图游戏。地图上每个小格子都有一个小机器人,机器人要把看到的信息传给别的机器人。地图越清晰,小格子越多,信息传递就越慢。Vision Mamba已经比某些传统方法省事,但超大地图还是会卡。

FastVim像是让一排机器人先开一个小队会议:大家把消息合成一份,再让这份消息快速穿过队伍。会议结束后,结果再发回每个机器人。下一轮不再开同样方向的会,而是换成另一方向。这样,机器人不会只知道自己的横向邻居,也能慢慢获得纵向和全地图的信息。

关键是它没有把地图永久缩小,只是在传递消息时暂时压缩,所以模型仍然输出原来的格子数量。实验很酷:在ImageNet-1K上,FastVim-Base达到82.6%准确率;处理2048×2048图片时,整体推理最多快72.5%。这就像游戏画面保持清晰,但后台通信更高效。

当然,平均消息也可能漏掉小细节,比如草丛里的一只小猫。FastVim更擅长理解大范围背景,未来可以加入专门保存细节的办法。它还能帮助分析细胞显微图、病理图和卫星图像,特别是那些通道很多、图片特别大的场景。

术语表

State Space Model(状态空间模型)

用隐藏状态记录历史信息,并通过状态转移产生输出的序列模型。其基本形式为h_t=Ah_{t-1}+Bx_t,y_t=Ch_t+Dx_t。

FastVim使用Mamba形式的SSM作为视觉上下文模块。

Selective Scan(选择性扫描)

让B、C和Δ由输入token动态决定的扫描机制,可根据内容选择保留或更新信息。

FastVim只对池化后的token执行选择性扫描。

Parallel Scan(并行扫描)

把递归计算从L个串行步骤组织为约log(L)层并行操作。

池化使Vim的2log(h)并行深度降为log(h)。

Mean Pooling(平均池化)

沿一个空间维度求token平均值,以较少表示概括多个位置。它不引入额外参数。

FastVim默认在Conv1D后使用平均池化。

MAE(掩码自编码器)

遮住输入的大部分patch,再训练编码器和解码器重建缺失内容的自监督方法。

FastMaskVim在ImageNet-1K上采用75% masking ratio预训练。

Per-channel Tokenization(逐通道token化)

把不同成像通道分别形成token,而不是将所有通道合成一个token。它能保留显微镜和卫星数据的互补信息。

FastChannelVim用于JUMP-CP的8通道细胞图像。

开放问题 这项研究留下的未解疑问

  • 1 池化后repeat造成的细节损失如何量化?目前尚不清楚哪些目标尺寸、纹理频率或层数最适合稀疏交互。
  • 2 该方法为何在Vim有效、在ViT失败?需要理论分析Mamba递归状态与Transformer集合式注意力的上下文差异。
  • 3 真实部署中的收益受GPU内存访问、混合精度和算子融合影响,仍需更多硬件及端到端基准。

应用场景

近期应用

超高分辨率医学影像

病理切片或显微图像可用FastMaskVim处理掩码和空洞网格,用更短扫描序列降低推理延迟。部署前需实现稀疏索引、验证边界细节,并针对分割任务微调。

多通道细胞分析

药物扰动筛选平台可使用FastChannelVim保留8个以上成像通道的信息,在JUMP-CP上获得73.6%准确率。适合大规模筛选,但需固定通道顺序并处理HCS。

远期愿景

实时卫星与工业视觉

在高分辨率遥感、缺陷检测和机器人视觉中,FastVim可能降低端到端延迟与显存需求。实现价值取决于更高效的pool-repeat算子、局部细节保真和跨硬件优化。

原文摘要

State Space Models (SSMs) with selective scan (Mamba) have been adapted into efficient vision models. Mamba, unlike Vision Transformers, achieves linear complexity for token interactions through a recurrent hidden state process. This sequential processing is enhanced by a parallel scan algorithm, which reduces the computational time of recurrent steps from $L$ sequential steps to $log(L)$ parallel steps with respect to the number of input tokens ($L$). In this work, we propose Fast Vision Mamba (FastVim), that further reduces the computational time of the SSM block by reducing the number of recurrent steps in Vision Mamba models while still retaining model performance. By alternately pooling tokens along image dimensions across Mamba blocks, we obtain a 2$\times$ reduction in the number of parallel steps in SSM block. Our model offers up to $72.5\%$ speedup in inference speed compared to baseline Vision Mamba models on high resolution (2048$\times$2048) images. Our experiments demonstrate state-of-the-art performance with dramatically improved throughput in a range of tasks such as image classification, cell perturbation prediction, segmentation, and object detection. Code is made available at https://github.com/insitro/FastVim

cs.CV cs.AI