5 ms·
Is it fair to say JAX is what Google made when they looked at Pytorch/Autograd and thought, "oh damn, that's what we should have done?". If so, is this the beg
by jphoward 7y ago
Is it fair to say JAX is what Google made when they looked at Pytorch/Autograd and thought, "oh damn, that's what we should have done?".
If so, is this the beginning of the end of Tensorflow? I know Tensorflow is still top for production, but it is certainly rapidly losing followings in the research field, and Pytorch and now starting to focus on deployment as they know this is their weakness.
- cl3misch 7y agoFYI, the core team behind JAX are the autograd people.
- 317070 7y agoYes, I think Jax is indeed the nail for tensorflow. It's not there yet, but the part of the research community that did not go to pytorch is going to jax now. And jax is made by the autograd people.
- nestorD 7y agoYes but its not mature enough to kill Tensorflow in the short term. I still see people prefering Tensorflow over Pytorch because they have the feeling that it is more mature for production use. Meanwhile Jax has not converged on a recommended deep neural network framework (it has the low level pieces). At the moment its a great building block that researcher should probably know.
- kragen 7y agoIt sounds like JAX is necessarily storing your whole calculation in memory, so it will necessarily use more memory for automatic differentiation of heavily iterative calculations, while other implementations of backward-mode automatic differentiation can instead restart your calculation from checkpoints to avoid storing the whole thing. This could be an advantage of several orders of magnitude for some calculations: using twice the CPU or GPU time in exchange for one thousandth or one ten-thousandth of the memory.
- mattjjatgoogle 7y agoRematerialization 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/pull/1749 https://github.com/google/jax/pull/1749) which lets you control reverse-mode checkpointing yourself; you can use it recursively to implement sophisticated checkpointing strategies (see Example 5 in that PR, which is the classic strategy for getting memory cost to scale like log(N) for iteration count N but requiring log(N) times as much computational work). It'd be interesting to experiment with heuristics for deploying those strategies automatically (e.g. given a program in JAX's jaxpr IR) but one of JAX's core philosophies is to keep things explicit and give users control through composable APIs. Automatic heuristics can be built on top. Another goal is to make JAX a great system for playing with things like this!
- kragen 7y agoThank you for the correction! I should have checked out the software before posting my incorrect surmises from the blog post. It sounds awesome!
- mattjjatgoogle 7y agoNo 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- https://openreview.net/forum?id=BkYYXJ9i-). I'm hoping as users push JAX on new applications, especially in simulation and scientific computing, we'll learn a lot and be able to improve! There's also "cross-country optimization" (https://www-sop.inria.fr/tropics/slides/EdfCea05.pdf https://www-sop.inria.fr/tropics/slides/EdfCea05.pdf) for mixing some forward-mode into reverse-mode to improve memory efficiency. Analogously to jax.checkpoint, we've only experimented with exposing that manually (in jax.jarrett, named because of https://arxiv.org/abs/1810.08297 https://arxiv.org/abs/1810.08297), and even then only for a special case. There's a lot to learn about, experiment with, and build!