Skip to content

Support bidirectional decoder backbones for text embeddings - #475

Closed
dkuku wants to merge 2 commits into
elixir-nx:mainfrom
dkuku:feat/bidirectional-decoder-embeddings
Closed

dkuku wants to merge 2 commits into
elixir-nx:mainfrom
dkuku:feat/bidirectional-decoder-embeddings

Conversation

@dkuku

@dkuku dkuku commented Oct 2, 2026

Copy link
Copy Markdown
Contributor

Decoder embedding checkpoints need their declared bidirectional attention mode to be honored. This adds that support to Llama, Mistral, and Qwen3, and corrects Gemma3Text bidirectional window handling, while retaining causal defaults.

Depends on #474. This draft targets upstream main, so its current diff includes the prerequisite RoPE commit. The embedding-only changes are in commit 11d6863. Merge the prerequisite first, then rebase this branch before merging.

  • Load attention-mode configuration with model-appropriate precedence and reject cache/generation use in bidirectional mode.
  • Honor Qwen3 sliding-window configuration and per-layer attention types; keep model-specific window rules in each model.
  • Add attention-boundary, padding, causal-cache, and generation regressions, plus 12 deterministic local checkpoints comparing complete hidden states with Python Transformers.
  • Document loading embedding backbones, RoPE options, and the additional projection/pooling requirements of EmbeddingGemma.

Validation: 83 targeted tests passed, including existing Llama, Mistral, Qwen3, Gemma3Text, and Phi3 tests. Formatting, compilation with --warnings-as-errors, unused-dependency checks, and diff whitespace checks pass. All 12 fixture checkpoints were regenerated using PyTorch 2.7.1 and Transformers commit 3693f8d26311305e914735a6373fb03468d6aaa0; their configurations, weights, and Python reference outputs reproduce exactly. The generator is included under test/fixtures/embedding_models/.

The full test suite and opt-in real EmbeddingGemma integration were not run. This supports the transformer backbone, not the complete trained EmbeddingGemma projection pipeline. The dependency reference above links a PR and does not close an issue.

AI disclosure: this contribution was developed with substantial LLM/coding-agent assistance, including implementation, tests, and review. Validation above was executed by the coding agent.

@dkuku dkuku closed this Oct 2, 2026
@dkuku

dkuku commented Oct 2, 2026

Copy link
Copy Markdown
Contributor Author

Sorry, agent went crazy - I won't ask for a review of 9k lines of code.

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