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

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?

I tried a few things on a Colab GPU:


I ran a few checks on an NVIDIA L4 with Qwen/Qwen3-0.6B, mainly around the GPU/TPU and bf16 questions.

The short version is:

  • the CPU observation from #11 did not transfer to the L4 in a simple “JIT bf16 is worse than eager bf16” way;
  • on GPU, the fp32 reference comparison was very sensitive to JAX matmul precision;
  • long-context and decision-level checks caught things that RMS-only / short-prompt checks missed;
  • for one small RMSNorm fixture I could reproduce a concrete JIT/XLA rounding-boundary effect, but it does not seem sufficient to explain the whole model by itself;
  • for scan-over-layers, L4 gave a real compile-time win but a small warm-decode regression;
  • for a next architecture, FLAN-T5-small looks interesting to me from a coverage perspective, because it adds encoder-decoder + cross-attention + relative-position behavior rather than another decoder-only variant.

For accelerator verification, my default matrix after these experiments would probably be:

axis minimal useful cases
reference HF fp32 + HF bf16 on the same accelerator where possible
JAX fp32 default + highest matmul precision
JAX bf16 eager/op-by-op + whole JIT
input short prompt + one longer prompt
metrics RMS/max + rank/top-k + greedy decision
generation full forward and cached decoding separately when generation matters

The same-device part seems important: otherwise framework/port differences and accelerator/backend arithmetic differences get mixed together.

JAX explicitly documents that fp32 dot products can use reduced-precision arithmetic internally on accelerators (TF32 on recent NVIDIA GPUs), while highest requests true fp32 on GPU:

JAX matmul precision

In my L4 run, that distinction was not subtle.

L4 setup and the first parity result

The main comparison used:

  • GPU: NVIDIA L4
  • model: Qwen/Qwen3-0.6B
  • eqx-zoo commit: 329b21b342b5213a6d0d15eb381517cd9d00d1e4
  • JAX / jaxlib: 0.11.1
  • Equinox: 0.13.8
  • Transformers: 5.18.0
  • HF reference attention path: eager
  • HF fp32 reference: CUDA TF32 disabled
  • short prompt: 5 tokens
  • long prompt: 256 tokens

For fp32, comparing eqx-zoo with HF fp32 on the same L4:

input JAX fp32 default RMS JAX fp32 highest RMS
short ~2.43e-3 ~8.66e-6
long ~4.04e-3 ~1.21e-5

So highest reduced the difference by roughly 280x–330x in this setup.

That makes me think device + dtype alone is not quite enough metadata for accelerator parity. I would also record at least:

  • JAX / jaxlib version
  • matmul precision
  • attention backend on the reference side
  • eager vs JIT
  • relevant XLA flags

JAX also makes the useful distinction that precision is not dtype: storage/activation dtype and dot-product arithmetic precision are separate controls.

For bf16, the result was more surprising: the L4 did not reproduce the CPU direction from #11.

Relative to HF’s own bf16-vs-fp32 RMS error, eqx-zoo eager was around 1.23x on the short input and 1.08x on the long input, while JIT was around 0.71x on both.

So on this L4, JIT was actually closer to HF fp32 by RMS.

I would not interpret that as “JIT is more correct”, though. Once I looked at token decisions, RMS and behavioral parity stopped moving monotonically together.

That seems like a useful reason to keep more than one verification metric.

Why I would keep long-context and decision-level checks

I tried a deliberately bad control related to #12: removing the fp32 upcast before attention softmax.

With the short prompt, the mutation did not change the top-1 decision.

With a 256-token input, it did: in one of the comparisons the top-1 agreement dropped from 100% to about 98.8%, even though RMS did not simply move in the same “worse” direction.

That made the long prompt look less like a generic stress test and more like coverage for numerical paths that short prompts barely exercise.

There is a close practical precedent in MaxText’s HF/golden-logit checker:

MaxText forward-pass logit checker

It uses several complementary signals, including:

  • numerical logit tolerance
  • KL divergence
  • top-k token comparison
  • longer prompts for bugs that are invisible on short inputs

The source even contains a long-prompt case specifically because a partial-RoPE bug was not detected by short prompts.

Levanter has a similar general philosophy when porting HF models: align the weights/input and compare the ported implementation against HF at meaningful model/module boundaries rather than treating “loads and runs” as sufficient verification:

Levanter model-porting guide

For eqx-zoo, a compact test ladder that seems useful to me would be:

  1. parameter/config mapping
  2. fp32 intermediate parity
  3. bf16 numerical parity
  4. top-k / rank parity
  5. greedy token parity
  6. cached-generation parity
  7. one longer input that exercises attention/position/cache paths

Not every model needs every item in CI, but having the distinction available seems useful.

One concrete bf16/JIT mechanism I could reproduce on the L4

I also followed #11 down into one real Qwen3-0.6B RMSNorm fixture.

This is intentionally a local finding, not a claim that it explains all of #11.

For the same bf16 input and weights:

  • JAX eager/op-by-op matched the equivalent PyTorch/HF RMSNorm arithmetic exactly at the final bf16 output;
  • whole-JIT differed;
  • about 17,545 / 65,536 output elements differed;
  • RMS difference was about 5.98e-4;
  • max absolute difference was 0.015625.

The first tiny eager/JIT difference appeared around mean(x²), but that was not enough to explain the final discrepancy.

The more interesting part showed up in optimized HLO.

The source-level structure was effectively:

norm in fp32
→ cast normalized value to bf16
→ multiply by bf16 RMSNorm weight

In the default whole-JIT optimized program, the computation instead appeared effectively as:

norm in fp32
→ multiply by weight at higher precision
→ cast the result to bf16

So the source-level bf16 rounding boundary was no longer in the same place.

Two independent diagnostics made the final RMSNorm result exactly match the eager/PyTorch result again:

  1. starting the process with
XLA_FLAGS=--xla_allow_excess_precision=false
  1. or putting jax.lax.optimization_barrier() immediately after the source-level bf16 cast.

The latter is consistent with JAX’s documented semantics: an optimization barrier prevents operations from being moved across it and prevents compiler fusion across the barrier:

jax.lax.optimization_barrier

The lowered StableHLO was unchanged by the excess-precision flag; the change appeared in the optimized program. That makes the compiler optimization stage a fairly strong localization for this exact fixture/environment.

I would still avoid turning either the barrier or the XLA flag into a recommended production fix from this result alone.

In particular, the model-wide experiments below show that RMSNorm is not the whole story.

Also, one subtle debugging point: jax.disable_jit() is useful for separating large compiled regions, but JAX documents that individual primitive operations are still compiled by XLA during eager/op-by-op execution:

jax.disable_jit

So I think “op-by-op XLA vs larger compiled region” is a more precise interpretation than “XLA vs no XLA”.

Model-wide excess precision: RMS and token decisions did not agree

I then repeated the comparison at model level.

The especially interesting control was global:

XLA_FLAGS=--xla_allow_excess_precision=false

Again, I mean this as a diagnostic, not a proposed default.

On the fixed 16-token greedy sequence:

condition tokens matching HF fp32
default model 7 / 16
global no-excess-precision 16 / 16
RMSNorm-only barrier 7 / 16

So the local RMSNorm mechanism was real, but RMSNorm-only barriers did not account for the model-wide decision difference.

There was another useful twist: the global no-excess condition did not necessarily give the smallest logit RMS against HF fp32. In some comparisons the default JIT logits were closer by RMS, while no-excess gave the exact HF greedy sequence.

That is probably the clearest result from these experiments for me:

lower aggregate logit error and identical token decisions are not the same objective.

For near-tied tokens, a small structured change in a few logits can matter more than a larger diffuse error over the vocabulary.

So if the goal is verification of a generative port, I would probably treat RMS/max error as one layer of the contract rather than the final oracle.

Full forward and cached decoding also behaved differently

I also separated fixed-input teacher-forced decisions from actual cached greedy decoding.

With the default compiled model:

  • teacher-forced HF path: 15 / 16 token decisions matched
  • cached greedy generation: 7 / 16 matched
  • the first mismatch was at the same early position, after which generation naturally cascaded

Global no-excess precision gave:

  • teacher-forced: 16 / 16
  • cached greedy: 16 / 16

I then tried selective optimization_barrier probes at a few semantic boundaries:

  • attention output → residual add
  • MLP output → residual add
  • final hidden state → tied embedding projection

The surprising one was the MLP residual boundary:

  • teacher-forced stayed 15 / 16
  • cached greedy became 16 / 16

But putting barriers at all three boundaries moved cached generation back to the default-style result rather than improving it further.

So the barrier effect was non-monotonic.

I would not read that as “the MLP residual is the root cause”. A barrier changes compiler optimization/fusion around the boundary, so it can perturb the compiled graph beyond the one arithmetic expression we are looking at.

What I do think this supports is keeping cached decoding as its own verification path instead of assuming that full-sequence forward parity automatically covers it.

That is especially relevant here because eqx-zoo generation has a prompt prefill followed by one-token cached decoding inside lax.scan; those are different shapes/programs from a full teacher-forced forward.

Scan-over-layers on the L4

I also tried the experiment from #15.

One detail: I did not compare the old experiment/scan-layers branch directly against current main, because that branch had accumulated unrelated historical distance from current main.

Instead I compared the scan tip with its merge-base, so the A/B mostly isolates the scan change itself.

For Qwen3-0.6B, prompt length 128 and 64 generated tokens on the L4:

  • estimated generation compile time dropped by roughly 57% in fp32
  • and roughly 63% in bf16
  • warm decode throughput was about 10% slower

So at least on this L4 the tradeoff looked like:

scan-over-layers
    ├─ compile latency: clearly better
    └─ warm decode throughput: slightly worse

That is different from the CPU magnitude in #15, but not a simple “scan wins on GPU” result either.

It makes me think compile latency and steady-state decode throughput should probably remain separate benchmark axes.

Equinox itself documents scan-over-layers specifically as a technique for improving compilation speed:

Equinox: improve compilation speed with scan-over-layers

So an accelerator-specific or optional scan path might still be attractive if compile latency matters enough, even if decode is not faster.

A possible small accelerator test matrix

If you want a compact matrix that is cheap enough for contributors to report, I think this would already distinguish a lot:

model/revision:
eqx-zoo commit:
device:
jax/jaxlib:

HF fp32, same device:
  attention backend:
  reduced-fp32 mode / TF32 setting:

JAX fp32:
  default matmul precision
  highest matmul precision

JAX bf16:
  eager/op-by-op
  whole JIT

inputs:
  short
  ~256+ token long case

metrics:
  RMS
  max abs
  top-1 / top-k
  greedy token parity

if generation is relevant:
  full-sequence / teacher-forced
  cached decode

I would probably start there before asking someone to dump HLO or bisect compiler behavior.

It gives fairly high information gain without turning every accelerator report into a compiler investigation.

Next model

For the “what Hub model next?” question, my vote from a coverage perspective would be google/flan-t5-small.

Not because I know it is the most requested model, but because it would exercise several new implementation boundaries at relatively small scale:

  • encoder + decoder rather than decoder-only
  • decoder cross-attention
  • encoder/decoder masks
  • relative position bias
  • encoder state reuse during generation
  • different cache semantics

The Transformers T5 docs expose the encoder/decoder and cross-attention structure, and T5’s relative-position behavior would add a meaningfully different verification surface.

So if the priority is architecture diversity per unit implementation effort, FLAN-T5-small looks attractive to me.

If the priority is instead actual user demand, I would keep that as a separate question rather than treating this suggestion as a popularity ranking.

Overall, the verification-first direction of eqx-zoo looks useful to me. The GPU run mostly convinced me that accelerator verification is not just “rerun the CPU thresholds on CUDA”: device math mode, compilation boundaries, context length, and cached decoding can all change what the useful oracle is.

The good news is that most of those dimensions seem coverable with a relatively small deterministic matrix rather than a large benchmark suite.

This is actually fantastic, thank you for taking the time to run all of this on an L4 and to write it up so nicely.

A few things really change my picture:

  • The JIT-vs-eager direction flipping on the L4 suggests #11 is more of a CPU/XLA-backend effect than a general one. I’ll add your numbers there.
  • The matmul-precision point is the most urgent one. If the default fp32 path uses TF32 on GPU, the fp32 parity tests aren’t really testing fp32 there. I’ll make the parity tests request highest precision and document it.
  • The long-context and decision-level results line up with #12 and the scan-over-layers numbers are exactly what #15 needed.
  • Your accelerator matrix is a great template; I’d like to turn it into an issue template for GPU/TPU reports, with credit to you.

FLAN-T5-small as a coverage choice makes a lot of sense too; I’ll open an issue for it.

If you’re up for it, I’d love to see the details: a notebook or opening issues / PRs directly. The RMSNorm rounding-boundary reproduction in particular would be really valuable on #11.

Quick follow-up now that I’ve read the expanded sections properly (they were collapsed when I first replied, sorry for asking for details you’d already given).

Your results overturned one of my conclusions. On CPU I’d tried --xla_allow_excess_precision=false, seen no improvement in max error and written excess precision off. Your RMSNorm mechanism and the 7/16 → 16/16 greedy result show that max error was the wrong metric, so it’s now the leading hypothesis. I’ve posted a correction on #11.

Everything is now tracked, with credit to you:

  • #11: the RMSNorm rounding-boundary mechanism, the excess-precision results and the teacher-forced vs cached split
  • #12: the softmax-upcast mutation that only a 256-token input caught, plus your test ladder as the issue’s scope
  • #15: the scan compile/decode numbers
  • #36: requesting highest matmul precision for fp32 parity on accelerators
  • #37: an issue template based on your accelerator matrix
  • #38: FLAN-T5 as the first encoder-decoder model

If you’d be willing to share the notebook or scripts, I’d love to reproduce the #11 results on CPU with RMS and greedy agreement rather than max error. And if any of #36, #37 or #12 appeals to you, PRs are very welcome. Thanks again, this was genuinely the most useful feedback the project has had.