Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
test_shape_base.py814 linesDownload Raw Back to tests
1import functools
2import sys
3
4import pytest
5
6import numpy as np
7from numpy import (
8    apply_along_axis,
9    apply_over_axes,
10    array_split,
11    column_stack,
12    dsplit,
13    dstack,
14    expand_dims,
15    hsplit,
16    kron,
17    put_along_axis,
18    split,
19    take_along_axis,
20    tile,
21    vsplit,
22)
23from numpy.exceptions import AxisError
24from numpy.testing import assert_, assert_array_equal, assert_equal, assert_raises
25
26IS_64BIT = sys.maxsize > 2**32
27
28
29def _add_keepdims(func):
30    """ hack in keepdims behavior into a function taking an axis """
31    @functools.wraps(func)
32    def wrapped(a, axis, **kwargs):
33        res = func(a, axis=axis, **kwargs)
34        if axis is None:
35            axis = 0  # res is now a scalar, so we can insert this anywhere
36        return np.expand_dims(res, axis=axis)
37    return wrapped
38
39
40class TestTakeAlongAxis:
41    def test_argequivalent(self):
42        """ Test it translates from arg<func> to <func> """
43        from numpy.random import rand
44        a = rand(3, 4, 5)
45
46        funcs = [
47            (np.sort, np.argsort, {}),
48            (_add_keepdims(np.min), _add_keepdims(np.argmin), {}),
49            (_add_keepdims(np.max), _add_keepdims(np.argmax), {}),
50            #(np.partition, np.argpartition, dict(kth=2)),
51        ]
52
53        for func, argfunc, kwargs in funcs:
54            for axis in list(range(a.ndim)) + [None]:
55                a_func = func(a, axis=axis, **kwargs)
56                ai_func = argfunc(a, axis=axis, **kwargs)
57                assert_equal(a_func, take_along_axis(a, ai_func, axis=axis))
58
59    def test_invalid(self):
60        """ Test it errors when indices has too few dimensions """
61        a = np.ones((10, 10))
62        ai = np.ones((10, 2), dtype=np.intp)
63
64        # sanity check
65        take_along_axis(a, ai, axis=1)
66
67        # not enough indices
68        assert_raises(ValueError, take_along_axis, a, np.array(1), axis=1)
69        # bool arrays not allowed
70        assert_raises(IndexError, take_along_axis, a, ai.astype(bool), axis=1)
71        # float arrays not allowed
72        assert_raises(IndexError, take_along_axis, a, ai.astype(float), axis=1)
73        # invalid axis
74        assert_raises(AxisError, take_along_axis, a, ai, axis=10)
75        # invalid indices
76        assert_raises(ValueError, take_along_axis, a, ai, axis=None)
77
78    def test_empty(self):
79        """ Test everything is ok with empty results, even with inserted dims """
80        a = np.ones((3, 4, 5))
81        ai = np.ones((3, 0, 5), dtype=np.intp)
82
83        actual = take_along_axis(a, ai, axis=1)
84        assert_equal(actual.shape, ai.shape)
85
86    def test_broadcast(self):
87        """ Test that non-indexing dimensions are broadcast in both directions """
88        a = np.ones((3, 4, 1))
89        ai = np.ones((1, 2, 5), dtype=np.intp)
90        actual = take_along_axis(a, ai, axis=1)
91        assert_equal(actual.shape, (3, 2, 5))
92
93
94class TestPutAlongAxis:
95    def test_replace_max(self):
96        a_base = np.array([[10, 30, 20], [60, 40, 50]])
97
98        for axis in list(range(a_base.ndim)) + [None]:
99            # we mutate this in the loop
100            a = a_base.copy()
101
102            # replace the max with a small value
103            i_max = _add_keepdims(np.argmax)(a, axis=axis)
104            put_along_axis(a, i_max, -99, axis=axis)
105
106            # find the new minimum, which should max
107            i_min = _add_keepdims(np.argmin)(a, axis=axis)
108
109            assert_equal(i_min, i_max)
110
111    def test_broadcast(self):
112        """ Test that non-indexing dimensions are broadcast in both directions """
113        a = np.ones((3, 4, 1))
114        ai = np.arange(10, dtype=np.intp).reshape((1, 2, 5)) % 4
115        put_along_axis(a, ai, 20, axis=1)
116        assert_equal(take_along_axis(a, ai, axis=1), 20)
117
118    def test_invalid(self):
119        """ Test invalid inputs """
120        a_base = np.array([[10, 30, 20], [60, 40, 50]])
121        indices = np.array([[0], [1]])
122        values = np.array([[2], [1]])
123
124        # sanity check
125        a = a_base.copy()
126        put_along_axis(a, indices, values, axis=0)
127        assert np.all(a == [[2, 2, 2], [1, 1, 1]])
128
129        # invalid indices
130        a = a_base.copy()
131        with assert_raises(ValueError) as exc:
132            put_along_axis(a, indices, values, axis=None)
133        assert "single dimension" in str(exc.exception)
134
135
136class TestApplyAlongAxis:
137    def test_simple(self):
138        a = np.ones((20, 10), 'd')
139        assert_array_equal(
140            apply_along_axis(len, 0, a), len(a) * np.ones(a.shape[1]))
141
142    def test_simple101(self):
143        a = np.ones((10, 101), 'd')
144        assert_array_equal(
145            apply_along_axis(len, 0, a), len(a) * np.ones(a.shape[1]))
146
147    def test_3d(self):
148        a = np.arange(27).reshape((3, 3, 3))
149        assert_array_equal(apply_along_axis(np.sum, 0, a),
150                           [[27, 30, 33], [36, 39, 42], [45, 48, 51]])
151
152    def test_preserve_subclass(self):
153        def double(row):
154            return row * 2
155
156        class MyNDArray(np.ndarray):
157            pass
158
159        m = np.array([[0, 1], [2, 3]]).view(MyNDArray)
160        expected = np.array([[0, 2], [4, 6]]).view(MyNDArray)
161
162        result = apply_along_axis(double, 0, m)
163        assert_(isinstance(result, MyNDArray))
164        assert_array_equal(result, expected)
165
166        result = apply_along_axis(double, 1, m)
167        assert_(isinstance(result, MyNDArray))
168        assert_array_equal(result, expected)
169
170    def test_subclass(self):
171        class MinimalSubclass(np.ndarray):
172            data = 1
173
174        def minimal_function(array):
175            return array.data
176
177        a = np.zeros((6, 3)).view(MinimalSubclass)
178
179        assert_array_equal(
180            apply_along_axis(minimal_function, 0, a), np.array([1, 1, 1])
181        )
182
183    def test_scalar_array(self, cls=np.ndarray):
184        a = np.ones((6, 3)).view(cls)
185        res = apply_along_axis(np.sum, 0, a)
186        assert_(isinstance(res, cls))
187        assert_array_equal(res, np.array([6, 6, 6]).view(cls))
188
189    def test_0d_array(self, cls=np.ndarray):
190        def sum_to_0d(x):
191            """ Sum x, returning a 0d array of the same class """
192            assert_equal(x.ndim, 1)
193            return np.squeeze(np.sum(x, keepdims=True))
194        a = np.ones((6, 3)).view(cls)
195        res = apply_along_axis(sum_to_0d, 0, a)
196        assert_(isinstance(res, cls))
197        assert_array_equal(res, np.array([6, 6, 6]).view(cls))
198
199        res = apply_along_axis(sum_to_0d, 1, a)
200        assert_(isinstance(res, cls))
201        assert_array_equal(res, np.array([3, 3, 3, 3, 3, 3]).view(cls))
202
203    def test_axis_insertion(self, cls=np.ndarray):
204        def f1to2(x):
205            """produces an asymmetric non-square matrix from x"""
206            assert_equal(x.ndim, 1)
207            return (x[::-1] * x[1:, None]).view(cls)
208
209        a2d = np.arange(6 * 3).reshape((6, 3))
210
211        # 2d insertion along first axis
212        actual = apply_along_axis(f1to2, 0, a2d)
213        expected = np.stack([
214            f1to2(a2d[:, i]) for i in range(a2d.shape[1])
215        ], axis=-1).view(cls)
216        assert_equal(type(actual), type(expected))
217        assert_equal(actual, expected)
218
219        # 2d insertion along last axis
220        actual = apply_along_axis(f1to2, 1, a2d)
221        expected = np.stack([
222            f1to2(a2d[i, :]) for i in range(a2d.shape[0])
223        ], axis=0).view(cls)
224        assert_equal(type(actual), type(expected))
225        assert_equal(actual, expected)
226
227        # 3d insertion along middle axis
228        a3d = np.arange(6 * 5 * 3).reshape((6, 5, 3))
229
230        actual = apply_along_axis(f1to2, 1, a3d)
231        expected = np.stack([
232            np.stack([
233                f1to2(a3d[i, :, j]) for i in range(a3d.shape[0])
234            ], axis=0)
235            for j in range(a3d.shape[2])
236        ], axis=-1).view(cls)
237        assert_equal(type(actual), type(expected))
238        assert_equal(actual, expected)
239
240    def test_subclass_preservation(self):
241        class MinimalSubclass(np.ndarray):
242            pass
243        self.test_scalar_array(MinimalSubclass)
244        self.test_0d_array(MinimalSubclass)
245        self.test_axis_insertion(MinimalSubclass)
246
247    def test_axis_insertion_ma(self):
248        def f1to2(x):
249            """produces an asymmetric non-square matrix from x"""
250            assert_equal(x.ndim, 1)
251            res = x[::-1] * x[1:, None]
252            return np.ma.masked_where(res % 5 == 0, res)
253        a = np.arange(6 * 3).reshape((6, 3))
254        res = apply_along_axis(f1to2, 0, a)
255        assert_(isinstance(res, np.ma.masked_array))
256        assert_equal(res.ndim, 3)
257        assert_array_equal(res[:, :, 0].mask, f1to2(a[:, 0]).mask)
258        assert_array_equal(res[:, :, 1].mask, f1to2(a[:, 1]).mask)
259        assert_array_equal(res[:, :, 2].mask, f1to2(a[:, 2]).mask)
260
261    def test_tuple_func1d(self):
262        def sample_1d(x):
263            return x[1], x[0]
264        res = np.apply_along_axis(sample_1d, 1, np.array([[1, 2], [3, 4]]))
265        assert_array_equal(res, np.array([[2, 1], [4, 3]]))
266
267    def test_empty(self):
268        # can't apply_along_axis when there's no chance to call the function
269        def never_call(x):
270            assert_(False)  # should never be reached
271
272        a = np.empty((0, 0))
273        assert_raises(ValueError, np.apply_along_axis, never_call, 0, a)
274        assert_raises(ValueError, np.apply_along_axis, never_call, 1, a)
275
276        # but it's sometimes ok with some non-zero dimensions
277        def empty_to_1(x):
278            assert_(len(x) == 0)
279            return 1
280
281        a = np.empty((10, 0))
282        actual = np.apply_along_axis(empty_to_1, 1, a)
283        assert_equal(actual, np.ones(10))
284        assert_raises(ValueError, np.apply_along_axis, empty_to_1, 0, a)
285
286    def test_with_iterable_object(self):
287        # from issue 5248
288        d = np.array([
289            [{1, 11}, {2, 22}, {3, 33}],
290            [{4, 44}, {5, 55}, {6, 66}]
291        ])
292        actual = np.apply_along_axis(lambda a: set.union(*a), 0, d)
293        expected = np.array([{1, 11, 4, 44}, {2, 22, 5, 55}, {3, 33, 6, 66}])
294
295        assert_equal(actual, expected)
296
297        # issue 8642 - assert_equal doesn't detect this!
298        for i in np.ndindex(actual.shape):
299            assert_equal(type(actual[i]), type(expected[i]))
300
301
302class TestApplyOverAxes:
303    def test_simple(self):
304        a = np.arange(24).reshape(2, 3, 4)
305        aoa_a = apply_over_axes(np.sum, a, [0, 2])
306        assert_array_equal(aoa_a, np.array([[[60], [92], [124]]]))
307
308
309class TestExpandDims:
310    def test_functionality(self):
311        s = (2, 3, 4, 5)
312        a = np.empty(s)
313        for axis in range(-5, 4):
314            b = expand_dims(a, axis)
315            assert_(b.shape[axis] == 1)
316            assert_(np.squeeze(b).shape == s)
317
318    def test_axis_tuple(self):
319        a = np.empty((3, 3, 3))
320        assert np.expand_dims(a, axis=(0, 1, 2)).shape == (1, 1, 1, 3, 3, 3)
321        assert np.expand_dims(a, axis=(0, -1, -2)).shape == (1, 3, 3, 3, 1, 1)
322        assert np.expand_dims(a, axis=(0, 3, 5)).shape == (1, 3, 3, 1, 3, 1)
323        assert np.expand_dims(a, axis=(0, -3, -5)).shape == (1, 1, 3, 1, 3, 3)
324
325    def test_axis_out_of_range(self):
326        s = (2, 3, 4, 5)
327        a = np.empty(s)
328        assert_raises(AxisError, expand_dims, a, -6)
329        assert_raises(AxisError, expand_dims, a, 5)
330
331        a = np.empty((3, 3, 3))
332        assert_raises(AxisError, expand_dims, a, (0, -6))
333        assert_raises(AxisError, expand_dims, a, (0, 5))
334
335    def test_repeated_axis(self):
336        a = np.empty((3, 3, 3))
337        assert_raises(ValueError, expand_dims, a, axis=(1, 1))
338
339    def test_subclasses(self):
340        a = np.arange(10).reshape((2, 5))
341        a = np.ma.array(a, mask=a % 3 == 0)
342
343        expanded = np.expand_dims(a, axis=1)
344        assert_(isinstance(expanded, np.ma.MaskedArray))
345        assert_equal(expanded.shape, (2, 1, 5))
346        assert_equal(expanded.mask.shape, (2, 1, 5))
347
348
349class TestArraySplit:
350    def test_integer_0_split(self):
351        a = np.arange(10)
352        assert_raises(ValueError, array_split, a, 0)
353
354    def test_integer_split(self):
355        a = np.arange(10)
356        res = array_split(a, 1)
357        desired = [np.arange(10)]
358        compare_results(res, desired)
359
360        res = array_split(a, 2)
361        desired = [np.arange(5), np.arange(5, 10)]
362        compare_results(res, desired)
363
364        res = array_split(a, 3)
365        desired = [np.arange(4), np.arange(4, 7), np.arange(7, 10)]
366        compare_results(res, desired)
367
368        res = array_split(a, 4)
369        desired = [np.arange(3), np.arange(3, 6), np.arange(6, 8),
370                   np.arange(8, 10)]
371        compare_results(res, desired)
372
373        res = array_split(a, 5)
374        desired = [np.arange(2), np.arange(2, 4), np.arange(4, 6),
375                   np.arange(6, 8), np.arange(8, 10)]
376        compare_results(res, desired)
377
378        res = array_split(a, 6)
379        desired = [np.arange(2), np.arange(2, 4), np.arange(4, 6),
380                   np.arange(6, 8), np.arange(8, 9), np.arange(9, 10)]
381        compare_results(res, desired)
382
383        res = array_split(a, 7)
384        desired = [np.arange(2), np.arange(2, 4), np.arange(4, 6),
385                   np.arange(6, 7), np.arange(7, 8), np.arange(8, 9),
386                   np.arange(9, 10)]
387        compare_results(res, desired)
388
389        res = array_split(a, 8)
390        desired = [np.arange(2), np.arange(2, 4), np.arange(4, 5),
391                   np.arange(5, 6), np.arange(6, 7), np.arange(7, 8),
392                   np.arange(8, 9), np.arange(9, 10)]
393        compare_results(res, desired)
394
395        res = array_split(a, 9)
396        desired = [np.arange(2), np.arange(2, 3), np.arange(3, 4),
397                   np.arange(4, 5), np.arange(5, 6), np.arange(6, 7),
398                   np.arange(7, 8), np.arange(8, 9), np.arange(9, 10)]
399        compare_results(res, desired)
400
401        res = array_split(a, 10)
402        desired = [np.arange(1), np.arange(1, 2), np.arange(2, 3),
403                   np.arange(3, 4), np.arange(4, 5), np.arange(5, 6),
404                   np.arange(6, 7), np.arange(7, 8), np.arange(8, 9),
405                   np.arange(9, 10)]
406        compare_results(res, desired)
407
408        res = array_split(a, 11)
409        desired = [np.arange(1), np.arange(1, 2), np.arange(2, 3),
410                   np.arange(3, 4), np.arange(4, 5), np.arange(5, 6),
411                   np.arange(6, 7), np.arange(7, 8), np.arange(8, 9),
412                   np.arange(9, 10), np.array([])]
413        compare_results(res, desired)
414
415    def test_integer_split_2D_rows(self):
416        a = np.array([np.arange(10), np.arange(10)])
417        res = array_split(a, 3, axis=0)
418        tgt = [np.array([np.arange(10)]), np.array([np.arange(10)]),
419                   np.zeros((0, 10))]
420        compare_results(res, tgt)
421        assert_(a.dtype.type is res[-1].dtype.type)
422
423        # Same thing for manual splits:
424        res = array_split(a, [0, 1], axis=0)
425        tgt = [np.zeros((0, 10)), np.array([np.arange(10)]),
426               np.array([np.arange(10)])]
427        compare_results(res, tgt)
428        assert_(a.dtype.type is res[-1].dtype.type)
429
430    def test_integer_split_2D_cols(self):
431        a = np.array([np.arange(10), np.arange(10)])
432        res = array_split(a, 3, axis=-1)
433        desired = [np.array([np.arange(4), np.arange(4)]),
434                   np.array([np.arange(4, 7), np.arange(4, 7)]),
435                   np.array([np.arange(7, 10), np.arange(7, 10)])]
436        compare_results(res, desired)
437
438    def test_integer_split_2D_default(self):
439        """ This will fail if we change default axis
440        """
441        a = np.array([np.arange(10), np.arange(10)])
442        res = array_split(a, 3)
443        tgt = [np.array([np.arange(10)]), np.array([np.arange(10)]),
444                   np.zeros((0, 10))]
445        compare_results(res, tgt)
446        assert_(a.dtype.type is res[-1].dtype.type)
447        # perhaps should check higher dimensions
448
449    @pytest.mark.skipif(not IS_64BIT, reason="Needs 64bit platform")
450    def test_integer_split_2D_rows_greater_max_int32(self):
451        a = np.broadcast_to([0], (1 << 32, 2))
452        res = array_split(a, 4)
453        chunk = np.broadcast_to([0], (1 << 30, 2))
454        tgt = [chunk] * 4
455        for i in range(len(tgt)):
456            assert_equal(res[i].shape, tgt[i].shape)
457
458    def test_index_split_simple(self):
459        a = np.arange(10)
460        indices = [1, 5, 7]
461        res = array_split(a, indices, axis=-1)
462        desired = [np.arange(0, 1), np.arange(1, 5), np.arange(5, 7),
463                   np.arange(7, 10)]
464        compare_results(res, desired)
465
466    def test_index_split_low_bound(self):
467        a = np.arange(10)
468        indices = [0, 5, 7]
469        res = array_split(a, indices, axis=-1)
470        desired = [np.array([]), np.arange(0, 5), np.arange(5, 7),
471                   np.arange(7, 10)]
472        compare_results(res, desired)
473
474    def test_index_split_high_bound(self):
475        a = np.arange(10)
476        indices = [0, 5, 7, 10, 12]
477        res = array_split(a, indices, axis=-1)
478        desired = [np.array([]), np.arange(0, 5), np.arange(5, 7),
479                   np.arange(7, 10), np.array([]), np.array([])]
480        compare_results(res, desired)
481
482
483class TestSplit:
484    # The split function is essentially the same as array_split,
485    # except that it test if splitting will result in an
486    # equal split.  Only test for this case.
487
488    def test_equal_split(self):
489        a = np.arange(10)
490        res = split(a, 2)
491        desired = [np.arange(5), np.arange(5, 10)]
492        compare_results(res, desired)
493
494    def test_unequal_split(self):
495        a = np.arange(10)
496        assert_raises(ValueError, split, a, 3)
497
498
499class TestColumnStack:
500    def test_non_iterable(self):
501        assert_raises(TypeError, column_stack, 1)
502
503    def test_1D_arrays(self):
504        # example from docstring
505        a = np.array((1, 2, 3))
506        b = np.array((2, 3, 4))
507        expected = np.array([[1, 2],
508                             [2, 3],
509                             [3, 4]])
510        actual = np.column_stack((a, b))
511        assert_equal(actual, expected)
512
513    def test_2D_arrays(self):
514        # same as hstack 2D docstring example
515        a = np.array([[1], [2], [3]])
516        b = np.array([[2], [3], [4]])
517        expected = np.array([[1, 2],
518                             [2, 3],
519                             [3, 4]])
520        actual = np.column_stack((a, b))
521        assert_equal(actual, expected)
522
523    def test_generator(self):
524        with pytest.raises(TypeError, match="arrays to stack must be"):
525            column_stack(np.arange(3) for _ in range(2))
526
527
528class TestDstack:
529    def test_non_iterable(self):
530        assert_raises(TypeError, dstack, 1)
531
532    def test_0D_array(self):
533        a = np.array(1)
534        b = np.array(2)
535        res = dstack([a, b])
536        desired = np.array([[[1, 2]]])
537        assert_array_equal(res, desired)
538
539    def test_1D_array(self):
540        a = np.array([1])
541        b = np.array([2])
542        res = dstack([a, b])
543        desired = np.array([[[1, 2]]])
544        assert_array_equal(res, desired)
545
546    def test_2D_array(self):
547        a = np.array([[1], [2]])
548        b = np.array([[1], [2]])
549        res = dstack([a, b])
550        desired = np.array([[[1, 1]], [[2, 2, ]]])
551        assert_array_equal(res, desired)
552
553    def test_2D_array2(self):
554        a = np.array([1, 2])
555        b = np.array([1, 2])
556        res = dstack([a, b])
557        desired = np.array([[[1, 1], [2, 2]]])
558        assert_array_equal(res, desired)
559
560    def test_generator(self):
561        with pytest.raises(TypeError, match="arrays to stack must be"):
562            dstack(np.arange(3) for _ in range(2))
563
564
565# array_split has more comprehensive test of splitting.
566# only do simple test on hsplit, vsplit, and dsplit
567class TestHsplit:
568    """Only testing for integer splits.
569
570    """
571    def test_non_iterable(self):
572        assert_raises(ValueError, hsplit, 1, 1)
573
574    def test_0D_array(self):
575        a = np.array(1)
576        try:
577            hsplit(a, 2)
578            assert_(0)
579        except ValueError:
580            pass
581
582    def test_1D_array(self):
583        a = np.array([1, 2, 3, 4])
584        res = hsplit(a, 2)
585        desired = [np.array([1, 2]), np.array([3, 4])]
586        compare_results(res, desired)
587
588    def test_2D_array(self):
589        a = np.array([[1, 2, 3, 4],
590                  [1, 2, 3, 4]])
591        res = hsplit(a, 2)
592        desired = [np.array([[1, 2], [1, 2]]), np.array([[3, 4], [3, 4]])]
593        compare_results(res, desired)
594
595
596class TestVsplit:
597    """Only testing for integer splits.
598
599    """
600    def test_non_iterable(self):
601        assert_raises(ValueError, vsplit, 1, 1)
602
603    def test_0D_array(self):
604        a = np.array(1)
605        assert_raises(ValueError, vsplit, a, 2)
606
607    def test_1D_array(self):
608        a = np.array([1, 2, 3, 4])
609        try:
610            vsplit(a, 2)
611            assert_(0)
612        except ValueError:
613            pass
614
615    def test_2D_array(self):
616        a = np.array([[1, 2, 3, 4],
617                  [1, 2, 3, 4]])
618        res = vsplit(a, 2)
619        desired = [np.array([[1, 2, 3, 4]]), np.array([[1, 2, 3, 4]])]
620        compare_results(res, desired)
621
622
623class TestDsplit:
624    # Only testing for integer splits.
625    def test_non_iterable(self):
626        assert_raises(ValueError, dsplit, 1, 1)
627
628    def test_0D_array(self):
629        a = np.array(1)
630        assert_raises(ValueError, dsplit, a, 2)
631
632    def test_1D_array(self):
633        a = np.array([1, 2, 3, 4])
634        assert_raises(ValueError, dsplit, a, 2)
635
636    def test_2D_array(self):
637        a = np.array([[1, 2, 3, 4],
638                  [1, 2, 3, 4]])
639        try:
640            dsplit(a, 2)
641            assert_(0)
642        except ValueError:
643            pass
644
645    def test_3D_array(self):
646        a = np.array([[[1, 2, 3, 4],
647                   [1, 2, 3, 4]],
648                  [[1, 2, 3, 4],
649                   [1, 2, 3, 4]]])
650        res = dsplit(a, 2)
651        desired = [np.array([[[1, 2], [1, 2]], [[1, 2], [1, 2]]]),
652                   np.array([[[3, 4], [3, 4]], [[3, 4], [3, 4]]])]
653        compare_results(res, desired)
654
655
656class TestSqueeze:
657    def test_basic(self):
658        from numpy.random import rand
659
660        a = rand(20, 10, 10, 1, 1)
661        b = rand(20, 1, 10, 1, 20)
662        c = rand(1, 1, 20, 10)
663        assert_array_equal(np.squeeze(a), np.reshape(a, (20, 10, 10)))
664        assert_array_equal(np.squeeze(b), np.reshape(b, (20, 10, 20)))
665        assert_array_equal(np.squeeze(c), np.reshape(c, (20, 10)))
666
667        # Squeezing to 0-dim should still give an ndarray
668        a = [[[1.5]]]
669        res = np.squeeze(a)
670        assert_equal(res, 1.5)
671        assert_equal(res.ndim, 0)
672        assert_equal(type(res), np.ndarray)
673
674
675class TestKron:
676    def test_basic(self):
677        # Using 0-dimensional ndarray
678        a = np.array(1)
679        b = np.array([[1, 2], [3, 4]])
680        k = np.array([[1, 2], [3, 4]])
681        assert_array_equal(np.kron(a, b), k)
682        a = np.array([[1, 2], [3, 4]])
683        b = np.array(1)
684        assert_array_equal(np.kron(a, b), k)
685
686        # Using 1-dimensional ndarray
687        a = np.array([3])
688        b = np.array([[1, 2], [3, 4]])
689        k = np.array([[3, 6], [9, 12]])
690        assert_array_equal(np.kron(a, b), k)
691        a = np.array([[1, 2], [3, 4]])
692        b = np.array([3])
693        assert_array_equal(np.kron(a, b), k)
694
695        # Using 3-dimensional ndarray
696        a = np.array([[[1]], [[2]]])
697        b = np.array([[1, 2], [3, 4]])
698        k = np.array([[[1, 2], [3, 4]], [[2, 4], [6, 8]]])
699        assert_array_equal(np.kron(a, b), k)
700        a = np.array([[1, 2], [3, 4]])
701        b = np.array([[[1]], [[2]]])
702        k = np.array([[[1, 2], [3, 4]], [[2, 4], [6, 8]]])
703        assert_array_equal(np.kron(a, b), k)
704
705    def test_return_type(self):
706        class myarray(np.ndarray):
707            __array_priority__ = 1.0
708
709        a = np.ones([2, 2])
710        ma = myarray(a.shape, a.dtype, a.data)
711        assert_equal(type(kron(a, a)), np.ndarray)
712        assert_equal(type(kron(ma, ma)), myarray)
713        assert_equal(type(kron(a, ma)), myarray)
714        assert_equal(type(kron(ma, a)), myarray)
715
716    @pytest.mark.parametrize(
717        "array_class", [np.asarray, np.asmatrix]
718    )
719    def test_kron_smoke(self, array_class):
720        a = array_class(np.ones([3, 3]))
721        b = array_class(np.ones([3, 3]))
722        k = array_class(np.ones([9, 9]))
723
724        assert_array_equal(np.kron(a, b), k)
725
726    def test_kron_ma(self):
727        x = np.ma.array([[1, 2], [3, 4]], mask=[[0, 1], [1, 0]])
728        k = np.ma.array(np.diag([1, 4, 4, 16]),
729                mask=~np.array(np.identity(4), dtype=bool))
730
731        assert_array_equal(k, np.kron(x, x))
732
733    @pytest.mark.parametrize(
734        "shape_a,shape_b", [
735            ((1, 1), (1, 1)),
736            ((1, 2, 3), (4, 5, 6)),
737            ((2, 2), (2, 2, 2)),
738            ((1, 0), (1, 1)),
739            ((2, 0, 2), (2, 2)),
740            ((2, 0, 0, 2), (2, 0, 2)),
741        ])
742    def test_kron_shape(self, shape_a, shape_b):
743        a = np.ones(shape_a)
744        b = np.ones(shape_b)
745        normalised_shape_a = (1,) * max(0, len(shape_b) - len(shape_a)) + shape_a
746        normalised_shape_b = (1,) * max(0, len(shape_a) - len(shape_b)) + shape_b
747        expected_shape = np.multiply(normalised_shape_a, normalised_shape_b)
748
749        k = np.kron(a, b)
750        assert np.array_equal(
751                k.shape, expected_shape), "Unexpected shape from kron"
752
753
754class TestTile:
755    def test_basic(self):
756        a = np.array([0, 1, 2])
757        b = [[1, 2], [3, 4]]
758        assert_equal(tile(a, 2), [0, 1, 2, 0, 1, 2])
759        assert_equal(tile(a, (2, 2)), [[0, 1, 2, 0, 1, 2], [0, 1, 2, 0, 1, 2]])
760        assert_equal(tile(a, (1, 2)), [[0, 1, 2, 0, 1, 2]])
761        assert_equal(tile(b, 2), [[1, 2, 1, 2], [3, 4, 3, 4]])
762        assert_equal(tile(b, (2, 1)), [[1, 2], [3, 4], [1, 2], [3, 4]])
763        assert_equal(tile(b, (2, 2)), [[1, 2, 1, 2], [3, 4, 3, 4],
764                                       [1, 2, 1, 2], [3, 4, 3, 4]])
765
766    def test_tile_one_repetition_on_array_gh4679(self):
767        a = np.arange(5)
768        b = tile(a, 1)
769        b += 2
770        assert_equal(a, np.arange(5))
771
772    def test_empty(self):
773        a = np.array([[[]]])
774        b = np.array([[], []])
775        c = tile(b, 2).shape
776        d = tile(a, (3, 2, 5)).shape
777        assert_equal(c, (2, 0))
778        assert_equal(d, (3, 2, 0))
779
780    def test_kroncompare(self):
781        from numpy.random import randint
782
783        reps = [(2,), (1, 2), (2, 1), (2, 2), (2, 3, 2), (3, 2)]
784        shape = [(3,), (2, 3), (3, 4, 3), (3, 2, 3), (4, 3, 2, 4), (2, 2)]
785        for s in shape:
786            b = randint(0, 10, size=s)
787            for r in reps:
788                a = np.ones(r, b.dtype)
789                large = tile(b, r)
790                klarge = kron(a, b)
791                assert_equal(large, klarge)
792
793
794class TestMayShareMemory:
795    def test_basic(self):
796        d = np.ones((50, 60))
797        d2 = np.ones((30, 60, 6))
798        assert_(np.may_share_memory(d, d))
799        assert_(np.may_share_memory(d, d[::-1]))
800        assert_(np.may_share_memory(d, d[::2]))
801        assert_(np.may_share_memory(d, d[1:, ::-1]))
802
803        assert_(not np.may_share_memory(d[::-1], d2))
804        assert_(not np.may_share_memory(d[::2], d2))
805        assert_(not np.may_share_memory(d[1:, ::-1], d2))
806        assert_(np.may_share_memory(d2[1:, ::-1], d2))
807
808
809# Utility
810def compare_results(res, desired):
811    """Compare lists of arrays."""
812    for x, y in zip(res, desired, strict=False):
813        assert_array_equal(x, y)
814 
codekingpro/portable-devtools · Team Ai