I hadn't uploaded the models that I trained using JAX to the Hugging Face Hub because
Transformers has been PyTorch-only since version 5
(though they say they're working to add interoperability with JAX in the future), so it
would have been tough to get them working natively with AutoModelForCausalLM
and the like.
But then it dawned on me that I'd already written a conversion script that could take my JAX safetensors files and convert them into ones compatible with my PyTorch code. It's actually those converted models that I use for my evals!
So, I've now uploaded PyTorch-compatible versions of all of my JAX-trained models:
"[Writing an LLM from scratch, part 34b -- from bigrams to GPT-2, one component at a time (in JAX)](/2026/07/llm-from-scratch-34b-building-and-training-gpt-2-small-in-jax)"
`gpjt/jax-no-mha-bias-no-dropout`
`gpjt/jax-no-mha-bias-with-dropout`
`gpjt/jax-with-mha-bias-no-dropout`
"[Why do OpenAI's GPT-2 weights beat mine? Part three: testing overtraining](/2026/07/why-do-openai-gpt2-weights-beat-mine-3-overtraining)"
`gpjt/jax-with-mha-bias-no-dropout-extended`
`gpjt/jax-with-mha-bias-no-dropout-2-epoch`
"[A quick(ish) Chinchilla check](/2026/08/chinchilla-check)"
`gpjt/jax-with-mha-bias-larger-chinchilla-1`
slightly-larger
model.gpjt/jax-with-mha-bias-larger-chinchilla-2
slightly-smaller
model.I've also added links to the posts in question.