A from-scratch PyTorch reproduction of Mamba-3 with complex-valued SSD state spaces.
~434M params · 8.0B Chinchilla-optimal tokens · 12–15 h on a single A100 80GB · N=64 complex64 states







Architecture · Headline metric · Quick start · Documentation · References
📖 Overview#
Mamba-3-Lite is a from-scratch PyTorch implementation of the Mamba-3 architecture (Dao & Gu, 2025) at Chinchilla-optimal scale. It succeeds Mamba-2 with three architectural breakthroughs that are implemented end-to-end in pure PyTorch — no mamba-ssm, no custom CUDA kernels (one sanctioned, opt-in Triton kernel covers the SSD hot path, see AGENTS.md §1):
- Complex-Valued SSD state spaces. State dimension is halved (N=128 → N=64) by promoting the recurrence into the complex plane (
complex64). Two real sub-states are packed into one complex state, achieving parity perplexity with Mamba-2 at double the state size. - MIMO (Multi-Input Multi-Output) head mixing. A fully-connected mixer across SSM heads replaces the classical SISO (single-input single-output) constraint, giving the model cross-head communication for free.
- Zero causal convolution. The memory-bound
causal_conv1dpass is eliminated in favor of a purely chunked linear projection — saving memory bandwidth and simplifying the block.
Why does this exist? Mamba-3's complex SSD extension is the key contribution that breaks the "real SSM only" paradigm. This repo implements the algorithm faithfully, tests the math against a naive reference, and benchmarks it on a single A100.
How it compares to the rest of the portfolio#
| Project | Backbone | State | Mixer | Causal conv |
|---|---|---|---|---|
| GPT-2 (From Scratch) | Transformer | — | — | — |
| LLaMA-3-Lite | Transformer + GQA | — | — | — |
| DeepSeek-v3-Lite | MLA + MoE | — | — | — |
| HyMo | GDN + MLA hybrid | real SSM (in GDN blocks) | — | — |
| Mamba-3-Lite | Pure complex SSD | N=64, complex64 | ✅ MIMO | ❌ none |
🏆 Headline metric#
Mamba-3-Lite: 50% smaller complex state (N=64, complex64) achieves parity loss with Mamba-2 at N=128 on the same 8.0B-token Chinchilla run (single A100 80GB, ~10–12 h wall time).
The complex recurrence h_t = exp((A_real + i·A_imag)·dt) · h_{t-1} + (B_real + i·B_imag)·x_t packs two real eigenvalues (one decay, one rotation) into a single complex state, doubling the expressive capacity per parameter. Verified by tests/test_ssd.py::test_chunkwise_matches_naive_complex; the derivation is in docs/concepts/ssd-theory.md.
🗺️ Visual Architecture Atlas#
Explore the full Interactive Visual Systems Guide: three verified Archify showcase maps, live complex-state SSD simulation, interactive parameter calculator, and verification receipts.
<div align="center">
<a href="docs/diagrams/mamba3_interactive_guide.html">
<img src="docs/diagrams/mamba3_architecture.visual-check.1440x900.dark.png" alt="Mamba-3 Architecture Overview" width="100%" style="border-radius: 8px; box-shadow: 0 4px 20px rgba(0,0,0,0.3);" />
</a>
<p><em>Figure 1: Mamba-3-Lite Architecture Map — 28-layer complex SSD state-space model with $N=64$ complex64 states, MIMO head mixing, and chunkwise recurrence. Click image to open interactive guide.</em></p>
</div>
Interactive Architecture & Systems Diagrams#
| Diagram | Description | Interactive HTML | Visual Preview |
|---|---|---|---|
| Complex SSD Architecture | 28-layer Mamba-3 block, complex64 SSD recurrence, MIMO head mixing matrix, and SwiGLU FFN | Open Map ↗ | PNG |
| Data Pipeline | 8.0B-token universal pipeline, GPT-2 BPE tokenizer, binary chunk sharding, and memory-mapped PretrainDataset | Open Map ↗ | PNG |
| Training Workflow | End-to-end pretraining loop, chunked cross-entropy (chunk=4096), AdamW optimizer, and state checkpointing | Open Map ↗ | PNG |
🏗 Architecture#
Input tokens (vocab = 50,257, GPT-2 BPE)
│
▼
Embedding (d_model=1024) ← weight-tied with output head
│
▼
28 × Mamba-3 Blocks (gradient checkpointing enabled globally):
┌──────────────────────────────────────────────────────────────┐
│ RMSNorm → in_proj → Chunkwise SSD (complex64) │
│ → MIMO mixer → out_proj → Residual │
│ RMSNorm → SwiGLU FFN (intermediate=2048) → Residual │
└──────────────────────────────────────────────────────────────┘
│
▼
Final RMSNorm → Linear head → Chunked Cross-Entropy (chunk=4096)Per-block components#
| Component | Spec | Purpose |
|---|---|---|
| Input projection | in_proj: d_model → n_heads × head_dim × 2 (real + imag packed) | One projection instead of separate x/B |
| Complex SSD | N=64, complex64, chunk=64 | State-space scan with complex eigenvalues |
| MIMO mixer | n_heads × head_dim → n_heads × head_dim (fully connected) | Cross-head information flow |
| Output projection | n_heads × head_dim → d_model | Aggregate heads back to model dim |
| FFN | SwiGLU, ffn_dim=2048 (not 4096) | Gated MLP, matches Mamba-2 design |
| Normalization | RMSNorm, pre-norm, eps=1e-5 | |
| Weight tying | Embed ↔ output head | Saves ~52M params |
| Causal conv | None | Pure chunked linear projection |
⚙️ Configuration#
The canonical config is configs/pretrain_a100_400m.yaml:
Model#
| Parameter | Value |
|---|---|
vocab_size | 50,257 (GPT-2 BPE) |
d_model | 1,024 |
n_layers | 28 |
n_heads | 16 (SSM heads) |
head_dim | 64 (D) |
state_dim | 64 (N, complex64) |
chunk_size | 64 (SSD tunable) |
ffn_dim | 2,048 (SwiGLU intermediate) |
max_seq_len | 2,048 |
weight_tying | true |
init_std | 0.02 |
| Total params | ~434M |
Training#
| Parameter | Value |
|---|---|
micro_batch_size | 16 |
gradient_accumulation_steps | 2 |
total_steps | 256,000 (~8.0B tokens) |
warmup_steps | 2,000 (linear) |
lr | 3.0 × 10⁻⁴ |
min_lr_ratio | 0.05 (cosine decay) |
weight_decay | 0.1 |
beta1 / beta2 | 0.9 / 0.95 |
grad_clip | 1.0 |
grad_checkpoint | true (uniform across all blocks) |
compile_mode | max-autotune |
nan_guard_max_consecutive | 5 (with checkpoint rollback) |
data_mix | fineweb-edu 0.50 / fineweb 0.20 / the-stack-python 0.15 / openmath-instruct-2 0.10 / arxiv 0.05 (spec annotation only — pretrain.py reads no data: key except train_data_path; see docs/training.md) |
🚀 Quick start#
1. Install#
git clone https://github.com/atandra2000/Mamba-3-Lite.git
cd Mamba-3-Lite
pip install -r requirements.txt2. Verify the SSD math (CPU-friendly)#
python3 -m pytest tests/ -v37 tests cover the complex chunkwise SSD (vs naive scan oracle), MIMO mixer identity init (in the class and inside the full model), transformer forward, grad-checkpoint wiring, one-step training on dummy data, the Triton kernel reference + dispatch guards, an autograd gradcheck of the kernel's backward plumbing, and the doc↔code alignment checker. 32 pass on CPU in <3s; 5 GPU-gated tests skip.
3. Launch a full pretraining run#
python3 training/pretrain.py --config configs/pretrain_a100_400m.yaml4. Resume from checkpoint#
python3 training/pretrain.py \
--config configs/pretrain_a100_400m.yaml \
--resume 80000🧠 Why complex-valued SSD?#
The classical Mamba-2 recurrence is real:
h_t = exp(A · dt) · h_{t-1} + B · x_t (A, B, h, x ∈ ℝ)
y_t = C · h_tMamba-3 promotes everything to the complex plane:
h_t = exp((A_real + i·A_imag) · dt) · h_{t-1} + (B_real + i·B_imag) · x_t
y_t = (C_real + i·C_imag) · h_t (A, B, C, h, x ∈ ℂ)This is not just "use complex64 tensors" — it's a genuine representational upgrade:
| Aspect | Real SSD (Mamba-2) | Complex SSD (Mamba-3) |
|---|---|---|
| Eigenvalues | Real scalars (decay only) | Complex (decay + rotation) |
| State expressive power | 1 real dimension | 2 real dimensions (1 complex) |
| State size for parity | N=128 | N=64 (50% smaller) |
| Memory (per layer, BF16) | N·D·2 bytes | N·D·4 bytes for complex, but half the N |
| Net KV-equivalent cost | — | Lower at same effective capacity |
The complex exponential exp(α + iβ) = exp(α)·(cos β + i·sin β) natively captures both decay (α) and oscillation (β), which is impossible in real SSMs without doubling the state.
📖 Full math deep-dive: see docs/concepts/ssd-theory.md for the chunkwise algorithm derivation (and its section on state-space duality for the connection to self-attention).
🔬 Why MIMO (no SISO)?#
Classical SSMs are Single-Input Single-Output per head: head i sees only its own channel. Mamba-3 inserts a fully-connected mixer across heads after the SSD scan:
y_mixed = y.view(B, T, n_heads, head_dim) # (B, T, H, D)
y_mixed = y_mixed.transpose(1, 2) # (B, H, T, D)
y_mixed = y_mixed.reshape(B, T, n_heads * head_dim) # merge into channels
y_mixed = mimo_linear(y_mixed) # (B, T, n_heads * head_dim)
y = out_proj(y_mixed)This is the same role cross-attention plays in transformers but at zero extra sequence cost.
🧪 Purity#
This repo intentionally avoids:
- ❌
mamba-ssmpackage - ❌
causal_conv1dpackage - ❌ Custom CUDA kernels
- ❌ HuggingFace Trainer / PyTorch Lightning
- ❌ Pickle checkpoints (uses
safetensors+ atomic writes)
The single sanctioned exception: the opt-in per_chunk_ssd_triton kernel (models/ssd_triton.py, gated behind ssd_dispatch='triton' + ENABLE_TRITON_KERNELS=1 — see AGENTS.md §1). Everything else is pure PyTorch (torch.*matmul, torch.*einsum, torch.*fft where applicable). This makes the code:
- Auditable — every line is plain tensor ops.
- Hardware-portable — runs on CPU, MPS, CUDA, AMD ROCm, TPU.
- Educational — the SSD math is the algorithm, not a hidden kernel.
📂 Project structure#
Mamba-3-Lite/
├── assets/
│ ├── style.css # doc portal design system
│ └── portal.js # interactive hero + mechanism demos
├── .github/workflows/
│ └── deploy-docs.yml # auto-deploy docs to GitHub Pages
├── configs/
│ └── pretrain_a100_400m.yaml
├── models/
│ ├── ssd_complex.py # ★ complex-valued chunkwise SSD
│ ├── ssd_triton.py # ★ sanctioned fused Triton kernel (opt-in)
│ ├── mimo.py # ★ inter-head mixer (identity-init)
│ ├── mamba_block.py # block wiring (no causal conv)
│ └── transformer.py # top-level Mamba-3
├── training/
│ └── pretrain.py # full training loop + resume
├── utils/
│ ├── checkpoint.py # atomic safetensors
│ └── logging.py # WandB-capable logger
├── data/
│ ├── prepare_data.py # shim over the shared 8.0B-token pipeline
│ └── data_config.yaml # materialised by the shim (GPT-2 vocab)
├── scripts/
│ ├── build_docs_html.py # HTML docs generator for GitHub Pages
│ └── launch_a100.sh
├── tests/
│ ├── test_ssd.py # ★ chunk vs naive equivalence
│ ├── test_ssd_triton.py # kernel reference, dispatch guards, GPU parity
│ ├── test_doc_refs.py # doc↔code alignment checker (docs CI gate)
│ ├── test_mimo.py
│ ├── test_grad_checkpoint.py
│ ├── test_train_step.py
│ ├── test_transformer.py
│ └── e2e_gpu_smoke.py # 8-check GPU pipeline smoke (CUDA + triton)
├── docs/ # ★ full documentation tree
│ ├── README.md # doc map + reading paths
│ ├── concepts/ # from-scratch concept building
│ ├── references/ # symbol-anchored API docs
│ ├── guides/ # task-oriented runbooks
│ └── training.md # data pipeline + dataset path
├── AGENTS.md
├── SKILLS.md
├── LICENSE # Apache 2.0
├── requirements.txt
└── pytest.iniTest suite status.tests/contains 37 tests: the complex chunkwise SSD (vs naive scan oracle), MIMO mixer identity init (class-level and inside the full model), the Triton kernel reference + dispatch guards + autograd gradcheck, transformer forward, grad-checkpoint wiring, one-step training on dummy data, and the doc↔code alignment checker. 32 pass on CPU; 5 are GPU-gated. Seedocs/concepts/ssd-theory.mdfor the full mathematical derivation.
📖 Documentation#
Live docs on GitHub Pages — auto-deployed from main via GitHub Actions on every push.
The full doc tree lives in docs/, machine-checked for doc↔code alignment by tests/test_doc_refs.py (--coverage --links):
| Area | Where |
|---|---|
| SSD theory (foundations → duality → complex states → chunkwise algorithm) | docs/concepts/ |
| API references (config, SSD/kernel, model, training) | docs/references/ |
| How-to guides (quickstart, runbook, tuning, extending, pretrain CLI) | docs/guides/ |
| Data pipeline + dataset path | docs/training.md |
🧪 Verification#
The Mamba-3 SSD math is verified by the test suite in tests/ and by inline assertions in models/ssd_complex.py. Manual smoke checks:
python3 -c "
import torch
from models.transformer import Mamba3Transformer, ModelConfig
cfg = ModelConfig(vocab_size=100, d_model=64, n_layers=2, n_heads=4,
head_dim=16, state_dim=8, chunk_size=4, ffn_dim=128,
max_seq_len=32, weight_tying=True)
m = Mamba3Transformer(cfg)
x = torch.randint(0, 100, (2, 16))
y = m(x)
assert y.shape == (2, 16, 100), y.shape
print('forward ok, param count:', sum(p.numel() for p in m.parameters()))
"
# 2. Headline equivalence (chunkwise SSD vs naive O(T) recurrence)
# See docs/concepts/ssd-theory.md for the derivation. The math is exercised by every
# forward pass — if it regressed, training loss would diverge.🤝 Contributing#
PRs welcome for:
- New chunkwise algorithms (e.g., parallel prefix-scan variants).
- Selective vs static A/B/C parameterizations.
- Hybrid attention + Mamba blocks (e.g., 1-in-N global attention).
- New data mixes with documented perplexity deltas.
Please:
- Read
docs/concepts/ssd-theory.mdbefore touchingmodels/ssd_complex.py. - Run
python3 -m pytest tests/ -v— all must pass. - Do not add attention layers, MoE, or MTP — this is a pure SSM repo (avoids overlap with the rest of the portfolio).
- Do not add
mamba-ssmorcausal_conv1ddependencies.
⚠️ Known caveats#
- Full 8B-token pretraining run not yet started (no GPU on dev machine). The inline assertions validate all primitives on CPU + tiny shapes.
- Complex SSD has 2× element bandwidth vs real SSD (complex64 = 2× float32) — the per-state size halving must offset this. Theoretical analysis in
docs/concepts/block-and-stability.md; will be measured at full scale. - No causal conv = slightly weaker local-pattern bias. Mamba-3 trades a small amount of inductive bias for memory bandwidth and simplicity.
📚 References#
- Mamba-3 — Dao & Gu, 2025 (arXiv:2603.15569)
- Mamba-2 / SSD — Dao & Gu, 2024 (arXiv:2405.21060)
- S4 — Gu et al., 2021 (arXiv:2111.00396)
- S6 (selective state spaces) — Gu & Dao, 2023 (arXiv:2312.00752)
- H3 — Fu et al., 2022 (arXiv:2212.14052)
- RetNet — Sun et al., 2023 (arXiv:2307.08621)
- RWKV — Peng et al., 2023 (arXiv:2305.13048)
- Chinchilla scaling laws — Hoffmann et al., arXiv:2203.15556
📄 License#
Apache 2.0. See LICENSE.
<div align="center">
⭐ Star this repo if you find it useful · Part of the CoreProjects portfolio
</div>