mirror of
https://github.com/qdrant/fastembed.git
synced 2026-10-03 19:37:46 -05:00
* Raise ValueError for non-positive batch_size in iter_batch iter_batch passed batch_size straight into islice with no lower-bound check. A batch_size of 0 made every islice call return an empty list immediately, so embed and rerank returned an empty result with no error across every modality that shares this helper. A negative batch_size hit islice's own argument validation and raised a cryptic ValueError instead of a clear one. Validate batch_size >= 1 once in iter_batch itself, since every embed and rerank entrypoint across dense text, sparse text, image, late-interaction and cross-encoder rerank funnels through it. Added test_iter_batch_rejects_non_positive_size and test_iter_batch_accepts_positive_size to tests/test_common.py. The new rejection test fails on unmodified code (batch_size=0 returns silently instead of raising) and passes with the fix. Fixes #719 * fix: fix batch size check for parallel execution --------- Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
75 lines
2.5 KiB
Python
75 lines
2.5 KiB
Python
import numpy as np
|
|
import pytest
|
|
|
|
from fastembed import (
|
|
TextEmbedding,
|
|
SparseTextEmbedding,
|
|
ImageEmbedding,
|
|
LateInteractionMultimodalEmbedding,
|
|
LateInteractionTextEmbedding,
|
|
)
|
|
from fastembed.common.utils import iter_batch, last_token_pooling
|
|
|
|
|
|
def test_text_list_supported_models():
|
|
for model_type in [
|
|
TextEmbedding,
|
|
SparseTextEmbedding,
|
|
ImageEmbedding,
|
|
LateInteractionMultimodalEmbedding,
|
|
LateInteractionTextEmbedding,
|
|
]:
|
|
supported_models = model_type.list_supported_models()
|
|
assert isinstance(supported_models, list)
|
|
description = supported_models[0]
|
|
assert isinstance(description, dict)
|
|
|
|
assert "model" in description and description["model"]
|
|
if model_type != SparseTextEmbedding:
|
|
assert "dim" in description and description["dim"]
|
|
assert "license" in description and description["license"]
|
|
assert "size_in_GB" in description and description["size_in_GB"]
|
|
assert "model_file" in description and description["model_file"]
|
|
assert "sources" in description and description["sources"]
|
|
assert "hf" in description["sources"] or "url" in description["sources"]
|
|
|
|
|
|
def test_last_token_pooling():
|
|
token_embeddings = np.array(
|
|
[
|
|
[[1.0, 1.0], [2.0, 2.0], [9.0, 9.0], [9.0, 9.0]], # 2 real tokens, then padding
|
|
[[3.0, 3.0], [4.0, 4.0], [5.0, 5.0], [6.0, 6.0]], # no padding
|
|
]
|
|
)
|
|
attention_mask = np.array([[1, 1, 0, 0], [1, 1, 1, 1]], dtype=np.int64)
|
|
|
|
pooled = last_token_pooling(token_embeddings, attention_mask)
|
|
|
|
assert np.allclose(pooled, [[2.0, 2.0], [6.0, 6.0]])
|
|
|
|
|
|
def test_last_token_pooling_with_left_padding():
|
|
token_embeddings = np.array(
|
|
[
|
|
[[9.0, 9.0], [9.0, 9.0], [1.0, 1.0], [2.0, 2.0]], # padding, then 2 real tokens
|
|
[[3.0, 3.0], [4.0, 4.0], [5.0, 5.0], [6.0, 6.0]], # no padding
|
|
]
|
|
)
|
|
attention_mask = np.array([[0, 0, 1, 1], [1, 1, 1, 1]], dtype=np.int64)
|
|
|
|
pooled = last_token_pooling(token_embeddings, attention_mask)
|
|
|
|
assert np.allclose(pooled, [[2.0, 2.0], [6.0, 6.0]])
|
|
|
|
|
|
def test_iter_batch_accepts_positive_size():
|
|
assert list(iter_batch([1, 2, 3, 4, 5], 3)) == [[1, 2, 3], [4, 5]]
|
|
|
|
|
|
def test_iter_batch_rejects_non_positive_size():
|
|
with pytest.raises(ValueError, match="batch_size must be >= 1, got 0"):
|
|
iter_batch([1, 2, 3], 0)
|
|
|
|
with pytest.raises(ValueError, match="batch_size must be >= 1, got -1"):
|
|
iter_batch([1, 2, 3], -1)
|