Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 11 additions & 4 deletions fastembed/common/preprocessor_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,12 +47,15 @@ def _valid_context(value: Any) -> int | None:
return value


def _resolve_max_context(tokenizer_config: dict[str, Any], model_dir: Path) -> int:
def _resolve_max_context(
tokenizer_config: dict[str, Any], model_dir: Path, default_max_length: int | None = None
) -> int:
"""Pick the truncation limit, preferring the stricter of the two tokenizer config keys.

`config.json:max_position_embeddings` deliberately is not used as a fallback: it is the size
of the position table, not the usable context, and the two differ per architecture, e.g.
roberta reports 514 for a usable 512.
roberta reports 514 for a usable 512. `default_max_length` is a limit the model class
declares for an export whose tokenizer config carries none; the config keys still win.
"""
candidates = [
context
Expand All @@ -62,6 +65,8 @@ def _resolve_max_context(tokenizer_config: dict[str, Any], model_dir: Path) -> i
)
if context is not None
]
if not candidates and default_max_length is not None:
candidates = [default_max_length]
if not candidates:
raise ValueError(
f"Could not determine the maximum context length for {model_dir}. Set a positive "
Expand All @@ -71,7 +76,9 @@ def _resolve_max_context(tokenizer_config: dict[str, Any], model_dir: Path) -> i
return min(candidates)


def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
def load_tokenizer(
model_dir: Path, default_max_length: int | None = None
) -> tuple[Tokenizer, dict[str, int]]:
tokenizer_path = model_dir / "tokenizer.json"
if not tokenizer_path.exists():
raise ValueError(f"Could not find tokenizer.json in {model_dir}")
Expand All @@ -90,7 +97,7 @@ def load_tokenizer(model_dir: Path) -> tuple[Tokenizer, dict[str, int]]:
with open(str(tokenizer_config_path)) as tokenizer_config_file:
tokenizer_config = json.load(tokenizer_config_file)

max_context = _resolve_max_context(tokenizer_config, model_dir)
max_context = _resolve_max_context(tokenizer_config, model_dir, default_max_length)

tokens_map = load_special_tokens(model_dir)

Expand Down
64 changes: 64 additions & 0 deletions fastembed/text/builtin_sentence_embedding.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,21 @@
from pathlib import Path
from typing import Any, Iterable, Type

import numpy as np

from fastembed.common.types import NumpyArray
from fastembed.common.onnx_model import OnnxOutputContext
from fastembed.common.preprocessor_utils import load_tokenizer
from fastembed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
from fastembed.common.model_description import DenseModelDescription, ModelSource

# Context length for exports whose tokenizer_config.json carries transformers' "unknown"
# sentinel instead of a usable `model_max_length`.
DEFAULT_MAX_LENGTHS: dict[str, int] = {
"google/embeddinggemma-2": 8192,
"google/embeddinggemma-2-q": 8192,
}


supported_builtin_sentence_embedding_models: list[DenseModelDescription] = [
DenseModelDescription(
Expand All @@ -24,6 +34,39 @@
model_file="onnx/model.onnx",
additional_files=["onnx/model.onnx_data"],
),
DenseModelDescription(
model="google/embeddinggemma-2",
dim=768,
description=(
"Text embeddings, Multimodal model used text-only, multilingual, 8192 input tokens "
"truncation, Prefixes for queries/documents: `task: search result | query: {content}` "
"for query, `title: {title | 'none'} | text: {content}` for documents, 2026 year."
),
license="apache-2.0",
size_in_GB=1.08,
sources=ModelSource(
hf="onnx-community/embeddinggemma-2-ONNX",
),
model_file="onnx/model.onnx",
additional_files=["onnx/model.onnx_data"],
),
DenseModelDescription(
model="google/embeddinggemma-2-Q",
dim=768,
description=(
"Text embeddings, Multimodal model used text-only, multilingual, 8192 input tokens "
"truncation, Prefixes for queries/documents: `task: search result | query: {content}` "
"for query, `title: {title | 'none'} | text: {content}` for documents, int8 weights, "
"2026 year."
),
license="apache-2.0",
size_in_GB=0.31,
sources=ModelSource(
hf="onnx-community/embeddinggemma-2-ONNX",
),
model_file="onnx/model_quantized.onnx",
additional_files=["onnx/model_quantized.onnx_data"],
),
DenseModelDescription(
model="ibm-granite/granite-embedding-small-english-r2",
dim=384,
Expand Down Expand Up @@ -58,6 +101,27 @@ def _list_supported_models(cls) -> list[DenseModelDescription]:
"""
return supported_builtin_sentence_embedding_models

def _load_tokenizer(self, model_dir: Path) -> None:
self.tokenizer, self.special_token_to_id = load_tokenizer(
model_dir=model_dir,
default_max_length=DEFAULT_MAX_LENGTHS.get(self.model_name.lower()),
)

def _preprocess_onnx_input(
self, onnx_input: dict[str, NumpyArray], **kwargs: Any
) -> dict[str, NumpyArray]:
"""Feed empty modality inputs to multimodal graphs used for text only.

The embeddinggemma-2 export is a single graph that also declares `image_features`,
`video_features` and `audio_features` inputs. For text, each is a zero-row tensor.
"""
for node in self.model.get_inputs(): # type: ignore[union-attr]
if node.name in onnx_input or not node.name.endswith("_features"):
continue
width = node.shape[-1] if isinstance(node.shape[-1], int) else 0
onnx_input[node.name] = np.zeros((0, width), dtype=np.float32)
return onnx_input

def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[NumpyArray]:
Expand Down
21 changes: 21 additions & 0 deletions tests/test_preprocessor_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -279,6 +279,27 @@ def test_unusable_max_context_raises(make_model_dir, model_max_length, max_lengt
load_tokenizer(model_dir)


@pytest.mark.parametrize(
"model_max_length,max_length,expected",
[
(HF_SENTINEL, None, 8192), # onnx-community/embeddinggemma-2-ONNX
(None, None, 8192),
(512, None, 512), # a usable config key still wins over the declared default
(HF_SENTINEL, 256, 256),
],
)
def test_default_max_length_fills_unusable_config(
make_model_dir, model_max_length, max_length, expected
) -> None:
model_dir = make_model_dir(
tokenizer_config={"model_max_length": model_max_length, "max_length": max_length},
)

tokenizer, _ = load_tokenizer(model_dir, default_max_length=8192)

assert tokenizer.truncation["max_length"] == expected


def test_absent_max_context_keys_raise(make_model_dir) -> None:
model_dir = make_model_dir(
drop_from_tokenizer_config=("model_max_length", "max_length"),
Expand Down
16 changes: 16 additions & 0 deletions tests/test_text_onnx_embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,12 @@
"google/embeddinggemma-300m": np.array(
[-0.08181356, 0.0214127, 0.05120273, -0.03690156, -0.0254504]
),
"google/embeddinggemma-2": np.array(
[-0.03013361, 0.01577923, 0.05849417, -0.00649917, -0.03439925]
),
"google/embeddinggemma-2-Q": np.array(
[-0.03013361, 0.01577923, 0.05849417, -0.00649917, -0.03439925]
),
"Qwen/Qwen3-Embedding-0.6B": np.array(
[-0.01476084, 0.01723184, -0.01195498, -0.07275258, 0.00281229]
),
Expand Down Expand Up @@ -113,16 +119,26 @@

DOC_PREFIXES = {
"google/embeddinggemma-300m": "title: none | text: ",
"google/embeddinggemma-2": "title: none | text: ",
"google/embeddinggemma-2-Q": "title: none | text: ",
}
QUERY_PREFIXES = {
"google/embeddinggemma-300m": "task: search result | query: ",
"google/embeddinggemma-2": "task: search result | query: ",
"google/embeddinggemma-2-Q": "task: search result | query: ",
"Qwen/Qwen3-Embedding-0.6B": QWEN3_INSTRUCT_PREFIX,
"Qwen/Qwen3-Embedding-0.6B-Q": QWEN3_INSTRUCT_PREFIX,
}
CANONICAL_QUERY_VECTOR_VALUES = {
"google/embeddinggemma-300m": np.array(
[-0.22990295, 0.03311195, 0.04290345, -0.03558498, -0.01399477]
),
"google/embeddinggemma-2": np.array(
[-0.02504013, 0.05100445, 0.05460444, -0.02423813, -0.04161748]
),
"google/embeddinggemma-2-Q": np.array(
[-0.02504013, 0.05100445, 0.05460444, -0.02423813, -0.04161748]
),
"Qwen/Qwen3-Embedding-0.6B": np.array(
[-0.01908712, 0.01635596, -0.00356586, -0.03947155, -0.01387356]
),
Expand Down