Y
HN Search
Hacker News Search
new
|
comments
|
top
|
jobs
mattjjatgoogle
searching PlanetScale…
1.
▲
2.
▲
3.
▲
4.
▲
5.
▲
6.
▲
4 ms
·
1.
▲
by
mattjjatgoogle
2y ago
An author's tweet thread: https://x.com/jacobaustin132/status/1886844716446007300
2.
▲
How to scale your model: A systems view of LLMs on TPUs
(jax-ml.github.io)
185 points
by
mattjjatgoogle
2y ago
|
30 comments
3.
▲
by
mattjjatgoogle
3y ago
Actually that never changed. The README has always had an example of differentiating through native Python control flow: https://github.com/google/jax/commit/948a8db0adf233f333f3e5f... The constraints on cont
4.
▲
by
mattjjatgoogle
3y ago
Actually, that's never been a constraint for JAX autodiff. JAX grew out of the original Autograd ( https://github.com/hips/autograd ), so differentiating through Python control flow always worked. It's jax.jit
5.
▲
by
mattjjatgoogle
3y ago
You're right! Maybe we should revise that... I made https://github.com/google/jax/pull/17851 , comments welcome!
6.
▲
by
mattjjatgoogle
3y ago
Have you seen JAX MD? https://github.com/jax-md/jax-md
7.
▲
by
mattjjatgoogle
4y ago
You're right that downstream libraries have often tended to introduce magic (some more than others), and moreover one library's magic is typically incompatible with other libraries'. It's something that we're workin
8.
▲
by
mattjjatgoogle
4y ago
Thanks for taking the time to explain these. > 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'
9.
▲
by
mattjjatgoogle
4y ago
If you have any particular examples in mind, and time to share them on https://github.com/google/jax/issues , we'd love to try to improve them. Improving error messages is a priority. About introspection tools
10.
▲
by
mattjjatgoogle
4y ago
Thanks for the kind words! We've been doing a lot more work in this direction too, for both compiler-based automatic parallelization [0] and a work-in-progress pmap successor for 'manual' parallelism (per-device code and expl
11.
▲
by
mattjjatgoogle
4y ago
I'm sure there's a lot of good material around, but here are some links that are conceptually very close to the linked Autodidax. (Disclaimer: I wrote Autodidax and some of these other materials.) There's Autodidact [0], a pr
12.
▲
by
mattjjatgoogle
6y ago
The difference is that in TF1 you had to use tf.cond, tf.while_loop etc for differentiable control flow. In JAX you can differentiate Python control flow directly, e.g.: In [1]: from jax import grad In [2]: def f(x): ...:
13.
▲
by
mattjjatgoogle
7y ago
No worries! I didn't mean it as a correction so much as just a discussion; I'm sure it's true that other autodiff systems have very sophisticated automatic remat (like https://openreview.net/forum?id=BkYYXJ9i-
14.
▲
by
mattjjatgoogle
7y ago
Rematerialization in autodiff is super interesting! XLA does rematerialization optimizations, so you get those automatically under jax.jit. There's also the jax.checkpoint decorator ( https://github.com/google/jax&#