mirror of
https://github.com/qdrant/fastembed.git
synced 2026-10-03 11:27:40 -05:00
Raise ValueError for non-positive batch_size in iter_batch (#721)
* 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>
This commit is contained in:
co-authored by
George Panchuk
parent
9ab3ffe9bc
commit
9667d11037
@@ -43,16 +43,23 @@ def last_token_pooling(input_array: NumpyArray, attention_mask: NDArray[np.int64
|
||||
|
||||
|
||||
def iter_batch(iterable: Iterable[T], size: int) -> Iterable[list[T]]:
|
||||
"""
|
||||
"""Validate the batch size immediately and consume the iterable lazily.
|
||||
|
||||
>>> list(iter_batch([1,2,3,4,5], 3))
|
||||
[[1, 2, 3], [4, 5]]
|
||||
"""
|
||||
source_iter = iter(iterable)
|
||||
while source_iter:
|
||||
b = list(islice(source_iter, size))
|
||||
if len(b) == 0:
|
||||
break
|
||||
yield b
|
||||
if size < 1:
|
||||
raise ValueError(f"batch_size must be >= 1, got {size}")
|
||||
|
||||
def batches() -> Iterable[list[T]]:
|
||||
source_iter = iter(iterable)
|
||||
while source_iter:
|
||||
b = list(islice(source_iter, size))
|
||||
if len(b) == 0:
|
||||
break
|
||||
yield b
|
||||
|
||||
return batches()
|
||||
|
||||
|
||||
def define_cache_dir(cache_dir: str | None = None) -> Path:
|
||||
|
||||
+14
-1
@@ -1,4 +1,5 @@
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from fastembed import (
|
||||
TextEmbedding,
|
||||
@@ -7,7 +8,7 @@ from fastembed import (
|
||||
LateInteractionMultimodalEmbedding,
|
||||
LateInteractionTextEmbedding,
|
||||
)
|
||||
from fastembed.common.utils import last_token_pooling
|
||||
from fastembed.common.utils import iter_batch, last_token_pooling
|
||||
|
||||
|
||||
def test_text_list_supported_models():
|
||||
@@ -59,3 +60,15 @@ def test_last_token_pooling_with_left_padding():
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user