Using JAX in 2022
assemblyai.com
assemblyai.com
I was planning on doing a more thorough introductory tutorial or deep dive into Transformations, so let me know if you (all) think that would be instructive!
> In short - speed.
For me personally, the magic of JAX is that it able to have this performance, while being as close as possible to having first class differentiation in Python. The latter is a far more important reason to use JAX. It can really change how you think about programming and ML. Rather than implementing a specific model, you can write up the parameterized solution to a problem then solve it.
However first class differentiation ultimately isn't really useful unless you happen to also solve the speed problem. That is what makes JAX incredible. From the programming perspective JAX is to differentiable programming what Prolog is to logic programming, however Prolog has always been limited ultimately by performance problems where JAX is not.
Also, jit() is where you'll likely see your speed increases.
As for autodiff - not every application will use it. What exactly are you trying to optimize? If you can parameterize your model somehow (e.g. parameterize edge weights) and then figure out some way to measure how "bad" your model is, you can use autodiff to tune your edge weights to minimize that metric. Not too familiar with DAGs but the first step is figuring out how to parameterize your model in a way that tuning the params can lead to your goal (and how to measure how close you are to that goal)
I like your comment on thinking about implementing models/thinking about them as parameterized solutions
I remember JAX (the XML parser) really saving my bacon when I was parsing larger-than-memory XML files ages ago.
I bet it was just a matter of the wrong tool for the job, but I (with some weird fondness, I might add), remember finding it extremely overengineered, even for Java standards.
It's a somewhat fair comparison; in my experience, highly optimized JAX matches highly optimized Tensorflow. However, non-optimized (but JITted) JAX beats non-optimized Tensorflow, as Tensorflow requires a lot of architectural changes to make it perform well. JAX, on the other hand, tends to perform well as long as you just JIT it. So it's much easier to get to, say, 90% of optimal performance. In Tensorflow, it's much harder (in my experience- maybe I'm just bad at Tensorflow).
The JIT compilation that JAX does is really, really good, as it combines operations together in a highly performant way.
IIRC, you guys standardized TensorFlow several years back. What does the current split look like between JAX and TF internally? Do some people use TF and some use JAX, or do you use JAX for specific tasks?
In [9]: x = np.random.randn(10000,10000).astype('f')
In [10]: %timeit -n5 fn(x)
623 ms ± 9.31 ms per loop (mean ± std. dev. of 7 runs, 5 loops each)
With Jax In [17]: %timeit -n5 jax_fn(x).block_until_ready()
98.4 ms ± 1.55 ms per loop (mean ± std. dev. of 7 runs, 5 loops each)
It's still 6x improvement but not as large that the article claims. I am on the latest Intel MacIf you fix the benchmarks then looks like this
5 loops, best of 5: 99.2 ms per loop
10 loops, best of 5: 114 ms per loop
10 loops, best of 5: 20.2 ms per loop
5x faster is to be expected as there are 5 pointwise operations (that are bandwidth bound) that can be fused.
The leading comparison is also quite misleading, imo, since I think it's comparing Numpy on CPU vs. Jax on an accelerator.
As for the other part about the leading comparison - I was trying to highlight just how much faster JAX could be in the best-case scenario. Beyond the accelerator and JIT, the function itself lends to being expedited significantly when JITted. I posted benchmarks with a comparison of JAX vs NumPy both on CPU, and then with JAX on TPU further down to control more variables. (reposted from reddit)
"(n.b. JAX is using TPU and NumPy is using CPU in order to highlight that JAX's speed ceiling is much higher than NumPy's)"
A much more interesting comparison would be CuPy or Pytorch code vs Jax code running on A100.
Obviously, the ability to use an accelerator makes JAX faster, but even without that on CPU it was faster. This is in part because of JIT, and in fairness the calculation in question does lend itself well to being expedited by JIT.
I actually have some preliminary benchmarks for a follow up specifically on just NumPy vs JAX, and it has become clear so far that NumPy is better in certain cases, especially for small operations where the overhead of JAX is not worth it.
In the article I mention this briefly, along with how JAX hasn't been focused on being optimized on CPU because they have bigger fish to fry, so to speak. I also link to the JAX documentation that has some comments comparing the two!
Relating to TF - I don't actually use TF at any point, but I did use PyTorch for Hessian calculation. TF and PT obviously do both work on GPU, but JAX has the benefit of being able to JIT more and implement everything in terms of XLA (although TF obviously has XLA support as well, and PT kind of does but just to get PT working on TPU).
Thinking about doing another article on a direct comparison of JAX with PT and TF - let me know if that's something you'd like to see!
If you are running your benchmark on a single Numpy operation then I would expect Numpy to have equal or better speed (you are not paying for JIT compilation). However, when you are doing several operations, Numpy will do a loop on array elements for each operation while JAX can fuse everything, that can end up making a big difference.
It's clear so far that NumPy can outperform JAX on CPU for small computations (unsurprisingly). Please let me know if you'd be interested in seeing a more thorough analysis!