# Eqx-zoo: Hub models in JAX/Equinox, verified against transformers and what verifying bf16 taught me

> Source: <https://discuss.huggingface.co/t/eqx-zoo-hub-models-in-jax-equinox-verified-against-transformers-and-what-verifying-bf16-taught-me/182905#post_1>
> Published: 2026-10-04 22:54:04+00:00

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](https://huggingface.co/collections/xquantize/verified-in-eqx-zoo-6ac2d731fd1ef3e109541bf7)

**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:

``` python
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](https://github.com/xquantize/eqx-zoo)

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?
