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.
317 karma · joined March 13, 2022
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.
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.
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.)
Instead, I've come to prefer a style I took from Julia: every class is either (a) abstract, or (b) concrete and final.
Abstract classes exist to declare interfaces.
__init__ methods only exist on concrete classes. After that it should be thought of as unsubclassable, and concerns about inheritance and diamond dependencies etc just don't exist.
(If you do need to extend some functionality: prefer composition over inheritance.)
And if you're doing implicit methods, then you may be interested to hear that today's release of Diffrax now includes IMEX solvers! Sil3, KenCarp3, KenCarp4, KenCarp5.
(Although technically that works by tracing and extracting all constants from the jaxpr, rather than introspecting the function's closure cells -- it sounds like your trick is the latter.)
When you discuss dynamic allocation, I'm guessing you're mainly referring to not being able to backprop through `jax.lax.while_loop`. If so, you might find `equinox.internal.while_loop` interesting, which is an unbounded while loop that you can backprop through! The secret sauce is to use a treeverse-style checkpointing scheme.
https://github.com/patrick-kidger/equinox/blob/f95a8ba13fb35...
https://github.com/patrick-kidger/equinox
(Now on 1.1k stars so it's achieved some popularity!)
This makes model-building elegant (IMO), without any new abstractions to learn. Quite a PyTorch-like experience overall.
It also supports numpy/pytorch/etc.
At a feature level there are things like descriptors and metaclasses, which are complex and rarely-used.
There's a huge number of obscure magic methods/attributes: __origin__, __orig_bases__, __prepare__, __mro_entries__, etc.
It's full of gotchas: methods are looked up on the instance, except magic methods which are looked up on the class, except except __getattr__ and __dir__ on module types. There are both __getattr__ and __getattribute__ and they do different things. Most `typing.Foo` (and also modern `list[int]`) aren't `isinstance`-able. The use of += with tuples. The lack of any coherent numeric tower. Etc. etc.
I think it's very clear that Python has grown organically, and is now struggling under the weight of all its bolted-on extra bits, or sheer unnecessary complexity. I would love to throw it all out and start again.
(Preferably a proper version with thumb clusters.)
At least for me, that would make this laptop a must-buy.
https://docs.helix-editor.com/master/configuration.html#edit...
(FWIW I use neither, and format via pre-commit hooks instead of in-editor.)
That said, I must admit that zellij doesn't add a whole lot to my workflow relative to tmux. I think I'm just easily distracted by the new-and-shiny.
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.
Geron is good but now a bit out-of-date. No transformers (just CNNs/RNNs/etc.) and the coding component is all in scikit-learn and TensorFlow (rather than PyTorch or JAX).
FWIW I did ask this question recently over at https://twitter.com/PatrickKidger/status/1602776438159339521, in case any of the responses are helpful.
But what you do need for specifically a PhD? I argue that "knowing stuff" is what is necessary -- and that indeed it's essentially the purpose of the whole academic institutiom.
(Hands-on machine learning, by Geron, back when I made the jump math->ML.)
For example another commentor mentions low-latency concerns in finance, and that's something I have zero experience with.
HPC has often also meant writing a lot of C++ to do e.g. MD or something.
These days, I consider myself HPC-adjacent -- I write scientific ML software, often for use on pretty beefy hardware (TPU pods etc.) So at least for that, here's an off-the-cuff list of a few items that come to mind:
- Know JAX. Really, really well: its internals, how its transforms work. It's definitely a bit bumpy in places, but it's still one of the best things we have for easily scaling programs, e.g. through `jax.pmap`, being able to test on CPU and then run on TPU, etc.
- Triton! New(-ish) kid on the block for GPU programming.
- How CPUs work: L1/L2/L3 caches, branch prediction, etc. Parallelism via OpenMP.
- How GPUs work: warps etc.
- How BLAS works (e.g. tiling)..
- Compiler theory. Inlining functions, argment aliasing, NRVO, ...
- Know autodiff well. E.g. have a read of the Dex paper, and the concerns with doing autodiff through index operations. Modern scientific computing is moving towards a ubiquitously autodifferentiable future.
- ... plus loads more, haha. Probably I'm still missing ten different things that another reader considers crucial.
When I wrote this list I certainly wasn't expecting/recommending all of this to stay in the reader's head forever.
Rather: if you've worked with something deeply at one time, then -- even if you've forgotten the details -- you still can pattern-match on it later. And then look up whatever you've forgotten!
So one thing I learned a lot of in my PhD (for literally a whole year), that I literally never needed, was functional analytic methods for PDEs. Stuff like Moser iterations / the De Giorgi-Nash-Moser theorem, etc.
The finer details of Turing machines have never really helped me, although in my case that's probably the exception as I imagine that's still pretty important.
On a more ML note, I have literally never needed SVMs. (And hope I never get asked about them, I've forgotten everything about them haha.)
I think there's a lot of other stuff I could add to the "just-don't-know-stuff" list!
(And to answer your last question: this list is curated, and based on the criteria of (a) is it useful, and (b) is it widely applicable.)
FWIW being able to "program like nobody's business" is still really really valuable. It's why I dedicate such a large chunk of the post to software dev skills. :)
JAX introduced a lot of cool concepts (e.g. autobatching (vmap), autoparallel (pmap)) and supported a lot of things that PyTorch didn't (e.g. forward mode autodiff).
And at least for my applications (scientific computing), it was much faster (~100x) due to a much better JIT compiler and reduced Python overhead.
...but! PyTorch has worked hard to introduce all of the former, and the recent PyTorch 2 announcement was primarily about a better JIT compiler for PyTorch. (I don't think anyone has done serious non-ML benchmarks for this though, so it remains to be seen how this holds up.)
There are still a few differences. E.g. JAX has a better differential equation solving ecosystem. PyTorch has a better protein language model ecosystem. JAX offers some better power-user features like custom vmap rules. PyTorch probably has a lower barrier to entry.
(FWIW I don't know how either hold up specifically for DSP.)
I'd honestly suggest just trying both; always nice to have a broader selection of tools available.
This definitely isn't true. On any benchmark I've tried, JAX and Julia basically match each other. Usually I find JAX to be a bit faster, but that might just be that I'm a bit more skilled at optimising that framework.
Anyway I'm not going to try and debunk things point-by-point, I'd rather avoid yet another unpleasant Julia flame-war.