VMamba:基于Mamba的计算机视觉通用骨干网络项目

VMamba: Visual State Space Models,code is based on mamba

分支2Tags15
文件最后提交记录最后更新时间
2 年前
2 年前
2 年前
2 年前
2 年前
2 年前
2 年前
2 年前
2 年前
2 年前
2 年前
2 年前

VMamba

VMamba: 视觉状态空间模型

Yue Liu1,Yunjie Tian1,Yuzhong Zhao1, Hongtian Yu1, Lingxi Xie2, Yaowei Wang3, Qixiang Ye1, Yunfan Liu1

1 中国科学院大学, 2 华为公司, 3 鹏城实验室.

论文: (arXiv 2401.10166)

✅ 更新

  • 2024年6月14日: 更新: 我们清理了代码,使其更易于阅读;我们增加了对 mamba2 的支持。
  • 2024年5月26日: 更新: 我们发布了 VMambav2 的更新权重,以及新的 arxiv 论文。
  • 2024年5月7日: 更新: 重要! 在下游任务中使用 torch.backends.cudnn.enabled=True 可能会非常慢。如果你发现 vmamba 在你的机器上运行很慢,请在 vmamba.py 中禁用它,否则忽略此信息。
  • ...

详情请见 detailed_updates.md

摘要

设计计算高效的网络架构在计算机视觉中持续成为必要。本文中,我们将 Mamba,一种状态空间语言模型,移植到 VMamba,一种线性时间复杂度的视觉骨干网络。VMamba 的核心是一组视觉状态空间 (VSS) 块,其中包含 2D 选择性扫描 (SS2D) 模块。通过沿四条扫描路径遍历,SS2D 帮助弥合了 1D 选择性扫描的有序性与 2D 视觉数据的非顺序结构之间的差距,从而促进了从各种来源和角度收集上下文信息。基于 VSS 块,我们开发了一系列 VMamba 架构,并通过一系列架构和实现优化加速它们。广泛的实验展示了 VMamba 在各种视觉感知任务中的出色表现,突显了其在输入缩放效率方面的优势,优于现有的基准模型。

概述

  • VMamba 作为计算机视觉的通用骨干网络。

architecture

  • VMamba 的 2D 选择性扫描

arch

  • VMamba 具有全局有效感受野

erf

  • VMamba 在激活图上类似于基于 Transformer 的方法

attn

activation

主要结果

📖 详情请见 performance.md.

ImageNet-1K 上的分类

名称 预训练 分辨率 acc@1 参数数量 FLOPs TP. 训练 TP. 配置/日志/检查点
Swin-T ImageNet-1K 224x224 81.2 28M 4.5G 1244 987 --
Swin-S ImageNet-1K 224x224 83.2 50M 8.7G 718 642 --
Swin-B ImageNet-1K 224x224 83.5 88M 15.4G 458 496 --
VMamba-S[s2l15] ImageNet-1K 224x224 83.6 50M 8.7G 877 314 配置/日志/检查点
VMamba-B[s2l15] ImageNet-1K 224x224 83.9 89M 15.4G 646 247 配置/日志/检查点
VMamba-T[s1l8] ImageNet-1K 224x224 82.6 30M 4.9G 1686 571 配置/日志/检查点
  • 本节中的模型是从头开始训练的,使用随机或手动初始化。超参数继承自 Swin,除了 drop_path_rateEMA。所有模型都使用 EMA 训练,除了 Vanilla-VMamba-T
  • TP.(吞吐量)训练 TP.(训练吞吐量) 在 A100 GPU 和 AMD EPYC 7542 CPU 上评估,批量大小为 128。训练 TP. 使用混合分辨率测试,不包括优化器的时间消耗。
  • FLOPs参数 现在包含 head(在之前的版本中,它们不包括 head,因此数字略有上升)。
  • 我们使用 @albertgu 提供的算法 计算 FLOPs,这将比之前的计算(基于 selective_scan_ref 函数,忽略了硬件感知算法)更大。

COCO 上的目标检测

骨干网络 参数数量 FLOPs 检测器 bboxAP bboxAP50 bboxAP75 segmAP segmAP50 segmAP75 配置/日志/检查点
Swin-T 48M 267G MaskRCNN@1x 42.7 65.2 46.8 39.3 62.2 42.2 --
Swin-S 69M 354G MaskRCNN@1x 44.8 66.6 48.9 40.9 63.4 44.2 --
Swin-B 107M 496G MaskRCNN@1x 46.9 -- -- 42.3 -- -- --
VMamba-S[s2l15] 70M 384G MaskRCNN@1x 48.7 70.0 53.4 43.7 67.3 47.0 配置/日志/检查点
VMamba-B[s2l15] 108M 485G MaskRCNN@1x 49.2 71.4 54.0 44.1 68.3 47.7 配置/日志/检查点
VMamba-B[s2l15] 108M 485G MaskRCNN@1x[bs8] 49.2 70.9 53.9 43.9 67.7 47.6 配置/日志/检查点
VMamba-T[s1l8] 50M 271G MaskRCNN@1x 47.3 69.3 52.0 42.7 66.4 45.9 配置/日志/检查点
:---: :---: :---: :---: :---: :---: :---: :---: :---: :---: :---:
Swin-T 48M 267G MaskRCNN@3x 46.0 68.1 50.3 41.6 65.1 44.9 --
Swin-S 69M 354G MaskRCNN@3x 48.2 69.8 52.8 43.2 67.0 46.1 --
VMamba-S[s2l15] 70M 384G MaskRCNN@3x 49.9 70.9 54.7 44.20 68.2 47.7 配置/日志/检查点
VMamba-T[s1l8] 50M 271G MaskRCNN@3x 48.8 70.4 53.50 43.7 67.4 47.0 配置/日志/检查点
  • 本节中的模型从 分类 中训练的模型初始化。
  • 我们现在使用 @albertgu 提供的算法 计算 FLOPs,这将比之前的计算(基于 selective_scan_ref 函数,忽略了硬件感知算法)更大。

ADE20K 语义分割

骨干网络 输入尺寸 参数量 FLOPs 分割器 mIoU(SS) mIoU(MS) 配置/日志/日志(ms)/模型
Swin-T 512x512 60M 945G UperNet@160k 44.4 45.8 --
Swin-S 512x512 81M 1039G UperNet@160k 47.6 49.5 --
Swin-B 512x512 121M 1188G UperNet@160k 48.1 49.7 --
VMamba-S[s2l15] 512x512 82M 1028G UperNet@160k 50.6 51.2 配置/日志/日志(ms)/模型
VMamba-B[s2l15] 512x512 122M 1170G UperNet@160k 51.0 51.6 配置/日志/日志(ms)/模型
VMamba-T[s1l8] 512x512 62M 949G UperNet@160k 47.9 48.8 配置/日志/日志(ms)/模型
  • 本节中的模型是从 classfication 中训练的模型初始化的。
  • 我们目前使用 @albertgu 提供的算法计算 FLOPs,这会比之前的计算结果更大(之前的计算基于 selective_scan_ref 函数,忽略了硬件感知的算法)。

开始使用

安装

步骤 1: 克隆 VMamba 仓库:

首先,克隆 VMamba 仓库并导航到项目目录:

git clone https://github.com/MzeroMiko/VMamba.git
cd VMamba

步骤 2: 环境设置:

VMamba 推荐通过 conda 创建环境并通过 pip 安装依赖。使用以下命令设置环境: 我们还推荐使用 pytorch>=2.0, cuda>=11.8。但较低版本的 pytorch 和 CUDA 也支持。

创建并激活新的 conda 环境

conda create -n vmamba
conda activate vmamba

安装依赖

pip install -r requirements.txt
cd kernels/selective_scan && pip install .

检查 Selective Scan(可选)

  • 如果你想检查与 mamba_ssm 相比的模块,首先安装 mamba_ssm

  • 如果你想检查我们的 selective scan 实现是否与 mamba_ssm 相同,selective_scan/test_selective_scan.py 可以帮你。在 selective_scan/test_selective_scan.py 中将 MODE = "mamba_ssm_sscore",然后运行 pytest selective_scan/test_selective_scan.py

  • 如果你想检查我们的 selective scan 实现是否与参考代码 (selective_scan_ref) 相同,在 selective_scan/test_selective_scan.py 中将 MODE = "sscore",然后运行 pytest selective_scan/test_selective_scan.py

  • MODE = "mamba_ssm" 用于检查 mamba_ssm 的结果是否接近 selective_scan_ref,而 "sstest" 保留用于开发。

  • 如果你发现 mamba_ssm (selective_scan_cuda) 或 selective_scan (selctive_scan_cuda_core) 与 selective_scan_ref 不够接近,并且测试失败,不用担心。检查 mamba_ssmselective_scan 是否足够接近 instead

  • 如果你对 selective scan 感兴趣,可以查看 mamba, mamba-mini, mamba.py mamba-minimal 了解更多信息。

DetectionSegmentation 的依赖(可选)

pip install mmengine==0.10.1 mmcv==2.1.0 opencv-python-headless ftfy regex
pip install mmdet==3.3.0 mmsegmentation==1.2.2 mmpretrain==1.2.0

模型训练与推理

分类

要在 ImageNet 上训练 VMamba 分类模型,使用以下命令进行不同配置:

python -m torch.distributed.launch --nnodes=1 --node_rank=0 --nproc_per_node=8 --master_addr="127.0.0.1" --master_port=29501 main.py --cfg </path/to/config> --batch-size 128 --data-path </path/of/dataset> --output /tmp

如果你只想测试性能(包括参数和 FLOPs):

python -m torch.distributed.launch --nnodes=1 --node_rank=0 --nproc_per_node=1 --master_addr="127.0.0.1" --master_port=29501 main.py --cfg </path/to/config> --batch-size 128 --data-path </path/of/dataset> --output /tmp --pretrained </path/of/checkpoint>

更多详情请参考 modelcard

检测与分割

使用 mmdetectionmmsegmentation 进行评估:

bash ./tools/dist_test.sh </path/to/config> </path/to/checkpoint> 1

使用 --tta 获取分割中的 mIoU(ms)

使用 mmdetectionmmsegmentation 进行训练:

bash ./tools/dist_train.sh </path/to/config> 8

有关检测和分割任务的更多信息,请参考 mmdetectionmmsegmentation 的手册。记得在 configs 目录中使用适当的骨干网络配置。

分析工具

VMamba 包含用于可视化 mamba "attention" 和有效感受野、分析吞吐量和训练吞吐量的工具。使用以下命令进行分析:

# 可视化 Mamba "Attention"
CUDA_VISIBLE_DEVICES=0 python analyze/attnmap.py

# 分析有效感受野
CUDA_VISIBLE_DEVICES=0 python analyze/erf.py

# 分析吞吐量和训练吞吐量
CUDA_VISIBLE_DEVICES=0 python analyze/tp.py

我们还包含了其他可能在此项目中使用的分析工具。感谢所有为这些工具做出贡献的人。

Star 历史

Star 历史图表

引用

@article{liu2024vmamba,
  title={VMamba: 视觉状态空间模型},
  author={刘越, 田云杰, 赵宇中, 余洪天, 谢凌曦, 王耀伟, 叶启祥, 刘云帆},
  journal={arXiv预印本 arXiv:2401.10166},
  year={2024}
}

致谢

本项目基于Mamba(论文代码),Swin-Transformer(论文代码),ConvNeXt(论文代码),OpenMMLab, 以及analyze/get_erf.py采用了replknet的代码,感谢他们的优秀工作。

  • 我们最近发布了Fast-iTPN,据我们所知,它在ImageNet-1K的Tiny/Small/Base级别模型中报告了最佳性能。(Tiny-24M-86.5%,Small-40M-87.8%,Base-85M-88.75%)

项目介绍

VMamba: Visual State Space Models,code is based on mamba

定制我的领域
153.22 K239访问 GitHub