41 minutes ago · Tech · hide · 0 comments

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)" gpjt/jax-no-mha-bias-no-dropout -- the first full LLM trained in the post, in the "Adding LayerNorm" section. gpjt/jax-no-mha-bias-with-dropout -- the second full LLM trained in the post, in the "Dropout" section. gpjt/jax-with-mha-bias-no-dropout -- the third full…

No comments yet. Log in to reply on the Fediverse. Comments will appear here.