cd /news/natural-language-processing/sentence-embeddings-in-jax-after-tra… · home › topics › natural-language-processing › article
[ARTICLE · art-147913] src=dev.to ↗ pub= topic=natural-language-processing verified=true sentiment=· neutral

Sentence embeddings in JAX, after Transformers v5 dropped it

A developer released eqx-zoo, an open-source library that loads Hugging Face sentence-transformer checkpoints as plain Equinox modules to compute sentence embeddings in JAX, after Hugging Face Transformers v5 removed its TensorFlow and JAX code in favor of PyTorch. The library reads safetensors files directly without PyTorch and reproduces sentence-transformers embeddings to within float32 rounding, supporting English and multilingual models including Qwen3-Embedding-0.6B.

by read6 min views1 publishedOct 8, 2026

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.

── more in #natural-language-processing 4 stories · sorted by recency
── more on @hugging face 3 stories trending now
sponsored brought to you by zahid.host 4,200+ EU-deployed projects
reading about agents? ship yours in a single git push.

Run your AI side-project on zahid.host

EU-based hosting, git-push deploys, automatic HTTPS, no cold starts. Free tier with a custom domain — perfect for shipping the agent you just read about.

$git push zahid main
→ Live at https://your-agent.zahid.host ✓
Get free account → Pricing
from €0/mo · no card required
LIVE [news/sentence-embeddings-…] indexed:0 read:6min 2026-10-08 · —