HNHacker News
TopNewBestAskShowJobs

patrickkidger

317 karma · joined March 13, 2022

submissionscomments
patrickkidger··on JAX – NumPy on the CPU, GPU, and TPU
Indeed, it is much better to use jit(grad(f)) in general.

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.

patrickkidger··on JAX – NumPy on the CPU, GPU, and TPU
I'm going to disagree here! Classes and functional programming can go very well together, just don't expect to do in-place mutation. (I.e. OO-style programming.)

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.

patrickkidger··on JAX – NumPy on the CPU, GPU, and TPU
I feel like I see the opposite -- that everything scientific computing is getting rewritten in something autodifferentiable! Whether that's JAX or something else.

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.)

patrickkidger··on Python Type Hints – *args and **kwargs (2021)
FWIW, I've come to regard this (cooperative multiple inheritance) as a failed experiment. It's just been too confusing, and hasn't seen adoption.

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.)

patrickkidger··on Pytrees
Incidentally `diffrax.PIDController` also has a `dtmin` argument that could probably be used here instead of re-running things. :)
patrickkidger··on Pytrees
This is awesome to hear!

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.

patrickkidger··on Pytrees
You're thinking of `jax.closure_convert`. :)

(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...

patrickkidger··on Pytrees
Shameless advert -- Equinox is a neural network library for JAX based entirely around pytrees:

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.

patrickkidger··on Writing Python like it's Rust
Try using jaxtyping: https://github.com/google/jaxtyping.

It also supports numpy/pytorch/etc.

patrickkidger··on How big should a programming language be?
I think of Python as being a pretty big language.

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.

patrickkidger··on Framework announces AMD, new Intel gen, 16“ laptop and more
+1 to the sibling comments for a split keyboard option.

(Preferably a proper version with thumb clusters.)

At least for me, that would make this laptop a must-buy.

patrickkidger··on AstroNvim is an aesthetic and feature-rich Neovim config
Heads-up that soft-wrap is now supported at HEAD!

https://docs.helix-editor.com/master/configuration.html#edit...

patrickkidger··on AstroNvim is an aesthetic and feature-rich Neovim config
There is the option to `auto-format` on save, if that helps.

(FWIW I use neither, and format via pre-commit hooks instead of in-editor.)

patrickkidger··on AstroNvim is an aesthetic and feature-rich Neovim config
FWIW, Helix does have tabs. (Set `editor.bufferline = True` in `config.toml` in order to have them listed.)
patrickkidger··on Be Careful Using Tmux and Environment Variables
This can be changed: I use zellij and just rewrote its config file to place almost all the keys in a single mode, basically 1-1 with tmux.

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.

patrickkidger··on Training Deep Networks with Data Parallelism in Jax
It sounds like you're concerned about how downstream libraries tend to wrap JAX transformations to handle their own thing? (E.g. `haiku.grad`.)

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.

patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
Aha, I have probably read that sentence literally hundreds of times, and never spotted the typo. I will never be able to unsee that.

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.

patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
I think those are all things you need for life.

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.

patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
As the author of this article... I have read maybe one textbook cover-to-cover in my life. :D

(Hands-on machine learning, by Geron, back when I made the jump math->ML.)

patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
Right! This was the situation for me.
patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
So HPC means like ten different things.

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.

patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
Completely agreed!
patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
Haha, so I actually have an atrocious memory! Famously so amongst my friends, I never remember what we've discussed.

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!

patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
Ooh, that's a great suggestion!

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.)

patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
This is a great description, thank you! What you've said is precisely the reason I emphasised knowing a bit of foundational math, e.g. topology.
patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
I'm sorry it came across this way for you! Rather, I'm just outlining why folks seem to keep asking me this question. :)
patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
I'm sorry to hear that things didn't go so well for you!

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. :)

patrickkidger··on Just know stuff (or, how to achieve success in a machine learning PhD)
This is a great point I didn't cover! "Just know stuff" tends to follow naturally from "care about stuff".
patrickkidger··on Numba: A High Performance Python Compiler
Honestly, the two are now incredibly close.

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.

patrickkidger··on PyTorch 2.0
> If you care about performance

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.

← PreviousPage 2 of 3Next →