Skip to content
Closed
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
7 changes: 7 additions & 0 deletions fastembed/image/transform/operators.py
Original file line number Diff line number Diff line change
Expand Up @@ -325,6 +325,11 @@ def _get_convert_to_rgb(transforms: list[Transform], config: dict[str, Any]) ->

@classmethod
def _get_resize(cls, transforms: list[Transform], config: dict[str, Any]) -> None:
"""Append resize operations selected by the image processor configuration.

ConvNeXT defaults to resizing when do_resize is omitted. Disabling it
skips both resize and its coupled crop while later preprocessing continues.
"""
mode = config.get("image_processor_type", "CLIPImageProcessor")
if mode in ("CLIPImageProcessor", "SiglipImageProcessor"):
if config.get("do_resize", False):
Expand All @@ -344,6 +349,8 @@ def _get_resize(cls, transforms: list[Transform], config: dict[str, Any]) -> Non
)
)
elif mode == "ConvNextFeatureExtractor":
if not config.get("do_resize", True):
return
if "size" in config and "shortest_edge" not in config["size"]:
raise ValueError(
f"Size dictionary must contain 'shortest_edge' key. Got {config['size'].keys()}"
Expand Down
7 changes: 6 additions & 1 deletion fastembed/parallel_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -112,6 +112,11 @@ def __init__(
self.num_active_workers: BaseValue | None = None

def start(self, **kwargs: Any) -> None:
"""Start workers with copies of the supplied initialization options.

The pool's cuda setting takes precedence over a cuda option in kwargs.
When device IDs are provided, assign them to workers in round-robin order.
"""
self.emergency_shutdown = False
self.input_queue = self.ctx.Queue(self.queue_size)
# An emergency shutdown unblocks the feeder thread with EPIPE (see semi_ordered_map), let it
Expand All @@ -125,10 +130,10 @@ def start(self, **kwargs: Any) -> None:

for worker_id in range(0, self.num_workers):
worker_kwargs = deepcopy(kwargs)
worker_kwargs["cuda"] = self.cuda
if self.device_ids:
device_id = self.device_ids[worker_id % len(self.device_ids)]
worker_kwargs["device_id"] = device_id
worker_kwargs["cuda"] = self.cuda

assert hasattr(self.ctx, "Process")
process = self.ctx.Process(
Expand Down
16 changes: 11 additions & 5 deletions fastembed/sparse/bm25.py
Original file line number Diff line number Diff line change
Expand Up @@ -342,6 +342,10 @@ def _term_frequency(self, tokens: list[str]) -> dict[int, float]:

Returns:
dict[int, float]: The token_id to term frequency mapping.

Counts and document length refer to the processed tokens supplied here.
This experimental collision policy adds separate lexical-token weights
at a shared ID; it does not merge counts before applying the BM25 formula.
"""
tf_map: dict[int, float] = {}
counter: defaultdict[str, int] = defaultdict(int)
Expand All @@ -352,10 +356,9 @@ def _term_frequency(self, tokens: list[str]) -> dict[int, float]:
for stemmed_token in counter:
token_id = self.compute_token_id(stemmed_token)
num_occurrences = counter[stemmed_token]
tf_map[token_id] = num_occurrences * (self.k + 1)
tf_map[token_id] /= num_occurrences + self.k * (
1 - self.b + self.b * doc_len / self.avg_len
)
weight = num_occurrences * (self.k + 1)
weight /= num_occurrences + self.k * (1 - self.b + self.b * doc_len / self.avg_len)
tf_map[token_id] = tf_map.get(token_id, 0.0) + weight
return tf_map

@classmethod
Expand All @@ -365,6 +368,9 @@ def compute_token_id(cls, token: str) -> int:
def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[SparseEmbedding]:
"""To emulate BM25 behaviour, we don't need to use weights in the query, and
it's enough to just hash the tokens and assign a weight of 1.0 to them.

Store unique token IDs as int64 to include the absolute signed-hash
boundary value 2147483648 without changing established IDs.
"""
if isinstance(query, str):
query = [query]
Expand All @@ -375,7 +381,7 @@ def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[Spa
stemmed_tokens = self._stem(tokens)
token_ids = np.array(
list(set(self.compute_token_id(token) for token in stemmed_tokens)),
dtype=np.int32,
dtype=np.int64,
)
values = np.ones_like(token_ids)
yield SparseEmbedding(indices=token_ids, values=values)
Expand Down
7 changes: 6 additions & 1 deletion fastembed/sparse/sparse_embedding_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,8 +25,13 @@ def as_dict(self) -> dict[int, float]:

@classmethod
def from_dict(cls, data: dict[int, float]) -> "SparseEmbedding":
"""Build aligned indices and values from an index-to-weight mapping.

Mapping iteration order is retained. Empty input produces empty values
and int64 indices, so the indices remain valid for NumPy indexing.
"""
if len(data) == 0:
return cls(values=np.array([]), indices=np.array([]))
return cls(values=np.array([]), indices=np.array([], dtype=np.int64))
indices, values = zip(*data.items())
return cls(values=np.array(values), indices=np.array(indices))

Expand Down
7 changes: 6 additions & 1 deletion fastembed/sparse/splade_pp.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,10 +38,15 @@ class SpladePP(SparseTextEmbeddingBase, OnnxTextModel[SparseEmbedding]):
def _post_process_onnx_output(
self, output: OnnxOutputContext, **kwargs: Any
) -> Iterable[SparseEmbedding]:
"""Yield sparse scores from log-ReLU weights and masked maximum pooling.

Use log1p to retain tiny positive logits that adding to one would lose.
An attention mask is required; dimensions with zero pooled weight are omitted.
"""
if output.attention_mask is None:
raise ValueError("attention_mask must be provided for document post-processing")

relu_log = np.log(1 + np.maximum(output.model_output, 0))
relu_log = np.log1p(np.maximum(output.model_output, 0))

weighted_log = relu_log * np.expand_dims(output.attention_mask, axis=-1)

Expand Down
135 changes: 135 additions & 0 deletions tests/test_bm25_hash_collisions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,135 @@
"""Regression tests for order-independent BM25 weights when token hashes collide."""

from collections import Counter
from pathlib import Path

import mmh3
import pytest

from fastembed.sparse.bm25 import Bm25


COLLIDING_TOKENS = ("kitchens", "prostaglandins")
COLLIDING_TOKEN_ID = 358434922


@pytest.fixture
def model() -> Bm25:
"""Create only the BM25 state required to calculate term-frequency weights."""
instance = Bm25.__new__(Bm25)
instance.k = 1.2
instance.b = 0.75
instance.avg_len = 256.0
return instance


def _assert_bag_weights_are_order_independent(model: Bm25, tokens: list[str]) -> None:
"""Check that permutations preserve hash IDs and weights within rounding tolerance."""
expected_keys = {model.compute_token_id(token) for token in tokens}
reference = model._term_frequency(tokens)
assert set(reference) == expected_keys
for reordered in (
list(reversed(tokens)),
sorted(tokens),
sorted(tokens, reverse=True),
tokens[1:] + tokens[:1],
):
result = model._term_frequency(reordered)
assert set(result) == expected_keys
assert result == pytest.approx(reference, rel=1e-12, abs=1e-12)


def test_real_tokens_collide_after_absolute_value_of_signed_hash() -> None:
"""Verify a real token pair with opposite signed hashes and the same sparse ID."""
assert mmh3.hash("kitchens") == COLLIDING_TOKEN_ID
assert mmh3.hash("prostaglandins") == -COLLIDING_TOKEN_ID
assert {Bm25.compute_token_id(token) for token in COLLIDING_TOKENS} == {COLLIDING_TOKEN_ID}


@pytest.mark.parametrize("counts", [(1, 2), (3, 1), (2, 5), (7, 3)])
def test_real_collision_weights_are_independent_of_token_order(
model: Bm25, counts: tuple[int, int]
) -> None:
"""Keep unequal colliding token counts independent of first-appearance order."""
tokens = [COLLIDING_TOKENS[0]] * counts[0] + [COLLIDING_TOKENS[1]] * counts[1]

_assert_bag_weights_are_order_independent(model, tokens)


def test_real_collision_with_other_tokens_preserves_keys_and_order_independence(
model: Bm25,
) -> None:
"""Preserve sparse IDs and order independence when collisions mix with other tokens."""
tokens = ["kitchens", "hello", "kitchens", "prostaglandins", "hello", "kitchens"]

_assert_bag_weights_are_order_independent(model, tokens)


@pytest.mark.parametrize("counts", [(1, 2, 3), (2, 4, 1), (5, 1, 2)])
def test_three_distinct_colliding_tokens_are_order_independent(
model: Bm25, monkeypatch: pytest.MonkeyPatch, counts: tuple[int, int, int]
) -> None:
"""Cover controlled three-way collisions without fixing a collision scoring policy."""
colliding = ("alpha", "beta", "gamma")
original_hash = model.compute_token_id

def controlled_hash(token: str) -> int:
"""Force selected tokens to collide while preserving ordinary token hashes."""
return COLLIDING_TOKEN_ID if token in colliding else original_hash(token)

monkeypatch.setattr(model, "compute_token_id", controlled_hash)
tokens = [token for token, count in zip(colliding, counts) for _ in range(count)]
tokens += ["hello", "hello"]

_assert_bag_weights_are_order_independent(model, tokens)


@pytest.mark.parametrize(
"tokens",
[
["hello"],
["hello", "hello", "hello"],
["hello", "world"],
["hello", "world", "hello", "third", "world", "world"],
],
)
def test_noncolliding_tokens_preserve_bm25_term_frequency_formula(
model: Bm25, tokens: list[str]
) -> None:
"""Preserve the standard BM25 formula for ordinary and repeated noncolliding tokens."""
counts = Counter(tokens)
token_ids = {token: model.compute_token_id(token) for token in counts}
assert len(set(token_ids.values())) == len(counts)
expected = {
token_ids[token]: count
* (model.k + 1)
/ (count + model.k * (1 - model.b + model.b * len(tokens) / model.avg_len))
for token, count in counts.items()
}

assert model._term_frequency(tokens) == pytest.approx(expected, rel=1e-12, abs=1e-12)


def test_empty_tokens_have_no_term_frequencies(model: Bm25) -> None:
"""Return an empty weight mapping when there are no document tokens."""
assert model._term_frequency([]) == {}


def test_document_embeddings_of_the_same_colliding_token_bag_are_order_independent(
tmp_path: Path,
) -> None:
"""Keep actual document embeddings stable when colliding tokens are reordered."""
model = Bm25(
"Qdrant/bm25",
cache_dir=str(tmp_path),
specific_model_path=str(tmp_path),
disable_stemmer=True,
local_files_only=True,
)
embeddings = list(
model.embed(["kitchens kitchens prostaglandins", "prostaglandins kitchens kitchens"])
)

assert len(embeddings) == 2
assert set(embeddings[0].as_dict()) == set(embeddings[1].as_dict()) == {COLLIDING_TOKEN_ID}
assert embeddings[0].as_dict() == pytest.approx(embeddings[1].as_dict(), rel=1e-12, abs=1e-12)
104 changes: 104 additions & 0 deletions tests/test_bm25_query_index_overflow.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
"""Regression tests for BM25 query token IDs beyond the signed int32 range."""

from pathlib import Path

import mmh3
import numpy as np
import pytest

from fastembed.sparse.bm25 import Bm25


BOUNDARY_TOKEN = "ad1u66pi"
BOUNDARY_TOKEN_ID = 2**31


@pytest.fixture
def model(tmp_path: Path) -> Bm25:
"""Create a BM25 instance with local paths and stemming disabled."""
return Bm25(
"Qdrant/bm25",
cache_dir=str(tmp_path),
specific_model_path=str(tmp_path),
disable_stemmer=True,
local_files_only=True,
)


def test_boundary_token_hash_exceeds_signed_int32_after_absolute_value() -> None:
"""Verify that a real token's absolute hash reaches the int32 overflow boundary."""
assert mmh3.hash(BOUNDARY_TOKEN) == -(2**31)
assert Bm25.compute_token_id(BOUNDARY_TOKEN) == BOUNDARY_TOKEN_ID


def test_query_embedding_preserves_boundary_token_id(model: Bm25) -> None:
"""Keep the boundary token's exact ID and unit query weight in int64 indices."""
embedding = list(model.query_embed(BOUNDARY_TOKEN))[0]

assert embedding.indices.dtype == np.int64
assert embedding.indices.tolist() == [BOUNDARY_TOKEN_ID]
assert embedding.values.tolist() == [1]


def test_document_and_query_embeddings_use_same_boundary_token_id(model: Bm25) -> None:
"""Use the same boundary token ID for document and query embeddings."""
document_embedding = list(model.embed(BOUNDARY_TOKEN))[0]
query_embedding = list(model.query_embed(BOUNDARY_TOKEN))[0]

assert (
document_embedding.indices.tolist()
== query_embedding.indices.tolist()
== [BOUNDARY_TOKEN_ID]
)


@pytest.mark.parametrize("query", ["AD1U66PI", "(ad1u66pi)!", "ad1u66pi ad1u66pi"])
def test_query_normalization_preserves_boundary_token_id(model: Bm25, query: str) -> None:
"""Preserve boundary IDs through case folding, punctuation removal, and deduplication."""
embedding = list(model.query_embed(query))[0]

assert embedding.indices.tolist() == [BOUNDARY_TOKEN_ID]
assert embedding.values.tolist() == [1]


@pytest.mark.parametrize("as_generator", [False, True])
def test_query_iterables_handle_mixed_ordinary_and_boundary_tokens(
model: Bm25, as_generator: bool
) -> None:
"""Accept lists and generators containing ordinary, boundary, and empty queries."""
queries = ["hello", BOUNDARY_TOKEN, f"hello {BOUNDARY_TOKEN}", ""]
query_input = (query for query in queries) if as_generator else queries

embeddings = list(model.query_embed(query_input))

hello_id = model.compute_token_id("hello")
expected_indices = [{hello_id}, {BOUNDARY_TOKEN_ID}, {hello_id, BOUNDARY_TOKEN_ID}, set()]
assert len(embeddings) == len(expected_indices)
for embedding, expected in zip(embeddings, expected_indices):
assert np.issubdtype(embedding.indices.dtype, np.integer)
assert set(embedding.indices.tolist()) == expected
assert embedding.values.tolist() == [1] * len(expected)


@pytest.mark.parametrize("query", ["hello", "hello world hello"])
def test_ordinary_query_tokens_have_unit_weights_and_are_deduplicated(
model: Bm25, query: str
) -> None:
"""Keep ordinary token hashing, deduplication, and unit query weights unchanged."""
embedding = list(model.query_embed(query))[0]
expected_indices = {model.compute_token_id(token) for token in query.split()}

assert np.issubdtype(embedding.indices.dtype, np.integer)
assert set(embedding.indices.tolist()) == expected_indices
assert embedding.values.tolist() == [1] * len(expected_indices)


def test_empty_queries_have_empty_integer_indices_and_values(model: Bm25) -> None:
"""Produce empty arrays with integer indices for queries without usable tokens."""
embeddings = list(model.query_embed(["", "!!!"]))

assert len(embeddings) == 2
for embedding in embeddings:
assert np.issubdtype(embedding.indices.dtype, np.integer)
assert embedding.indices.shape == (0,)
assert embedding.values.shape == (0,)
Loading