It runs about as fast as any of the other popular machine learning frameworks, occasionally faster.
Disclaimer: I work for Google and use JAX, although I'm not on the Jax team.
I've yet to see anything get "a lot faster" because of XLA. It's a ton of complicated code, but then you end up spending the vast majority of time in NVIDIA's cuDNN anyway, so any benefits you might have hoped for will be marginal at best.
Almost double speedup for FP16 Resnet-50.
In fact, also seems to be outperformed by plain PyTorch using a single V100: https://github.com/NVIDIA/DeepLearningExamples/tree/master/P...
Nvidia might have eliminated any potential data pipeline bottlenecks (with careful DALI tuning), but I'd still expect a lot less speedup. Maybe they compiled pytorch with certain tricks, and used newer CUDA/CuDNN code, idk.