Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_overrides.py801 linesDownload Raw Back to tests
1import inspect
2import os
3import pickle
4import sys
5import tempfile
6from io import StringIO
7from unittest import mock
8
9import pytest
10
11import numpy as np
12from numpy._core.overrides import (
13    _get_implementing_args,
14    array_function_dispatch,
15    verify_matching_signatures,
16)
17from numpy.testing import assert_, assert_equal, assert_raises, assert_raises_regex
18from numpy.testing.overrides import get_overridable_numpy_array_functions
19
20
21def _return_not_implemented(self, *args, **kwargs):
22    return NotImplemented
23
24
25# need to define this at the top level to test pickling
26@array_function_dispatch(lambda array: (array,))
27def dispatched_one_arg(array):
28    """Docstring."""
29    return 'original'
30
31
32@array_function_dispatch(lambda array1, array2: (array1, array2))
33def dispatched_two_arg(array1, array2):
34    """Docstring."""
35    return 'original'
36
37
38class TestGetImplementingArgs:
39
40    def test_ndarray(self):
41        array = np.array(1)
42
43        args = _get_implementing_args([array])
44        assert_equal(list(args), [array])
45
46        args = _get_implementing_args([array, array])
47        assert_equal(list(args), [array])
48
49        args = _get_implementing_args([array, 1])
50        assert_equal(list(args), [array])
51
52        args = _get_implementing_args([1, array])
53        assert_equal(list(args), [array])
54
55    def test_ndarray_subclasses(self):
56
57        class OverrideSub(np.ndarray):
58            __array_function__ = _return_not_implemented
59
60        class NoOverrideSub(np.ndarray):
61            pass
62
63        array = np.array(1).view(np.ndarray)
64        override_sub = np.array(1).view(OverrideSub)
65        no_override_sub = np.array(1).view(NoOverrideSub)
66
67        args = _get_implementing_args([array, override_sub])
68        assert_equal(list(args), [override_sub, array])
69
70        args = _get_implementing_args([array, no_override_sub])
71        assert_equal(list(args), [no_override_sub, array])
72
73        args = _get_implementing_args(
74            [override_sub, no_override_sub])
75        assert_equal(list(args), [override_sub, no_override_sub])
76
77    def test_ndarray_and_duck_array(self):
78
79        class Other:
80            __array_function__ = _return_not_implemented
81
82        array = np.array(1)
83        other = Other()
84
85        args = _get_implementing_args([other, array])
86        assert_equal(list(args), [other, array])
87
88        args = _get_implementing_args([array, other])
89        assert_equal(list(args), [array, other])
90
91    def test_ndarray_subclass_and_duck_array(self):
92
93        class OverrideSub(np.ndarray):
94            __array_function__ = _return_not_implemented
95
96        class Other:
97            __array_function__ = _return_not_implemented
98
99        array = np.array(1)
100        subarray = np.array(1).view(OverrideSub)
101        other = Other()
102
103        assert_equal(_get_implementing_args([array, subarray, other]),
104                     [subarray, array, other])
105        assert_equal(_get_implementing_args([array, other, subarray]),
106                     [subarray, array, other])
107
108    def test_many_duck_arrays(self):
109
110        class A:
111            __array_function__ = _return_not_implemented
112
113        class B(A):
114            __array_function__ = _return_not_implemented
115
116        class C(A):
117            __array_function__ = _return_not_implemented
118
119        class D:
120            __array_function__ = _return_not_implemented
121
122        a = A()
123        b = B()
124        c = C()
125        d = D()
126
127        assert_equal(_get_implementing_args([1]), [])
128        assert_equal(_get_implementing_args([a]), [a])
129        assert_equal(_get_implementing_args([a, 1]), [a])
130        assert_equal(_get_implementing_args([a, a, a]), [a])
131        assert_equal(_get_implementing_args([a, d, a]), [a, d])
132        assert_equal(_get_implementing_args([a, b]), [b, a])
133        assert_equal(_get_implementing_args([b, a]), [b, a])
134        assert_equal(_get_implementing_args([a, b, c]), [b, c, a])
135        assert_equal(_get_implementing_args([a, c, b]), [c, b, a])
136
137    def test_too_many_duck_arrays(self):
138        namespace = {'__array_function__': _return_not_implemented}
139        types = [type('A' + str(i), (object,), namespace) for i in range(65)]
140        relevant_args = [t() for t in types]
141
142        actual = _get_implementing_args(relevant_args[:64])
143        assert_equal(actual, relevant_args[:64])
144
145        with assert_raises_regex(TypeError, 'distinct argument types'):
146            _get_implementing_args(relevant_args)
147
148
149class TestNDArrayArrayFunction:
150
151    def test_method(self):
152
153        class Other:
154            __array_function__ = _return_not_implemented
155
156        class NoOverrideSub(np.ndarray):
157            pass
158
159        class OverrideSub(np.ndarray):
160            __array_function__ = _return_not_implemented
161
162        array = np.array([1])
163        other = Other()
164        no_override_sub = array.view(NoOverrideSub)
165        override_sub = array.view(OverrideSub)
166
167        result = array.__array_function__(func=dispatched_two_arg,
168                                          types=(np.ndarray,),
169                                          args=(array, 1.), kwargs={})
170        assert_equal(result, 'original')
171
172        result = array.__array_function__(func=dispatched_two_arg,
173                                          types=(np.ndarray, Other),
174                                          args=(array, other), kwargs={})
175        assert_(result is NotImplemented)
176
177        result = array.__array_function__(func=dispatched_two_arg,
178                                          types=(np.ndarray, NoOverrideSub),
179                                          args=(array, no_override_sub),
180                                          kwargs={})
181        assert_equal(result, 'original')
182
183        result = array.__array_function__(func=dispatched_two_arg,
184                                          types=(np.ndarray, OverrideSub),
185                                          args=(array, override_sub),
186                                          kwargs={})
187        assert_equal(result, 'original')
188
189        with assert_raises_regex(TypeError, 'no implementation found'):
190            np.concatenate((array, other))
191
192        expected = np.concatenate((array, array))
193        result = np.concatenate((array, no_override_sub))
194        assert_equal(result, expected.view(NoOverrideSub))
195        result = np.concatenate((array, override_sub))
196        assert_equal(result, expected.view(OverrideSub))
197
198    def test_no_wrapper(self):
199        # Regular numpy functions have wrappers, but do not presume
200        # all functions do (array creation ones do not): check that
201        # we just call the function in that case.
202        array = np.array(1)
203        func = lambda x: x * 2
204        result = array.__array_function__(func=func, types=(np.ndarray,),
205                                          args=(array,), kwargs={})
206        assert_equal(result, array * 2)
207
208    def test_wrong_arguments(self):
209        # Check our implementation guards against wrong arguments.
210        a = np.array([1, 2])
211        with pytest.raises(TypeError, match="args must be a tuple"):
212            a.__array_function__(np.reshape, (np.ndarray,), a, (2, 1))
213        with pytest.raises(TypeError, match="kwargs must be a dict"):
214            a.__array_function__(np.reshape, (np.ndarray,), (a,), (2, 1))
215
216
217class TestArrayFunctionDispatch:
218
219    def test_pickle(self):
220        for proto in range(2, pickle.HIGHEST_PROTOCOL + 1):
221            roundtripped = pickle.loads(
222                    pickle.dumps(dispatched_one_arg, protocol=proto))
223            assert_(roundtripped is dispatched_one_arg)
224
225    def test_name_and_docstring(self):
226        assert_equal(dispatched_one_arg.__name__, 'dispatched_one_arg')
227        if sys.flags.optimize < 2:
228            assert_equal(dispatched_one_arg.__doc__, 'Docstring.')
229
230    def test_interface(self):
231
232        class MyArray:
233            def __array_function__(self, func, types, args, kwargs):
234                return (self, func, types, args, kwargs)
235
236        original = MyArray()
237        (obj, func, types, args, kwargs) = dispatched_one_arg(original)
238        assert_(obj is original)
239        assert_(func is dispatched_one_arg)
240        assert_equal(set(types), {MyArray})
241        # assert_equal uses the overloaded np.iscomplexobj() internally
242        assert_(args == (original,))
243        assert_equal(kwargs, {})
244
245    def test_not_implemented(self):
246
247        class MyArray:
248            def __array_function__(self, func, types, args, kwargs):
249                return NotImplemented
250
251        array = MyArray()
252        with assert_raises_regex(TypeError, 'no implementation found'):
253            dispatched_one_arg(array)
254
255    def test_where_dispatch(self):
256
257        class DuckArray:
258            def __array_function__(self, ufunc, method, *inputs, **kwargs):
259                return "overridden"
260
261        array = np.array(1)
262        duck_array = DuckArray()
263
264        result = np.std(array, where=duck_array)
265
266        assert_equal(result, "overridden")
267
268
269class TestVerifyMatchingSignatures:
270
271    def test_verify_matching_signatures(self):
272
273        verify_matching_signatures(lambda x: 0, lambda x: 0)
274        verify_matching_signatures(lambda x=None: 0, lambda x=None: 0)
275        verify_matching_signatures(lambda x=1: 0, lambda x=None: 0)
276
277        with assert_raises(RuntimeError):
278            verify_matching_signatures(lambda a: 0, lambda b: 0)
279        with assert_raises(RuntimeError):
280            verify_matching_signatures(lambda x: 0, lambda x=None: 0)
281        with assert_raises(RuntimeError):
282            verify_matching_signatures(lambda x=None: 0, lambda y=None: 0)
283        with assert_raises(RuntimeError):
284            verify_matching_signatures(lambda x=1: 0, lambda y=1: 0)
285
286    def test_array_function_dispatch(self):
287
288        with assert_raises(RuntimeError):
289            @array_function_dispatch(lambda x: (x,))
290            def f(y):
291                pass
292
293        # should not raise
294        @array_function_dispatch(lambda x: (x,), verify=False)
295        def f(y):
296            pass
297
298
299def _new_duck_type_and_implements():
300    """Create a duck array type and implements functions."""
301    HANDLED_FUNCTIONS = {}
302
303    class MyArray:
304        def __array_function__(self, func, types, args, kwargs):
305            if func not in HANDLED_FUNCTIONS:
306                return NotImplemented
307            if not all(issubclass(t, MyArray) for t in types):
308                return NotImplemented
309            return HANDLED_FUNCTIONS[func](*args, **kwargs)
310
311    def implements(numpy_function):
312        """Register an __array_function__ implementations."""
313        def decorator(func):
314            HANDLED_FUNCTIONS[numpy_function] = func
315            return func
316        return decorator
317
318    return (MyArray, implements)
319
320
321class TestArrayFunctionImplementation:
322
323    def test_one_arg(self):
324        MyArray, implements = _new_duck_type_and_implements()
325
326        @implements(dispatched_one_arg)
327        def _(array):
328            return 'myarray'
329
330        assert_equal(dispatched_one_arg(1), 'original')
331        assert_equal(dispatched_one_arg(MyArray()), 'myarray')
332
333    def test_optional_args(self):
334        MyArray, implements = _new_duck_type_and_implements()
335
336        @array_function_dispatch(lambda array, option=None: (array,))
337        def func_with_option(array, option='default'):
338            return option
339
340        @implements(func_with_option)
341        def my_array_func_with_option(array, new_option='myarray'):
342            return new_option
343
344        # we don't need to implement every option on __array_function__
345        # implementations
346        assert_equal(func_with_option(1), 'default')
347        assert_equal(func_with_option(1, option='extra'), 'extra')
348        assert_equal(func_with_option(MyArray()), 'myarray')
349        with assert_raises(TypeError):
350            func_with_option(MyArray(), option='extra')
351
352        # but new options on implementations can't be used
353        result = my_array_func_with_option(MyArray(), new_option='yes')
354        assert_equal(result, 'yes')
355        with assert_raises(TypeError):
356            func_with_option(MyArray(), new_option='no')
357
358    def test_not_implemented(self):
359        MyArray, implements = _new_duck_type_and_implements()
360
361        @array_function_dispatch(lambda array: (array,), module='my')
362        def func(array):
363            return array
364
365        array = np.array(1)
366        assert_(func(array) is array)
367        assert_equal(func.__module__, 'my')
368
369        with assert_raises_regex(
370                TypeError, "no implementation found for 'my.func'"):
371            func(MyArray())
372
373    @pytest.mark.parametrize("name", ["concatenate", "mean", "asarray"])
374    def test_signature_error_message_simple(self, name):
375        func = getattr(np, name)
376        try:
377            # all of these functions need an argument:
378            func()
379        except TypeError as e:
380            exc = e
381
382        assert exc.args[0].startswith(f"{name}()")
383
384    def test_signature_error_message(self):
385        # The lambda function will be named "<lambda>", but the TypeError
386        # should show the name as "func"
387        def _dispatcher():
388            return ()
389
390        @array_function_dispatch(_dispatcher)
391        def func():
392            pass
393
394        try:
395            func._implementation(bad_arg=3)
396        except TypeError as e:
397            expected_exception = e
398
399        try:
400            func(bad_arg=3)
401            raise AssertionError("must fail")
402        except TypeError as exc:
403            if exc.args[0].startswith("_dispatcher"):
404                # We replace the qualname currently, but it used `__name__`
405                # (relevant functions have the same name and qualname anyway)
406                pytest.skip("Python version is not using __qualname__ for "
407                            "TypeError formatting.")
408
409            assert exc.args == expected_exception.args
410
411    @pytest.mark.parametrize("value", [234, "this func is not replaced"])
412    def test_dispatcher_error(self, value):
413        # If the dispatcher raises an error, we must not attempt to mutate it
414        error = TypeError(value)
415
416        def dispatcher():
417            raise error
418
419        @array_function_dispatch(dispatcher)
420        def func():
421            return 3
422
423        try:
424            func()
425            raise AssertionError("must fail")
426        except TypeError as exc:
427            assert exc is error  # unmodified exception
428
429    def test_properties(self):
430        # Check that str and repr are sensible
431        func = dispatched_two_arg
432        assert str(func) == str(func._implementation)
433        repr_no_id = repr(func).split("at ")[0]
434        repr_no_id_impl = repr(func._implementation).split("at ")[0]
435        assert repr_no_id == repr_no_id_impl
436
437    @pytest.mark.parametrize("func", [
438            lambda x, y: 0,  # no like argument
439            lambda like=None: 0,  # not keyword only
440            lambda *, like=None, a=3: 0,  # not last (not that it matters)
441        ])
442    def test_bad_like_sig(self, func):
443        # We sanity check the signature, and these should fail.
444        with pytest.raises(RuntimeError):
445            array_function_dispatch()(func)
446
447    def test_bad_like_passing(self):
448        # Cover internal sanity check for passing like as first positional arg
449        def func(*, like=None):
450            pass
451
452        func_with_like = array_function_dispatch()(func)
453        with pytest.raises(TypeError):
454            func_with_like()
455        with pytest.raises(TypeError):
456            func_with_like(like=234)
457
458    def test_too_many_args(self):
459        # Mainly a unit-test to increase coverage
460        objs = []
461        for i in range(80):
462            class MyArr:
463                def __array_function__(self, *args, **kwargs):
464                    return NotImplemented
465
466            objs.append(MyArr())
467
468        def _dispatch(*args):
469            return args
470
471        @array_function_dispatch(_dispatch)
472        def func(*args):
473            pass
474
475        with pytest.raises(TypeError, match="maximum number"):
476            func(*objs)
477
478
479class TestNDArrayMethods:
480
481    def test_repr(self):
482        # gh-12162: should still be defined even if __array_function__ doesn't
483        # implement np.array_repr()
484
485        class MyArray(np.ndarray):
486            def __array_function__(*args, **kwargs):
487                return NotImplemented
488
489        array = np.array(1).view(MyArray)
490        assert_equal(repr(array), 'MyArray(1)')
491        assert_equal(str(array), '1')
492
493
494class TestNumPyFunctions:
495
496    def test_set_module(self):
497        assert_equal(np.sum.__module__, 'numpy')
498        assert_equal(np.char.equal.__module__, 'numpy.char')
499        assert_equal(np.fft.fft.__module__, 'numpy.fft')
500        assert_equal(np.linalg.solve.__module__, 'numpy.linalg')
501
502    def test_inspect_sum(self):
503        signature = inspect.signature(np.sum)
504        assert_('axis' in signature.parameters)
505
506    def test_override_sum(self):
507        MyArray, implements = _new_duck_type_and_implements()
508
509        @implements(np.sum)
510        def _(array):
511            return 'yes'
512
513        assert_equal(np.sum(MyArray()), 'yes')
514
515    def test_sum_on_mock_array(self):
516
517        # We need a proxy for mocks because __array_function__ is only looked
518        # up in the class dict
519        class ArrayProxy:
520            def __init__(self, value):
521                self.value = value
522
523            def __array_function__(self, *args, **kwargs):
524                return self.value.__array_function__(*args, **kwargs)
525
526            def __array__(self, *args, **kwargs):
527                return self.value.__array__(*args, **kwargs)
528
529        proxy = ArrayProxy(mock.Mock(spec=ArrayProxy))
530        proxy.value.__array_function__.return_value = 1
531        result = np.sum(proxy)
532        assert_equal(result, 1)
533        proxy.value.__array_function__.assert_called_once_with(
534            np.sum, (ArrayProxy,), (proxy,), {})
535        proxy.value.__array__.assert_not_called()
536
537    def test_sum_forwarding_implementation(self):
538
539        class MyArray(np.ndarray):
540
541            def sum(self, axis, out):
542                return 'summed'
543
544            def __array_function__(self, func, types, args, kwargs):
545                return super().__array_function__(func, types, args, kwargs)
546
547        # note: the internal implementation of np.sum() calls the .sum() method
548        array = np.array(1).view(MyArray)
549        assert_equal(np.sum(array), 'summed')
550
551
552class TestArrayLike:
553    def _create_MyArray(self):
554        class MyArray:
555            def __init__(self, function=None):
556                self.function = function
557
558            def __array_function__(self, func, types, args, kwargs):
559                assert func is getattr(np, func.__name__)
560                try:
561                    my_func = getattr(self, func.__name__)
562                except AttributeError:
563                    return NotImplemented
564                return my_func(*args, **kwargs)
565
566        return MyArray
567
568    def _create_MyNoArrayFunctionArray(self):
569        class MyNoArrayFunctionArray:
570            def __init__(self, function=None):
571                self.function = function
572
573        return MyNoArrayFunctionArray
574
575    def _create_MySubclass(self):
576        class MySubclass(np.ndarray):
577            def __array_function__(self, func, types, args, kwargs):
578                result = super().__array_function__(func, types, args, kwargs)
579                return result.view(self.__class__)
580
581        return MySubclass
582
583    def add_method(self, name, arr_class, enable_value_error=False):
584        def _definition(*args, **kwargs):
585            # Check that `like=` isn't propagated downstream
586            assert 'like' not in kwargs
587
588            if enable_value_error and 'value_error' in kwargs:
589                raise ValueError
590
591            return arr_class(getattr(arr_class, name))
592        setattr(arr_class, name, _definition)
593
594    def func_args(*args, **kwargs):
595        return args, kwargs
596
597    def test_array_like_not_implemented(self):
598        MyArray = self._create_MyArray()
599        self.add_method('array', MyArray)
600
601        ref = MyArray.array()
602
603        with assert_raises_regex(TypeError, 'no implementation found'):
604            array_like = np.asarray(1, like=ref)
605
606    _array_tests = [
607        ('array', *func_args((1,))),
608        ('asarray', *func_args((1,))),
609        ('asanyarray', *func_args((1,))),
610        ('ascontiguousarray', *func_args((2, 3))),
611        ('asfortranarray', *func_args((2, 3))),
612        ('require', *func_args((np.arange(6).reshape(2, 3),),
613                               requirements=['A', 'F'])),
614        ('empty', *func_args((1,))),
615        ('full', *func_args((1,), 2)),
616        ('ones', *func_args((1,))),
617        ('zeros', *func_args((1,))),
618        ('arange', *func_args(3)),
619        ('frombuffer', *func_args(b'\x00' * 8, dtype=int)),
620        ('fromiter', *func_args(range(3), dtype=int)),
621        ('fromstring', *func_args('1,2', dtype=int, sep=',')),
622        ('loadtxt', *func_args(lambda: StringIO('0 1\n2 3'))),
623        ('genfromtxt', *func_args(lambda: StringIO('1,2.1'),
624                                  dtype=[('int', 'i8'), ('float', 'f8')],
625                                  delimiter=',')),
626    ]
627
628    def test_nep35_functions_as_array_functions(self,):
629        all_array_functions = get_overridable_numpy_array_functions()
630        like_array_functions_subset = {
631            getattr(np, func_name) for func_name, *_ in self.__class__._array_tests
632        }
633        assert like_array_functions_subset.issubset(all_array_functions)
634
635        nep35_python_functions = {
636            np.eye, np.fromfunction, np.full, np.genfromtxt,
637            np.identity, np.loadtxt, np.ones, np.require, np.tri,
638        }
639        assert nep35_python_functions.issubset(all_array_functions)
640
641        nep35_C_functions = {
642            np.arange, np.array, np.asanyarray, np.asarray,
643            np.ascontiguousarray, np.asfortranarray, np.empty,
644            np.frombuffer, np.fromfile, np.fromiter, np.fromstring,
645            np.zeros,
646        }
647        assert nep35_C_functions.issubset(all_array_functions)
648
649    @pytest.mark.parametrize('function, args, kwargs', _array_tests)
650    @pytest.mark.parametrize('numpy_ref', [True, False])
651    def test_array_like(self, function, args, kwargs, numpy_ref):
652        MyArray = self._create_MyArray()
653        self.add_method('array', MyArray)
654        self.add_method(function, MyArray)
655        np_func = getattr(np, function)
656        my_func = getattr(MyArray, function)
657
658        if numpy_ref is True:
659            ref = np.array(1)
660        else:
661            ref = MyArray.array()
662
663        like_args = tuple(a() if callable(a) else a for a in args)
664        array_like = np_func(*like_args, **kwargs, like=ref)
665
666        if numpy_ref is True:
667            assert type(array_like) is np.ndarray
668
669            np_args = tuple(a() if callable(a) else a for a in args)
670            np_arr = np_func(*np_args, **kwargs)
671
672            # Special-case np.empty to ensure values match
673            if function == "empty":
674                np_arr.fill(1)
675                array_like.fill(1)
676
677            assert_equal(array_like, np_arr)
678        else:
679            assert type(array_like) is MyArray
680            assert array_like.function is my_func
681
682    @pytest.mark.parametrize('function, args, kwargs', _array_tests)
683    @pytest.mark.parametrize('ref', [1, [1], "MyNoArrayFunctionArray"])
684    def test_no_array_function_like(self, function, args, kwargs, ref):
685        MyNoArrayFunctionArray = self._create_MyNoArrayFunctionArray()
686        self.add_method('array', MyNoArrayFunctionArray)
687        self.add_method(function, MyNoArrayFunctionArray)
688        np_func = getattr(np, function)
689
690        # Instantiate ref if it's the MyNoArrayFunctionArray class
691        if ref == "MyNoArrayFunctionArray":
692            ref = MyNoArrayFunctionArray.array()
693
694        like_args = tuple(a() if callable(a) else a for a in args)
695
696        with assert_raises_regex(TypeError,
697                'The `like` argument must be an array-like that implements'):
698            np_func(*like_args, **kwargs, like=ref)
699
700    @pytest.mark.parametrize('function, args, kwargs', _array_tests)
701    def test_subclass(self, function, args, kwargs):
702        MySubclass = self._create_MySubclass()
703        ref = np.array(1).view(MySubclass)
704        np_func = getattr(np, function)
705        like_args = tuple(a() if callable(a) else a for a in args)
706        array_like = np_func(*like_args, **kwargs, like=ref)
707        assert type(array_like) is MySubclass
708        if np_func is np.empty:
709            return
710        np_args = tuple(a() if callable(a) else a for a in args)
711        np_arr = np_func(*np_args, **kwargs)
712        assert_equal(array_like.view(np.ndarray), np_arr)
713
714    @pytest.mark.parametrize('numpy_ref', [True, False])
715    def test_array_like_fromfile(self, numpy_ref):
716        MyArray = self._create_MyArray()
717        self.add_method('array', MyArray)
718        self.add_method("fromfile", MyArray)
719
720        if numpy_ref is True:
721            ref = np.array(1)
722        else:
723            ref = MyArray.array()
724
725        data = np.random.random(5)
726
727        with tempfile.TemporaryDirectory() as tmpdir:
728            fname = os.path.join(tmpdir, "testfile")
729            data.tofile(fname)
730
731            array_like = np.fromfile(fname, like=ref)
732            if numpy_ref is True:
733                assert type(array_like) is np.ndarray
734                np_res = np.fromfile(fname, like=ref)
735                assert_equal(np_res, data)
736                assert_equal(array_like, np_res)
737            else:
738                assert type(array_like) is MyArray
739                assert array_like.function is MyArray.fromfile
740
741    def test_exception_handling(self):
742        MyArray = self._create_MyArray()
743        self.add_method('array', MyArray, enable_value_error=True)
744
745        ref = MyArray.array()
746
747        with assert_raises(TypeError):
748            # Raises the error about `value_error` being invalid first
749            np.array(1, value_error=True, like=ref)
750
751    @pytest.mark.parametrize('function, args, kwargs', _array_tests)
752    def test_like_as_none(self, function, args, kwargs):
753        MyArray = self._create_MyArray()
754        self.add_method('array', MyArray)
755        self.add_method(function, MyArray)
756        np_func = getattr(np, function)
757
758        like_args = tuple(a() if callable(a) else a for a in args)
759        # required for loadtxt and genfromtxt to init w/o error.
760        like_args_exp = tuple(a() if callable(a) else a for a in args)
761
762        array_like = np_func(*like_args, **kwargs, like=None)
763        expected = np_func(*like_args_exp, **kwargs)
764        # Special-case np.empty to ensure values match
765        if function == "empty":
766            array_like.fill(1)
767            expected.fill(1)
768        assert_equal(array_like, expected)
769
770
771def test_function_like():
772    # We provide a `__get__` implementation, make sure it works
773    assert type(np.mean) is np._core._multiarray_umath._ArrayFunctionDispatcher
774
775    class MyClass:
776        def __array__(self, dtype=None, copy=None):
777            # valid argument to mean:
778            return np.arange(3)
779
780        func1 = staticmethod(np.mean)
781        func2 = np.mean
782        func3 = classmethod(np.mean)
783
784    m = MyClass()
785    assert m.func1([10]) == 10
786    assert m.func2() == 1  # mean of the arange
787    with pytest.raises(TypeError, match="unsupported operand type"):
788        # Tries to operate on the class
789        m.func3()
790
791    # Manual binding also works (the above may shortcut):
792    bound = np.mean.__get__(m, MyClass)
793    assert bound() == 1
794
795    bound = np.mean.__get__(None, MyClass)  # unbound actually
796    assert bound([10]) == 10
797
798    bound = np.mean.__get__(MyClass)  # classmethod
799    with pytest.raises(TypeError, match="unsupported operand type"):
800        bound()
801 
codekingpro/portable-devtools · Team Ai