mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
fix(core): make abatch_iterate consistent with batch_iterate for None and zero size (#39367)
This commit is contained in:
1 parent
162f9d9582
commit
f9ee55d94c
4 files changed
+50
-7
No files matched your search
@@ -323,25 +323,35 @@ class aclosing(AbstractAsyncContextManager[Any]): # noqa: N801
|
||||
|
||||
|
||||
async def abatch_iterate(
|
||||
size: int, iterable: AsyncIterable[T]
|
||||
size: int | None, iterable: AsyncIterable[T]
|
||||
) -> AsyncIterator[list[T]]:
|
||||
"""Utility batching function for async iterables.
|
||||
|
||||
Args:
|
||||
size: The size of the batch.
|
||||
|
||||
If `None`, returns a single batch.
|
||||
iterable: The async iterable to batch.
|
||||
|
||||
Yields:
|
||||
The batches.
|
||||
|
||||
Raises:
|
||||
ValueError: If `size` is not `None` and is not a positive integer.
|
||||
"""
|
||||
if size is None:
|
||||
single_batch = [el async for el in iterable]
|
||||
if single_batch:
|
||||
yield single_batch
|
||||
return
|
||||
if size <= 0:
|
||||
msg = f"Batch size must be a positive integer, got {size}."
|
||||
raise ValueError(msg)
|
||||
batch: list[T] = []
|
||||
async for element in iterable:
|
||||
if len(batch) < size:
|
||||
batch.append(element)
|
||||
|
||||
batch.append(element)
|
||||
if len(batch) >= size:
|
||||
yield batch
|
||||
batch = []
|
||||
|
||||
if batch:
|
||||
yield batch
|
||||
@@ -214,7 +214,13 @@ def batch_iterate(size: int | None, iterable: Iterable[T]) -> Iterator[list[T]]:
|
||||
|
||||
Yields:
|
||||
The batches of the iterable.
|
||||
|
||||
Raises:
|
||||
ValueError: If `size` is not `None` and is not a positive integer.
|
||||
"""
|
||||
if size is not None and size <= 0:
|
||||
msg = f"Batch size must be a positive integer, got {size}."
|
||||
raise ValueError(msg)
|
||||
it = iter(iterable)
|
||||
while True:
|
||||
chunk = list(islice(it, size))
|
||||
|
||||
@@ -12,10 +12,14 @@ from langchain_core.utils.aiter import abatch_iterate
|
||||
(3, [10, 20, 30, 40, 50], [[10, 20, 30], [40, 50]]),
|
||||
(1, [100, 200, 300], [[100], [200], [300]]),
|
||||
(4, [], []),
|
||||
(None, [1, 2, 3], [[1, 2, 3]]),
|
||||
(None, [], []),
|
||||
],
|
||||
)
|
||||
async def test_abatch_iterate(
|
||||
input_size: int, input_iterable: list[str], expected_output: list[list[str]]
|
||||
input_size: int | None,
|
||||
input_iterable: list[str],
|
||||
expected_output: list[list[str]],
|
||||
) -> None:
|
||||
"""Test batching function."""
|
||||
|
||||
@@ -29,3 +33,15 @@ async def test_abatch_iterate(
|
||||
|
||||
output = [el async for el in iterator_]
|
||||
assert output == expected_output
|
||||
|
||||
|
||||
@pytest.mark.parametrize("input_size", [0, -1])
|
||||
async def test_abatch_iterate_invalid_size(input_size: int) -> None:
|
||||
"""Non-positive sizes should raise instead of silently discarding data."""
|
||||
|
||||
async def _to_async_iterable(iterable: list[int]) -> AsyncIterator[int]:
|
||||
for item in iterable:
|
||||
yield item
|
||||
|
||||
with pytest.raises(ValueError, match="positive integer"):
|
||||
_ = [el async for el in abatch_iterate(input_size, _to_async_iterable([1, 2]))]
|
||||
@@ -10,10 +10,21 @@ from langchain_core.utils.iter import batch_iterate
|
||||
(3, [10, 20, 30, 40, 50], [[10, 20, 30], [40, 50]]),
|
||||
(1, [100, 200, 300], [[100], [200], [300]]),
|
||||
(4, [], []),
|
||||
(None, [1, 2, 3], [[1, 2, 3]]),
|
||||
(None, [], []),
|
||||
],
|
||||
)
|
||||
def test_batch_iterate(
|
||||
input_size: int, input_iterable: list[str], expected_output: list[list[str]]
|
||||
input_size: int | None,
|
||||
input_iterable: list[str],
|
||||
expected_output: list[list[str]],
|
||||
) -> None:
|
||||
"""Test batching function."""
|
||||
assert list(batch_iterate(input_size, input_iterable)) == expected_output
|
||||
|
||||
|
||||
@pytest.mark.parametrize("input_size", [0, -1])
|
||||
def test_batch_iterate_invalid_size(input_size: int) -> None:
|
||||
"""Non-positive sizes should raise instead of silently discarding data."""
|
||||
with pytest.raises(ValueError, match="positive integer"):
|
||||
list(batch_iterate(input_size, [1, 2, 3]))
|
||||
Reference in new issue
Block a user