Using JAX to Accelerate Research
deepmind.com
deepmind.com
Everyone knows JAX is what Google realised Tensorflow should have been when they realised how much of a joy Pytorch was to use. I actually think JAX does offer some advantages, not least true numpy interoperability. However, not mentioning *torch a single time in the blog post seems a little disingenuous for a Google-owned deep learning enterprise.
A technical reason for "why not pytorch" is that JAX was also built in part to expose and leverage the power of the XLA compiler, which is at least for the moment a pretty uniquely powerful tool for producing efficient, highly-scalable accelerator code.
I should underline that this is a friendly community of peers though: there is a lot of respect for Pytorch, which in turn was certainly influenced by the original Autograd that many of the JAX devs also worked on. JAX (and its fancier sibling Dex) beyond being useful tools are also still research projects in and of themselves seeking to advance our ideas on how to write expressive, powerful numerical code on modern architectures.
OTOH I'm not sure most people know what tasks are GPU-worthy or not. I haven't the slightest idea of why MCMC/Variational Bayes is amenable to GPU speedups and Persistent Homology isn't.
Yes, that is in active development: https://pymc-devs.medium.com/the-future-of-pymc3-or-theano-i...
JAX enables using (parts of) existing numpy codebases in disciplines other than deep learning. Autodiff and compilation to GPUs are very useful for all kinds of algorithms and processing pipelines.
But, had a look in the code and jax has cublas and RoCm blas, and it looks like there is a flow where it uses the gpu directly, unless I'm missing something.
Definitely worth having a closer look. Autograd via function reflection should be faster than backprop. And if it's running on AMD GPUs then it's quite intriguing.
I haven't found a way that I like. The closest I've come is tftorch: https://twitter.com/theshawwn/status/1311925180126511104
I think it's important to have global scope names. "biggan.discriminator.3.conv1.kernel.b" is perfectly sensible: it's the bias value of the kernel for the convolution of the third block of your biggan discriminator.
Everyone tries to treat model variables as interchangeable nameless parts. I hate it. Every variable has a global name, conceptually.
A global name also has other benefits. It becomes far easier to create EMA weights, for example, since you can filter the variables by name.
Their is no clear best pattern at the moment. My favorite would be Flax and their recent Linen API which is a refined effort that pays off when using their framework.