Skip to content

Add Bumblebee.load_embedding_head for SentenceTransformers pipelines - #472

Closed
dkuku wants to merge 9 commits into
elixir-nx:mainfrom
dkuku:feature/embedding-gemma
Closed

dkuku wants to merge 9 commits into
elixir-nx:mainfrom
dkuku:feature/embedding-gemma

Conversation

@dkuku

@dkuku dkuku commented Sep 30, 2026 •

Copy link
Copy Markdown
Contributor

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

  1. 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.
    • Dense Projections: Supports activation functions Identity, Tanh, ReLU, GELU, and SiLU.
    • 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.
    • Passthrough Modules: Supports sentence_transformers.models.Dropout pass-through.
    • Normalization: Supports $L_2$ vector normalization.
    • 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.
  2. Dense Layer Fusion Optimization (fuse_dense option):

    • 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).
  3. 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.
  4. 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

Note: This PR was supported by Antigravity.

@dkuku
dkuku force-pushed the feature/embedding-gemma branch from e152d91 to 6ab4e05 Compare September 30, 2026 21:06
@dkuku
dkuku marked this pull request as draft September 30, 2026 21:19
@dkuku

dkuku commented Sep 30, 2026

Copy link
Copy Markdown
Contributor Author

Usage Examples

1. Standard Usage (e.g. google/embeddinggemma-300m or sentence-transformers/all-MiniLM-L6-v2)

repo = {:hf, "google/embeddinggemma-300m"}

{:ok, model_info} = Bumblebee.load_model(repo)
{:ok, model_info} = Bumblebee.load_embedding_head(repo, model_info)
{:ok, tokenizer} = Bumblebee.load_tokenizer(repo)

serving = Bumblebee.Text.text_embedding(model_info, tokenizer)
Nx.Serving.run(serving, "Hello world")
#=> %{embedding: #Nx.Tensor<f32[768] ...>}

2. Projection Head Fusion (fuse_dense: true)

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:

{:ok, model_info} = Bumblebee.load_model(repo)
{:ok, model_info} = Bumblebee.load_embedding_head(repo, model_info, fuse_dense: true)
{:ok, tokenizer} = Bumblebee.load_tokenizer(repo)

serving = Bumblebee.Text.text_embedding(model_info, tokenizer)
Nx.Serving.run(serving, "Hello world")

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:

repo = {:hf, "sentence-transformers/all-distilroberta-v1"}

{:ok, model_info} = Bumblebee.load_model(repo, architecture: :base)
{:ok, model_info} = Bumblebee.load_embedding_head(repo, model_info)
{:ok, tokenizer} = Bumblebee.load_tokenizer(repo)

serving = Bumblebee.Text.text_embedding(model_info, tokenizer)
Nx.Serving.run(serving, "Hello world")
#=> %{embedding: #Nx.Tensor<f32[768] ...>}

@dkuku

dkuku commented Sep 30, 2026

Copy link
Copy Markdown
Contributor Author

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.

Results Summary

Model Backbone Architecture Head Modules Output Dim Max Absolute Diff Mean Absolute Diff
unsloth/embeddinggemma-300m Gemma3Text Mean Pooling + 2x Dense + L2 Norm 768 1.04e-7 2.32e-8
sentence-transformers/all-MiniLM-L6-v2 Bert Mean Pooling + L2 Norm 384 1.12e-7 3.06e-8
sentence-transformers/distiluse-base-multilingual-cased-v1 DistilBert Mean Pooling + Dense (Tanh) 512 4.84e-8 1.14e-8
sentence-transformers/sentence-t5-base T5Encoder Mean Pooling + Dense + L2 Norm 768 6.44e-5 1.23e-5
sentence-transformers/all-distilroberta-v1 Roberta (architecture: :base) Mean Pooling + L2 Norm 768 1.10e-7 2.37e-8

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.


1. Python Reference Script (sentence-transformers 6.1.0)

from sentence_transformers import SentenceTransformer
import json

models = [
    "unsloth/embeddinggemma-300m",
    "sentence-transformers/all-MiniLM-L6-v2",
    "sentence-transformers/distiluse-base-multilingual-cased-v1",
    "sentence-transformers/sentence-t5-base",
    "sentence-transformers/all-distilroberta-v1"
]

results = {}
for m in models:
    model = SentenceTransformer(m)
    emb = model.encode("Hello world")
    results[m] = emb.tolist()

with open("/tmp/py_embeddings.json", "w") as f:
    json.dump(results, f)

2. Elixir Verification Script (Bumblebee)

py_data = File.read!("/tmp/py_embeddings.json") |> Jason.decode!()

models = [
  {"unsloth/embeddinggemma-300m", []},
  {"sentence-transformers/all-MiniLM-L6-v2", []},
  {"sentence-transformers/distiluse-base-multilingual-cased-v1", []},
  {"sentence-transformers/sentence-t5-base", []},
  {"sentence-transformers/all-distilroberta-v1", [architecture: :base]}
]

for {model_id, load_opts} <- models do
  IO.puts("\n=== #{model_id} ===")
  {:ok, model_info} = Bumblebee.load_model({:hf, model_id}, load_opts)
  {:ok, model_info} = Bumblebee.load_embedding_head({:hf, model_id}, model_info)
  {:ok, tokenizer} = Bumblebee.load_tokenizer({:hf, model_id})

  serving = Bumblebee.Text.text_embedding(model_info, tokenizer)
  res = Nx.Serving.run(serving, "Hello world")
  bb_emb = res.embedding

  py_emb = Nx.tensor(py_data[model_id], type: :f32)
  diff = Nx.abs(Nx.subtract(bb_emb, py_emb))
  max_diff = Nx.to_number(Nx.reduce_max(diff))
  mean_diff = Nx.to_number(Nx.mean(diff))

  IO.puts("Embedding shape: #{inspect(Nx.shape(bb_emb))}")
  IO.puts("Max absolute diff: #{max_diff}")
  IO.puts("Mean absolute diff: #{mean_diff}")
  IO.puts("Py first 5: #{inspect(Enum.take(Nx.to_flat_list(py_emb), 5))}")
  IO.puts("BB first 5: #{inspect(Enum.take(Nx.to_flat_list(bb_emb), 5))}")
end

Raw Output

=== unsloth/embeddinggemma-300m ===
Embedding shape: {768}
Max absolute diff: 1.043081283569336e-7
Mean absolute diff: 2.321415415451611e-8
Py first 5: [-0.20307935774326324, 0.03475954011082649, 0.06016656383872032, -0.016863875091075897, 0.006666663568466902]
BB first 5: [-0.20307935774326324, 0.03475956246256828, 0.060166530311107635, -0.016863878816366196, 0.006666661240160465]

=== sentence-transformers/all-MiniLM-L6-v2 ===
Embedding shape: {384}
Max absolute diff: 1.1175870895385742e-7
Mean absolute diff: 3.063611231368668e-8
Py first 5: [-0.03447720780968666, 0.031023211777210236, 0.006734978407621384, 0.026108982041478157, -0.039361998438835144]
BB first 5: [-0.034477267414331436, 0.031023208051919937, 0.006734946742653847, 0.026108991354703903, -0.03936201333999634]

=== sentence-transformers/distiluse-base-multilingual-cased-v1 ===
Embedding shape: {512}
Max absolute diff: 4.842877388000488e-8
Mean absolute diff: 1.144235284300521e-8
Py first 5: [0.03743688017129898, 2.512907958589494e-4, -0.04453359916806221, -0.006018291227519512, -0.004640689119696617]
BB first 5: [0.03743689879775047, 2.512825303710997e-4, -0.04453359544277191, -0.0060182781890034676, -0.0046407124027609825]

=== sentence-transformers/sentence-t5-base ===
Embedding shape: {768}
Max absolute diff: 6.43506646156311e-5
Mean absolute diff: 1.2338131455180701e-5
Py first 5: [0.0015573501586914062, -0.056060791015625, 0.0262908935546875, 0.058380126953125, 0.00701904296875]
BB first 5: [0.0015513345133513212, -0.05604176223278046, 0.026304110884666443, 0.05838458240032196, 0.007023626007139683]

=== sentence-transformers/all-distilroberta-v1 ===
Embedding shape: {768}
Max absolute diff: 1.0989606380462646e-7
Mean absolute diff: 2.373483631856743e-8
Py first 5: [0.027807772159576416, -0.024371134117245674, -0.02197802998125553, -0.05726218968629837, 0.034519851207733154]
BB first 5: [0.027807775884866714, -0.024371156468987465, -0.02197805419564247, -0.05726221576333046, 0.034519895911216736]

@dkuku
dkuku marked this pull request as ready for review October 1, 2026 05:14
@dkuku

dkuku commented Oct 1, 2026

Copy link
Copy Markdown
Contributor Author

I will extract the files logic first - then push the rest.

@dkuku dkuku closed this Oct 1, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant