RL and MARL with Flax NNX
-
Updated
Sep 7, 2026 - Python
RL and MARL with Flax NNX
The best ChatGPT-style model that $100 of TPU time can buy.
Educational, modular, and high-performance Diffusion Transformers (DiT) in JAX and Flax NNX.
Titans: Learning to Memorize at Test Time
Scientific machine learning for JAX/Flax NNX: neural operators (FNO family, DeepONet, PINO, UNO), physics-informed networks (PINN, FBPINN, XPINN), E(3)-equivariant atomistic potentials (SchNet, PaiNN, NequIP), differentiable Kohn-Sham DFT, SINDy equation discovery, uncertainty quantification (conformal, GPs, SBI), PDEBench benchmarking.
nanoGPT for Diffusion Language Models
Deep reinforcement learning algorithms and environments, written completely in JAX + Flax NNX for end-to-end GPU-accelerated training.
Differentiable data pipelines for JAX/Flax NNX: HuggingFace, TFDS and ArrayRecord sources, augmentation stages, scan-based epochs, exact mid-epoch resume.
JAX benchmarking, profiling and evaluation metrics for Flax NNX: a registry of pure-function metrics (regression, classification, calibration, uncertainty, forecasting, generative, image, text, audio, graph, fairness), XLA FLOP counting, roofline analysis, GPU and energy monitoring, regression detection, publication exports, W&B and MLflow.
Molecular active learning with JAX
End-to-end differentiable bioinformatics for JAX/Flax NNX: soft Smith-Waterman alignment, read mapping and assembly, variant calling, RNA-seq, single-cell, epigenomics, CRISPR, metabolomics, multi-omics, protein and RNA structure, molecular dynamics and drug-discovery operators composed into trainable pipelines on the Avitai stack.
Flax NNX implementation of common metrics.
Physics-informed, RL-aligned evaluation engine for autonomous driving: JAX/Flax NNX trajectory-diffusion world models on the Waymo Open Dataset, DPO and reward-guided adversarial steering toward the scenarios that break AV stacks, physics feasibility, occupancy flow, WOSAC-style realism metrics and sensor simulation.
KGE-JAXed: A simple knowledge graph embedding library created in JAX
Generative modeling for JAX/Flax NNX: VAEs, GANs, DDPM/DDIM/score/DiT/latent diffusion, normalizing flows, energy-based, autoregressive and geometric (SE(3), protein, point cloud, mesh) models; image, text, audio, molecular, protein, tabular and time-series modalities; a shared Trainer with callbacks and extensions; evaluation through calibrax.
A 2-in-1 notebook-based tutorial and implementation of Manifold-Constrained Hyper-Connections using JAX
JaxNN: Foundation Models in JAX/Flax
JAX/Flax NNX training infrastructure: device meshes and SPMD sharding, an Orbax checkpoint store, early stopping and callbacks, W&B and MLflow logging.
To associate your repository with the flax-nnx topic, visit your repo's landing page and select "manage topics."