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…
Hugging Face's Transformers v5 discontinued its TensorFlow and JAX implementations, concentrating on PyTorch instead. This alteration impacted those who employed FlaxBertModel or FlaxAutoModel for generating JAX-based sentence embeddings. To address this, the post introduces eqx-zoo, an open-source library capable of loading Hugging Face checkpoints as plain Equinox modules.
The embeddings produced by this method align with sentence-transformers to within the precision of float32 rounding. The article proceeds to discuss English and multilingual models, a compact semantic search, and Qwen3-Embedding, a contemporary embedding model constructed upon a language model. Installation can be accomplished using 'pip install eqx-zoo tokenizers'.
The library incorporates JAX and Equinox, while tokenizers, a rapid tokenizer library from Hugging Face, is utilized to convert text into token ids. The 'eqx-zoo' package reads the checkpoint's safetensors files directly, eliminating the need for PyTorch. The first example showcases the usage of the 'all-MiniLM-L6-v2' model, a compact and swift English model frequently utilized for semantic search.
The code imports the necessary libraries, loads the model and tokenizer, encodes the sentences, and computes the embeddings. The embeddings exhibit unit length, enabling the calculation of cosine similarities to ascertain the semantic relationships between the sentences. The dot products of the embeddings yield the cosine similarities, with the cat sentences scoring 0.56 with each other and approximately 0.05 with the stock-market sentence.
The post elucidates that the 'model.embed' function can process individual sentences, while 'jax.vmap' applies it to the entire batch. Furthermore, it emphasizes the importance of acknowledging that each checkpoint includes its unique method for transforming per-token outputs into a single vector, encapsulated in a 'modules.json' file and a pooling configuration.
The 'eqx-zoo' library reads these configurations, allowing the 'embed' function to adhere to the specific pooling method defined by each checkpoint. A mistake in this aspect will result in embeddings that may seem plausible but are not the ones the model was trained to produce. The post provides a comprehensive comparison of pooling strategies for the 'all-MiniLM-L6-v2', 'bge-small-en-v1.5', and 'Qwen3-Embedding-0.6B' models.
If a checkpoint specifies a pooling method that 'eqx-zoo' does not support, an 'NotImplementedError' is raised, notifying the user of the incompatible pooling method. The text concludes with a demonstration of constructing a miniature semantic search system utilizing the 'eqx-zoo' library, JAX, and a small set of documents.
Written by urgent.news from Dev.to's reporting — not their text. Machine-written — may contain errors; check the original before relying on it.