From 5f13901bb5666a8359187fc211d30cfadff52971 Mon Sep 17 00:00:00 2001 From: mika <211269698+mikamikasuki@users.noreply.github.com> Date: Fri, 2 Oct 2026 19:49:48 -0700 Subject: [PATCH] Support NumPy vector initials in accumulate --- toolz/itertoolz.py | 3 ++- toolz/tests/test_itertoolz.py | 28 ++++++++++++++++++++++++++++ 2 files changed, 30 insertions(+), 1 deletion(-) diff --git a/toolz/itertoolz.py b/toolz/itertoolz.py index 05c36cca..7272384b 100644 --- a/toolz/itertoolz.py +++ b/toolz/itertoolz.py @@ -55,7 +55,8 @@ def accumulate(binop, seq, initial=no_default): itertools.accumulate : In standard itertools for Python 3.2+ """ seq = iter(seq) - if initial == no_default: + if (isinstance(initial, (str, collections.UserString)) + and initial == no_default): try: result = next(seq) except StopIteration: diff --git a/toolz/tests/test_itertoolz.py b/toolz/tests/test_itertoolz.py index d20da2b8..a66a49f3 100644 --- a/toolz/tests/test_itertoolz.py +++ b/toolz/tests/test_itertoolz.py @@ -1,4 +1,6 @@ import itertools +import pytest +from collections import UserString from itertools import starmap from toolz.utils import raises from functools import partial @@ -325,6 +327,32 @@ def test_accumulate_works_on_consumable_iterables(): assert list(accumulate(add, iter((1, 2, 3)))) == [1, 3, 6] +def test_accumulate_numpy_initial(): + np = pytest.importorskip('numpy') + values = [np.array([1., 2.]), np.array([3., 4.])] + initial = np.zeros(2) + result = list(accumulate(np.add, values, initial)) + expected = list(itertools.accumulate(values, np.add, initial=initial)) + + assert len(result) == len(expected) + assert result[0] is initial + for actual, wanted in zip(result, expected): + np.testing.assert_array_equal(actual, wanted) + + +def test_accumulate_none_initial(): + assert list(accumulate(lambda acc, item: item, [1, 2], None)) == [None, 1, 2] + + +def test_accumulate_user_string_sentinel(): + assert list(accumulate(add, [1, 2], UserString(no_default2))) == [1, 3] + + initial = UserString('prefix') + result = list(accumulate(add, ['a', 'b'], initial)) + assert result[0] is initial + assert result == ['prefix', 'prefixa', 'prefixab'] + + def test_sliding_window(): assert list(sliding_window(2, [1, 2, 3, 4])) == [(1, 2), (2, 3), (3, 4)] assert list(sliding_window(3, [1, 2, 3, 4])) == [(1, 2, 3), (2, 3, 4)]