JAX – NumPy on the CPU, GPU, and TPU
jax.readthedocs.io
jax.readthedocs.io
When I was in astronomy (about a decade ago) I did large scale simulations of gravitational interactions. But at the time all these simulations were done on CPU. Some of the really big efforts used more specialized chips, but it was a huge effort to write the code for it.
But today with Jax, if you want to write an N-body simulation of a globular cluster, you can just code it up in numpy and it'll run on a GPU for free and be about 1000x faster. From what I can tell though, very few people in the sciences have caught on yet.
And, more importantly, things that cannot be expressed that way tend to not be a good fit for GPU computing anyway (independently of the language / framework you are using).
[0]: `array` is a shortcut here, JAX is not limited to operations on arrays.
I'll have to disagree with you a little bit here. SIMT model of GPUs are quiet a bit more expressive than the numpy's SIMD model. As an obvious example, you'll have to manually maintain a mask to implement if/else i.e. code path divergence in SIMD. GPUs automatically does this and many more to make your life easier. And frankly, I find it lot more easier to reason about what should happen to one data point than a bunch of them together.
An interesting article I read recently that has some relevance to this discussion. https://pharr.org/matt/blog/2018/04/18/ispc-origins
To be fair though, modern GPUs are pretty good at branching and latency hiding, while numpy-style code has poor data locality unless you have a magic compiler.
Having access to high performance explicit loops and ifs/masks allows one to focus on the hard parts of the algorithms, rather than on the purely incidental puzzle how to best avoid spending time in the Python runtime.
When you chain a sequence of vectorized operations on arrays, loop fusion would save you from allocating memory for each intermediate variable, and the round trip time of moving it from RAM to CPU multiple times. I don’t know how good JAX’s JITted loop fusion is on CPU, but I’ve been very very impressed by Julia.
Eg: I had some Numpy code that took hours (and needed terabyte RAM) that was very straightforward to code in Julia, and needed only a few GB to finish in a few seconds — on my laptop.
I want to be able to think in arrays, but to also not have to materialize the arrays as much as possible.
Depening on the workload this is no problem. If you have many cheap iterations, you will notice the overhead.
I am not sure if they are working on fusing scan and what's the current status.
Back in the tensorflow days, I had this issue and submitted a patch that gave a ~50x speedup for my usecase. It's always better to optimize the base function rather than have 100 people all manually working around the same performance issue.
just sharing for those who want to learn more
you seem to have very precise knowledge of the SLOC at a point in time - just curious is there any tooling you used to do that? that can be pretty nifty to pull out on occasion
I am also continously surprised how little adoption JIT and autodiff libraries have gotten in scientific computing. A lot of my colleagues somehow really like coding cost function gradients and fine-tuned GPU code by hand. I guess using something like JAX can reduce your standing in the group, because it can make it seem like coding algorithms is pretty easy.
My experience might be biased though: shameless advert for Equinox (https://github.com/patrick-kidger/equinox, 1.4k GitHub stars), which is now the foundation of quite a lot of SciComp in JAX. (Both internal and open-source.)
...if that SciComp uses machine learning, I guess? In my "physics of biomedical imaging" bubble, people are hardly doing state-of-the-art ML, but rather expensive forward models for which computing a gradient is cumbersome.
But I know that e.g. Stephan Hoyer is a physicist and you are a mathematician originally -- I have read a lot of your JAX issues and libraries ;-) maybe it just depends on the "mini-bubble' aka. the indiviual research group and not only the field of science.
Not necessarily! It's perfectly possible (and quite common) to e.g. write down a traditional parameterised ODE, and then optimise its parameters via gradient descent. Compute the gradients wrt parameters using autodiff through the numerical ODE solver. All without a single neural network in sight! ;)
My usual spiel is that autodiff+autoparallel are really useful for any kind of numerical computation -- of which ML is a (popular, well funded) special case.
At least in my mini bubble, these kinds of "scipy but autodifferentiable" use-cases are fairly common.
> I have read a lot of your JAX issues and libraries ;-)
Haha, that's fun to hear though! Thank you for sharing that.
In general, I very much agree that "autodiff+autoparallel are really useful for any kind of numerical computation". And the use cases are also really common in my bubble. It's just that (imho) most people have not realized this.
I’d appreciate any pointers to the literature; curious to see the kinds of models people work with. Thanks!
In all conferences like NeurIPS, in Google ML Community days, etc., whenever there is a JAX workshop/tutorial/talk, it is always touted as a numerical computation library. And it was developed as such. Sure the focus is in ML, but everyone involved in it always have said that this is a general purpose scientific computing library.
Flax, Haiku, etc. are Deep Learning libraries.
> JAX is Autograd and XLA, brought together for high-performance machine learning research.
That does not really convey the generality of it that well.
Are there any benchmarks for that? Running on GPU never comes for free. You have to transfer data back and forth which has a cost, for instance.
Is anyone making any serious progress in fast GPU based computational tools for other faster languages? I'm looking for something that also works on the GPU on windows (unlike JAX)
It’s because most of the people doing these computations don’t have the capacity to become experts in multiple fields. They understand the math and analytics very well, and they expend all their time thinking about that, not about type systems, memory management, etc. Python lets them code without having to think about a lot of that stuff so they can focus on the things they care about. These aren’t computer scientists or programmers, they’re meteorologists, astronomers, oil and gas analysts, investment bankers etc. That’s why some truly great computer scientists and programmers invested their time into building these tools for python vs other languages.
I'm guessing this was mostly Fast Multipole Method? I don't think it ports that easily GPU since there's so much communication involved and the leaves don't do a whole lot
x[0] = 10
Instead I have to do: y = x.at[0].set(10)
Of course this has advantages, but you can't then go and claim that JAX is a drop in replacement for numpy, because this such a fundamental change to how numpy developers think (and in this regard, PyTorch is closer to numpy than JAX).Why can't you do the first in functional programming (not in this specific case because it's just how it is, but in general)?
And even if you can't do so for any reasonable reason in functional (again, in general), what stops us to just add syntactic sugar to equal it to the second to make programmer's life easier?
We could indeed introduce syntactic sugar (`y= (x[0]:=10)` maybe), but you'll still need to introduce a new variable to hold the modified list.
- higher order functions (lambdas, currying, closures, etc.)
- pure functions, immutability by default, side effects are pushed to the top level and marked clearly
The first aspect of functional programming has been already accepted by most OOP languages (even C++ has lambdas and closures).
The second aspect of functional programming is what makes it useful on GPU (because GPU architecture that makes it so powerful requires no interactions between code fragments that are run in parallel on 1000s of cores). So you can easily run pure functional code on GPU, but you can't easily run imperative code on GPU.
You can introduce side effects to functional programming, but then it ceases to be any more useful for GPU (and other parallel programming) than imperative/OOP.
https://github.com/explosion/thinc/blob/master/thinc/backend...
Though I guess the question is why one would still use NumPy when there are good libraries for CPU and GPU. Maybe for interop with other libraries, but DLPack works pretty well for converting arrays.
e.g, `x[0] = 10` is the same as `x.__set_item__(0, 10)`, so there shouldn't be any technical limitation to using `x[0]` (says the guy who never even imported jax)
I somehow completely missed the assignment part of the second example.
Thank you for the clarification.
If no, you're just passing things between functions, then go ahead with Jax! But converting larger codebases with classes is just significantly better with PyTorch even if they use different method names etc.
You might like Equinox (https://github.com/patrick-kidger/equinox ; 1.4k GitHub stars) which deliberately offers a very PyTorch-like feel for JAX.
Regarding speed, I would strongly recommend JAX over PyTorch for SciComp. The XLA compiler seems to be much more effective for such use cases.
class JaxWrapper:
def __init__(self, arr):
self.arr = arr
def __setitem__(self, key, val):
return self.arr.at[key].set(val)
....The big issue I had was: I was developing on the CPU, then moved to running it on a GPU, and it wasn't as fast as I expected-- I started debugging, and saw there was still lots of communication between the CPU and GPU even tho it was all jit'd. I think PyTorch is a more user friendly for writing high performance models if you're not straying too far from the beaten path. But I really love JAX would like to play around with it more to understand these pits I'm falling into.
And another complaint is I can't run it on my Macbook M1 GPU... but I'm seeing this page now, so maybe that's not true anymore: https://developer.apple.com/metal/jax/
Naturally, I pay close attention to Jax and give it a look every now and then. So, I’ll focus my observations below on Jax’s Numpy API support.
At a glance, Jax code looks like regular Python, but it’s a very different style of programming. Two big differences I’ve found are:
- All Jax functions must be pure. You can’t pass references. - ndarrays cannot be created with dynamic shapes. You have to hardcode the shape tuples. One possible workaround can be to create a buffer much bigger than you need and return that along with actual shape.
Then there are many small things that are very well documented[1] by the Jax team.
Overall, if you are training ML models, the trouble might be worth it (Autograd). But for accelerating Numpy alone, it is no Numba replacement - which will happily work in the above mentioned use-cases.
[1] https://jax.readthedocs.io/en/latest/notebooks/Common_Gotcha...
This is not true. Rather, all shapes have to be known at compile time. That means that output shapes must not depend on input values, but may depend on input shapes -- also explicitely.
Furthermore, there are two useful additions:
1. You can use vanilla numpy for compile-time computations. An example would be computing an array of indices for some moving-window filter, depending on the input shape and a stride parameter.
2. You can mark function arguments as "static". Then their values may change shapes of the output, but accordingly the function is compiled for each value of those arguments.
We don't support Windows GPU because we haven't had the engineer bandwidth to support it well.
We recommend WSL2 for GPU on Windows at the moment because that is a compromise: it allows CUDA support, without us having to support another release variant.
But we welcome community contributions!
We release Windows CPU wheels (https://pypi.org/project/jaxlib/#files). So JAX on CPU works great on Windows.
We don't release Windows GPU wheels at the moment, but that's because we're a small team and none of us use Windows personally. We welcome contributions!
(I verified that the Windows CUDA GPU support built as recently as two weeks ago, but I don't have the ability to test that it works.)
We recommend WSL2 because that's just using our existing Linux CUDA release.
We felt that Windows CPU support was important so everyone can run JAX, even if it's not always the most-accelerated version of JAX. And we got some great PRs from the community that helped fix a few open issues.
[0]: https://developer.apple.com/metal/jax/
It's significantly faster than CPU. Something like 100x using sheet
It also covers things like functional purity in Deep Learning, and handling of random numbers in JAX.
[0]: Learn JAX: From Linear Regression to Neural Networks - https://www.kaggle.com/code/truthr/jax-0
For a lot of people I know whose main job was not to write code, switching from tensorflow to pytorch was something that saved them ten to hundreds of hours in the long run, even accounting for the initial learning time.
TF is not worth it anymore.
I'm working with sequences, e.g. speech recognition, machine translation, language modeling. This is a quite fundamental property for this type of models, that we have variable lengths sequences.
In those cases, for some example code, I have seen that training also used only fixed size dimensions. And at inference time, they had some non-JAX code for the loop over the sequence around the JAX code with fixed-size dimensions.
This seems like a quite fundamental issue to me? I wonder a bit that this is not an issue for others.
I went through the same hiring process and had a positive experience at every stage. I had a strong competing offer but went with the JAX team at NVIDIA.
I'll pass it along as feedback.
Just want to share that Ray (an open source project we're developing at Anyscale), can be used to scale Jax (e.g., across TPUs).
Some docs from Google on how to do this
https://cloud.google.com/tpu/docs/ray-guide
Alpa is an open source project scaling Jax on 1000+ GPUs
https://www.anyscale.com/blog/training-175b-parameter-langua...
Cohere uses Ray + Jax + TPUs to build their LLMs
https://www.youtube.com/watch?v=For8yLkZP5w
A demo from Matt Johnson on the Jax team
PSPP was a long-standing issue and to see a fairly new computational tool used to significantly aid in the process of "solving" it speaks greatly towards its general utility in the sciences.
nvidia dropped cuda support for perfectly good gpu's, showing the perils and waste of being locked-in in a profit-maximazing monopoly.
For now I'm happy with Pytorch->ONNX and then running the ONNX model directly. But as I said, that means I can't easily train using JAX :-(
Nice. When did they make this change?
Here is the old way in the docs, where you needed to define functions for the if-true branch and the if-false branch, and feed them to a conditional function, to get the normal if-then-else conditional.
https://jax.readthedocs.io/en/latest/notebooks/Common_Gotcha...
For vanilla "if", the condition must be known at compile time. For runtime, you have to use "cond", "where", or "select" (which may be analogous).
https://github.com/google/jax/commit/948a8db0adf233f333f3e5f...
The constraints on control flow expressions come from jax.jit (because Python control flow can't be staged out) and jax.vmap (because we can't take multiple branches of Python control flow, which we might need to do for different batch elements). But autodiff of Python-native control flow works fine!
Don’t know about usage and uptake though.
What’s the point of having to explicitly call grad(jit(f))? Why doesn’t grad just call jit internally? Is there a usecase where you want the grad without jit?
Supporting the opposite composition is still useful in some edge cases -- for example when debugging, you want to step through a computation without jit, and simply not crash when differentiating any inner functions also decorated with jit.
To add some colour to my answer. When writing a library, it's typical to a put a JIT statement on everything in the public API. This means you get the benefits of JIT compilation even when you're just hacking around in the REPL, and mitigates the new-user-footgun in which they forget to use JIT themselves.
Meanwhile, good practice is always to JIT your whole computation.
Combined, this mean that it's fairly common to go jit (at the top level) -> grad (of your operation) -> jit (of some library call).
When debugging your code, the JIT'd library call is _probably_ not the culprit. So you only want to disable the top-level JIT when stepping through, and still take advantage of JIT compilation where you can. Overall one obtains a composition of the form grad(jit(...)).
TL;DR: even if use case doesn't come up super frequently, it's more user-friendly to support grad(jit(...)) than it is to just crash.
edit: Made comparison more fair.
In terms of is it worth using it - that depends on what you're doing. If you just want to start with ML training probably not. If you have something already and you want to take it to next level (e.g. influence how training and inference work) than it's a good choice. You might be interested in looking into flax or haiku instead of using vanilla Jax. These are closer to pytorch.
also: