mamba-rs
Mamba SSM implementation in Rust with optional CUDA GPU acceleration. Supports Mamba-1 and Mamba-3 SISO.
Full inference and training pipelines with BPTT through recurrent SSM state. Custom CUDA kernels with CUDA Graph capture for minimal-latency GPU inference.
Features
- Two architectures — Mamba SSM (Gu & Dao, 2023) and Mamba-3 SISO (Lahoti et al., ICLR 2026)
- CPU inference — zero-allocation single-step recurrent forward pass with SIMD + BLAS
- GPU inference — CUDA kernels with optional CUDA Graph capture (~1.6x speedup)
- CPU training — full backward pass with BPTT, parallel batch training via Rayon
- GPU training — custom CUDA forward + backward kernels (47 for M3, 12 for M1)
- Serialization — safetensors format (HuggingFace compatible)
- Standalone — no framework dependency (no PyTorch, no Burn, no Candle)
- f32 — native single precision, TF32 Tensor Cores on Ampere/Hopper
Quick Start — Mamba SSM
use ;
let cfg = default; // d_model=128, 3 layers
let weights = init;
let mut state = zeros;
let mut scratch = new;
let mut output = vec!;
mamba_step;
state.reset; // episode boundary
Quick Start — Mamba-3 SISO
use Mamba3Config;
use ;
use Mamba3State;
use Mamba3Weights;
let cfg = default;
let weights = init;
let mut state = zeros;
let mut scratch = new;
let mut output = vec!;
mamba3_step;
GPU Inference (CUDA)
[]
= { = "0.2", = ["cuda"] }
Mamba SSM
use GpuMambaBackbone;
let mut gpu = new?;
gpu.capture_graph?; // optional ~2x speedup
gpu.step?;
gpu.reset?;
Mamba-3
use GpuMamba3Backbone;
let mut gpu = new?;
gpu.capture_graph?; // optional ~1.6x speedup
gpu.step?;
gpu.reset?;
Requires NVIDIA GPU + CUDA toolkit. Kernels compiled at runtime via NVRTC.
Weight Serialization
use serialize;
// Mamba-1
save?;
let = load?;
// Mamba-3
use ;
save_mamba3?;
let = load_mamba3?;
Performance (RTX 6000 Ada)
| Mamba SSM | Mamba-3 SISO | |
|---|---|---|
| GPU Inference B=1 (CUDA Graph) | 79 us | 86 us |
| GPU Training Fwd+Bwd (T=32) | 1,640 us | 2,169 us |
| CPU Inference B=1 | 84 us | 65 us |
| CPU Training Fwd+Bwd (T=32) | 15,874 us | 3,609 us |
Zero heap allocations per inference step. See detailed results:
Documentation
- Mamba SSM architecture — pipeline, modular API, weight layout
- Mamba-3 architecture — trapezoidal SSM, RoPE, BCNorm, CUDA kernels
- Mamba SSM benchmarks — GPU/CPU inference + training numbers
- Mamba-3 benchmarks — GPU/CPU inference + training numbers
Citation
License
Dual-licensed under MIT or Apache-2.0.