VMamba: Visual State Space Models,code is based on mamba
| 文件 | 最后提交记录 | 最后更新时间 |
|---|---|---|
| 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 作为计算机视觉的通用骨干网络。
- VMamba 的 2D 选择性扫描
- VMamba 具有全局有效感受野
- VMamba 在激活图上类似于基于 Transformer 的方法
主要结果
📖 详情请见 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_rate和EMA。所有模型都使用 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_ssm和selective_scan是否足够接近 instead。 -
如果你对 selective scan 感兴趣,可以查看 mamba, mamba-mini, mamba.py mamba-minimal 了解更多信息。
Detection 和 Segmentation 的依赖(可选)
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。
检测与分割
使用 mmdetection 或 mmsegmentation 进行评估:
bash ./tools/dist_test.sh </path/to/config> </path/to/checkpoint> 1
使用 --tta 获取分割中的 mIoU(ms)
使用 mmdetection 或 mmsegmentation 进行训练:
bash ./tools/dist_train.sh </path/to/config> 8
有关检测和分割任务的更多信息,请参考 mmdetection 和 mmsegmentation 的手册。记得在 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 历史
引用
@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%)