DEV Community

Zane Neave
Zane Neave

Posted on

Sentence embeddings in JAX, after Transformers v5 dropped it

Hugging Face Transformers v5 removed its TensorFlow and JAX code to focus on PyTorch. If you used FlaxBertModel or FlaxAutoModel to compute sentence embeddings in JAX, those classes are gone.

This post shows how to compute sentence embeddings in JAX with eqx-zoo, an open-source library that loads Hugging Face checkpoints as plain Equinox modules. The embeddings match sentence-transformers to within float32 rounding, and the post ends with how that's verified.

We'll cover English and multilingual models, a small semantic search, and Qwen3-Embedding, a modern embedding model built on a language model.

Install

pip install eqx-zoo tokenizers
Enter fullscreen mode Exit fullscreen mode

eqx-zoo brings in JAX and Equinox; tokenizers is Hugging Face's fast tokenizer library, which we'll use to turn text into token ids. PyTorch isn't needed: eqx-zoo reads the checkpoint's safetensors files directly.

Your first embeddings

all-MiniLM-L6-v2 is a small, fast English model that's a common default for semantic search:

import jax
import jax.numpy as jnp
from tokenizers import Tokenizer

from eqx_zoo import Encoder

repo = "sentence-transformers/all-MiniLM-L6-v2"
tokenizer = Tokenizer.from_pretrained(repo)
tokenizer.enable_padding()
model = Encoder.from_pretrained(repo)

sentences = ["The cat sits on the mat.", "A feline rests on a rug.", "Stock markets fell today."]
batch = tokenizer.encode_batch(sentences)
ids = jnp.array([e.ids for e in batch])
mask = jnp.array([e.attention_mask for e in batch])

embeddings = jax.vmap(model.embed)(ids, mask)
print(embeddings @ embeddings.T)
Enter fullscreen mode Exit fullscreen mode

The embeddings have unit length, so their dot products are cosine similarities:

[[0.9999997  0.55840427 0.05399836]
 [0.55840427 1.0000001  0.06003189]
 [0.05399836 0.06003189 0.99999946]]
Enter fullscreen mode Exit fullscreen mode

The two cat sentences score 0.56 with each other, and about 0.05 with the stock-market one, even though the only word they share is "on".

Two things to notice, model.embed works on one sentence, and jax.vmap maps it over the batch. And mask marks which tokens are real: the shorter sentences are padded to the longest, and padding is excluded from the embedding.

Each checkpoint brings its own recipe

An embedding model is more than its transformer. sentence-transformers checkpoints also say how to turn per-token outputs into one vector, in a modules.json file and a pooling config. eqx-zoo's from_pretrained reads those, so embed follows each checkpoint's own recipe:

Checkpoint Pooling Normalised
all-MiniLM-L6-v2 Mean over real tokens Yes
bge-small-en-v1.5 The first token, [CLS] Yes
Qwen3-Embedding-0.6B The last real token Yes

This matters because the wrong pooling gives embeddings that look plausible but aren't the ones the model was trained to produce. If a checkpoint asks for something eqx-zoo doesn't implement yet, such as another pooling mode or an extra projection layer, loading fails with a clear NotImplementedError instead of quietly computing something different.

You can see what was read with model.pooling and model.normalize, and use model(ids, mask) to get the per-token hidden states instead.

A tiny semantic search

Models are ordinary JAX pytrees, so the usual transformations apply. Here's a search over a handful of documents, with the embedding step JIT-compiled:

import equinox as eqx

@eqx.filter_jit
def embed_batch(model, ids, mask):
    return jax.vmap(model.embed)(ids, mask)

def encode(texts):
    batch = tokenizer.encode_batch(texts)
    ids = jnp.array([e.ids for e in batch])
    mask = jnp.array([e.attention_mask for e in batch])
    return embed_batch(model, ids, mask)

documents = [
    "The lighthouse keeper wrote in the logbook every night.",
    "Interest rates rose for the third month in a row.",
    "The ferry to the island leaves at nine.",
    "A new recipe for lemon cake.",
]
doc_embeddings = encode(documents)

query = encode(["When does the boat depart?"])[0]
scores = doc_embeddings @ query
for i in jnp.argsort(-scores):
    print(f"{scores[i]:.3f}  {documents[i]}")
Enter fullscreen mode Exit fullscreen mode
0.533  The ferry to the island leaves at nine.
0.207  The lighthouse keeper wrote in the logbook every night.
0.065  A new recipe for lemon cake.
0.052  Interest rates rose for the third month in a row.
Enter fullscreen mode Exit fullscreen mode

The ferry timetable comes out on top, though the only word it shares with the query is "the".

eqx.filter_jit compiles the function once per input shape. With enable_padding(), each batch is padded to its own longest sentence, so a new length means a new compilation. For steady throughput, pad to a fixed length instead, for example tokenizer.enable_padding(length=128).

Multilingual embeddings

The multilingual-e5 models cover about 100 languages and use the same Encoder API. Swapping the repository is the only code change:

repo = "intfloat/multilingual-e5-base"
tokenizer = Tokenizer.from_pretrained(repo)
tokenizer.enable_padding()
model = Encoder.from_pretrained(repo)

texts = [
    "query: Where is the lighthouse?",
    "passage: Le phare se trouve au bout du port.",
    "passage: Die Zinsen sind im dritten Monat gestiegen.",
]
Enter fullscreen mode Exit fullscreen mode

Tokenize and embed them exactly as before. Against the English query, the French lighthouse passage scores 0.741, and the German one about interest rates 0.670.

Two details come from the model card, and both matter:

  • Every input starts with query: or passage:. That's how the model was trained, and leaving the prefixes out degrades results. For similarity between texts of the same kind, the card recommends query: for both.
  • Scores sit high, mostly around 0.7 to 1.0, because of how the model was trained. What matters is the order of the scores, not their absolute values.

The base model is an XLM-RoBERTa, and multilingual-e5-small is a BERT with a multilingual vocabulary; eqx-zoo supports both architectures, so you don't need to know which is which.

Qwen3-Embedding: a language model as an embedder

Qwen3-Embedding works differently from BERT-style models. It's a decoder, like a chat model, with causal attention: each token only sees the tokens before it. So the embedding is the hidden state of the last token, the only one that has seen the whole input. eqx-zoo loads it with DecoderEmbedder, which has the same embed method:

from eqx_zoo import DecoderEmbedder

repo = "Qwen/Qwen3-Embedding-0.6B"
tokenizer = Tokenizer.from_pretrained(repo)
tokenizer.enable_padding()
model = DecoderEmbedder.from_pretrained(repo)

prompt = "Instruct: Given a web search query, retrieve relevant passages that answer the query\nQuery:"
texts = [
    prompt + "What is the capital of France?",
    "Paris is the capital of France.",
    "Cats sleep a lot.",
]
batch = tokenizer.encode_batch(texts)
ids = jnp.array([e.ids for e in batch])
mask = jnp.array([e.attention_mask for e in batch])

query, passage, unrelated = jax.vmap(model.embed)(ids, mask)
print(query @ passage, query @ unrelated)  # 0.736 0.117
Enter fullscreen mode Exit fullscreen mode

As with e5, there's a convention to follow: queries get an instruction prompt, and documents don't. The prompt comes from the checkpoint's sentence-transformers configuration, and you can change the task description to suit your search.

The tokenizer also appends an end-of-text token to every input, and that's the token whose hidden state becomes the embedding. The standalone tokenizers library adds it automatically, just as sentence-transformers does.

How we know it's right

A port that loads and runs isn't necessarily correct: a missing bias or the wrong pooling still produces vectors that look reasonable. So every supported checkpoint is checked against the reference implementations on every pull request:

  • Layer by layer: each layer's output matches Hugging Face Transformers in float32.
  • End to end: embeddings match sentence-transformers itself, on a padded batch of sentences of different lengths. For the examples in this post, the largest difference is 1.7e-7 for all-MiniLM-L6-v2, and 4.7e-7 for Qwen3-Embedding with its query prompt, which is float32 rounding.
  • In bfloat16: the error against float32 must be within 2x the reference library's own bf16 error. Measured ratios are 0.84 to 1.51 for encoders and up to 1.38 for Qwen3-Embedding.
  • Block by block: precision-sensitive pieces such as LayerNorm are tested directly in bf16, because a subtle bug there barely moves a whole model's output but is off by thousands of rounding steps at the block level.
  • Tiny random models: small randomly initialised versions of each architecture cover code paths no single checkpoint exercises, and run in seconds.

The verified checkpoints so far are all-MiniLM-L6-v2, bge-small-en-v1.5, multilingual-e5-small, multilingual-e5-base and Qwen3-Embedding-0.6B. They're collected on the Hugging Face Hub in Verified in eqx-zoo. Other checkpoints with the same architectures (BERT, RoBERTa, XLM-RoBERTa, Qwen3) load the same way.

Limitations, and trying it

eqx-zoo is a young project, so a few honest caveats:

  • The test suite runs on CPU. JAX runs the same code on GPUs and TPUs, but the numerical checks haven't been run there systematically yet. GPU and TPU reports are very welcome, and the repository has an issue template for them.
  • Embedding models are limited to the architectures above. MPNet-based models such as all-mpnet-base-v2, and checkpoints that need custom code, aren't supported yet.
  • Tokenization is up to you. eqx-zoo takes token ids, so you choose how to tokenize, truncate and pad. The tokenizers library, as used here, matches what sentence-transformers does.

To try it:

pip install eqx-zoo tokenizers
Enter fullscreen mode Exit fullscreen mode

The code is on GitHub at xquantize/eqx-zoo, along with Llama, Qwen and Qwen3-MoE language models verified the same way. Issues and pull requests are welcome, from bug reports to new checkpoints.

Which embedding model would you most like to use from JAX? Let me know in the comments.

Top comments (1)

Some comments may only be visible to logged-in visitors. Sign in to view all comments.