1 of 19

From stateful code to purified JAX:�how to build your neural net framework

Sabrina J. Mielke • @sjmielke • 2021-07-01

2 of 19

JAX as “numpy + jit()”

“Tired of Topic Models? Clusters of Pretrained Word Embeddings Make for Fast and Good Topics too!”

Sia, Dalmia, Mielke. EMNLP 2020.

Speed and coolness!

3 of 19

JAX as “numpy + jit()”

omg Jeff Dean!

be still my beating heart

4 of 19

JAX as “numpy + grad()”

Reversible jump MCMC requires the Jacobian determinant of a diffeomorphism:

jit() + vmap()

= speeeeeed

5 of 19

Tape-based autograd, e.g. PyTorch:

6 of 19

Pure transformation-based autograd: JAX

7 of 19

But how to “train” a big neural network?

How can weights and gradients be stateful if our Tensors can’t be mutated?

8 of 19

Basis for most of this talk:�https://sjmielke.com/jax-purify.htm�...a blogpost from March 2020!

Is it still relevant?

Good news: it never was*! :)�(* or it will never not be)

9 of 19

What will we accomplish?

No SotA models or frameworks.

No “idiomatic” elegance.

No batching (vmap), distributing (pmap), and XLA-compiling (jit).

  1. Linear Regression, statefully
  2. Moving to JAX & Purity
  3. Aside: classes as PyTree nodes?
  4. A pure-ifying framework sketch

10 of 19

A “module” in PyTorch

modules own their weights...

...and use them when building

the execution graph/tape

11 of 19

Using a PyTorch module: mutation!

mutate gradients and weights!

12 of 19

...off to the notebook!

13 of 19

pure-ification

Let the user write stateful methods and the framework convert such functions into actually pure functions that we can call grad() on!

14 of 19

The purification scheme, visually

15 of 19

This is what I want it to look like:

not quite self.w

but close enough.

use as before :)

16 of 19

Let’s implement this first...

...off to the notebook!

17 of 19

ours

dm-haiku

randomness handling!

18 of 19

ours

flax.linen

19 of 19

More details and complications:�https://sjmielke.com/jax-purify.htm

Follow me on Twitter @sjmielke :)