Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion toolz/itertoolz.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
28 changes: 28 additions & 0 deletions toolz/tests/test_itertoolz.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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)]
Expand Down