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. Hugging Face Transformers https://huggingface.co/docs/transformers/en/index 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 https://github.com/xquantize/eqx-zoo , an open-source library that loads Hugging Face checkpoints as plain Equinox https://github.com/patrick-kidger/equinox modules. The embeddings match sentence-transformers https://www.sbert.net/ 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 https://huggingface.co/Qwen/Qwen3-Embedding-0.6B , a modern embedding model built on a language model. pip install eqx-zoo tokenizers eqx-zoo https://github.com/xquantize/eqx-zoo brings in JAX and Equinox https://github.com/patrick-kidger/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 https://github.com/xquantize/eqx-zoo reads the checkpoint's safetensors files directly. all-MiniLM-L6-v2 https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2 is a small, fast English model that's a common default for semantic search: python 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 https://www.sbert.net/ checkpoints also say how to turn per-token outputs into one vector, in a modules.json file and a pooling config. eqx-zoo's https://github.com/xquantize/eqx-zoo from pretrained reads those, so embed follows each checkpoint's own recipe: | Checkpoint | Pooling | Normalised | |---|---|---| | all-MiniLM-L6-v2 https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2 | Mean over real tokens | Yes | | bge-small-en-v1.5 https://huggingface.co/BAAI/bge-small-en-v1.5 | The first token, CLS | Yes | | Qwen3-Embedding-0.6B https://huggingface.co/Qwen/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 https://github.com/xquantize/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. 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: python 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 https://huggingface.co/intfloat/multilingual-e5-base 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 https://huggingface.co/intfloat/multilingual-e5-base , 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 https://huggingface.co/docs/transformers/en/model doc/xlm-roberta , and multilingual-e5-small https://huggingface.co/intfloat/multilingual-e5-small is a BERT https://huggingface.co/docs/transformers/en/model doc/bert with a multilingual vocabulary; eqx-zoo https://github.com/xquantize/eqx-zoo supports both architectures, so you don't need to know which is which. Qwen3-Embedding https://huggingface.co/Qwen/Qwen3-Embedding-0.6B works differently from BERT-style models https://huggingface.co/docs/transformers/en/model doc/bert . 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 https://github.com/xquantize/eqx-zoo loads it with DecoderEmbedder , which has the same embed method: python 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 https://www.sbert.net/ 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 https://www.sbert.net/ 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 https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2 , bge-small-en-v1.5 https://huggingface.co/BAAI/bge-small-en-v1.5 , multilingual-e5-small https://huggingface.co/intfloat/multilingual-e5-small , multilingual-e5-base https://huggingface.co/intfloat/multilingual-e5-base and Qwen3-Embedding-0.6B https://huggingface.co/Qwen/Qwen3-Embedding-0.6B . They're collected on the Hugging Face Hub in Verified in eqx-zoo https://github.com/xquantize/eqx-zoo . Other checkpoints with the same architectures BERT https://huggingface.co/docs/transformers/en/model doc/bert , RoBERTa https://huggingface.co/docs/transformers/en/model doc/roberta , XLM-RoBERTa https://huggingface.co/docs/transformers/en/model doc/xlm-roberta , Qwen3 https://huggingface.co/collections/Qwen/qwen3 load the same way. eqx-zoo https://github.com/xquantize/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 https://github.com/xquantize/eqx-zoo , along with Llama https://huggingface.co/docs/transformers/en/model doc/llama , Qwen https://huggingface.co/Qwen and Qwen3-MoE https://huggingface.co/docs/transformers/en/model doc/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.