Mamba2-Jax; Mamba2 implemented in pure Jax/Flax
Mamba2-Jax; Mamba2 implemented in pure Jax/Flax
I ported the Mamba2 state-space model (“Transformers are SSMs”) to pure JAX/Flax as mamba2-jax. – Core Mamba2 block with LM (Mamba2ForCausalLM) and time-series (Mamba2Forecaster) heads – Pure JAX/Flax (no Triton/custom CUDA), runs on CPU / CUDA / TPU via standard JAX backends – Small CPU-only parity test vs mamba2-torch: similar loss curves, final MSE diff ≈ 0.012, prediction correlation ≈ 0.99; after JIT warmup JAX was ≈ 2× faster per step I’d really appreciate feedback on: – API design, especially for streaming/stateful inference – Performance gotchas you hit if you try it – Any hooks you’d want exposed for research use PyPI: https://pypi.org/project/mamba2-jax/ Thanks, Cosmo
Share cardActual performance
Launch Intel predictions
Analyze your own launch →Correct prediction on native model
Similar products
A pure Tcl JPEG decoder
A pure-Ruby implementation of systemd's sd_notify(3)
Pointfree FP-Interpreter with Combinators and Pure Functions
Goredo – redo implementation on pure Go
Rtpstalk – pure Java RTPS implementation
Pure C Implementation of the Chatarjee – Xi Correlation Coefficient
6502 in pure Lua, no dependencies
IRCd in Pure Bash
Go-osc – OSC Packet Implementation for Golang. Implemented in Pure Go
Pure functional lenses in Racket