Fetching the paper…

Equinox: neural networks in JAX via callable PyTrees and filtered transformations · Around