Flax
Centered on flexibility and clarity.
JAX - A curated list of resources https://github.com/google/jax
This page lists names, links and short descriptions. The original list on GitHub is the source and belongs to its authors.
Focused on simplicity, created by the authors of Sonnet at DeepMind.
Has an object oriented design similar to PyTorch.
A High Level API for Deep Learning in JAX. Supports Flax, Haiku, and Optax.
"Batteries included" deep learning library focused on providing solutions for common workloads.
Lightweight graph neural network library.
High-level API for specifying neural networks of both finite and infinite width.
Ecosystem of pretrained Transformers for a wide range of natural language tasks (Flax).
Callable PyTrees and filtered JIT/grad transformations => neural networks in JAX.
A Jax Library for Computer Vision Research and Beyond.
Prioritizes legibility, visualization, and easy editing of neural network models with composable tools and a simple mental model.
Legible, Scalable, Reproducible Foundation Models with Named Tensors and JAX.
LLMs made easy: Pre-training, finetuning, evaluating and serving LLMs in JAX/Flax.
Probabilistic programming based on the Pyro library.
Utilities to write and test reliable JAX code.
Gradient processing and optimization library.
Library for implementing reinforcement learning agents.
Accelerated, differential molecular dynamics.
Turn RL papers into code, the easy way.
Reimplementation of TensorFlow Probability, containing probability distributions and bijectors.
Construct differentiable convex optimization layers.
Tensor learning made simple.
Machine Learning toolbox for Quantum Physics.
AWS library for Uncertainty Quantification in Deep Learning.
Library of samplers for JAX.
Probabilistic state space models.
Federated learning in JAX, built on Optax and Haiku.
Construct equivariant neural network layers.
Implementations and checkpoints for ResNet variants in Flax.
JAX/Flax port of the RAFT optical flow estimator.
Immutable Torch Modules for JAX.
Root finding, minimisation, fixed points, and least squares.
Hardware accelerated (GPU/TPU), batchable and differentiable optimizers in JAX.
Library implementing the UniRep model for protein machine learning applications.
Distributions and normalizing flows built as equinox modules.
Framework and Library for building and training Diffusion models in multi-node multi-device distributed settings (TPUs)
Normalizing flows in JAX.
scikit-learn kernel matrices using JAX.
Differentiable cosmology library.
Exponential Families in JAX.
Combine MPI operations with your Jax code on CPUs and GPUs.
Image augmentations and transformations.
Flax version of TorchVision.
Probabilistic programming language based on program transformations.
Toolbox that bundles utilities to solve optimal transport problems.
A photovoltaic simulator with automatic differentation.
Lie theory library for rigid body transformations and optimization.
Differentiable physics engine to simulate environments along with learning algorithms to train agents for these environments.
Pretrained models for Jax/Flax.
XLA accelerated algorithms for sparse representations and compressive sensing.
Automatic differentiable spectrum modeling of exoplanets/brown dwarfs compatible to JAX.
PIX is an image processing library in JAX, for JAX.
Bayesian Optimization powered by JAX.
Framework for differentiable simulators with arbitrary discretizations.
Convert functions that operate on arrays into functions that operate on PyTrees.
Implementations of research papers originally without code or code written with frameworks other than JAX.
A framework for building discrete Probabilistic Graphical Models (PGM's) and running inference inference on them via JAX.
Hardware-Accelerated Neuroevolution
JAX-Based Evolution Strategies
Symbolic CPU/GPU/TPU programming.
Express & compile probabilistic programs for performant inference.
DSL-based reshaping library for JAX and other frameworks.
Open-source library for distributed matrix factorization using Alternating Least Squares, more info in ALX: Large Scale Matrix Factorization on TPUs.
Numerical differential equation solvers in JAX.
The tiniest of Gaussian process libraries in JAX.
Reinforcement Learning Environments with the well-known gym API.
Monte Carlo tree search algorithms in native JAX.
Second Order Optimization with Approximate Curvature for NNs.
Convert functions/graphs to JAX functions.
A library for differentiable acoustic simulations
Gaussian processes in JAX.
A Suite of Industry-Driven Hardware-Accelerated RL Environments written in JAX.
Equinox version of Torchvision.
Accelerated curve fitting library for nonlinear least-squares problems (see arXiv paper).
Solve macroeconomic models with hetereogeneous agents using JAX.
A domain-specific compiler and runtime suite to run JAX code with MPC(Secure Multi-Party Computation).
Add a tqdm progress bar to JAX scans and loops.
Serialize JAX, Flax, Haiku, or Objax model params with 🤗safetensors.
Differentiable stencil decorators in JAX.
A simple, performant and scalable Jax LLM written in pure Python/Jax and targeting Google Cloud TPUs.
A Jax-based machine learning framework for training large scale models.
The layer library for Pax with a goal to be usable by other JAX-based ML projects.
Vectorisable, end-to-end RL algorithms in JAX.
Automatically apply LoRA to JAX models (Flax, Haiku, etc.)
Scientific computational imaging in JAX.
Spiking Neural Networks in JAX for machine learning on neuromorphic hardware.
Brain Dynamics Programming in Python.
Physical units and unit-aware mathematical system in JAX.
Dendritic Modeling in JAX.
State-based Transformation System for Program Compilation and Augmentation.
Leveraging Taichi Lang to customize brain dynamics operators.
Optimal transport tools in JAX.
Quality Diversity optimization in Jax.
Nightly CI and optimized examples for JAX on NVIDIA GPUs using libraries such as T5x, Paxml, and Transformer Engine.
Vectorized board game environments for RL with an AlphaZero example.
EasyDeL 🔮 is an OpenSource Library to make your training faster and more Optimized With cool Options for training and serving (Llama, MPT, Mixtral, Falcon, etc) in JAX
A Differentiable Massively Parallel Lattice Boltzmann Library in Python for Physics-Based Machine Learning.
High-performance and differentiable simulations of quantum systems with JAX.
Agent-Based modelling framework in JAX.
Vectorized calculation of optical properties in thin-film structures using JAX. Swiss Army knife tool for thin-film optics research
Algorithms for finding coresets to compress large datasets while retaining their statistical properties.
A reimplementation of MiniGrid, a Reinforcement Learning environment, in JAX
Finite-Difference Time-Domain Electromagnetic Simulations in JAX
Differentiable Ray Tracing toolbox for Radio Propagation powered by the JAX ecosystem.
Plasma physics simulations using a PIC (Particle-in-Cell) method to self-consistently solve for electron and ion dynamics in electromagnetic fields
A FlashAttention implementation for JAX with support for efficient document mask computation and context parallelism.
differentiable (magneto)hydrodynamics for astrophysics in JAX
Fluid-structure interaction simulations using Immersed Boundary-Lattice Boltzmann Method.
High-performance tomographic reconstruction.
torchax is a library for Jax to interoperate with model code written in PyTorch.
Official implementation of Fourier Features Let Networks Learn High Frequency Functions in Low Dimensional Domains.
Approximate inference for Markov (i.e., temporal) Gaussian processes using iterated Kalman filtering and smoothing.
Nested sampling in JAX.
Open-source library for distributed matrix factorization using Alternating Least Squares, more info in ALX: Large Scale Matrix Factorization on TPUs.
Collection of LLMs implemented in JAX & Flax
Flax implementation of DeepSeek-R1 1.5B distilled reasoning LLM.
Open-source library for distributed matrix factorization using Alternating Least Squares, more info in ALX: Large Scale Matrix Factorization on TPUs.
Official implementation of Mip-NeRF: A Multiscale Representation for Anti-Aliasing Neural Radiance Fields.
Implementation of NeuS: Learning Neural Implicit Surfaces by Volume Rendering for Multi-view Reconstruction
Implementation of Big Transfer (BiT): General Visual Representation Learning.
Implementations of reinforcement learning algorithms.
Implementation of Pay Attention to MLPs.
Minimal implementation of MLP-Mixer: An all-MLP Architecture for Vision.
Official implementation of Aggregating Nested Transformers.
Official implementation of Cross-Modal Contrastive Learning for Text-to-Image Generation.
Official implementation of An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale.
Port of mseitzer/pytorch-fid to Flax.
A JAX/Flax implementation of the Sharpened Cosine Similarity layer.
A JAX + Flax implementation of Combinatorial Optimization with Physics-Inspired Graph Neural Networks.
Flax implementation of DETR: End-to-end Object Detection with Transformers using Sinkhorn solver and parallel bipartite matching.
Implementation of the inference pipeline of AlphaFold v2.0, presented in Highly accurate protein structure prediction with AlphaFold.
Reference code for Uncovering the Limits of Adversarial Training against Norm-Bounded Adversarial Examples and Fixing Data Augmentation to Improve Adversarial Robustness.
Normalizing flows with JAX.
Open-source library for distributed matrix factorization using Alternating Least Squares, more info in ALX: Large Scale Matrix Factorization on TPUs.
JAX implementation of the paper Auction learning as a two-player game.
"Batteries included" deep learning library focused on providing solutions for common workloads.
Official implementation of Bayesian inverse optimal control for linear-quadratic Gaussian problems from the paper Putting perception into action with inverse optimal control for continuous psychophysics
Official tutorial and implementation from the paper Towards Generative Ray Path Sampling for Faster Point-to-Point Ray Tracing.
JAX, its use at DeepMind, and discussion between engineers, scientists, and JAX core team.
Simple neural network from scratch in JAX.
JAX's core design, how it's powering new research, and how you can start using it.
Introduction to Bayesian modelling using NumPyro.
JAX intro presentation in Program Transformations for Machine Learning workshop.
Presentation of TPU host access with demo.
Tutorial created by Zico Kolter, David Duvenaud, and Matt Johnson with Colab notebooks avaliable in Deep Implicit Layers.
A four part YouTube tutorial series with Colab notebooks that starts with Jax fundamentals and moves up to training with a data parallel approach on a v3-32 TPU Pod slice.
Ecosystem of pretrained Transformers for a wide range of natural language tasks (Flax).
White paper describing an early version of JAX, detailing how computation is traced and compiled.
Introduces JAX, M.D., a differentiable physics library which includes simulation environments, interaction potentials, neural networks, and more.
Uses JAX's JIT and VMAP to achieve faster differentially private than existing libraries.
White paper describing the XLB library: benchmarks, validations, and more details about the library.
Describes the state of JAX and the JAX ecosystem at DeepMind.
Neural network building blocks from scratch with the basic JAX operators.
A gentle introduction to JAX and using it to implement Linear and Logistic Regression, and Neural Network models and using them to solve real world problems.
Learn how to create a simple convolutional network with the Linen API by Flax and train it to recognize handwritten digits.
Compares Flax, Haiku, and Objax on the Kaggle flower classification challenge.
Introduction to both JAX and Meta-Learning.
Concise implementation of RealNVP.
Tutorial on implementing path tracing.
Ensemble nets are a method of representing an ensemble of models as one single logical model.
Implements different methods for OOD detection.
Understand how autodiff works using JAX.
Showcases how to go from a PyTorch-like style of coding to a more Functional-style of coding.
Tutorial demonstrating the infrastructure required to provide custom ops in JAX.
Explores how JAX can power the next generation of scalable neuroevolution algorithms.
Demonstrates how to use JAX to perform inner-loss optimization with SGD and Momentum, outer-loss optimization with gradients, and outer-loss optimization using evolutionary strategies.
Walk through of implementing automatic differentiation variational inference (ADVI) easily and cleanly with JAX.
Trains a classification model robust to different combinations of input channels at different resolutions, then uses a genetic algorithm to decide the best combination for a particular loss.
Colab that introduces various aspects of the language and applies them to simple ML problems.
Tutorial on the different ways to write an MCMC sampler in JAX along with speed benchmarks.
Tutorial on how to add a progress bar to compiled loops in JAX using the host_callback module.
A series of notebooks and videos going from zero JAX knowledge to building neural networks in Haiku.
A tutorial on writing a simple end-to-end training and evaluation pipeline in JAX, Flax and Optax.
A tutorial on 3D volumetric rendering of scenes represented by Neural Radiance Fields in JAX.
A series of notebooks explaining various deep learning concepts, from basics (e.g. intro to JAX/Flax, activiation functions) to recent advances (e.g., Vision Transformers, SimCLR), with translations to PyTorch.
A blog post on how JAX can massively speedup RL training through vectorisation.
A simple example of solving the advection-diffusion equations with JAX and using it in a constrained optimization problem to find initial conditions that yield desired result.
A hands-on guide to using JAX for deep learning and other mathematically-intensive applications.
| Python | - Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more
hesreallyhim/awesome-claude-code
A hand-picked collection of the finest of resources for the most awesome of agents, Claude Code, the undisputed champion of coding companions, from the unstoppable team…
VoltAgent/awesome-agent-skills
A curated collection of 1000+ agent skills from official dev teams and the community, compatible with Claude Code, Codex, Gemini CLI, Cursor, and more.
josephmisiti/awesome-machine-learning
A curated list of awesome Machine Learning frameworks, libraries and software.
EthicalML/awesome-production-machine-learning
A curated list of awesome open source libraries to deploy, monitor, version and scale your machine learning
academic/awesome-datascience
:memo: An awesome Data Science repository to learn and apply for real world problems.
analysis-tools-dev/static-analysis
⚙️ A curated list of static analysis (SAST) tools and linters for all programming languages, config files, build tools, and more. The focus is on tools which improve…