Project Overview (README)

§ README.md 2,322 words τ ~12 min read

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

![Python 3.10+](https://www.python.org/)

![PyTorch 2.1+](https://pytorch.org/)

![License: Apache 2.0](LICENSE)

![GPU: A100 80GB](#-hardware)

![No custom CUDA](#-purity)

![Code style: black](https://github.com/psf/black)

![Docs](https://atandra2000.github.io/Mamba-3-Lite/)

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):

  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.
  2. 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.
  3. Zero causal convolution. The memory-bound causal_conv1d pass 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#

ProjectBackboneStateMixerCausal conv
GPT-2 (From Scratch)Transformer———
LLaMA-3-LiteTransformer + GQA———
DeepSeek-v3-LiteMLA + MoE———
HyMoGDN + MLA hybridreal SSM (in GDN blocks)——
Mamba-3-LitePure complex SSDN=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#

DiagramDescriptionInteractive HTMLVisual Preview
Complex SSD Architecture28-layer Mamba-3 block, complex64 SSD recurrence, MIMO head mixing matrix, and SwiGLU FFNOpen Map ↗PNG
Data Pipeline8.0B-token universal pipeline, GPT-2 BPE tokenizer, binary chunk sharding, and memory-mapped PretrainDatasetOpen Map ↗PNG
Training WorkflowEnd-to-end pretraining loop, chunked cross-entropy (chunk=4096), AdamW optimizer, and state checkpointingOpen Map ↗PNG

🏗 Architecture#

Per-block components#

ComponentSpecPurpose
Input projectionin_proj: d_model → n_heads × head_dim × 2 (real + imag packed)One projection instead of separate x/B
Complex SSDN=64, complex64, chunk=64State-space scan with complex eigenvalues
MIMO mixern_heads × head_dim → n_heads × head_dim (fully connected)Cross-head information flow
Output projectionn_heads × head_dim → d_modelAggregate heads back to model dim
FFNSwiGLU, ffn_dim=2048 (not 4096)Gated MLP, matches Mamba-2 design
NormalizationRMSNorm, pre-norm, eps=1e-5
Weight tyingEmbed ↔ output headSaves ~52M params
Causal convNonePure chunked linear projection

⚙️ Configuration#

The canonical config is configs/pretrain_a100_400m.yaml:

Model#

ParameterValue
vocab_size50,257 (GPT-2 BPE)
d_model1,024
n_layers28
n_heads16 (SSM heads)
head_dim64 (D)
state_dim64 (N, complex64)
chunk_size64 (SSD tunable)
ffn_dim2,048 (SwiGLU intermediate)
max_seq_len2,048
weight_tyingtrue
init_std0.02
Total params~434M

Training#

ParameterValue
micro_batch_size16
gradient_accumulation_steps2
total_steps256,000 (~8.0B tokens)
warmup_steps2,000 (linear)
lr3.0 × 10⁻⁴
min_lr_ratio0.05 (cosine decay)
weight_decay0.1
beta1 / beta20.9 / 0.95
grad_clip1.0
grad_checkpointtrue (uniform across all blocks)
compile_modemax-autotune
nan_guard_max_consecutive5 (with checkpoint rollback)
data_mixfineweb-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#

bash
git clone https://github.com/atandra2000/Mamba-3-Lite.git
cd Mamba-3-Lite
pip install -r requirements.txt

2. Verify the SSD math (CPU-friendly)#

bash
python3 -m pytest tests/ -v

37 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#

bash
python3 training/pretrain.py --config configs/pretrain_a100_400m.yaml

4. Resume from checkpoint#

bash
python3 training/pretrain.py \
    --config configs/pretrain_a100_400m.yaml \
    --resume 80000

🧠 Why complex-valued SSD?#

The classical Mamba-2 recurrence is real:

code
h_t = exp(A · dt) · h_{t-1} + B · x_t        (A, B, h, x ∈ ℝ)
y_t = C · h_t

Mamba-3 promotes everything to the complex plane:

code
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:

AspectReal SSD (Mamba-2)Complex SSD (Mamba-3)
EigenvaluesReal scalars (decay only)Complex (decay + rotation)
State expressive power1 real dimension2 real dimensions (1 complex)
State size for parityN=128N=64 (50% smaller)
Memory (per layer, BF16)N·D·2 bytesN·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:

python
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-ssm package
  • ❌ causal_conv1d package
  • ❌ 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#

Test 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. See docs/concepts/ssd-theory.md for 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):

AreaWhere
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 pathdocs/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:


🤝 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:

  1. Read docs/concepts/ssd-theory.md before touching models/ssd_complex.py.
  2. Run python3 -m pytest tests/ -v — all must pass.
  3. Do not add attention layers, MoE, or MTP — this is a pure SSM repo (avoids overlap with the rest of the portfolio).
  4. Do not add mamba-ssm or causal_conv1d dependencies.

⚠️ 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>