With transformers v5 focusing on PyTorch, I’ve been building eqx-zoo, which loads Hub checkpoints directly (safetensors, by repo name, no conversion) as plain Equinox modules and verifies each model against transformers itself.
What’s supported: Llama, Qwen2, Qwen3 and Qwen3-MoE for generation (KV cache, batching) and BERT, RoBERTa and XLM-RoBERTa embedding models, with each checkpoint’s sentence-transformers pooling and normalisation read from its config. The verified checkpoints are in this collection: Verified in eqx-zoo - a xquantize Collection
How verification works
- float32: every layer’s output is compared with transformers’ (captured with forward hooks) and greedy generation must reproduce transformers’ output token for token
- embeddings must match sentence-transformers on padded batches
- tiny randomly initialised models cover code paths no single checkpoint exercises, with every parameter randomised (default init sets norm weights to 1 and biases to 0, which can hide a dropped scale or bias)
What bfloat16 taught me
- Two bf16 implementations round differently, so comparing ours with transformers’ bf16 directly is the wrong test: they’re often further apart than either is from float32. Instead, ours is compared with float32, relative to transformers’ own bf16 error.
- Exact greedy agreement in bf16 isn’t a sound criterion: transformers’ own bf16 drifted from its float32 after 2 of 20 tokens on Llama 3.2 1B.
- RMS error is a much more stable statistic than max error, which is dominated by single unlucky roundings.
- JIT-compiling a whole decoder in XLA made bf16 measurably less accurate than eager execution (up to 2.4x transformers’ own bf16 error), and compiling a single layer alone reproduced it. Encoders, which are post-norm, didn’t show this at all.
- A subtle bug (LayerNorm statistics in bf16) shifted model-level error less than the difference between x86 and ARM did, so model-level tests couldn’t catch it reliably. At the block level, the same bug was off by thousands of bf16 steps, so it’s now tested there directly.
Quick example:
from eqx_zoo import CausalLM, generate
model = CausalLM.from_pretrained("Qwen/Qwen3-0.6B")
tokens = generate(model, prompt_ids, max_new_tokens=30)
pip install eqx-zoo — GitHub - xquantize/eqx-zoo: Verified Equinox ports of pretrained models, numerically matched against Hugging Face · GitHub
Everything is tested on CPU so far. I’d love feedback from anyone running JAX on GPU or TPU and I’m curious: which Hub models would be most useful to have in JAX next?