Putting my Jax-trained models on the Hugging Face Hub

Putting my Jax-trained models on the Hugging Face Hub 图片 1

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 could use my existing PyTorch script to upload them.

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)"

"Why do OpenAI's GPT-2 weights beat mine? Part three: testing overtraining"

"A quick(ish) Chinchilla check"

I've also added links to the posts in question.

添加评论
点赞收藏
点踩分享查看原文
评论
?
参与讨论