3 ms·
(I'm an engineer at DeepMind, and I work with JAX daily) It's a somewhat fair comparison; in my experience, highly optimized JAX matches highly optimized Tenso
by fnbr 5y ago
(I'm an engineer at DeepMind, and I work with JAX daily)
It's a somewhat fair comparison; in my experience, highly optimized JAX matches highly optimized Tensorflow. However, non-optimized (but JITted) JAX beats non-optimized Tensorflow, as Tensorflow requires a lot of architectural changes to make it perform well. JAX, on the other hand, tends to perform well as long as you just JIT it. So it's much easier to get to, say, 90% of optimal performance. In Tensorflow, it's much harder (in my experience- maybe I'm just bad at Tensorflow).
The JIT compilation that JAX does is really, really good, as it combines operations together in a highly performant way.
- SleekEagle 5y agoThanks for this comment, I probably could've stressed JIT more. Random question - I noted in the article that you all at DeepMind announced that you're using JAX to accelerate your research. IIRC, you guys standardized TensorFlow several years back. What does the current split look like between JAX and TF internally? Do some people use TF and some use JAX, or do you use JAX for specific tasks?
- sivakon 5y agoThe first benchmark is wrong. Numpy random by default uses float64, and jax uses float32. In [9]: x = np.random.randn(10000,10000).astype('f') In [10]: %timeit -n5 fn(x) 623 ms ± 9.31 ms per loop (mean ± std. dev. of 7 runs, 5 loops each) With Jax In [17]: %timeit -n5 jax_fn(x).block_until_ready() 98.4 ms ± 1.55 ms per loop (mean ± std. dev. of 7 runs, 5 loops each) It's still 6x improvement but not as large that the article claims. I am on the latest Intel Mac