Skip to content

Add bidirectional attention support to Gemma 3 - #468

Merged
jonatanklosko merged 1 commit into
elixir-nx:mainfrom
dkuku:feature/gemma3-bidirectional
Sep 30, 2026
Merged

jonatanklosko merged 1 commit into
elixir-nx:mainfrom
dkuku:feature/gemma3-bidirectional

Conversation

@dkuku

@dkuku dkuku commented Sep 29, 2026

Copy link
Copy Markdown

This PR adds support for bidirectional attention in Bumblebee.Text.Gemma3Text.

Motivation

In huggingface/transformers, Gemma3TextConfig introduces the use_bidirectional_attention boolean configuration option. When set to true, the model uses non-causal (bidirectional) attention across all layers rather than causal masking. This is required for embedding and encoder tasks built on top of the Gemma 3 backbone (such as EmbeddingGemma).

Changes

  • Added :use_bidirectional_attention option (defaulting to false) to Bumblebee.Text.Gemma3Text.
  • Configured causal: not spec.use_bidirectional_attention when building transformer blocks.
  • Mapped use_bidirectional_attention in Bumblebee.HuggingFace.Transformers.Config.load/2.
  • Added integration test in test/bumblebee/text/gemma3_text_test.exs with reference values generated from Python transformers.

Reference verification (Python)

Output verified against PyTorch transformers with atol: 1.0e-4 using bumblebee-testing/tiny-random-Gemma3TextModel:

import torch
from transformers import AutoModel, AutoConfig

model_id = "bumblebee-testing/tiny-random-Gemma3TextModel"
config = AutoConfig.from_pretrained(model_id)
config.use_bidirectional_attention = True
model = AutoModel.from_pretrained(model_id, config=config)
model.eval()

inputs = {
    "input_ids": torch.tensor([[10, 20, 30, 40, 50, 60, 70, 80, 0, 0]]),
    "attention_mask": torch.tensor([[1, 1, 1, 1, 1, 1, 1, 1, 0, 0]])
}

with torch.no_grad():
    out = model(**inputs)
    print(out.last_hidden_state[:, 1:4, 1:4])
    # => tensor([[[ 0.3755, -0.6418,  1.9320],
    # =>          [-0.5731,  0.3630,  2.5736],
    # =>          [-0.0530,  0.5178,  1.3310]]])

@jonatanklosko

Copy link
Copy Markdown
Member

Please rebase on top of main, and then we can merge this.

@dkuku
dkuku force-pushed the feature/gemma3-bidirectional branch from d78e616 to dc96970 Compare September 30, 2026 17:27
@dkuku

dkuku commented Sep 30, 2026

Copy link
Copy Markdown
Author

馃憤馃徎 Done

@jonatanklosko
jonatanklosko merged commit 5d02fe4 into elixir-nx:main Sep 30, 2026
2 checks passed
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.

2 participants