Jax and Flax LLMs – Transformer Implementations Optimized for TPUs
Jax and Flax LLMs – Transformer Implementations Optimized for TPUs
I've open-sourced awesome-jax-flax-llms, a curated collection of large language model (LLM) implementations built from scratch using JAX and Flax. The repo is designed for high-performance training on TPUs/GPUs, making it ideal for researchers, ML engineers, and curious tinkerers looking to explore or extend modern transformer models. Key Features: Modular, readable, and extensible codebase Implementations of GPT-2 and LLaMA 3 in pure JAX/Flax Accelerated training with XLA + Optax Google Colab support (TPU-ready) Hugging Face dataset integration Upcoming support for fine-tuning, Mistral, and DeepSeek-R This is primarily an educational resource, but it's written with performance in mind and can be adapted for more serious use. Contributions are welcome — whether you’re improving performance, adding new models, or experimenting with different attention mechanisms.
Share cardActual performance
Launch Intel predictions
Analyze your own launch →Incorrect prediction on native model
Similar products
(bits) of a Libc, Optimized for Wasm
GPTCache – Redis for LLMs
prompttest – pytest for LLMs
Mobile optimized Arguman
GruntStart (Grunt+H5BP+jQuery+Modernizr) - Optimized Development
Optimized Travel
KraspAI Kompass – keep up with new LLMs
pry-rescue — workflow-optimized debugging for ruby
How Optimized Threaded Pagination Works
An optimized and ergonomic image annotation platform