4 ms·
I find JAX really exciting. The idea of numpy with autograd is exactly what Pythonistas want. The elephant in the room, though, is “why not Pytorch?” Everyone
by jphoward 6y ago
I find JAX really exciting. The idea of numpy with autograd is exactly what Pythonistas want. The elephant in the room, though, is “why not Pytorch?”
Everyone knows JAX is what Google realised Tensorflow should have been when they realised how much of a joy Pytorch was to use. I actually think JAX does offer some advantages, not least true numpy interoperability. However, not mentioning *torch a single time in the blog post seems a little disingenuous for a Google-owned deep learning enterprise.
- orbifold 6y agoWhat is more some of the deep learning libraries on top of JAX (like flax, linen) are close to a one-to-one copy of PyTorch. However the underlying implementation is in some ways more interesting than the one of PyTorch and naturally accomodates higher derivatives and also the way that optimization and state initialization is implemented is more functional and "theoretically sound".
- riyadparvez 6y agoCan you explain how JAX compares to pytorch? AFAIK pytorch also closely resembles numpy API.
- nestorD 6y agoFor me the main point is that JAX is slightly lower level than Pytorch and has nicer abstraction (no need to worry about tensors that might not store a gradient or wetehr you are on the GPU: it eliminates lots of newcomers bug) which makes it a great fit to build Deep learning frameworks, but also simulations, on top of it.
- alevskaya 6y agoJAX and Pytorch have somewhat different scopes though... JAX is concerned about much more than the kinds of neural nets we write today. It's a general system for expressing and transforming numerical programs, and the devs are as genuinely excited about e.g. scientific programming, probabilistic modeling, etc. as they are about NNs. A technical reason for "why not pytorch" is that JAX was also built in part to expose and leverage the power of the XLA compiler, which is at least for the moment a pretty uniquely powerful tool for producing efficient, highly-scalable accelerator code. I should underline that this is a friendly community of peers though: there is a lot of respect for Pytorch, which in turn was certainly influenced by the original Autograd that many of the JAX devs also worked on. JAX (and its fancier sibling Dex) beyond being useful tools are also still research projects in and of themselves seeking to advance our ideas on how to write expressive, powerful numerical code on modern architectures.
- prionassembly 6y agoPytorch might be suboptimal but it's what's already there e.g. for Pyro. (PyMC3 uses Theano; can Theano drive JAX?) OTOH I'm not sure most people know what tasks are GPU-worthy or not. I haven't the slightest idea of why MCMC/Variational Bayes is amenable to GPU speedups and Persistent Homology isn't.
- ihnorton 6y ago> can Theano drive JAX? Yes, that is in active development: https://pymc-devs.medium.com/the-future-of-pymc3-or-theano-is-dead-long-live-theano-d8005f8a0e9b https://pymc-devs.medium.com/the-future-of-pymc3-or-theano-i...
- cl3misch 6y ago> why not Pytorch? JAX enables using (parts of) existing numpy codebases in disciplines other than deep learning. Autodiff and compilation to GPUs are very useful for all kinds of algorithms and processing pipelines.