You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
Add Bumblebee.load_embedding_head for SentenceTransformers pipelines - #472
This PR introduces Bumblebee.load_embedding_head/3 to support SentenceTransformers composite pipelines (such as google/embeddinggemma-300m, sentence-transformers/all-MiniLM-L6-v2, sentence-transformers/distiluse-base-multilingual-cased-v1, sentence-transformers/sentence-t5-base, and sentence-transformers/all-distilroberta-v1) in Bumblebee.
Following the discussion in #468, #469, and #471, instead of implementing a standalone EmbeddingGemma model, we reuse existing Bumblebee base transformer backbones (like Gemma3Text, Bert, DistilBert, Roberta, T5) and parse modules.json to attach the embedding head (pooling, dense projection, normalization) onto model_info.model.
What's Included
Bumblebee.load_embedding_head/3:
Parses modules.json from the repository.
Pooling: Supports all 5 SentenceTransformers pooling modes:
cls_token
max_tokens
mean_tokens
mean_sqrt_len_tokens
last_token (supports both pooling_mode_lasttoken and pooling_mode_last_token keys)
Preserves f16/bf16 precision and prevents underflow on lower precision types.
Matches PyTorch concatenation order when multiple pooling strategies are combined.
Weights Formats: Supports both Safetensors (model.safetensors) and PyTorch checkpoints (pytorch_model.bin).
Automatic Type & Backend Propagation: Automatically inherits :type (e.g. {:f, 16}) and :backend (e.g. EXLA.Backend) from the base model parameters if not explicitly overridden in opts.
Attaches onto base model hidden_state and returns %{embedding: ..., pooled_state: ..., hidden_state: ...} with updated Axon.ModelState parameter tracking.
Provides clear, actionable error messages if a backbone model defaults to language modeling and lacks :hidden_state.
Original (fuse_dense: false, default): Stacks linear projection layers separately, matching the Hugging Face checkpoint structure.
Fused (fuse_dense: true): Automatically pre-multiplies consecutive linear projection layers without bias ($W_{\text{fused}} = W_1 \times W_2 \in \mathbb{R}^{d_{in} \times d_{out}}$).
In models like EmbeddingGemma (which has consecutive 768 $\to$ 3072 and 3072 $\to$ 768 layers), fusion provides an 8x reduction in projection parameters (0.59M vs 4.7M) and a 1.72x speedup on the projection head with exact numerical identity ($1.7 \times 10^{-7}$ max diff).
Subdirectory Tree Listing:
Adds ?recursive=true to HuggingFace.Hub.file_listing_url and recursive wildcard in local repo file discovery so files in module subdirectories (e.g. 2_Dense/model.safetensors) are discovered and loaded.
Serving Integration:
Bumblebee.Text.text_embedding/3 automatically extracts the :embedding or :pooled_state attribute.
How to Test
# Run unit tests (module parsing, pooling modes, dense activations, fused layers, error handling)
mix test test/bumblebee/huggingface/sentence_transformers_test.exs
# Run slow end-to-end integration tests with real checkpoints (all 5 models)
mix test test/bumblebee/huggingface/sentence_transformers_test.exs --only slow
For models with consecutive linear projection layers without bias (such as EmbeddingGemma: 768 $\to$ 3072 $\to$ 768), fuse_dense: true merges them at load time into a single matrix multiplication ($768 \times 768$), reducing projection head parameters by 8x:
3. Models Configured for Masked LM by Default (e.g. all-distilroberta-v1)
If a repository's config.json lists a masked LM architecture by default, specify architecture: :base so Bumblebee.load_model loads the base transformer encoder outputs (hidden_state) instead of masked LM logits:
Numerical Parity Verification Across Architectures
We evaluated numerical parity for "Hello world" between PyTorch sentence-transformers 6.1.0 and Bumblebee across 5 different model families and architectures.
All Dense models match reference outputs within float32 numerical precision ($\le 1.12 \times 10^{-7}$). Sentence-T5 matches within $6.44 \times 10^{-5}$ due to floating point accumulation in relative position bias computation.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
This PR introduces
Bumblebee.load_embedding_head/3to support SentenceTransformers composite pipelines (such asgoogle/embeddinggemma-300m,sentence-transformers/all-MiniLM-L6-v2,sentence-transformers/distiluse-base-multilingual-cased-v1,sentence-transformers/sentence-t5-base, andsentence-transformers/all-distilroberta-v1) in Bumblebee.Following the discussion in #468, #469, and #471, instead of implementing a standalone
EmbeddingGemmamodel, we reuse existing Bumblebee base transformer backbones (likeGemma3Text,Bert,DistilBert,Roberta,T5) and parsemodules.jsonto attach the embedding head (pooling, dense projection, normalization) ontomodel_info.model.What's Included
Bumblebee.load_embedding_head/3:modules.jsonfrom the repository.cls_tokenmax_tokensmean_tokensmean_sqrt_len_tokenslast_token(supports bothpooling_mode_lasttokenandpooling_mode_last_tokenkeys)f16/bf16precision and prevents underflow on lower precision types.Identity,Tanh,ReLU,GELU, andSiLU.model.safetensors) and PyTorch checkpoints (pytorch_model.bin).:type(e.g.{:f, 16}) and:backend(e.g.EXLA.Backend) from the base model parameters if not explicitly overridden inopts.sentence_transformers.models.Dropoutpass-through.hidden_stateand returns%{embedding: ..., pooled_state: ..., hidden_state: ...}with updatedAxon.ModelStateparameter tracking.:hidden_state.Dense Layer Fusion Optimization (
fuse_denseoption):fuse_dense: false, default): Stacks linear projection layers separately, matching the Hugging Face checkpoint structure.fuse_dense: true): Automatically pre-multiplies consecutive linear projection layers without bias (EmbeddingGemma(which has consecutive 768Subdirectory Tree Listing:
?recursive=truetoHuggingFace.Hub.file_listing_urland recursive wildcard in local repo file discovery so files in module subdirectories (e.g.2_Dense/model.safetensors) are discovered and loaded.Serving Integration:
Bumblebee.Text.text_embedding/3automatically extracts the:embeddingor:pooled_stateattribute.How to Test
Note: This PR was supported by Antigravity.