jaxns by Joshuaalbert

Probabilistic Programming and Nested sampling in JAX

created at July 31, 2020, 6 p.m.

Python

5 +0

125 +4

8 +0

GitHub
tf2jax by deepmind

None

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

Python

7 +0

96 +0

8 +0

GitHub
lqg by RothkopfLab

Inverse optimal control for continuous psychophysics

created at Aug. 24, 2021, 7:38 a.m.

Jupyter Notebook

2 +0

24 +0

8 +0

GitHub
jaxdf by ucl-bug

A JAX-based research framework for writing differentiable numerical simulators with arbitrary discretizations

created at Sept. 8, 2021, 4:38 p.m.

Python

7 +0

107 +0

7 +0

GitHub
tree-math by google

Mathematical operations for JAX pytrees

created at Dec. 18, 2021, 1:56 a.m.

Python

10 -1

167 +0

7 +0

GitHub
einshape by deepmind

None

created at July 22, 2021, 4:32 p.m.

Python

7 +0

89 +0

5 +0

GitHub
jax-fid by matthias-wright

FID computation in Jax/Flax.

created at July 29, 2021, 4:05 a.m.

Python

3 +0

21 +0

5 +0

GitHub
SymJAX by SymJAX

Documentation:

created at Sept. 4, 2019, 2:23 p.m.

Python

8 +0

118 +1

5 +0

GitHub
sklearn-jax-kernels by ExpectationMax

Composable kernels for scikit-learn implemented in JAX.

created at March 10, 2020, 6:59 p.m.

Python

5 +0

40 -1

4 +0

GitHub
lorax by davisyoshida

LoRA for arbitrary JAX models and functions

created at April 21, 2023, 11:50 a.m.

Python

3 +0

117 +0

4 +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

55 +0

4 +0

GitHub
NuX by Information-Fusion-Lab-Umass

Normalizing Flows using JAX

created at March 9, 2020, 12:38 a.m.

Python

9 +0

82 +0

4 +0

GitHub
parallax by srush

None

created at May 19, 2020, 1:49 a.m.

Python

6 +0

157 +0

4 +0

GitHub
kernex by ASEM000

Stencil computations in JAX

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

Python

1 +0

59 +0

3 +0

GitHub
efax by NeilGirdhar

Exponential families for JAX

created at March 24, 2020, 5:44 a.m.

Python

4 +0

50 +0

3 +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 +1

GitHub
imax by 4rtemi5

Image augmentation library for Jax

created at Feb. 9, 2021, 1:49 a.m.

Python

3 +0

34 +0

3 +0

GitHub
JAX-Flax-Tutorial-Image-Classification-with-Linen by 8bitmp3

How to use the Flax Linen API to build a convolutional neural network model and train it for image classification (using TensorFlow Datasets).

created at Dec. 24, 2020, 3 a.m.

Jupyter Notebook

2 +0

22 +0

3 +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
flaxvision by rolandgvc

A selection of neural network models ported from torchvision for JAX & Flax.

created at June 14, 2020, 4:34 p.m.

Python

4 +0

44 +0

2 +0

GitHub