Training Deep Networks with Data Parallelism in Jax
mishalaskin.com
mishalaskin.com
I've become a big believer that it would be beneficial for PL research in ML which makes heavy use of program transformations to provide small JAX-based implementations. There's really no other system which allows you to express interpreter-based transformations with the benefits that JAX provides (maybe `functorch` in a few months? I have some doubts of transformation composition with systems like torchdynamo - but I don't know much about it)
Edit: note this is coming from a long time Julia stan, take that for what it is worth :)
There’s another mechanism in Julia - generated functions. These allow method body specialization given type knowledge about the signature of the function — so a user can write code which is generated for the method body when inference determines the signature (and the inferred signature is tight enough) which depends on the inferred types.
All of Julia’s program transformation based AD packages are based on the latter transformation — most of them do terrible things to the compiler, including massively blowing up the size of code before optimization.
The only package which is more promising is Diffractor — but I’m not convinced it is more than a research prototype at its current level of development. That may change. This was written by one of the compiler devs, and uses lower level hooks into the compiler, developed to support its transformation.
The big issue in general: Julia doesn’t let you write transformations on its typed IR from user space, unless you want to ignore Julia’s native execution engine. There are hooks that someone can work with — but they aren’t user-facing (for all but the most advanced users) — and they break pass composability with code generation using the native engine (this may have changed since I last looked at this!) I would know, because I’ve spent several attempts trying to do stuff like this, and making crappy, unstable packages :)
Separately: macros are one level of reflection -> code generation. JAX supports a different form — you can’t emit data which represents generic expressions — it’s not quite like Lisp in that sense. It’s better to think about JAX as a “two-level” language system — where you have a meta-level which is Python, and there’s a statically typed array language which is the object level. JAX supports a stage -like operation which allows transforming compat subset of Python to the statically typed array language. But you can write interpreters in the full dynamism of Python — as long as the “active paths” (under tracers) are in that compat set, you can then stage out applications of the interpreters on Python functions, etc.
JAX provides one solution to the composable transformation problem, and they’ve done it in an elegant way - that's ~pretty~ easy to understand (c.f. the Autodidax tutorial). With my current knowledge of things, I can’t effectively argue that Julia supports the same right now (caveat: things may have changed since I last had a look). This is an area where a lot of stuff seems to be going on in Julia, so I doubt it will remain this way forever.
> Julia doesn’t let you write transformations on its typed IR from user space, unless you want to ignore Julia’s native execution engine. There are hooks [...] but [...] they break pass composability with code generation using the native engine
Can you elaborate on that? What is pass composability? Are we talking about LLVM's IR or is there another specific to Julia?
When I discussed writing passes, I was referring to interacting with these two phases of the compiler (abstract interpreter and optimization). In practice, these two phases are interlinked.
Shuffling data between these two phases make a lot of assumptions which are mostly opaque to users. Like I said, you can hack it — but it’s hard to learn what you need to know, and there’s not a stable interface or a nice “this is how you write a transformation to operate on this IR” or “this is how you write a custom opt”.
In any case, I’m not totally convinced that it’s a good idea to expose this stuff to user libraries. Or, at least, it needs to be carefully thought about.
See some of the complaints about “magic” in this post for some of that. I’m just fascinated by this stuff for some weird reason.
> What is pass composability?
Pass composability is something that comes up in julia a lot because our custom compiler passes are often done in the user space and have all sorts of interesting applications. The idea is just that we want to have multiple program transformations occuring at once.
I.e. suppose I'm using some sort of program transformation to turn regular code into derivative code with automatic differentiation (AD(, but suppose I *also* want to performing a program transformation in order to generate say GPU code, or I want to perform a program transformation that replaces all my heap allocations with allocations onto a Bump Allocator, or something else. One has to take care to make sure these different transformations can cooperate with eachother. Hell, it can even occur when one wants to do higher order AD that you have to stack two AD passes on top of eachother.
One problem here is that layering passes on top of eachother can cause a combinatorial explosion in the amount of generated code if things aren't being pruned or optimized between passes.
_________________________________
> Are we talking about LLVM's IR or is there another specific to Julia?
The person you were talking to was referring to Julia's own untyped and typed IR's respectively. Julia programs go through quite a few different forms of representation before they end up getting run. The pipeline looks like this:
1) String: Just a regular, unparsed string of text.
2) Expressions: this is a user facing representation of parsed code that our macros operate on. At this level, all that's really been done is parsing and a bit of canonicalization. There is no name or scope resolution done at this level, and everything is in terms of trees.
3) Untyped IR: This is a not-so-user-facing intermediate representation of julia code that is produced after an Expression tree gets linearized into SSA form. This has had name and scope resolution performed on it, but no type inference or optimization passes passes performed on it. Generated functions and various user-level compiler pass injection techniques are able to operate on this level of julia representation.
4) Typed IR: This is actually the same object as untyped IR, just with slots that used to be empty filled in. It has had type inference performed on it, and many of our custom julia optimization passes performed on it. The types here still correspond to julia level types. Ideally, we'd be doing user level pass injections on this level of IR where types are resolved, performing optimizations using those types to prune down the amount of code, and then performing the next program transformation, and so on.
5) LLVM IR: The next step after we're done with the typed IR is to translate it down to LLVM IR. This involves replacing julia types with LLVM types, and a bunch of other stuff. LLVM will then perform its own optimization passes (of our choice) on this IR. Some packages do program transformations on this level of code, for instance Enzyme.jl. One advantage of this is that the work can be easily shared with other LLVM backed languages.
6) Assembly code: The LLVM IR then gets compiled to assembly with involves yet more optimization and translation passes.
All I can see about program transformations in Jax is [1], and it appears to me there are 4, grad, jit, vmap and pmap. It seems to me you are implying that there are ways to create custom transformations, and this has actual use cases.
Would you mind giving some more details? Or maybe some links? I can't help but be excited by your enthusiasm, it feels like Jax could be the ultimate programming language.
Note that JAX enforces certain limitations, so you should be careful when considering JAX to be "the perfect language" - in general I don't think this is true. It's quite good at what it's designed for.
[0] https://jax.readthedocs.io/en/latest/notebooks/Distributed_a...
[1] https://jax.readthedocs.io/en/latest/jep/14273-shard-map.htm...
About introspection tools, at least for runtime value debugging there is to some extent a fundamental challenge: since jax.jit stages computation out of Python (though jax.grad and jax.vmap don't), it means standard Python runtime value inspection tools, like printing and pdb, can't work under a jax.jit as the values aren't available as the Python code is executing. You can always remove the jax.jit while debugging (or use `with jax.disable_jit(): ...`), but that's not always convenient, and we need jax.jit for good performance.
We recently added some runtime value debugging tools which work even with jax.jit-staged-out code (even in automatically parallelized code!), though they're not the standard introspection tools: see `jax.debug.print` and `jax.debug.breakpoint` on https://jax.readthedocs.io/en/latest/debugging/index.html and https://jax.readthedocs.io/en/latest/debugging/print_breakpo....
If you were thinking about other kinds of introspection tooling, I'd love to hear about it!
That's handy, and I hadn't seen it before, thanks.
It's been a bit, but I think the most frustrating errors were around mapping pytrees (like this issue https://github.com/google/jax/issues/9928). I'm not sure the exact solution, but the axis juggling and specifications were where I remember a lot of pain, and the docs (though extensive) were unclear. At times it feels like improvements are punted on in the hopes that xmap eventually fixes everything (and xmap has been in experimental for far longer than I expected).
Also the barriers where I couldn't disable jit. IIRC pmap automatically jits, so there was no way to avoid staging that part out. When it came to doing some complex jax.lax.ppermute, it felt more difficult than it needed to be to debug.
Next time I encounter something particularly opaque, I'll share on the github issue tracker.
> It's been a bit, but I think the most frustrating errors were around mapping pytrees (like this issue https://github.com/google/jax/issues/9928).
We've improved some of these pytree error messages but it seems that vmap one is still not great. Thanks for the ping on it.
> Also the barriers where I couldn't disable jit. IIRC pmap automatically jits, so there was no way to avoid staging that part out.
That was indeed a longstanding issue in pmap's implementation. And since people came to expect jit to be "built in" to pmap, it wasn't easy to revise.
However, we recently (https://github.com/google/jax/pull/11854) made `jax.disable_jit()` work with pmap, in the sense that it makes pmap execute eagerly, so that you can print/pdb/etc to your heart's content. (The pmap successor, shard_map (https://jax.readthedocs.io/en/latest/jep/14273-shard-map.htm...), is eager by default. Also it has uniformly good error messages from the start!)
> Next time I encounter something particularly opaque, I'll share on the github issue tracker.
Thank you for the constructive feedback!
Higher order functions are difficult in general, and it would be fantastic to have core patterns or tools for breaking them open.
If so, then allow me to make my usual advert here for Equinox:
https://github.com/patrick-kidger/equinox
This actually works with JAX's native transformations. (There's no `equinox.vmap` for example.)
On higher-order functions more generally, Equinox offers a way to control these quite carefully, by making ubiquitous use of callables that are also pytrees. E.g. a neural network is both a callable in that it has a forward pass, and a pytree in that it records its parameters in its tree structure.
1. as you say, exposing patterns and tools for library authors to implement transformations/higher-order primitives using JAX's machinery rather than requiring each library to introduce bespoke magic to do the same;
2. adding JAX core infrastructure which directly solves the common problems that libraries tend to solve independently (and with bespoke magic).
As a first impression the tower of abstraction (Python - NumPy - functional subset JIT) looks a bit too multilayered for a sustainable foundation.
The beauty of JAX is that basic usage is basically a single function: `grad`.
You just write whatever Python function you want and can get the derivative/gradient of it trivially. It gets a bit trickier when you need more sophisticated numeric tools like numpy/scipy, but in those cases it's just about swapping out with a JAX version of those.
In this sense JAX is the spiritual success to Autograd. However the really amazing thing about JAX is that not only do you get the autodiff for basically free, you also get very good performance, and basically GPU parallelism without needing to think about it at all.
PyTorch is an awesome library, but largely focus on building Neural Networks specifically. JAX should be thought of a tool that basically any Python programmer can just throw in there whenever they come across a problem that benefits from having differentiable code (which is a lot of cases once you start thinking about differentiation as a first class feature).
I don't get the point of this distinction - JAX was developed specifically for ML. What else is it being used for right now?
It’s probably also useful in population and metaheuristic scenarios where the optimisation objective can be described mathematically, allowing you to make use of GPGPUs, and if possible first and second order derivatives.
Although, I've mainly heard about Julia in that context, not Jax.
But in general, I would suspect youth.
Certain inference serving solutions like Nvidia Triton Inference Server will even take an ONNX model and then do TRT compilation (with cache!) on the actual inference hardware dynamically at model load time. This is really nice because you can deploy a standard ONNX model across instances and varying GPU hardware and always get TRT optimized and compatible with Compute Capability, etc. Really handy and basically comes down to a few lines of config in the model configuration.
I'm not terribly familiar with JAX but I have to imagine there's ONNX export or straight to TRT export somewhere.
There's some effort going into systems for saving and restoring the computation graph for Jax programs, which will help a lot. I'm surprised it didn't happen sooner, as it seems like quite a natural fit with the jax architecture.