So why does there seem to be no published metrics showing performance of various common ML models on common hardware with OpenXLA vs other frameworks/compilers?
So why does there seem to be no published metrics showing performance of various common ML models on common hardware with OpenXLA vs other frameworks/compilers?
But people like Google offering ML compute as a service want to keep the actual optimization modules secret and closed source. That way they can make their hardware perform superspeed while competitors hardware looks slow.
For simple stuff, we can compare JAX to PyTorch on a 4090, and JAX seems faster by 10-50%. It's way way way faster on TPU.
That said, JAX is apparently a pain to work with in comparison to PyTorch (disclaimer I haven't used JAX enough to opine yet).
I'd also want a comparison between JAX and Taichi-lang for parallel stuff, even though the problem domains aren't exactly aligned.
What benchmarks are you looking at here?
Not super scientific, Im sorry to say, but methodologically benchmarking this is really hard.