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.
pip install eqx-zoo tokenizers
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.
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)
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]]
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.
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, 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.
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]}")
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.
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).
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.",
]
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:
query: 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 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 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
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.
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:
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.
eqx-zoo is a young project, so a few honest caveats:
To try it:
pip install eqx-zoo tokenizers
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.