Jax – Composable transformations of Python and NumPy programs
github.com
github.com
We'd like to take this opportunity to give a shout out to some of the awesome projects folks are building on top of JAX, e.g.,
* Flax, a neural network library for JAX (https://github.com/google/flax)
* Haiku, a neural network library for JAX inspired by Sonnet (https://github.com/deepmind/dm-haiku)
* RLax, a library for building reinforcement learning agents (https://github.com/deepmind/rlax)
* NumPyro, a probabilistic programming library on top of JAX (https://github.com/pyro-ppl/numpyro)
* JAX-MD, a differentiable molecular dynamics package built on top of JAX (https://github.com/google/jax-md)
Interesting that googlers who are supposed to use Tensorflow are now actively developing a new autograd engine and at least three new DL frameworks on top of it. What do you think about this segmentation?
Comparisons are hard in general and I don't have a good answer for you right now, but keep in mind most of these libraries are from researchers openly sharing the codebases they develop for their own work. We see the role of JAX as analogous to NumPy, that is, a common substrate on which folks can build these sorts of tools.
The advantages JAX brings are
* Numpy-adjacent interface
* Auto-vectorization
* Great tie-in with XLA
* Higher-order derivatives (critical for some applications)
* Simplified, functional interface
The disadvantage is that it's younger and not trying to be a fully-fledged competitive neural network framework and thus is behind and less resourced compared with other libraries implementing auto-differentiation.I think it's really exciting to see more people learning about JAX and beginning to use it for serious projects.
Perhaps this is an unpopular opinion, but to me this is an advantage. JAX is a library upon which you can build a neural-network framework -- or a framework for something else.
It seems clear that we're still figuring out the Right Way to write code that defines a neural network. The fact that JAX lets you write different, competing libraries/APIs/DSLs -- however you want to think of it -- lets us innovate more freely.
Also, Zygote is much more ambitious than Jax. Zygote aims to support all of the Julia language whereas Jax is limited to a subset of Python. I wonder if the Zygote folks are biting off too much here? Though currently Zygote doesn't support mutation.
Since Python is not compiled like Julia, the approach usually involves using custom types that store a separate intermediate representation from python (a custom one or the MLIR for example) and overloading functions (and perhaps other more sophisticated ways of metaprogramming) to map all the required operations while interpreting the Python code before compiling. That's why Zygote says it works on Zygote unaware libraries, as it doesn't store the graph within any of the types or directly overload any of the methods (it does need to know how to reverse them though), and also why it can directly operate on control structures since they are part of the full IR even when you can't overload them.
And while Zygote aims to fully support the Julia language, they'll certainly work on getting the most important operations working well enough before focusing on the more complicated stuff like mutation.
you want to work with the code reduced down to Single assignment form, and work with code blocks and the control flow graph. for a number of reasons: 1. all control flow looks the same now (no 2 different kinds of loops + GOTOs) 2. one expression per line, no need to untangle larger experessions 3. Host of techniques from compiler world can come out to be used.
This is how Zgyote works, it used code that has been lowered into this form
A year or two ago, it at least seemed to be a much more vibrant situation where there was more effort going into lots of different approaches, but then Zygote really picked up steam and it seemed like it was going to solve all our problems and cure cancer so development on the other packages slowed down, but now the 'finish line' for Zygote doesn't seem to be as close as we initially thought.
Also, I think performance is very important because Zygote is used in Flux (Julia's main deep learning framework at this point.) and if it's keeping Flux's performance from matching, say, PyTorch's then that's going to limit adoption of Flux and Julia. One of Julia's main claims to fame is performance so people coming to kick the tires are going to be disappointed if it's actually slower for this domain.
Besides, any implementation of AD is going to have downsides. Having many different, hot swappable implementations is quite nice because you can tailor towards your needs.
One of my on going projects is ChainRules (http://www.juliadiff.org/ChainRulesCore.jl/dev/) which will unite them under one set of custom senstitivities and more generally make it easier to mix and match them
I need to update the JuliaDiff website, I want to list all of them in a table with some some key points.
Forward mode:
- [ForwardDiff](https://github.com/JuliaDiff/ForwardDiff.jl)
- [ForwardDiff2](https://github.com/YingboMa//ForwardDiff2.jl)
Reverse Mode:
- [Nabla](https://github.com/invenia/Nabla.jl/)
- [Tracker](https://github.com/FluxML/Tracker.jl)
- [Yota](https://github.com/dfdx/Yota.jl)
- [Zygote](https://github.com/FluxML/Zygote.jl)
- [ReverseDiff](https://github.com/JuliaDiff/ReverseDiff.jl)
- [AutoGrad.jl](https://github.com/denizyuret/AutoGrad.jl)
- [NiLang](https://github.com/GiggleLiu/NiLang.jl) (arguably not reverse mode)
Symbolic:
- [ModelingToolKit](https://github.com/JuliaDiffEq/ModelingToolkit.jl)
- [XGrad.jl](https://github.com/dfdx/XGrad.jl)
Finite Differencing:
- [Calculus](https://github.com/JuliaMath/Calculus.jl) (please stop)
- [FiniteDifferences](https://github.com/JuliaDiff/FiniteDifferences.jl)
- [FiniteDiff](https://github.com/JuliaDiff/FiniteDiff.jl)
https://news.ycombinator.com/item?id=18636054
JAX is pretty neat because it is effectively a derivatives compiler: it can automatically differentiate a function and JIT compile the result. This makes training in machine learning both fast and easy because gradient descent no longer has to be written by hand.
I thought that PyTorch, Tensorflow and similar already do that.
These methods are much faster than perturbation-based derivatives and much more applicable than symbolic methods (which cannot be automatically extracted from a program).
Although, honestly, I misspoke. The difference between AD and symbolic differentiation is more subtle. Really AD is profiting because it uses AST representations to keep a graph of intermediate values while symbolic methods can blow up exponentially (or require clever, difficult to generalize tricks to reconstruct that graph).
Nobody's been writing derivatives by hand for 5+ years. All major frameworks (PyTorch, Tensorflow, MXNet, autodiff, Chainer, Theano, etc.) have decent to great automatic differentiation.
The differences and improvements are more subtle (easy parallelization/vectorization, higher-order gradients, good XLA support).
Automatic differentiation allows for great flexibility and composability but the performance is still far from good, even with the various JITs available. Jax seems to be one of the most flexible and optimized for many use cases for now however.
But I have some almost-reasonably-performant pytorch that I'd rather not just use as a cash burning machine, so it looks like it might be time to dive into CUDA :-\
So yes, if need a new primitive to add an efficient CUDA kernel, you will probably also have to write its derivative manually too. JAX has a few shortcuts that occasionally make this easier but fundamentally it has the same challenge as any auto-diff system.
Next to none of the frameworks are yet able to JIT you a performant RNN, yet RNNs only use very standard components[1]. OpenAI had a massive speed and memory usage boost for attention by implementing what amounts to a few standard primitives together[2].
There are massive gaps in the optimizations that existing ML compilers provide. The landscape is starting to get better but it's still filled with many pitholes.
[1]: https://twitter.com/stanfordnlp/status/1224106217192087552
import jax.numpy as np
from jax import grad, jit, vmap
from jax import random
def relu(x):
return np.where(x>0, x, np.zeros_like(x))
def identity(x):
return relu(x) - relu(-x)
derivative_identity = grad(identity)
derivative_identity(0.0)
returns 1.0 or 0.0? It currently returns 0.0. (Edit: typo in sign)Subgradients are only applicable when summing convex functions. Here relu(x) - relu(-x) is a sum of a convex function and a concave function.
Critical kinks are those that affect the geometry of the gradient in systematic ways. For instance, a model with a mixture of discrete and continuous parameters. These are serious blockers and require more complex methods to solve such as Rao-Blackwellization (marginalizing out the discrete parameters). Generally this appears as model bias or substantially increased, often fatal variance in loss curves.
I've yet to see anything get "a lot faster" because of XLA. It's a ton of complicated code, but then you end up spending the vast majority of time in NVIDIA's cuDNN anyway, so any benefits you might have hoped for will be marginal at best.
Almost double speedup for FP16 Resnet-50.
In fact, also seems to be outperformed by plain PyTorch using a single V100: https://github.com/NVIDIA/DeepLearningExamples/tree/master/P...
Nvidia might have eliminated any potential data pipeline bottlenecks (with careful DALI tuning), but I'd still expect a lot less speedup. Maybe they compiled pytorch with certain tricks, and used newer CUDA/CuDNN code, idk.
It runs about as fast as any of the other popular machine learning frameworks, occasionally faster.
Disclaimer: I work for Google and use JAX, although I'm not on the Jax team.
Support for general n-th total derivatives is rather good :)
JAX core is an extensible system for transforming numerical Python functions. This core is used to implement automatic differentiation, translation to TF XLA, etc.
Numba does not have such a generic function transformation framework - it just supports a single transformation, that is from numerical Python function to machine code.
(I'm not from either jax or numba, but a keen jax user for non-ML research.)
We haven't tried combining them yet, but we think it would be fun to explore (https://github.com/google/jax/issues/1870). For example, you could use Numba to hand write a numerical kernel that then participates in a machine learning model that uses JAX automatic differentiation.