Minimal library to train LLMs on TPU in JAX with pjit(). - View it on GitHub
Star
279
Rank
111503