3 ms·
It's not quite a fair comparison, since Numpy is running with float64 while Jax is running with float32. If you fix the benchmarks then looks like this 5 loop
by chillee 5y ago
It's not quite a fair comparison, since Numpy is running with float64 while Jax is running with float32.
If you fix the benchmarks then looks like this
5 loops, best of 5: 99.2 ms per loop
10 loops, best of 5: 114 ms per loop
10 loops, best of 5: 20.2 ms per loop
5x faster is to be expected as there are 5 pointwise operations (that are bandwidth bound) that can be fused.
The leading comparison is also quite misleading, imo, since I think it's comparing Numpy on CPU vs. Jax on an accelerator.
- SleekEagle 5y agoThanks very much for posting this - I forgot JAX defaults to float32 - I'll fix that soon. As for the other part about the leading comparison - I was trying to highlight just how much faster JAX could be in the best-case scenario. Beyond the accelerator and JIT, the function itself lends to being expedited significantly when JITted. I posted benchmarks with a comparison of JAX vs NumPy both on CPU, and then with JAX on TPU further down to control more variables. (reposted from reddit)
- brrrrrm 5y agore: leading comparison, that makes sense. I recommend labeling it a bit better: JAX on TPU vs NumPy on CPU
- SleekEagle 5y agoUpdated! Thanks for the feedback. Added this note in the figure description "(n.b. JAX is using TPU and NumPy is using CPU in order to highlight that JAX's speed ceiling is much higher than NumPy's)"
- p1esk 5y agoComparing Jax on TPU vs Numpy on CPU does not make any sense. Of course a GEMM hardware accelerator will be much faster than a general purpose CPU. What is the point of this comparison? You either run the same code on different hardware, or different code on the same hardware. A much more interesting comparison would be CuPy or Pytorch code vs Jax code running on A100.