I'm a huge fan of Jax. The Jax team is incredibly strong!
Just want to share that Ray (an open source project we're developing at Anyscale), can be used to scale Jax (e.g., across TPUs).
Some docs from Google on how to do this
https://cloud.google.com/tpu/docs/ray-guide
Alpa is an open source project scaling Jax on 1000+ GPUs
https://www.anyscale.com/blog/training-175b-parameter-langua...
Cohere uses Ray + Jax + TPUs to build their LLMs
https://www.youtube.com/watch?v=For8yLkZP5w
A demo from Matt Johnson on the Jax team