{"slug": "sentence-embeddings-in-jax-after-transformers-v5-dropped-it", "title": "Sentence embeddings in JAX, after Transformers v5 dropped it", "summary": "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.", "body_md": "[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.\n\nThis 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.\n\nWe'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.\n\n```\npip install eqx-zoo tokenizers\n```\n\n[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.\n\n[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:\n\n``` python\nimport jax\nimport jax.numpy as jnp\nfrom tokenizers import Tokenizer\n\nfrom eqx_zoo import Encoder\n\nrepo = \"sentence-transformers/all-MiniLM-L6-v2\"\ntokenizer = Tokenizer.from_pretrained(repo)\ntokenizer.enable_padding()\nmodel = Encoder.from_pretrained(repo)\n\nsentences = [\"The cat sits on the mat.\", \"A feline rests on a rug.\", \"Stock markets fell today.\"]\nbatch = tokenizer.encode_batch(sentences)\nids = jnp.array([e.ids for e in batch])\nmask = jnp.array([e.attention_mask for e in batch])\n\nembeddings = jax.vmap(model.embed)(ids, mask)\nprint(embeddings @ embeddings.T)\n```\n\nThe embeddings have unit length, so their dot products are cosine similarities:\n\n```\n[[0.9999997  0.55840427 0.05399836]\n [0.55840427 1.0000001  0.06003189]\n [0.05399836 0.06003189 0.99999946]]\n```\n\nThe 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\".\n\nTwo 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.\n\nAn 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:\n\n| Checkpoint | Pooling | Normalised | \n|---|---|---|\n| [all-MiniLM-L6-v2](https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2) | Mean over real tokens | Yes | \n| [bge-small-en-v1.5](https://huggingface.co/BAAI/bge-small-en-v1.5) | The first token, `[CLS]` | Yes | \n| [Qwen3-Embedding-0.6B](https://huggingface.co/Qwen/Qwen3-Embedding-0.6B) | The last real token | Yes | \n\nThis 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.\n\nYou can see what was read with `model.pooling` and `model.normalize`, and use `model(ids, mask)` to get the per-token hidden states instead.\n\nModels are ordinary JAX pytrees, so the usual transformations apply. Here's a search over a handful of documents, with the embedding step JIT-compiled:\n\n``` python\nimport equinox as eqx\n\n@eqx.filter_jit\ndef embed_batch(model, ids, mask):\n    return jax.vmap(model.embed)(ids, mask)\n\ndef encode(texts):\n    batch = tokenizer.encode_batch(texts)\n    ids = jnp.array([e.ids for e in batch])\n    mask = jnp.array([e.attention_mask for e in batch])\n    return embed_batch(model, ids, mask)\n\ndocuments = [\n    \"The lighthouse keeper wrote in the logbook every night.\",\n    \"Interest rates rose for the third month in a row.\",\n    \"The ferry to the island leaves at nine.\",\n    \"A new recipe for lemon cake.\",\n]\ndoc_embeddings = encode(documents)\n\nquery = encode([\"When does the boat depart?\"])[0]\nscores = doc_embeddings @ query\nfor i in jnp.argsort(-scores):\n    print(f\"{scores[i]:.3f}  {documents[i]}\")\n0.533  The ferry to the island leaves at nine.\n0.207  The lighthouse keeper wrote in the logbook every night.\n0.065  A new recipe for lemon cake.\n0.052  Interest rates rose for the third month in a row.\n```\n\nThe ferry timetable comes out on top, though the only word it shares with the query is \"the\".\n\n`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)`.\n\nThe [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:\n\n```\nrepo = \"intfloat/multilingual-e5-base\"\ntokenizer = Tokenizer.from_pretrained(repo)\ntokenizer.enable_padding()\nmodel = Encoder.from_pretrained(repo)\n\ntexts = [\n    \"query: Where is the lighthouse?\",\n    \"passage: Le phare se trouve au bout du port.\",\n    \"passage: Die Zinsen sind im dritten Monat gestiegen.\",\n]\n```\n\nTokenize 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.\n\nTwo details come from the [model card](https://huggingface.co/intfloat/multilingual-e5-base), and both matter:\n\n`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.\n\n[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:\n\n``` python\nfrom eqx_zoo import DecoderEmbedder\n\nrepo = \"Qwen/Qwen3-Embedding-0.6B\"\ntokenizer = Tokenizer.from_pretrained(repo)\ntokenizer.enable_padding()\nmodel = DecoderEmbedder.from_pretrained(repo)\n\nprompt = \"Instruct: Given a web search query, retrieve relevant passages that answer the query\\nQuery:\"\ntexts = [\n    prompt + \"What is the capital of France?\",\n    \"Paris is the capital of France.\",\n    \"Cats sleep a lot.\",\n]\nbatch = tokenizer.encode_batch(texts)\nids = jnp.array([e.ids for e in batch])\nmask = jnp.array([e.attention_mask for e in batch])\n\nquery, passage, unrelated = jax.vmap(model.embed)(ids, mask)\nprint(query @ passage, query @ unrelated)  # 0.736 0.117\n```\n\nAs 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.\n\nThe 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.\n\nA 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:\n\nThe 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.\n\n[eqx-zoo](https://github.com/xquantize/eqx-zoo) is a young project, so a few honest caveats:\n\nTo try it:\n\n```\npip install eqx-zoo tokenizers\n```\n\nThe 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.\n\nWhich embedding model would you most like to use from JAX? Let me know in the comments.", "url": "https://wpnews.pro/news/sentence-embeddings-in-jax-after-transformers-v5-dropped-it", "canonical_source": "https://dev.to/xquantize/sentence-embeddings-in-jax-after-transformers-v5-dropped-it-137k", "published_at": "2026-10-08 23:05:00+00:00", "updated_at": "2026-10-08 23:18:17.125546+00:00", "lang": "en", "topics": ["natural-language-processing", "machine-learning", "ai-tools", "developer-tools", "large-language-models"], "entities": ["Hugging Face", "JAX", "Equinox", "eqx-zoo", "sentence-transformers", "Qwen3-Embedding", "all-MiniLM-L6-v2", "Transformers"], "also_reported_by": [], "alternates": {"html": "https://wpnews.pro/news/sentence-embeddings-in-jax-after-transformers-v5-dropped-it", "markdown": "https://wpnews.pro/news/sentence-embeddings-in-jax-after-transformers-v5-dropped-it.md", "text": "https://wpnews.pro/news/sentence-embeddings-in-jax-after-transformers-v5-dropped-it.txt", "jsonld": "https://wpnews.pro/news/sentence-embeddings-in-jax-after-transformers-v5-dropped-it.jsonld"}}