ott by ott-jax

Optimal transport tools implemented with the JAX framework, to get differentiable, parallel and jit-able computations.

created at Dec. 24, 2021, 7:28 a.m.

Python

10 +0

458 +0

78 +0

GitHub
jax-models by DarshanDeshpande

Unofficial JAX implementations of deep learning research papers

created at Jan. 8, 2022, 1:40 p.m.

Python

5 +0

140 +0

9 +0

GitHub
QDax by adaptive-intelligent-robotics

Accelerated Quality-Diversity

created at Feb. 11, 2022, 3:48 p.m.

Python

6 +0

243 +0

36 +0

GitHub
mctx by deepmind

Monte Carlo tree search in JAX

created at March 1, 2022, 5:26 p.m.

Python

29 +0

2,219 +4

171 -1

GitHub
tf2jax by deepmind

None

created at March 2, 2022, 8:22 p.m.

Python

7 +0

98 +2

8 +0

GitHub
kfac-jax by deepmind

Second Order Optimization and Curvature Estimation with K-FAC in JAX.

created at March 18, 2022, 10:19 a.m.

Python

10 +0

204 +2

15 +0

GitHub
levanter by stanford-crfm

Legible, Scalable, Reproducible Foundation Models with Named Tensors and Jax

created at May 24, 2022, 10:26 p.m.

Python

15 +0

454 +5

66 +0

GitHub
praxis by google

None

created at June 14, 2022, 4:04 p.m.

Python

8 +0

148 +0

39 +0

GitHub
paxml by google

Pax is a Jax-based machine learning framework for training large scale models. Pax allows for advanced and fully configurable experimentation and parallelization, and has demonstrated industry leading model flop utilization rates.

created at June 14, 2022, 4:04 p.m.

Python

16 +0

410 +4

61 +0

GitHub
kernex by ASEM000

Stencil computations in JAX

created at July 10, 2022, 10:01 a.m.

Python

1 +0

62 +1

3 +0

GitHub
eqxvision by paganpasta

A Python package of computer vision models for the Equinox ecosystem.

created at July 24, 2022, 10:02 p.m.

Python

4 +0

95 +0

10 +0

GitHub
jaxfit by Dipolar-Quantum-Gases

GPU/TPU accelerated nonlinear least-squares curve fitting using JAX

created at Aug. 8, 2022, 5 p.m.

Python

1 +0

43 +0

3 +0

GitHub
jumanji by instadeepai

🕹ī¸ A diverse suite of scalable reinforcement learning environments in JAX

created at Aug. 11, 2022, 7:34 a.m.

Python

10 +0

539 +4

67 +0

GitHub
fortuna by awslabs

A Library for Uncertainty Quantification.

created at Nov. 17, 2022, 1:11 p.m.

Python

12 +0

855 +1

45 +0

GitHub
EasyLM by young-geng

Large language models (LLMs) made easy, EasyLM is a one stop solution for pre-training, finetuning, evaluating and serving LLMs in JAX/Flax.

created at Nov. 22, 2022, 12:55 p.m.

Python

41 +0

2,257 +9

233 +0

GitHub
safejax by alvarobartt

Serialize JAX, Flax, Haiku, or Objax model params with 🤗`safetensors`

created at Dec. 21, 2022, 9:20 a.m.

Python

2 +0

38 +0

2 +0

GitHub
jax-tqdm by jeremiecoullon

Add a tqdm progress bar to your JAX scans and loops.

created at Jan. 16, 2023, 4:54 p.m.

Python

4 +0

56 +1

6 +0

GitHub
dynamiqs by dynamiqs

High-performance quantum systems simulation with JAX (GPU-accelerated & differentiable solvers).

created at Feb. 5, 2023, 5:04 p.m.

Python

6 +0

106 +1

11 +0

GitHub
purejaxrl by luchris429

Really Fast End-to-End Jax RL Implementations

created at Feb. 25, 2023, 3:38 p.m.

Python

12 +0

579 +1

47 +0

GitHub
maxtext by google

A simple, performant and scalable Jax LLM!

created at Feb. 28, 2023, 7:47 p.m.

Python

23 +1

1,300 +5

226 -2

GitHub