Skip to content

[Stacked on #468] Add EmbeddingGemma model with dense bottleneck projections - #469

Closed
dkuku wants to merge 1 commit into
elixir-nx:mainfrom
dkuku:feature/embedding-gemma
Closed

dkuku wants to merge 1 commit into
elixir-nx:mainfrom
dkuku:feature/embedding-gemma

Conversation

@dkuku

@dkuku dkuku commented Sep 29, 2026

Copy link
Copy Markdown
Contributor

Note

Stacked PR: This PR is stacked on top of #468 (feature/gemma3-bidirectional).
Please review #468 first. Once #468 is merged into main, the diff here will automatically update to only show the EmbeddingGemma additions.

This PR adds support for EmbeddingGemma (google/embeddinggemma-300m), a lightweight text embedding model based on Gemma 3.

Model Architecture

In the SentenceTransformers and HuggingFace ecosystem, EmbeddingGemma consists of:

  1. Bidirectional Gemma 3 backbone: Uses Gemma 3 with non-causal bidirectional attention (implemented in Add bidirectional attention support to Gemma 3聽#468).
  2. Attention-masked mean pooling: Aggregates token representations taking the attention mask into account.
  3. Dense bottleneck projections: Two sequential linear projection layers (2_Dense and 3_Dense), projecting $768 \to 3072 \to 768$.
  4. L2 Normalization: Output embeddings are normalized to unit length.

Changes

  • Implemented Bumblebee.Text.EmbeddingGemma with :base architecture producing :embedding, :pooled_state, and :hidden_state.
  • Extended PyTorch parameter loading to support {prefix, path} tuples in Bumblebee.Conversion.PyTorchParams so subdirectory weights (2_Dense/model.safetensors and 3_Dense/model.safetensors) load without key collisions.
  • Added extra_params_files/2 callback in Bumblebee.Text.EmbeddingGemma and handled directory loading in Bumblebee.load_model.
  • Registered EmbeddingGemma in Bumblebee.load_model/load_spec architecture mapping and registered embeddinggemma tokenizer mapping to :gemma.
  • Added unit tests for architecture, pooling, L2 normalization, params mapping, and an integration test verified against sentence-transformers with unsloth/embeddinggemma-300m.

Reference verification (Python)

Output verified against PyTorch sentence-transformers with atol: 1.0e-4 using unsloth/embeddinggemma-300m:

from sentence_transformers import SentenceTransformer

model = SentenceTransformer("unsloth/embeddinggemma-300m")
emb = model.encode("Hello world")

print([round(float(x), 6) for x in emb[:5]])
# => [-0.203079, 0.034760, 0.060167, -0.016864, 0.006667]

Bumblebee output:

res = Nx.Serving.run(serving, "Hello world")
# res.embedding[0..4] => [-0.203079, 0.034759, 0.060166, -0.016863, 0.006666]

@dkuku
dkuku marked this pull request as draft September 29, 2026 20:27
@dkuku
dkuku force-pushed the feature/embedding-gemma branch 2 times, most recently from d2276ec to 6d86b0d Compare September 29, 2026 20:34
pytorch_state
# Prepend path-specific prefix if configured. This avoids key collisions when
# loading multiple state dict files that share internal tensor names (e.g. "linear.weight")
case path_prefixes[path] do

@dkuku dkuku Sep 29, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Context on :path_prefixes for the reviewer:

Unlike models directly from huggingface/transformers where weights reside in a single root model.safetensors, EmbeddingGemma (google/embeddinggemma-300m) comes from SentenceTransformers and splits the architecture across subdirectories:

  • Root model.safetensors contains the Gemma 3 backbone.
  • 2_Dense/model.safetensors contains the 1st linear bottleneck projection (768 -> 3072).
  • 3_Dense/model.safetensors contains the 2nd linear bottleneck projection (3072 -> 768).

Because both 2_Dense and 3_Dense internally name their tensors linear.weight, merging them directly via Map.merge/2 would cause key collisions and overwrite the 1st dense layer. The :path_prefixes option prepends the folder name (2_Dense. / 3_Dense.) to disambiguate them.

If you have a preferred alternative design for handling modular / multi-directory checkpoints in Bumblebee, I'm happy to adapt or refactor this!

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