codekingpro/portable-devtools
114k
1import ctypes as ct
2import inspect
3import itertools
4import pickle
5import sys
6import warnings
7
8import pytest
9from pytest import param
10
11import numpy as np
12import numpy._core._operand_flag_tests as opflag_tests
13import numpy._core._rational_tests as _rational_tests
14import numpy._core._umath_tests as umt
15import numpy._core.umath as ncu
16import numpy.linalg._umath_linalg as uml
17from numpy.exceptions import AxisError
18from numpy.testing import (
19 HAS_REFCOUNT,
20 IS_PYPY,
21 IS_WASM,
22 assert_,
23 assert_allclose,
24 assert_almost_equal,
25 assert_array_almost_equal,
26 assert_array_equal,
27 assert_equal,
28 assert_no_warnings,
29 assert_raises,
30)
31from numpy.testing._private.utils import requires_memory
32
33UNARY_UFUNCS = [obj for obj in np._core.umath.__dict__.values()
34 if isinstance(obj, np.ufunc)]
35UNARY_OBJECT_UFUNCS = [uf for uf in UNARY_UFUNCS if "O->O" in uf.types]
36
37# Remove functions that do not support `floats`
38UNARY_OBJECT_UFUNCS.remove(np.bitwise_count)
39
40
41class TestUfuncKwargs:
42 def test_kwarg_exact(self):
43 assert_raises(TypeError, np.add, 1, 2, castingx='safe')
44 assert_raises(TypeError, np.add, 1, 2, dtypex=int)
45 assert_raises(TypeError, np.add, 1, 2, extobjx=[4096])
46 assert_raises(TypeError, np.add, 1, 2, outx=None)
47 assert_raises(TypeError, np.add, 1, 2, sigx='ii->i')
48 assert_raises(TypeError, np.add, 1, 2, signaturex='ii->i')
49 assert_raises(TypeError, np.add, 1, 2, subokx=False)
50 assert_raises(TypeError, np.add, 1, 2, wherex=[True])
51
52 def test_sig_signature(self):
53 assert_raises(TypeError, np.add, 1, 2, sig='ii->i',
54 signature='ii->i')
55
56 def test_sig_dtype(self):
57 assert_raises(TypeError, np.add, 1, 2, sig='ii->i',
58 dtype=int)
59 assert_raises(TypeError, np.add, 1, 2, signature='ii->i',
60 dtype=int)
61
62 def test_extobj_removed(self):
63 assert_raises(TypeError, np.add, 1, 2, extobj=[4096])
64
65
66class TestUfuncGenericLoops:
67 """Test generic loops.
68
69 The loops to be tested are:
70
71 PyUFunc_ff_f_As_dd_d
72 PyUFunc_ff_f
73 PyUFunc_dd_d
74 PyUFunc_gg_g
75 PyUFunc_FF_F_As_DD_D
76 PyUFunc_DD_D
77 PyUFunc_FF_F
78 PyUFunc_GG_G
79 PyUFunc_OO_O
80 PyUFunc_OO_O_method
81 PyUFunc_f_f_As_d_d
82 PyUFunc_d_d
83 PyUFunc_f_f
84 PyUFunc_g_g
85 PyUFunc_F_F_As_D_D
86 PyUFunc_F_F
87 PyUFunc_D_D
88 PyUFunc_G_G
89 PyUFunc_O_O
90 PyUFunc_O_O_method
91 PyUFunc_On_Om
92
93 Where:
94
95 f -- float
96 d -- double
97 g -- long double
98 F -- complex float
99 D -- complex double
100 G -- complex long double
101 O -- python object
102
103 It is difficult to assure that each of these loops is entered from the
104 Python level as the special cased loops are a moving target and the
105 corresponding types are architecture dependent. We probably need to
106 define C level testing ufuncs to get at them. For the time being, I've
107 just looked at the signatures registered in the build directory to find
108 relevant functions.
109
110 """
111 np_dtypes = [
112 (np.single, np.single), (np.single, np.double),
113 (np.csingle, np.csingle), (np.csingle, np.cdouble),
114 (np.double, np.double), (np.longdouble, np.longdouble),
115 (np.cdouble, np.cdouble), (np.clongdouble, np.clongdouble)]
116
117 @pytest.mark.parametrize('input_dtype,output_dtype', np_dtypes)
118 def test_unary_PyUFunc(self, input_dtype, output_dtype, f=np.exp, x=0, y=1):
119 xs = np.full(10, input_dtype(x), dtype=output_dtype)
120 ys = f(xs)[::2]
121 assert_allclose(ys, y)
122 assert_equal(ys.dtype, output_dtype)
123
124 def f2(x, y):
125 return x**y
126
127 @pytest.mark.parametrize('input_dtype,output_dtype', np_dtypes)
128 def test_binary_PyUFunc(self, input_dtype, output_dtype, f=f2, x=0, y=1):
129 xs = np.full(10, input_dtype(x), dtype=output_dtype)
130 ys = f(xs, xs)[::2]
131 assert_allclose(ys, y)
132 assert_equal(ys.dtype, output_dtype)
133
134 # class to use in testing object method loops
135 class foo:
136 def conjugate(self):
137 return np.bool(1)
138
139 def logical_xor(self, obj):
140 return np.bool(1)
141
142 def test_unary_PyUFunc_O_O(self):
143 x = np.ones(10, dtype=object)
144 assert_(np.all(np.abs(x) == 1))
145
146 def test_unary_PyUFunc_O_O_method_simple(self, foo=foo):
147 x = np.full(10, foo(), dtype=object)
148 assert_(np.all(np.conjugate(x) == True))
149
150 def test_binary_PyUFunc_OO_O(self):
151 x = np.ones(10, dtype=object)
152 assert_(np.all(np.add(x, x) == 2))
153
154 def test_binary_PyUFunc_OO_O_method(self, foo=foo):
155 x = np.full(10, foo(), dtype=object)
156 assert_(np.all(np.logical_xor(x, x)))
157
158 def test_binary_PyUFunc_On_Om_method(self, foo=foo):
159 x = np.full((10, 2, 3), foo(), dtype=object)
160 assert_(np.all(np.logical_xor(x, x)))
161
162 def test_python_complex_conjugate(self):
163 # The conjugate ufunc should fall back to calling the method:
164 arr = np.array([1 + 2j, 3 - 4j], dtype="O")
165 assert isinstance(arr[0], complex)
166 res = np.conjugate(arr)
167 assert res.dtype == np.dtype("O")
168 assert_array_equal(res, np.array([1 - 2j, 3 + 4j], dtype="O"))
169
170 @pytest.mark.parametrize("ufunc", UNARY_OBJECT_UFUNCS)
171 def test_unary_PyUFunc_O_O_method_full(self, ufunc):
172 """Compare the result of the object loop with non-object one"""
173 val = np.float64(np.pi / 4)
174
175 class MyFloat(np.float64):
176 def __getattr__(self, attr):
177 try:
178 return super().__getattr__(attr)
179 except AttributeError:
180 return lambda: getattr(np._core.umath, attr)(val)
181
182 # Use 0-D arrays, to ensure the same element call
183 num_arr = np.array(val, dtype=np.float64)
184 obj_arr = np.array(MyFloat(val), dtype="O")
185
186 with np.errstate(all="raise"):
187 try:
188 res_num = ufunc(num_arr)
189 except Exception as exc:
190 with assert_raises(type(exc)):
191 ufunc(obj_arr)
192 else:
193 res_obj = ufunc(obj_arr)
194 assert_array_almost_equal(res_num.astype("O"), res_obj)
195
196
197def _pickleable_module_global():
198 pass
199
200
201class TestUfunc:
202 def test_pickle(self):
203 for proto in range(2, pickle.HIGHEST_PROTOCOL + 1):
204 assert_(pickle.loads(pickle.dumps(np.sin,
205 protocol=proto)) is np.sin)
206
207 # Check that ufunc not defined in the top level numpy namespace
208 # such as numpy._core._rational_tests.test_add can also be pickled
209 res = pickle.loads(pickle.dumps(_rational_tests.test_add,
210 protocol=proto))
211 assert_(res is _rational_tests.test_add)
212
213 def test_pickle_withstring(self):
214 astring = (b"cnumpy.core\n_ufunc_reconstruct\np0\n"
215 b"(S'numpy._core.umath'\np1\nS'cos'\np2\ntp3\nRp4\n.")
216 assert_(pickle.loads(astring) is np.cos)
217
218 @pytest.mark.skipif(IS_PYPY, reason="'is' check does not work on PyPy")
219 def test_pickle_name_is_qualname(self):
220 # This tests that a simplification of our ufunc pickle code will
221 # lead to allowing qualnames as names. Future ufuncs should
222 # possible add a specific qualname, or a hook into pickling instead
223 # (dask+numba may benefit).
224 _pickleable_module_global.ufunc = umt._pickleable_module_global_ufunc
225
226 obj = pickle.loads(pickle.dumps(_pickleable_module_global.ufunc))
227 assert obj is umt._pickleable_module_global_ufunc
228
229 def test_reduceat_shifting_sum(self):
230 L = 6
231 x = np.arange(L)
232 idx = np.array(list(zip(np.arange(L - 2), np.arange(L - 2) + 2))).ravel()
233 assert_array_equal(np.add.reduceat(x, idx)[::2], [1, 3, 5, 7])
234
235 def test_all_ufunc(self):
236 """Try to check presence and results of all ufuncs.
237
238 The list of ufuncs comes from generate_umath.py and is as follows:
239
240 ===== ==== ============= =============== ========================
241 done args function types notes
242 ===== ==== ============= =============== ========================
243 n 1 conjugate nums + O
244 n 1 absolute nums + O complex -> real
245 n 1 negative nums + O
246 n 1 sign nums + O -> int
247 n 1 invert bool + ints + O flts raise an error
248 n 1 degrees real + M cmplx raise an error
249 n 1 radians real + M cmplx raise an error
250 n 1 arccos flts + M
251 n 1 arccosh flts + M
252 n 1 arcsin flts + M
253 n 1 arcsinh flts + M
254 n 1 arctan flts + M
255 n 1 arctanh flts + M
256 n 1 cos flts + M
257 n 1 sin flts + M
258 n 1 tan flts + M
259 n 1 cosh flts + M
260 n 1 sinh flts + M
261 n 1 tanh flts + M
262 n 1 exp flts + M
263 n 1 expm1 flts + M
264 n 1 log flts + M
265 n 1 log10 flts + M
266 n 1 log1p flts + M
267 n 1 sqrt flts + M real x < 0 raises error
268 n 1 ceil real + M
269 n 1 trunc real + M
270 n 1 floor real + M
271 n 1 fabs real + M
272 n 1 rint flts + M
273 n 1 isnan flts -> bool
274 n 1 isinf flts -> bool
275 n 1 isfinite flts -> bool
276 n 1 signbit real -> bool
277 n 1 modf real -> (frac, int)
278 n 1 logical_not bool + nums + M -> bool
279 n 2 left_shift ints + O flts raise an error
280 n 2 right_shift ints + O flts raise an error
281 n 2 add bool + nums + O boolean + is ||
282 n 2 subtract bool + nums + O boolean - is ^
283 n 2 multiply bool + nums + O boolean * is &
284 n 2 divide nums + O
285 n 2 floor_divide nums + O
286 n 2 true_divide nums + O bBhH -> f, iIlLqQ -> d
287 n 2 fmod nums + M
288 n 2 power nums + O
289 n 2 greater bool + nums + O -> bool
290 n 2 greater_equal bool + nums + O -> bool
291 n 2 less bool + nums + O -> bool
292 n 2 less_equal bool + nums + O -> bool
293 n 2 equal bool + nums + O -> bool
294 n 2 not_equal bool + nums + O -> bool
295 n 2 logical_and bool + nums + M -> bool
296 n 2 logical_or bool + nums + M -> bool
297 n 2 logical_xor bool + nums + M -> bool
298 n 2 maximum bool + nums + O
299 n 2 minimum bool + nums + O
300 n 2 bitwise_and bool + ints + O flts raise an error
301 n 2 bitwise_or bool + ints + O flts raise an error
302 n 2 bitwise_xor bool + ints + O flts raise an error
303 n 2 arctan2 real + M
304 n 2 remainder ints + real + O
305 n 2 hypot real + M
306 ===== ==== ============= =============== ========================
307
308 Types other than those listed will be accepted, but they are cast to
309 the smallest compatible type for which the function is defined. The
310 casting rules are:
311
312 bool -> int8 -> float32
313 ints -> double
314
315 """
316 pass
317
318 # from include/numpy/ufuncobject.h
319 size_inferred = 2
320 can_ignore = 4
321
322 def test_signature0(self):
323 # the arguments to test_signature are: nin, nout, core_signature
324 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
325 2, 1, "(i),(i)->()")
326 assert_equal(enabled, 1)
327 assert_equal(num_dims, (1, 1, 0))
328 assert_equal(ixs, (0, 0))
329 assert_equal(flags, (self.size_inferred,))
330 assert_equal(sizes, (-1,))
331
332 def test_signature1(self):
333 # empty core signature; treat as plain ufunc (with trivial core)
334 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
335 2, 1, "(),()->()")
336 assert_equal(enabled, 0)
337 assert_equal(num_dims, (0, 0, 0))
338 assert_equal(ixs, ())
339 assert_equal(flags, ())
340 assert_equal(sizes, ())
341
342 def test_signature2(self):
343 # more complicated names for variables
344 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
345 2, 1, "(i1,i2),(J_1)->(_kAB)")
346 assert_equal(enabled, 1)
347 assert_equal(num_dims, (2, 1, 1))
348 assert_equal(ixs, (0, 1, 2, 3))
349 assert_equal(flags, (self.size_inferred,) * 4)
350 assert_equal(sizes, (-1, -1, -1, -1))
351
352 def test_signature3(self):
353 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
354 2, 1, "(i1, i12), (J_1)->(i12, i2)")
355 assert_equal(enabled, 1)
356 assert_equal(num_dims, (2, 1, 2))
357 assert_equal(ixs, (0, 1, 2, 1, 3))
358 assert_equal(flags, (self.size_inferred,) * 4)
359 assert_equal(sizes, (-1, -1, -1, -1))
360
361 def test_signature4(self):
362 # matrix_multiply signature from _umath_tests
363 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
364 2, 1, "(n,k),(k,m)->(n,m)")
365 assert_equal(enabled, 1)
366 assert_equal(num_dims, (2, 2, 2))
367 assert_equal(ixs, (0, 1, 1, 2, 0, 2))
368 assert_equal(flags, (self.size_inferred,) * 3)
369 assert_equal(sizes, (-1, -1, -1))
370
371 def test_signature5(self):
372 # matmul signature from _umath_tests
373 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
374 2, 1, "(n?,k),(k,m?)->(n?,m?)")
375 assert_equal(enabled, 1)
376 assert_equal(num_dims, (2, 2, 2))
377 assert_equal(ixs, (0, 1, 1, 2, 0, 2))
378 assert_equal(flags, (self.size_inferred | self.can_ignore,
379 self.size_inferred,
380 self.size_inferred | self.can_ignore))
381 assert_equal(sizes, (-1, -1, -1))
382
383 def test_signature6(self):
384 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
385 1, 1, "(3)->()")
386 assert_equal(enabled, 1)
387 assert_equal(num_dims, (1, 0))
388 assert_equal(ixs, (0,))
389 assert_equal(flags, (0,))
390 assert_equal(sizes, (3,))
391
392 def test_signature7(self):
393 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
394 3, 1, "(3),(03,3),(n)->(9)")
395 assert_equal(enabled, 1)
396 assert_equal(num_dims, (1, 2, 1, 1))
397 assert_equal(ixs, (0, 0, 0, 1, 2))
398 assert_equal(flags, (0, self.size_inferred, 0))
399 assert_equal(sizes, (3, -1, 9))
400
401 def test_signature8(self):
402 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
403 3, 1, "(3?),(3?,3?),(n)->(9)")
404 assert_equal(enabled, 1)
405 assert_equal(num_dims, (1, 2, 1, 1))
406 assert_equal(ixs, (0, 0, 0, 1, 2))
407 assert_equal(flags, (self.can_ignore, self.size_inferred, 0))
408 assert_equal(sizes, (3, -1, 9))
409
410 def test_signature9(self):
411 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
412 1, 1, "( 3) -> ( )")
413 assert_equal(enabled, 1)
414 assert_equal(num_dims, (1, 0))
415 assert_equal(ixs, (0,))
416 assert_equal(flags, (0,))
417 assert_equal(sizes, (3,))
418
419 def test_signature10(self):
420 enabled, num_dims, ixs, flags, sizes = umt.test_signature(
421 3, 1, "( 3? ) , (3? , 3?) ,(n )-> ( 9)")
422 assert_equal(enabled, 1)
423 assert_equal(num_dims, (1, 2, 1, 1))
424 assert_equal(ixs, (0, 0, 0, 1, 2))
425 assert_equal(flags, (self.can_ignore, self.size_inferred, 0))
426 assert_equal(sizes, (3, -1, 9))
427
428 def test_signature_failure_extra_parenthesis(self):
429 with assert_raises(ValueError):
430 umt.test_signature(2, 1, "((i)),(i)->()")
431
432 def test_signature_failure_mismatching_parenthesis(self):
433 with assert_raises(ValueError):
434 umt.test_signature(2, 1, "(i),)i(->()")
435
436 def test_signature_failure_signature_missing_input_arg(self):
437 with assert_raises(ValueError):
438 umt.test_signature(2, 1, "(i),->()")
439
440 def test_signature_failure_signature_missing_output_arg(self):
441 with assert_raises(ValueError):
442 umt.test_signature(2, 2, "(i),(i)->()")
443
444 def test_get_signature(self):
445 assert_equal(np.vecdot.signature, "(n),(n)->()")
446
447 def test_forced_sig(self):
448 a = 0.5 * np.arange(3, dtype='f8')
449 assert_equal(np.add(a, 0.5), [0.5, 1, 1.5])
450 with assert_raises(TypeError):
451 np.add(a, 0.5, sig='i', casting='unsafe')
452 assert_equal(np.add(a, 0.5, sig='ii->i', casting='unsafe'), [0, 0, 1])
453 with assert_raises(TypeError):
454 np.add(a, 0.5, sig=('i4',), casting='unsafe')
455 assert_equal(np.add(a, 0.5, sig=('i4', 'i4', 'i4'),
456 casting='unsafe'), [0, 0, 1])
457
458 b = np.zeros((3,), dtype='f8')
459 np.add(a, 0.5, out=b)
460 assert_equal(b, [0.5, 1, 1.5])
461 b[:] = 0
462 with assert_raises(TypeError):
463 np.add(a, 0.5, sig='i', out=b, casting='unsafe')
464 assert_equal(b, [0, 0, 0])
465 np.add(a, 0.5, sig='ii->i', out=b, casting='unsafe')
466 assert_equal(b, [0, 0, 1])
467 b[:] = 0
468 with assert_raises(TypeError):
469 np.add(a, 0.5, sig=('i4',), out=b, casting='unsafe')
470 assert_equal(b, [0, 0, 0])
471 np.add(a, 0.5, sig=('i4', 'i4', 'i4'), out=b, casting='unsafe')
472 assert_equal(b, [0, 0, 1])
473
474 def test_signature_all_None(self):
475 # signature all None, is an acceptable alternative (since 1.21)
476 # to not providing a signature.
477 res1 = np.add([3], [4], sig=(None, None, None))
478 res2 = np.add([3], [4])
479 assert_array_equal(res1, res2)
480 res1 = np.maximum([3], [4], sig=(None, None, None))
481 res2 = np.maximum([3], [4])
482 assert_array_equal(res1, res2)
483
484 with pytest.raises(TypeError):
485 # special case, that would be deprecated anyway, so errors:
486 np.add(3, 4, signature=(None,))
487
488 def test_signature_dtype_type(self):
489 # Since that will be the normal behaviour (past NumPy 1.21)
490 # we do support the types already:
491 float_dtype = type(np.dtype(np.float64))
492 np.add(3, 4, signature=(float_dtype, float_dtype, None))
493
494 @pytest.mark.parametrize("get_kwarg", [
495 param(lambda dt: {"dtype": dt}, id="dtype"),
496 param(lambda dt: {"signature": (dt, None, None)}, id="signature")])
497 def test_signature_dtype_instances_allowed(self, get_kwarg):
498 # We allow certain dtype instances when there is a clear singleton
499 # and the given one is equivalent; mainly for backcompat.
500 int64 = np.dtype("int64")
501 int64_2 = pickle.loads(pickle.dumps(int64))
502 # Relies on pickling behavior, if assert fails just remove test...
503 assert int64 is not int64_2
504
505 assert np.add(1, 2, **get_kwarg(int64_2)).dtype == int64
506 td = np.timedelta64(2, "s")
507 assert np.add(td, td, **get_kwarg("m8")).dtype == "m8[s]"
508
509 msg = "The `dtype` and `signature` arguments to ufuncs"
510
511 with pytest.raises(TypeError, match=msg):
512 np.add(3, 5, **get_kwarg(np.dtype("int64").newbyteorder()))
513 with pytest.raises(TypeError, match=msg):
514 np.add(3, 5, **get_kwarg(np.dtype("m8[ns]")))
515 with pytest.raises(TypeError, match=msg):
516 np.add(3, 5, **get_kwarg("m8[ns]"))
517
518 @pytest.mark.parametrize("casting", ["unsafe", "same_kind", "safe"])
519 def test_partial_signature_mismatch(self, casting):
520 # If the second argument matches already, no need to specify it:
521 res = np.ldexp(np.float32(1.), np.int_(2), dtype="d")
522 assert res.dtype == "d"
523 res = np.ldexp(np.float32(1.), np.int_(2), signature=(None, None, "d"))
524 assert res.dtype == "d"
525
526 # ldexp only has a loop for long input as second argument, overriding
527 # the output cannot help with that (no matter the casting)
528 with pytest.raises(TypeError):
529 np.ldexp(1., np.uint64(3), dtype="d")
530 with pytest.raises(TypeError):
531 np.ldexp(1., np.uint64(3), signature=(None, None, "d"))
532
533 def test_partial_signature_mismatch_with_cache(self):
534 with pytest.raises(TypeError):
535 np.add(np.float16(1), np.uint64(2), sig=("e", "d", None))
536 # Ensure e,d->None is in the dispatching cache (double loop)
537 np.add(np.float16(1), np.float64(2))
538 # The error must still be raised:
539 with pytest.raises(TypeError):
540 np.add(np.float16(1), np.uint64(2), sig=("e", "d", None))
541
542 def test_use_output_signature_for_all_arguments(self):
543 # Test that providing only `dtype=` or `signature=(None, None, dtype)`
544 # is sufficient if falling back to a homogeneous signature works.
545 # In this case, the `intp, intp -> intp` loop is chosen.
546 res = np.power(1.5, 2.8, dtype=np.intp, casting="unsafe")
547 assert res == 1 # the cast happens first.
548 res = np.power(1.5, 2.8, signature=(None, None, np.intp),
549 casting="unsafe")
550 assert res == 1
551 with pytest.raises(TypeError):
552 # the unsafe casting would normally cause errors though:
553 np.power(1.5, 2.8, dtype=np.intp)
554
555 def test_signature_errors(self):
556 with pytest.raises(TypeError,
557 match="the signature object to ufunc must be a string or"):
558 np.add(3, 4, signature=123.) # neither a string nor a tuple
559
560 with pytest.raises(ValueError):
561 # bad symbols that do not translate to dtypes
562 np.add(3, 4, signature="%^->#")
563
564 with pytest.raises(ValueError):
565 np.add(3, 4, signature=b"ii-i") # incomplete and byte string
566
567 with pytest.raises(ValueError):
568 np.add(3, 4, signature="ii>i") # incomplete string
569
570 with pytest.raises(ValueError):
571 np.add(3, 4, signature=(None, "f8")) # bad length
572
573 with pytest.raises(UnicodeDecodeError):
574 np.add(3, 4, signature=b"\xff\xff->i")
575
576 def test_forced_dtype_times(self):
577 # Signatures only set the type numbers (not the actual loop dtypes)
578 # so using `M` in a signature/dtype should generally work:
579 a = np.array(['2010-01-02', '1999-03-14', '1833-03'], dtype='>M8[D]')
580 np.maximum(a, a, dtype="M")
581 np.maximum.reduce(a, dtype="M")
582
583 arr = np.arange(10, dtype="m8[s]")
584 np.add(arr, arr, dtype="m")
585 np.maximum(arr, arr, dtype="m")
586
587 @pytest.mark.parametrize("ufunc", [np.add, np.sqrt])
588 def test_cast_safety(self, ufunc):
589 """Basic test for the safest casts, because ufuncs inner loops can
590 indicate a cast-safety as well (which is normally always "no").
591 """
592 def call_ufunc(arr, **kwargs):
593 return ufunc(*(arr,) * ufunc.nin, **kwargs)
594
595 arr = np.array([1., 2., 3.], dtype=np.float32)
596 arr_bs = arr.astype(arr.dtype.newbyteorder())
597 expected = call_ufunc(arr)
598 # Normally, a "no" cast:
599 res = call_ufunc(arr, casting="no")
600 assert_array_equal(expected, res)
601 # Byte-swapping is not allowed with "no" though:
602 with pytest.raises(TypeError):
603 call_ufunc(arr_bs, casting="no")
604
605 # But is allowed with "equiv":
606 res = call_ufunc(arr_bs, casting="equiv")
607 assert_array_equal(expected, res)
608
609 # Casting to float64 is safe, but not equiv:
610 with pytest.raises(TypeError):
611 call_ufunc(arr_bs, dtype=np.float64, casting="equiv")
612
613 # but it is safe cast:
614 res = call_ufunc(arr_bs, dtype=np.float64, casting="safe")
615 expected = call_ufunc(arr.astype(np.float64)) # upcast
616 assert_array_equal(expected, res)
617
618 @pytest.mark.parametrize("ufunc", [np.add, np.equal])
619 def test_cast_safety_scalar(self, ufunc):
620 # We test add and equal, because equal has special scalar handling
621 # Note that the "equiv" casting behavior should maybe be considered
622 # a current implementation detail.
623 with pytest.raises(TypeError):
624 # this picks an integer loop, which is not safe
625 ufunc(3., 4., dtype=int, casting="safe")
626
627 with pytest.raises(TypeError):
628 # We accept python float as float64 but not float32 for equiv.
629 ufunc(3., 4., dtype="float32", casting="equiv")
630
631 # Special case for object and equal (note that equiv implies safe)
632 ufunc(3, 4, dtype=object, casting="equiv")
633 # Picks a double loop for both, first is equiv, second safe:
634 ufunc(np.array([3.]), 3., casting="equiv")
635 ufunc(np.array([3.]), 3, casting="safe")
636 ufunc(np.array([3]), 3, casting="equiv")
637
638 def test_cast_safety_scalar_special(self):
639 # We allow this (and it succeeds) via object, although the equiv
640 # part may not be important.
641 np.equal(np.array([3]), 2**300, casting="equiv")
642
643 def test_true_divide(self):
644 a = np.array(10)
645 b = np.array(20)
646 tgt = np.array(0.5)
647
648 for tc in 'bhilqBHILQefdgFDG':
649 dt = np.dtype(tc)
650 aa = a.astype(dt)
651 bb = b.astype(dt)
652
653 # Check result value and dtype.
654 for x, y in itertools.product([aa, -aa], [bb, -bb]):
655
656 # Check with no output type specified
657 if tc in 'FDG':
658 tgt = complex(x) / complex(y)
659 else:
660 tgt = float(x) / float(y)
661
662 res = np.true_divide(x, y)
663 rtol = max(np.finfo(res).resolution, 1e-15)
664 assert_allclose(res, tgt, rtol=rtol)
665
666 if tc in 'bhilqBHILQ':
667 assert_(res.dtype.name == 'float64')
668 else:
669 assert_(res.dtype.name == dt.name)
670
671 # Check with output type specified. This also checks for the
672 # incorrect casts in issue gh-3484 because the unary '-' does
673 # not change types, even for unsigned types, Hence casts in the
674 # ufunc from signed to unsigned and vice versa will lead to
675 # errors in the values.
676 for tcout in 'bhilqBHILQ':
677 dtout = np.dtype(tcout)
678 assert_raises(TypeError, np.true_divide, x, y, dtype=dtout)
679
680 for tcout in 'efdg':
681 dtout = np.dtype(tcout)
682 if tc in 'FDG':
683 # Casting complex to float is not allowed
684 assert_raises(TypeError, np.true_divide, x, y, dtype=dtout)
685 else:
686 tgt = float(x) / float(y)
687 rtol = max(np.finfo(dtout).resolution, 1e-15)
688 # The value of tiny for double double is NaN
689 with warnings.catch_warnings():
690 warnings.simplefilter('ignore', UserWarning)
691 if not np.isnan(np.finfo(dtout).tiny):
692 atol = max(np.finfo(dtout).tiny, 3e-308)
693 else:
694 atol = 3e-308
695 # Some test values result in invalid for float16
696 # and the cast to it may overflow to inf.
697 with np.errstate(invalid='ignore', over='ignore'):
698 res = np.true_divide(x, y, dtype=dtout)
699 if not np.isfinite(res) and tcout == 'e':
700 continue
701 assert_allclose(res, tgt, rtol=rtol, atol=atol)
702 assert_(res.dtype.name == dtout.name)
703
704 for tcout in 'FDG':
705 dtout = np.dtype(tcout)
706 tgt = complex(x) / complex(y)
707 rtol = max(np.finfo(dtout).resolution, 1e-15)
708 # The value of tiny for double double is NaN
709 with warnings.catch_warnings():
710 warnings.simplefilter('ignore', UserWarning)
711 if not np.isnan(np.finfo(dtout).tiny):
712 atol = max(np.finfo(dtout).tiny, 3e-308)
713 else:
714 atol = 3e-308
715 res = np.true_divide(x, y, dtype=dtout)
716 if not np.isfinite(res):
717 continue
718 assert_allclose(res, tgt, rtol=rtol, atol=atol)
719 assert_(res.dtype.name == dtout.name)
720
721 # Check booleans
722 a = np.ones((), dtype=np.bool)
723 res = np.true_divide(a, a)
724 assert_(res == 1.0)
725 assert_(res.dtype.name == 'float64')
726 res = np.true_divide(~a, a)
727 assert_(res == 0.0)
728 assert_(res.dtype.name == 'float64')
729
730 def test_sum_stability(self):
731 a = np.ones(500, dtype=np.float32)
732 assert_almost_equal((a / 10.).sum() - a.size / 10., 0, 4)
733
734 a = np.ones(500, dtype=np.float64)
735 assert_almost_equal((a / 10.).sum() - a.size / 10., 0, 13)
736
737 @pytest.mark.skipif(IS_WASM, reason="fp errors don't work in wasm")
738 def test_sum(self):
739 for dt in (int, np.float16, np.float32, np.float64, np.longdouble):
740 for v in (0, 1, 2, 7, 8, 9, 15, 16, 19, 127,
741 128, 1024, 1235):
742 # warning if sum overflows, which it does in float16
743 with warnings.catch_warnings(record=True) as w:
744 warnings.simplefilter("always", RuntimeWarning)
745
746 tgt = dt(v * (v + 1) / 2)
747 overflow = not np.isfinite(tgt)
748 assert_equal(len(w), 1 * overflow)
749
750 d = np.arange(1, v + 1, dtype=dt)
751
752 assert_almost_equal(np.sum(d), tgt)
753 assert_equal(len(w), 2 * overflow)
754
755 assert_almost_equal(np.sum(d[::-1]), tgt)
756 assert_equal(len(w), 3 * overflow)
757
758 d = np.ones(500, dtype=dt)
759 assert_almost_equal(np.sum(d[::2]), 250.)
760 assert_almost_equal(np.sum(d[1::2]), 250.)
761 assert_almost_equal(np.sum(d[::3]), 167.)
762 assert_almost_equal(np.sum(d[1::3]), 167.)
763 assert_almost_equal(np.sum(d[::-2]), 250.)
764 assert_almost_equal(np.sum(d[-1::-2]), 250.)
765 assert_almost_equal(np.sum(d[::-3]), 167.)
766 assert_almost_equal(np.sum(d[-1::-3]), 167.)
767 # sum with first reduction entry != 0
768 d = np.ones((1,), dtype=dt)
769 d += d
770 assert_almost_equal(d, 2.)
771
772 def test_sum_complex(self):
773 for dt in (np.complex64, np.complex128, np.clongdouble):
774 for v in (0, 1, 2, 7, 8, 9, 15, 16, 19, 127,
775 128, 1024, 1235):
776 tgt = dt(v * (v + 1) / 2) - dt((v * (v + 1) / 2) * 1j)
777 d = np.empty(v, dtype=dt)
778 d.real = np.arange(1, v + 1)
779 d.imag = -np.arange(1, v + 1)
780 assert_almost_equal(np.sum(d), tgt)
781 assert_almost_equal(np.sum(d[::-1]), tgt)
782
783 d = np.ones(500, dtype=dt) + 1j
784 assert_almost_equal(np.sum(d[::2]), 250. + 250j)
785 assert_almost_equal(np.sum(d[1::2]), 250. + 250j)
786 assert_almost_equal(np.sum(d[::3]), 167. + 167j)
787 assert_almost_equal(np.sum(d[1::3]), 167. + 167j)
788 assert_almost_equal(np.sum(d[::-2]), 250. + 250j)
789 assert_almost_equal(np.sum(d[-1::-2]), 250. + 250j)
790 assert_almost_equal(np.sum(d[::-3]), 167. + 167j)
791 assert_almost_equal(np.sum(d[-1::-3]), 167. + 167j)
792 # sum with first reduction entry != 0
793 d = np.ones((1,), dtype=dt) + 1j
794 d += d
795 assert_almost_equal(d, 2. + 2j)
796
797 def test_sum_initial(self):
798 # Integer, single axis
799 assert_equal(np.sum([3], initial=2), 5)
800
801 # Floating point
802 assert_almost_equal(np.sum([0.2], initial=0.1), 0.3)
803
804 # Multiple non-adjacent axes
805 assert_equal(np.sum(np.ones((2, 3, 5), dtype=np.int64), axis=(0, 2), initial=2),
806 [12, 12, 12])
807
808 def test_sum_where(self):
809 # More extensive tests done in test_reduction_with_where.
810 assert_equal(np.sum([[1., 2.], [3., 4.]], where=[True, False]), 4.)
811 assert_equal(np.sum([[1., 2.], [3., 4.]], axis=0, initial=5.,
812 where=[True, False]), [9., 5.])
813
814 def test_vecdot(self):
815 arr1 = np.arange(6).reshape((2, 3))
816 arr2 = np.arange(3).reshape((1, 3))
817
818 actual = np.vecdot(arr1, arr2)
819 expected = np.array([5, 14])
820
821 assert_array_equal(actual, expected)
822
823 actual2 = np.vecdot(arr1.T, arr2.T, axis=-2)
824 assert_array_equal(actual2, expected)
825
826 actual3 = np.vecdot(arr1.astype("object"), arr2)
827 assert_array_equal(actual3, expected.astype("object"))
828
829 def test_matvec(self):
830 arr1 = np.arange(6).reshape((2, 3))
831 arr2 = np.arange(3).reshape((1, 3))
832
833 actual = np.matvec(arr1, arr2)
834 expected = np.array([[5, 14]])
835
836 assert_array_equal(actual, expected)
837
838 actual2 = np.matvec(arr1.T, arr2.T, axes=[(-1, -2), -2, -1])
839 assert_array_equal(actual2, expected)
840
841 actual3 = np.matvec(arr1.astype("object"), arr2)
842 assert_array_equal(actual3, expected.astype("object"))
843
844 @pytest.mark.parametrize("vec", [
845 np.array([[1., 2., 3.], [4., 5., 6.]]),
846 np.array([[1., 2j, 3.], [4., 5., 6j]]),
847 np.array([[1., 2., 3.], [4., 5., 6.]], dtype=object),
848 np.array([[1., 2j, 3.], [4., 5., 6j]], dtype=object)])
849 @pytest.mark.parametrize("matrix", [
850 None,
851 np.array([[1. + 1j, 0.5, -0.5j],
852 [0.25, 2j, 0.],
853 [4., 0., -1j]])])
854 def test_vecmatvec_identity(self, matrix, vec):
855 """Check that (x†A)x equals x†(Ax)."""
856 mat = matrix if matrix is not None else np.eye(3)
857 matvec = np.matvec(mat, vec) # Ax
858 vecmat = np.vecmat(vec, mat) # x†A
859 if matrix is None:
860 assert_array_equal(matvec, vec)
861 assert_array_equal(vecmat.conj(), vec)
862 assert_array_equal(matvec, (mat @ vec[..., np.newaxis]).squeeze(-1))
863 assert_array_equal(vecmat, (vec[..., np.newaxis].mT.conj()
864 @ mat).squeeze(-2))
865 expected = np.einsum('...i,ij,...j', vec.conj(), mat, vec)
866 vec_matvec = (vec.conj() * matvec).sum(-1)
867 vecmat_vec = (vecmat * vec).sum(-1)
868 assert_array_equal(vec_matvec, expected)
869 assert_array_equal(vecmat_vec, expected)
870
871 @pytest.mark.parametrize("ufunc, shape1, shape2, conj", [
872 (np.vecdot, (3,), (3,), True),
873 (np.vecmat, (3,), (3, 1), True),
874 (np.matvec, (1, 3), (3,), False),
875 (np.matmul, (1, 3), (3, 1), False),
876 ])
877 def test_vecdot_matvec_vecmat_complex(self, ufunc, shape1, shape2, conj):
878 arr1 = np.array([1, 2j, 3])
879 arr2 = np.array([1, 2, 3])
880
881 actual1 = ufunc(arr1.reshape(shape1), arr2.reshape(shape2))
882 expected1 = np.array(((arr1.conj() if conj else arr1) * arr2).sum(),
883 ndmin=min(len(shape1), len(shape2)))
884 assert_array_equal(actual1, expected1)
885 # This would fail for conj=True, since matmul omits the conjugate.
886 if not conj:
887 assert_array_equal(arr1.reshape(shape1) @ arr2.reshape(shape2),
888 expected1)
889
890 actual2 = ufunc(arr2.reshape(shape1), arr1.reshape(shape2))
891 expected2 = np.array(((arr2.conj() if conj else arr2) * arr1).sum(),
892 ndmin=min(len(shape1), len(shape2)))
893 assert_array_equal(actual2, expected2)
894
895 actual3 = ufunc(arr1.reshape(shape1).astype("object"),
896 arr2.reshape(shape2).astype("object"))
897 expected3 = expected1.astype(object)
898 assert_array_equal(actual3, expected3)
899
900 def test_vecdot_subclass(self):
901 class MySubclass(np.ndarray):
902 pass
903
904 arr1 = np.arange(6).reshape((2, 3)).view(MySubclass)
905 arr2 = np.arange(3).reshape((1, 3)).view(MySubclass)
906 result = np.vecdot(arr1, arr2)
907 assert isinstance(result, MySubclass)
908
909 def test_vecdot_object_no_conjugate(self):
910 arr = np.array(["1", "2"], dtype=object)
911 with pytest.raises(AttributeError, match="conjugate"):
912 np.vecdot(arr, arr)
913
914 def test_vecdot_object_breaks_outer_loop_on_error(self):
915 arr1 = np.ones((3, 3)).astype(object)
916 arr2 = arr1.copy()
917 arr2[1, 1] = None
918 out = np.zeros(3).astype(object)
919 with pytest.raises(TypeError, match=r"\*: 'float' and 'NoneType'"):
920 np.vecdot(arr1, arr2, out=out)
921 assert out[0] == 3
922 assert out[1] == out[2] == 0
923
924 def test_broadcast(self):
925 msg = "broadcast"
926 a = np.arange(4).reshape((2, 1, 2))
927 b = np.arange(4).reshape((1, 2, 2))
928 assert_array_equal(np.vecdot(a, b), np.sum(a * b, axis=-1), err_msg=msg)
929 msg = "extend & broadcast loop dimensions"
930 b = np.arange(4).reshape((2, 2))
931 assert_array_equal(np.vecdot(a, b), np.sum(a * b, axis=-1), err_msg=msg)
932 # Broadcast in core dimensions should fail
933 a = np.arange(8).reshape((4, 2))
934 b = np.arange(4).reshape((4, 1))
935 assert_raises(ValueError, np.vecdot, a, b)
936 # Extend core dimensions should fail
937 a = np.arange(8).reshape((4, 2))
938 b = np.array(7)
939 assert_raises(ValueError, np.vecdot, a, b)
940 # Broadcast should fail
941 a = np.arange(2).reshape((2, 1, 1))
942 b = np.arange(3).reshape((3, 1, 1))
943 assert_raises(ValueError, np.vecdot, a, b)
944
945 # Writing to a broadcasted array with overlap should warn, gh-2705
946 a = np.arange(2)
947 b = np.arange(4).reshape((2, 2))
948 u, v = np.broadcast_arrays(a, b)
949 assert_equal(u.strides[0], 0)
950 x = u + v
951 with warnings.catch_warnings(record=True) as w:
952 warnings.simplefilter("always")
953 u += v
954 assert_equal(len(w), 1)
955 assert_(x[0, 0] != u[0, 0])
956
957 # Output reduction should not be allowed.
958 # See gh-15139
959 a = np.arange(6).reshape(3, 2)
960 b = np.ones(2)
961 out = np.empty(())
962 assert_raises(ValueError, np.vecdot, a, b, out)
963 out2 = np.empty(3)
964 c = np.vecdot(a, b, out2)
965 assert_(c is out2)
966
967 def test_out_broadcasts(self):
968 # For ufuncs and gufuncs (not for reductions), we currently allow
969 # the output to cause broadcasting of the input arrays.
970 # both along dimensions with shape 1 and dimensions which do not
971 # exist at all in the inputs.
972 arr = np.arange(3).reshape(1, 3)
973 out = np.empty((5, 4, 3))
974 np.add(arr, arr, out=out)
975 assert (out == np.arange(3) * 2).all()
976
977 # The same holds for gufuncs (gh-16484)
978 np.vecdot(arr, arr, out=out)
979 # the result would be just a scalar `5`, but is broadcast fully:
980 assert (out == 5).all()
981
982 @pytest.mark.parametrize(["arr", "out"], [
983 ([2], np.empty(())),
984 ([1, 2], np.empty(1)),
985 (np.ones((4, 3)), np.empty((4, 1)))],
986 ids=["(1,)->()", "(2,)->(1,)", "(4, 3)->(4, 1)"])
987 def test_out_broadcast_errors(self, arr, out):
988 # Output is (currently) allowed to broadcast inputs, but it cannot be
989 # smaller than the actual result.
990 with pytest.raises(ValueError, match="non-broadcastable"):
991 np.positive(arr, out=out)
992
993 with pytest.raises(ValueError, match="non-broadcastable"):
994 np.add(np.ones(()), arr, out=out)
995
996 def test_type_cast(self):
997 msg = "type cast"
998 a = np.arange(6, dtype='short').reshape((2, 3))
999 assert_array_equal(np.vecdot(a, a), np.sum(a * a, axis=-1),
1000 err_msg=msg)
1001 msg = "type cast on one argument"
1002 a = np.arange(6).reshape((2, 3))
1003 b = a + 0.1
1004 assert_array_almost_equal(np.vecdot(a, b), np.sum(a * b, axis=-1),
1005 err_msg=msg)
1006
1007 def test_endian(self):
1008 msg = "big endian"
1009 a = np.arange(6, dtype='>i4').reshape((2, 3))
1010 assert_array_equal(np.vecdot(a, a), np.sum(a * a, axis=-1),
1011 err_msg=msg)
1012 msg = "little endian"
1013 a = np.arange(6, dtype='<i4').reshape((2, 3))
1014 assert_array_equal(np.vecdot(a, a), np.sum(a * a, axis=-1),
1015 err_msg=msg)
1016
1017 # Output should always be native-endian
1018 Ba = np.arange(1, dtype='>f8')
1019 La = np.arange(1, dtype='<f8')
1020 assert_equal((Ba + Ba).dtype, np.dtype('f8'))
1021 assert_equal((Ba + La).dtype, np.dtype('f8'))
1022 assert_equal((La + Ba).dtype, np.dtype('f8'))
1023 assert_equal((La + La).dtype, np.dtype('f8'))
1024
1025 assert_equal(np.absolute(La).dtype, np.dtype('f8'))
1026 assert_equal(np.absolute(Ba).dtype, np.dtype('f8'))
1027 assert_equal(np.negative(La).dtype, np.dtype('f8'))
1028 assert_equal(np.negative(Ba).dtype, np.dtype('f8'))
1029
1030 def test_incontiguous_array(self):
1031 msg = "incontiguous memory layout of array"
1032 x = np.arange(64).reshape((2, 2, 2, 2, 2, 2))
1033 a = x[:, 0, :, 0, :, 0]
1034 b = x[:, 1, :, 1, :, 1]
1035 a[0, 0, 0] = -1
1036 msg2 = "make sure it references to the original array"
1037 assert_equal(x[0, 0, 0, 0, 0, 0], -1, err_msg=msg2)
1038 assert_array_equal(np.vecdot(a, b), np.sum(a * b, axis=-1), err_msg=msg)
1039 x = np.arange(24).reshape(2, 3, 4)
1040 a = x.T
1041 b = x.T
1042 a[0, 0, 0] = -1
1043 assert_equal(x[0, 0, 0], -1, err_msg=msg2)
1044 assert_array_equal(np.vecdot(a, b), np.sum(a * b, axis=-1), err_msg=msg)
1045
1046 def test_output_argument(self):
1047 msg = "output argument"
1048 a = np.arange(12).reshape((2, 3, 2))
1049 b = np.arange(4).reshape((2, 1, 2)) + 1
1050 c = np.zeros((2, 3), dtype='int')
1051 np.vecdot(a, b, c)
1052 assert_array_equal(c, np.sum(a * b, axis=-1), err_msg=msg)
1053 c[:] = -1
1054 np.vecdot(a, b, out=c)
1055 assert_array_equal(c, np.sum(a * b, axis=-1), err_msg=msg)
1056
1057 msg = "output argument with type cast"
1058 c = np.zeros((2, 3), dtype='int16')
1059 np.vecdot(a, b, c)
1060 assert_array_equal(c, np.sum(a * b, axis=-1), err_msg=msg)
1061 c[:] = -1
1062 np.vecdot(a, b, out=c)
1063 assert_array_equal(c, np.sum(a * b, axis=-1), err_msg=msg)
1064
1065 msg = "output argument with incontiguous layout"
1066 c = np.zeros((2, 3, 4), dtype='int16')
1067 np.vecdot(a, b, c[..., 0])
1068 assert_array_equal(c[..., 0], np.sum(a * b, axis=-1), err_msg=msg)
1069 c[:] = -1
1070 np.vecdot(a, b, out=c[..., 0])
1071 assert_array_equal(c[..., 0], np.sum(a * b, axis=-1), err_msg=msg)
1072
1073 @pytest.mark.parametrize("arg", ["array", "scalar", "subclass"])
1074 def test_output_ellipsis(self, arg):
1075 class subclass(np.ndarray):
1076 def __array_wrap__(self, obj, context=None, return_value=None):
1077 return super().__array_wrap__(obj, context, return_value)
1078
1079 if arg == "scalar":
1080 one = 1
1081 expected_type = np.ndarray
1082 elif arg == "array":
1083 one = np.array(1)
1084 expected_type = np.ndarray
1085 elif arg == "subclass":
1086 one = np.array(1).view(subclass)
1087 expected_type = subclass
1088
1089 assert type(np.add(one, 2, out=...)) is expected_type
1090 assert type(np.add.reduce(one, out=...)) is expected_type
1091 res1, res2 = np.divmod(one, 2, out=...)
1092 assert type(res1) is type(res2) is expected_type
1093
1094 def test_output_ellipsis_errors(self):
1095 with pytest.raises(TypeError,
1096 match=r"out=\.\.\. is only allowed as a keyword argument."):
1097 np.add(1, 2, ...)
1098
1099 with pytest.raises(TypeError,
1100 match=r"out=\.\.\. is only allowed as a keyword argument."):
1101 np.add.reduce(1, (), None, ...)
1102
1103 type_error = r"must use `\.\.\.` as `out=\.\.\.` and not per-operand/in a tuple"
1104 with pytest.raises(TypeError, match=type_error):
1105 np.negative(1, out=(...,))
1106
1107 with pytest.raises(TypeError, match=type_error):
1108 # We only allow out=... not individual args for now
1109 np.divmod(1, 2, out=(np.empty(()), ...))
1110
1111 with pytest.raises(TypeError, match=type_error):
1112 np.add.reduce(1, out=(...,))
1113
1114 def test_axes_argument(self):
1115 # vecdot signature: '(n),(n)->()'
1116 a = np.arange(27.).reshape((3, 3, 3))
1117 b = np.arange(10., 19.).reshape((3, 1, 3))
1118 # basic tests on inputs (outputs tested below with matrix_multiply).
1119 c = np.vecdot(a, b)
1120 assert_array_equal(c, (a * b).sum(-1))
1121 # default
1122 c = np.vecdot(a, b, axes=[(-1,), (-1,), ()])
1123 assert_array_equal(c, (a * b).sum(-1))
1124 # integers ok for single axis.
1125 c = np.vecdot(a, b, axes=[-1, -1, ()])
1126 assert_array_equal(c, (a * b).sum(-1))
1127 # mix fine
1128 c = np.vecdot(a, b, axes=[(-1,), -1, ()])
1129 assert_array_equal(c, (a * b).sum(-1))
1130 # can omit last axis.
1131 c = np.vecdot(a, b, axes=[-1, -1])
1132 assert_array_equal(c, (a * b).sum(-1))
1133 # can pass in other types of integer (with __index__ protocol)
1134 c = np.vecdot(a, b, axes=[np.int8(-1), np.array(-1, dtype=np.int32)])
1135 assert_array_equal(c, (a * b).sum(-1))
1136 # swap some axes
1137 c = np.vecdot(a, b, axes=[0, 0])
1138 assert_array_equal(c, (a * b).sum(0))
1139 c = np.vecdot(a, b, axes=[0, 2])
1140 assert_array_equal(c, (a.transpose(1, 2, 0) * b).sum(-1))
1141 # Check errors for improperly constructed axes arguments.
1142 # should have list.
1143 assert_raises(TypeError, np.vecdot, a, b, axes=-1)
1144 # needs enough elements
1145 assert_raises(ValueError, np.vecdot, a, b, axes=[-1])
1146 # should pass in indices.
1147 assert_raises(TypeError, np.vecdot, a, b, axes=[-1.0, -1.0])
1148 assert_raises(TypeError, np.vecdot, a, b, axes=[(-1.0,), -1])
1149 assert_raises(TypeError, np.vecdot, a, b, axes=[None, 1])
1150 # cannot pass an index unless there is only one dimension
1151 # (output is wrong in this case)
1152 assert_raises(AxisError, np.vecdot, a, b, axes=[-1, -1, -1])
1153 # or pass in generally the wrong number of axes
1154 assert_raises(AxisError, np.vecdot, a, b, axes=[-1, -1, (-1,)])
1155 assert_raises(AxisError, np.vecdot, a, b, axes=[-1, (-2, -1), ()])
1156 # axes need to have same length.
1157 assert_raises(ValueError, np.vecdot, a, b, axes=[0, 1])
1158
1159 # matrix_multiply signature: '(m,n),(n,p)->(m,p)'
1160 mm = umt.matrix_multiply
1161 a = np.arange(12).reshape((2, 3, 2))
1162 b = np.arange(8).reshape((2, 2, 2, 1)) + 1
1163 # Sanity check.
1164 c = mm(a, b)
1165 assert_array_equal(c, np.matmul(a, b))
1166 # Default axes.
1167 c = mm(a, b, axes=[(-2, -1), (-2, -1), (-2, -1)])
1168 assert_array_equal(c, np.matmul(a, b))
1169 # Default with explicit axes.
1170 c = mm(a, b, axes=[(1, 2), (2, 3), (2, 3)])
1171 assert_array_equal(c, np.matmul(a, b))
1172 # swap some axes.
1173 c = mm(a, b, axes=[(0, -1), (1, 2), (-2, -1)])
1174 assert_array_equal(c, np.matmul(a.transpose(1, 0, 2),
1175 b.transpose(0, 3, 1, 2)))
1176 # Default with output array.
1177 c = np.empty((2, 2, 3, 1))
1178 d = mm(a, b, out=c, axes=[(1, 2), (2, 3), (2, 3)])
1179 assert_(c is d)
1180 assert_array_equal(c, np.matmul(a, b))
1181 # Transposed output array
1182 c = np.empty((1, 2, 2, 3))
1183 d = mm(a, b, out=c, axes=[(-2, -1), (-2, -1), (3, 0)])
1184 assert_(c is d)
1185 assert_array_equal(c, np.matmul(a, b).transpose(3, 0, 1, 2))
1186 # Check errors for improperly constructed axes arguments.
1187 # wrong argument
1188 assert_raises(TypeError, mm, a, b, axis=1)
1189 # axes should be list
1190 assert_raises(TypeError, mm, a, b, axes=1)
1191 assert_raises(TypeError, mm, a, b, axes=((-2, -1), (-2, -1), (-2, -1)))
1192 # list needs to have right length
1193 assert_raises(ValueError, mm, a, b, axes=[])
1194 assert_raises(ValueError, mm, a, b, axes=[(-2, -1)])
1195 # list should not contain None, or lists
1196 assert_raises(TypeError, mm, a, b, axes=[None, None, None])
1197 assert_raises(TypeError,
1198 mm, a, b, axes=[[-2, -1], [-2, -1], [-2, -1]])
1199 assert_raises(TypeError,
1200 mm, a, b, axes=[(-2, -1), (-2, -1), [-2, -1]])
