JAX Mambas

GitLab (Main Repo) | GitHub (Mirror)

Summary

A project where I am doing JAX+Flax NNX implementation of Mamba-1, 2 and 3. This includes replicating some of their experiments (as my compute budget allows), and porting their CUDA kernels to either JAX/OpenXLA's FFI or (ideally) Pallas.

My objective is to provide both paper-accurate and PyTorch-repo-accurate implementations. This, in my opinion, is important as I found that the experiments conducted in the Mamba-1 paper were replicated by the paper-accurate version but not the repo accurate version if the repo accurate version was plugged in witht the paper's hyperparams.

Progress

Mamba

My work on this is complete.

Mamba-2

I've implemented a form of Structed State Space Duality, but turns out it's not quite what they did for the paper; my form is a bit more computationally expensive. I'm not gonna push anything until I actually have a proper SSD kernel.

Mamba-3

Work on this hasn't started yet