summaryrefslogtreecommitdiff
path: root/numpy/ma/tests/test_extras.py
diff options
context:
space:
mode:
authorJouke Witteveen <j.witteveen@gmail.com>2021-11-28 18:34:49 +0100
committerJouke Witteveen <j.witteveen@gmail.com>2022-05-10 18:19:21 +0200
commit4f1d95aa1fa0910b631e6aea91f5c2033593c11e (patch)
tree7f6b984305793db6d49b12917e9e192f154cf4da /numpy/ma/tests/test_extras.py
parentff3a9dae0f4b6e539b1170a9d334dcefe862d28f (diff)
downloadnumpy-4f1d95aa1fa0910b631e6aea91f5c2033593c11e.tar.gz
ENH: Add compressed= argument to ma.ndenumerate
Diffstat (limited to 'numpy/ma/tests/test_extras.py')
-rw-r--r--numpy/ma/tests/test_extras.py9
1 files changed, 9 insertions, 0 deletions
diff --git a/numpy/ma/tests/test_extras.py b/numpy/ma/tests/test_extras.py
index 5ea48ee51..74f744e09 100644
--- a/numpy/ma/tests/test_extras.py
+++ b/numpy/ma/tests/test_extras.py
@@ -1658,6 +1658,8 @@ class TestNDEnumerate:
list(ndenumerate(ordinary)))
assert_equal(list(ndenumerate(ordinary)),
list(ndenumerate(with_mask)))
+ assert_equal(list(ndenumerate(with_mask)),
+ list(ndenumerate(with_mask, compressed=False)))
def test_ndenumerate_allmasked(self):
a = masked_all(())
@@ -1665,7 +1667,11 @@ class TestNDEnumerate:
c = masked_all((2, 3, 4))
assert_equal(list(ndenumerate(a)), [])
assert_equal(list(ndenumerate(b)), [])
+ assert_equal(list(ndenumerate(b, compressed=False)),
+ list(zip(np.ndindex((100,)), 100 * [masked])))
assert_equal(list(ndenumerate(c)), [])
+ assert_equal(list(ndenumerate(c, compressed=False)),
+ list(zip(np.ndindex((2, 3, 4)), 2 * 3 * 4 * [masked])))
def test_ndenumerate_mixedmasked(self):
a = masked_array(np.arange(12).reshape((3, 4)),
@@ -1675,6 +1681,9 @@ class TestNDEnumerate:
items = [((1, 2), 6),
((2, 0), 8), ((2, 1), 9), ((2, 2), 10), ((2, 3), 11)]
assert_equal(list(ndenumerate(a)), items)
+ assert_equal(len(list(ndenumerate(a, compressed=False))), a.size)
+ for coordinate, value in ndenumerate(a, compressed=False):
+ assert_equal(a[coordinate], value)
class TestStack: