From 9667d11037d7c4e7c0d3541f2ca90ddc662da63c Mon Sep 17 00:00:00 2001 From: Aditya Nikam <118399608+adityaanikam@users.noreply.github.com> Date: Sat, 26 Sep 2026 23:00:07 +0530 Subject: [PATCH] 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 --- fastembed/common/utils.py | 21 ++++++++++++++------- tests/test_common.py | 15 ++++++++++++++- 2 files changed, 28 insertions(+), 8 deletions(-) diff --git a/fastembed/common/utils.py b/fastembed/common/utils.py index 60d229b..0e271fd 100644 --- a/fastembed/common/utils.py +++ b/fastembed/common/utils.py @@ -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: diff --git a/tests/test_common.py b/tests/test_common.py index f7cae5b..fe9626e 100644 --- a/tests/test_common.py +++ b/tests/test_common.py @@ -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)