引言:
本教程旨在让零基础的小白也能成功搭建mamba环境利用CUDA加速,也就是GPU版本

由于官方至今只发布了liunx版的mamba的whl,然后我在Ubuntu下安装了之后由于Vmware不能GPU直通,然后mamba要求必须使用CUDA,所以我在该系统版本下并没有成功,于是就只有在windows下寻找方法,然后就有了这篇教程,因为windows的mamba-ssm相关的编译包已经有大佬做出来了,所以经过我的尝试它们的whl来进行mamba环境搭建,并且成功运行测试mamba可以成功使用。
Ps:mamba相关的whl链接放在文章的最后了大家可以免费获取

环境:
系统:Windows11

工具:Anaconda
显卡:NVIDIA4060(显卡要求需要带CUDA)
废话不多少,直接开始!

第一步:创建anconda环境方便管理

win+R输入cmd进入shell,执行以下命令创建名为mamba的虚拟环境,指定python版本为3.10

conda create -n mamba python=3.10

第二步:进入创建好的虚拟环境安装mamba相关的包,以及其他基础环境

为了更方便的安装whl,直接找到whl所在目录,然后在该目录下输入cmd进入shell,或者在shell中cd到whl所在目录

激活我们刚刚创建的环境

conda activate mamba

左侧显示mamba虚拟环境名称就说明激活成功

现在依次安装下面的包
1.cudatoolkit (GPU加速环境)
2.torch + torchvision + torchaudio (深度学习主框架)
3.setuptools (负责正确安装底层扩展)
4.packaging (处理版本兼容性)
依次执行以下4个命令

conda install cudatoolkit==11.8  
pip install torch==2.1.1 torchvision==0.16.1 torchaudio==2.1.1 --index-url https://download.pytorch.org/whl/cu118
pip install setuptools==68.2.2
conda install packaging

现在介绍一下whl相关的文件,链接里面的文件如下,Readme有相关安装指导

mamba_ssm包是实现 Mamba 状态空间模型(Mamba SSM) 的 Python 库

Mamba 是一种用于序列建模的架构,旨在替代 Transformer,实现更高效的长序列处理。
主要功能包括:

  • 状态空间建模(State Space Models)

  • 高效并行计算内核(通常依赖 CUDA)

  • 推理和训练接口(PyTorch)

causal_conv1d实现了 一维因果卷积(Causal 1D Convolution),是 Mamba 模型在进行状态更新时用到的核心算子。

  • 用于高效实现时间序列的前向传播;

  • 替代标准卷积,以保证“因果性”(即不能看到未来输入);

  • 在 GPU 上优化过,性能很关键。

Ps: 没有这个包,mamba_ssm 运行时会报错(例如 ImportError: No module named 'causal_conv1d'

TritonOpenAI 出品的高性能 GPU kernel 编译框架
它允许开发者用 Python 写出比 PyTorch 更底层、但仍然简洁的 GPU 计算逻辑。

在这里,mamba_ssmcausal_conv1d 很可能使用 Triton 内核 来实现高效的 GPU 计算。

简单来说:

  • triton:提供 GPU 编译与加速的“引擎”;

  • causal_conv1d:定义底层的卷积算子(用 Triton 实现);

  • mamba_ssm:在高层调用这些算子,实现完整的 Mamba 模型。

为了保证安装成功建议按照我顺序来进行这3个包的安装
注意:先检查是否当前目录存在这些文件,如果没有这些文件可以返回第二步开头并激活环境再执行命令即可

先安装 triton

pip install triton-2.0.0-cp310-cp310-win_amd64.whl

再安装 causal_conv1d

pip install causal_conv1d-1.1.1-cp310-cp310-win_amd64.whl

最后安装 mamba_ssm

pip install mamba_ssm-1.1.3-cp310-cp310-win_amd64.whl

最后一步,验证安装是否有效

切换到conda的mamba虚拟环境,然后运行以下test.py代码可查看安装情况

import torch

print("检查 PyTorch GPU 状态...")
print(f"PyTorch 版本: {torch.__version__}")
print(f"CUDA 是否可用: {torch.cuda.is_available()}")
if torch.cuda.is_available():
    print(f"GPU 名称: {torch.cuda.get_device_name(0)}")
    print(f"CUDA 版本: {torch.version.cuda}")
    print(f"cuDNN 版本: {torch.backends.cudnn.version()}")

try:
    from mamba_ssm import Mamba
    print("\nMamba 模块导入成功。")
except ImportError:
    print("\nMamba 模块未安装或导入失败,请检查 mamba_ssm 是否正确安装。")
    exit()

try:
    print("\n正在测试 Mamba 在 GPU 上运行...")
    model = Mamba(d_model=64, d_state=16, d_conv=4, expand=2)
    model = model.cuda() if torch.cuda.is_available() else model
    x = torch.randn(1, 16, 64).cuda() if torch.cuda.is_available() else torch.randn(1, 16, 64)
    y = model(x)
    print(f"运行成功,输出张量形状: {y.shape}")
except Exception as e:
    print(f"\nMamba GPU 运行失败: {e}")

运行成功结果

到这里恭喜你大功告成是不是很简单啊!!!

常见报错解决方法

1.module 'torch.utils._pytree' has no attribute 'register_pytree_node'. Did you mean: '_register_pytree_node'?

原因:说明 PyTorch 与 Transformers(transformers 库)版本不兼容
解决方法:降低 transformers 版本 到mamba这个虚拟环境中执行下面的命令

# 卸载新版 transformers
pip uninstall transformers -y
# 安装兼容版本
pip install transformers==4.36.2 

2.其他错误可尝试numpy版本是否兼容,对numpy进行降级

pip install "numpy<2.0" 
#或者指定版本
pip install numpy==1.26.4

Ps:有其他问题可在评论区讨论哦~~~ zpa666~~~
whl链接
https://pan.baidu.com/s/1YQPN2fELYv_larfZtQLHIg?pwd=xqct 提取码: xqct 
PS:如果是50系的本文mamba版本暂不支持,大家可以去看该博主的解决方法,https://blog.csdn.net/yyywxk/article/details/146798627#t13

Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐