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:
Aditya Nikam
2026-09-27 00:30:07 +07:00
committed by GitHub
co-authored by George Panchuk
parent 9ab3ffe9bc
commit 9667d11037
2 changed files with 28 additions and 8 deletions
+14 -7
View File
@@ -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
View File
@@ -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)