codekingpro/portable-devtools
114k
1""" Test functions for linalg module
2
3"""
4import itertools
5import os
6import subprocess
7import sys
8import textwrap
9import threading
10import traceback
11import warnings
12
13import pytest
14
15import numpy as np
16from numpy import (
17 array,
18 asarray,
19 atleast_2d,
20 cdouble,
21 csingle,
22 dot,
23 double,
24 identity,
25 inf,
26 linalg,
27 matmul,
28 multiply,
29 single,
30)
31from numpy._core import swapaxes
32from numpy.exceptions import AxisError
33from numpy.linalg import LinAlgError, matrix_power, matrix_rank, multi_dot, norm
34from numpy.linalg._linalg import _multi_dot_matrix_chain_order
35from numpy.testing import (
36 HAS_LAPACK64,
37 IS_WASM,
38 NOGIL_BUILD,
39 assert_,
40 assert_allclose,
41 assert_almost_equal,
42 assert_array_equal,
43 assert_equal,
44 assert_raises,
45 assert_raises_regex,
46)
47
48try:
49 import numpy.linalg.lapack_lite
50except ImportError:
51 # May be broken when numpy was built without BLAS/LAPACK present
52 # If so, ensure we don't break the whole test suite - the `lapack_lite`
53 # submodule should be removed, it's only used in two tests in this file.
54 pass
55
56
57def consistent_subclass(out, in_):
58 # For ndarray subclass input, our output should have the same subclass
59 # (non-ndarray input gets converted to ndarray).
60 return type(out) is (type(in_) if isinstance(in_, np.ndarray)
61 else np.ndarray)
62
63
64old_assert_almost_equal = assert_almost_equal
65
66
67def assert_almost_equal(a, b, single_decimal=6, double_decimal=12, **kw):
68 if asarray(a).dtype.type in (single, csingle):
69 decimal = single_decimal
70 else:
71 decimal = double_decimal
72 old_assert_almost_equal(a, b, decimal=decimal, **kw)
73
74
75def get_real_dtype(dtype):
76 return {single: single, double: double,
77 csingle: single, cdouble: double}[dtype]
78
79
80def get_complex_dtype(dtype):
81 return {single: csingle, double: cdouble,
82 csingle: csingle, cdouble: cdouble}[dtype]
83
84
85def get_rtol(dtype):
86 # Choose a safe rtol
87 if dtype in (single, csingle):
88 return 1e-5
89 else:
90 return 1e-11
91
92
93# used to categorize tests
94all_tags = {
95 'square', 'nonsquare', 'hermitian', # mutually exclusive
96 'generalized', 'size-0', 'strided' # optional additions
97}
98
99
100class LinalgCase:
101 def __init__(self, name, a, b, tags=set()):
102 """
103 A bundle of arguments to be passed to a test case, with an identifying
104 name, the operands a and b, and a set of tags to filter the tests
105 """
106 assert_(isinstance(name, str))
107 self.name = name
108 self.a = a
109 self.b = b
110 self.tags = frozenset(tags) # prevent shared tags
111
112 def check(self, do):
113 """
114 Run the function `do` on this test case, expanding arguments
115 """
116 do(self.a, self.b, tags=self.tags)
117
118 def __repr__(self):
119 return f'<LinalgCase: {self.name}>'
120
121
122def apply_tag(tag, cases):
123 """
124 Add the given tag (a string) to each of the cases (a list of LinalgCase
125 objects)
126 """
127 assert tag in all_tags, "Invalid tag"
128 for case in cases:
129 case.tags = case.tags | {tag}
130 return cases
131
132
133#
134# Base test cases
135#
136
137np.random.seed(1234)
138
139CASES = []
140
141# square test cases
142CASES += apply_tag('square', [
143 LinalgCase("single",
144 array([[1., 2.], [3., 4.]], dtype=single),
145 array([2., 1.], dtype=single)),
146 LinalgCase("double",
147 array([[1., 2.], [3., 4.]], dtype=double),
148 array([2., 1.], dtype=double)),
149 LinalgCase("double_2",
150 array([[1., 2.], [3., 4.]], dtype=double),
151 array([[2., 1., 4.], [3., 4., 6.]], dtype=double)),
152 LinalgCase("csingle",
153 array([[1. + 2j, 2 + 3j], [3 + 4j, 4 + 5j]], dtype=csingle),
154 array([2. + 1j, 1. + 2j], dtype=csingle)),
155 LinalgCase("cdouble",
156 array([[1. + 2j, 2 + 3j], [3 + 4j, 4 + 5j]], dtype=cdouble),
157 array([2. + 1j, 1. + 2j], dtype=cdouble)),
158 LinalgCase("cdouble_2",
159 array([[1. + 2j, 2 + 3j], [3 + 4j, 4 + 5j]], dtype=cdouble),
160 array([[2. + 1j, 1. + 2j, 1 + 3j], [1 - 2j, 1 - 3j, 1 - 6j]], dtype=cdouble)),
161 LinalgCase("0x0",
162 np.empty((0, 0), dtype=double),
163 np.empty((0,), dtype=double),
164 tags={'size-0'}),
165 LinalgCase("8x8",
166 np.random.rand(8, 8),
167 np.random.rand(8)),
168 LinalgCase("1x1",
169 np.random.rand(1, 1),
170 np.random.rand(1)),
171 LinalgCase("nonarray",
172 [[1, 2], [3, 4]],
173 [2, 1]),
174])
175
176# non-square test-cases
177CASES += apply_tag('nonsquare', [
178 LinalgCase("single_nsq_1",
179 array([[1., 2., 3.], [3., 4., 6.]], dtype=single),
180 array([2., 1.], dtype=single)),
181 LinalgCase("single_nsq_2",
182 array([[1., 2.], [3., 4.], [5., 6.]], dtype=single),
183 array([2., 1., 3.], dtype=single)),
184 LinalgCase("double_nsq_1",
185 array([[1., 2., 3.], [3., 4., 6.]], dtype=double),
186 array([2., 1.], dtype=double)),
187 LinalgCase("double_nsq_2",
188 array([[1., 2.], [3., 4.], [5., 6.]], dtype=double),
189 array([2., 1., 3.], dtype=double)),
190 LinalgCase("csingle_nsq_1",
191 array(
192 [[1. + 1j, 2. + 2j, 3. - 3j], [3. - 5j, 4. + 9j, 6. + 2j]], dtype=csingle),
193 array([2. + 1j, 1. + 2j], dtype=csingle)),
194 LinalgCase("csingle_nsq_2",
195 array(
196 [[1. + 1j, 2. + 2j], [3. - 3j, 4. - 9j], [5. - 4j, 6. + 8j]], dtype=csingle),
197 array([2. + 1j, 1. + 2j, 3. - 3j], dtype=csingle)),
198 LinalgCase("cdouble_nsq_1",
199 array(
200 [[1. + 1j, 2. + 2j, 3. - 3j], [3. - 5j, 4. + 9j, 6. + 2j]], dtype=cdouble),
201 array([2. + 1j, 1. + 2j], dtype=cdouble)),
202 LinalgCase("cdouble_nsq_2",
203 array(
204 [[1. + 1j, 2. + 2j], [3. - 3j, 4. - 9j], [5. - 4j, 6. + 8j]], dtype=cdouble),
205 array([2. + 1j, 1. + 2j, 3. - 3j], dtype=cdouble)),
206 LinalgCase("cdouble_nsq_1_2",
207 array(
208 [[1. + 1j, 2. + 2j, 3. - 3j], [3. - 5j, 4. + 9j, 6. + 2j]], dtype=cdouble),
209 array([[2. + 1j, 1. + 2j], [1 - 1j, 2 - 2j]], dtype=cdouble)),
210 LinalgCase("cdouble_nsq_2_2",
211 array(
212 [[1. + 1j, 2. + 2j], [3. - 3j, 4. - 9j], [5. - 4j, 6. + 8j]], dtype=cdouble),
213 array([[2. + 1j, 1. + 2j], [1 - 1j, 2 - 2j], [1 - 1j, 2 - 2j]], dtype=cdouble)),
214 LinalgCase("8x11",
215 np.random.rand(8, 11),
216 np.random.rand(8)),
217 LinalgCase("1x5",
218 np.random.rand(1, 5),
219 np.random.rand(1)),
220 LinalgCase("5x1",
221 np.random.rand(5, 1),
222 np.random.rand(5)),
223 LinalgCase("0x4",
224 np.random.rand(0, 4),
225 np.random.rand(0),
226 tags={'size-0'}),
227 LinalgCase("4x0",
228 np.random.rand(4, 0),
229 np.random.rand(4),
230 tags={'size-0'}),
231])
232
233# hermitian test-cases
234CASES += apply_tag('hermitian', [
235 LinalgCase("hsingle",
236 array([[1., 2.], [2., 1.]], dtype=single),
237 None),
238 LinalgCase("hdouble",
239 array([[1., 2.], [2., 1.]], dtype=double),
240 None),
241 LinalgCase("hcsingle",
242 array([[1., 2 + 3j], [2 - 3j, 1]], dtype=csingle),
243 None),
244 LinalgCase("hcdouble",
245 array([[1., 2 + 3j], [2 - 3j, 1]], dtype=cdouble),
246 None),
247 LinalgCase("hempty",
248 np.empty((0, 0), dtype=double),
249 None,
250 tags={'size-0'}),
251 LinalgCase("hnonarray",
252 [[1, 2], [2, 1]],
253 None),
254 LinalgCase("matrix_b_only",
255 array([[1., 2.], [2., 1.]]),
256 None),
257 LinalgCase("hmatrix_1x1",
258 np.random.rand(1, 1),
259 None),
260])
261
262
263#
264# Gufunc test cases
265#
266def _make_generalized_cases():
267 new_cases = []
268
269 for case in CASES:
270 if not isinstance(case.a, np.ndarray):
271 continue
272
273 a = np.array([case.a, 2 * case.a, 3 * case.a])
274 if case.b is None:
275 b = None
276 elif case.b.ndim == 1:
277 b = case.b
278 else:
279 b = np.array([case.b, 7 * case.b, 6 * case.b])
280 new_case = LinalgCase(case.name + "_tile3", a, b,
281 tags=case.tags | {'generalized'})
282 new_cases.append(new_case)
283
284 a = np.array([case.a] * 2 * 3).reshape((3, 2) + case.a.shape)
285 if case.b is None:
286 b = None
287 elif case.b.ndim == 1:
288 b = np.array([case.b] * 2 * 3 * a.shape[-1])\
289 .reshape((3, 2) + case.a.shape[-2:])
290 else:
291 b = np.array([case.b] * 2 * 3).reshape((3, 2) + case.b.shape)
292 new_case = LinalgCase(case.name + "_tile213", a, b,
293 tags=case.tags | {'generalized'})
294 new_cases.append(new_case)
295
296 return new_cases
297
298
299CASES += _make_generalized_cases()
300
301
302#
303# Generate stride combination variations of the above
304#
305def _stride_comb_iter(x):
306 """
307 Generate cartesian product of strides for all axes
308 """
309
310 if not isinstance(x, np.ndarray):
311 yield x, "nop"
312 return
313
314 stride_set = [(1,)] * x.ndim
315 stride_set[-1] = (1, 3, -4)
316 if x.ndim > 1:
317 stride_set[-2] = (1, 3, -4)
318 if x.ndim > 2:
319 stride_set[-3] = (1, -4)
320
321 for repeats in itertools.product(*tuple(stride_set)):
322 new_shape = [abs(a * b) for a, b in zip(x.shape, repeats)]
323 slices = tuple(slice(None, None, repeat) for repeat in repeats)
324
325 # new array with different strides, but same data
326 xi = np.empty(new_shape, dtype=x.dtype)
327 xi.view(np.uint32).fill(0xdeadbeef)
328 xi = xi[slices]
329 xi[...] = x
330 xi = xi.view(x.__class__)
331 assert_(np.all(xi == x))
332 yield xi, "stride_" + "_".join(["%+d" % j for j in repeats])
333
334 # generate also zero strides if possible
335 if x.ndim >= 1 and x.shape[-1] == 1:
336 s = list(x.strides)
337 s[-1] = 0
338 xi = np.lib.stride_tricks.as_strided(x, strides=s)
339 yield xi, "stride_xxx_0"
340 if x.ndim >= 2 and x.shape[-2] == 1:
341 s = list(x.strides)
342 s[-2] = 0
343 xi = np.lib.stride_tricks.as_strided(x, strides=s)
344 yield xi, "stride_xxx_0_x"
345 if x.ndim >= 2 and x.shape[:-2] == (1, 1):
346 s = list(x.strides)
347 s[-1] = 0
348 s[-2] = 0
349 xi = np.lib.stride_tricks.as_strided(x, strides=s)
350 yield xi, "stride_xxx_0_0"
351
352
353def _make_strided_cases():
354 new_cases = []
355 for case in CASES:
356 for a, a_label in _stride_comb_iter(case.a):
357 for b, b_label in _stride_comb_iter(case.b):
358 new_case = LinalgCase(case.name + "_" + a_label + "_" + b_label, a, b,
359 tags=case.tags | {'strided'})
360 new_cases.append(new_case)
361 return new_cases
362
363
364CASES += _make_strided_cases()
365
366
367#
368# Test different routines against the above cases
369#
370class LinalgTestCase:
371 TEST_CASES = CASES
372
373 def check_cases(self, require=set(), exclude=set()):
374 """
375 Run func on each of the cases with all of the tags in require, and none
376 of the tags in exclude
377 """
378 for case in self.TEST_CASES:
379 # filter by require and exclude
380 if case.tags & require != require:
381 continue
382 if case.tags & exclude:
383 continue
384
385 try:
386 case.check(self.do)
387 except Exception as e:
388 msg = f'In test case: {case!r}\n\n'
389 msg += traceback.format_exc()
390 raise AssertionError(msg) from e
391
392
393class LinalgSquareTestCase(LinalgTestCase):
394
395 def test_sq_cases(self):
396 self.check_cases(require={'square'},
397 exclude={'generalized', 'size-0'})
398
399 def test_empty_sq_cases(self):
400 self.check_cases(require={'square', 'size-0'},
401 exclude={'generalized'})
402
403
404class LinalgNonsquareTestCase(LinalgTestCase):
405
406 def test_nonsq_cases(self):
407 self.check_cases(require={'nonsquare'},
408 exclude={'generalized', 'size-0'})
409
410 def test_empty_nonsq_cases(self):
411 self.check_cases(require={'nonsquare', 'size-0'},
412 exclude={'generalized'})
413
414
415class HermitianTestCase(LinalgTestCase):
416
417 def test_herm_cases(self):
418 self.check_cases(require={'hermitian'},
419 exclude={'generalized', 'size-0'})
420
421 def test_empty_herm_cases(self):
422 self.check_cases(require={'hermitian', 'size-0'},
423 exclude={'generalized'})
424
425
426class LinalgGeneralizedSquareTestCase(LinalgTestCase):
427
428 @pytest.mark.slow
429 def test_generalized_sq_cases(self):
430 self.check_cases(require={'generalized', 'square'},
431 exclude={'size-0'})
432
433 @pytest.mark.slow
434 def test_generalized_empty_sq_cases(self):
435 self.check_cases(require={'generalized', 'square', 'size-0'})
436
437
438class LinalgGeneralizedNonsquareTestCase(LinalgTestCase):
439
440 @pytest.mark.slow
441 def test_generalized_nonsq_cases(self):
442 self.check_cases(require={'generalized', 'nonsquare'},
443 exclude={'size-0'})
444
445 @pytest.mark.slow
446 def test_generalized_empty_nonsq_cases(self):
447 self.check_cases(require={'generalized', 'nonsquare', 'size-0'})
448
449
450class HermitianGeneralizedTestCase(LinalgTestCase):
451
452 @pytest.mark.slow
453 def test_generalized_herm_cases(self):
454 self.check_cases(require={'generalized', 'hermitian'},
455 exclude={'size-0'})
456
457 @pytest.mark.slow
458 def test_generalized_empty_herm_cases(self):
459 self.check_cases(require={'generalized', 'hermitian', 'size-0'},
460 exclude={'none'})
461
462
463def identity_like_generalized(a):
464 a = asarray(a)
465 if a.ndim >= 3:
466 r = np.empty(a.shape, dtype=a.dtype)
467 r[...] = identity(a.shape[-2])
468 return r
469 else:
470 return identity(a.shape[0])
471
472
473class SolveCases(LinalgSquareTestCase, LinalgGeneralizedSquareTestCase):
474 # kept apart from TestSolve for use for testing with matrices.
475 def do(self, a, b, tags):
476 x = linalg.solve(a, b)
477 if np.array(b).ndim == 1:
478 # When a is (..., M, M) and b is (M,), it is the same as when b is
479 # (M, 1), except the result has shape (..., M)
480 adotx = matmul(a, x[..., None])[..., 0]
481 assert_almost_equal(np.broadcast_to(b, adotx.shape), adotx)
482 else:
483 adotx = matmul(a, x)
484 assert_almost_equal(b, adotx)
485 assert_(consistent_subclass(x, b))
486
487
488class TestSolve(SolveCases):
489 @pytest.mark.parametrize('dtype', [single, double, csingle, cdouble])
490 def test_types(self, dtype):
491 x = np.array([[1, 0.5], [0.5, 1]], dtype=dtype)
492 assert_equal(linalg.solve(x, x).dtype, dtype)
493
494 def test_1_d(self):
495 class ArraySubclass(np.ndarray):
496 pass
497 a = np.arange(8).reshape(2, 2, 2)
498 b = np.arange(2).view(ArraySubclass)
499 result = linalg.solve(a, b)
500 assert result.shape == (2, 2)
501
502 # If b is anything other than 1-D it should be treated as a stack of
503 # matrices
504 b = np.arange(4).reshape(2, 2).view(ArraySubclass)
505 result = linalg.solve(a, b)
506 assert result.shape == (2, 2, 2)
507
508 b = np.arange(2).reshape(1, 2).view(ArraySubclass)
509 assert_raises(ValueError, linalg.solve, a, b)
510
511 def test_0_size(self):
512 class ArraySubclass(np.ndarray):
513 pass
514 # Test system of 0x0 matrices
515 a = np.arange(8).reshape(2, 2, 2)
516 b = np.arange(6).reshape(1, 2, 3).view(ArraySubclass)
517
518 expected = linalg.solve(a, b)[:, 0:0, :]
519 result = linalg.solve(a[:, 0:0, 0:0], b[:, 0:0, :])
520 assert_array_equal(result, expected)
521 assert_(isinstance(result, ArraySubclass))
522
523 # Test errors for non-square and only b's dimension being 0
524 assert_raises(linalg.LinAlgError, linalg.solve, a[:, 0:0, 0:1], b)
525 assert_raises(ValueError, linalg.solve, a, b[:, 0:0, :])
526
527 # Test broadcasting error
528 b = np.arange(6).reshape(1, 3, 2) # broadcasting error
529 assert_raises(ValueError, linalg.solve, a, b)
530 assert_raises(ValueError, linalg.solve, a[0:0], b[0:0])
531
532 # Test zero "single equations" with 0x0 matrices.
533 b = np.arange(2).view(ArraySubclass)
534 expected = linalg.solve(a, b)[:, 0:0]
535 result = linalg.solve(a[:, 0:0, 0:0], b[0:0])
536 assert_array_equal(result, expected)
537 assert_(isinstance(result, ArraySubclass))
538
539 b = np.arange(3).reshape(1, 3)
540 assert_raises(ValueError, linalg.solve, a, b)
541 assert_raises(ValueError, linalg.solve, a[0:0], b[0:0])
542 assert_raises(ValueError, linalg.solve, a[:, 0:0, 0:0], b)
543
544 def test_0_size_k(self):
545 # test zero multiple equation (K=0) case.
546 class ArraySubclass(np.ndarray):
547 pass
548 a = np.arange(4).reshape(1, 2, 2)
549 b = np.arange(6).reshape(3, 2, 1).view(ArraySubclass)
550
551 expected = linalg.solve(a, b)[:, :, 0:0]
552 result = linalg.solve(a, b[:, :, 0:0])
553 assert_array_equal(result, expected)
554 assert_(isinstance(result, ArraySubclass))
555
556 # test both zero.
557 expected = linalg.solve(a, b)[:, 0:0, 0:0]
558 result = linalg.solve(a[:, 0:0, 0:0], b[:, 0:0, 0:0])
559 assert_array_equal(result, expected)
560 assert_(isinstance(result, ArraySubclass))
561
562
563class InvCases(LinalgSquareTestCase, LinalgGeneralizedSquareTestCase):
564
565 def do(self, a, b, tags):
566 a_inv = linalg.inv(a)
567 assert_almost_equal(matmul(a, a_inv),
568 identity_like_generalized(a))
569 assert_(consistent_subclass(a_inv, a))
570
571
572class TestInv(InvCases):
573 @pytest.mark.parametrize('dtype', [single, double, csingle, cdouble])
574 def test_types(self, dtype):
575 x = np.array([[1, 0.5], [0.5, 1]], dtype=dtype)
576 assert_equal(linalg.inv(x).dtype, dtype)
577
578 def test_0_size(self):
579 # Check that all kinds of 0-sized arrays work
580 class ArraySubclass(np.ndarray):
581 pass
582 a = np.zeros((0, 1, 1), dtype=np.int_).view(ArraySubclass)
583 res = linalg.inv(a)
584 assert_(res.dtype.type is np.float64)
585 assert_equal(a.shape, res.shape)
586 assert_(isinstance(res, ArraySubclass))
587
588 a = np.zeros((0, 0), dtype=np.complex64).view(ArraySubclass)
589 res = linalg.inv(a)
590 assert_(res.dtype.type is np.complex64)
591 assert_equal(a.shape, res.shape)
592 assert_(isinstance(res, ArraySubclass))
593
594
595class EigvalsCases(LinalgSquareTestCase, LinalgGeneralizedSquareTestCase):
596
597 def do(self, a, b, tags):
598 ev = linalg.eigvals(a)
599 evalues, evectors = linalg.eig(a)
600 assert_almost_equal(ev, evalues)
601
602
603class TestEigvals(EigvalsCases):
604 @pytest.mark.parametrize('dtype', [single, double, csingle, cdouble])
605 def test_types(self, dtype):
606 x = np.array([[1, 0.5], [0.5, 1]], dtype=dtype)
607 assert_equal(linalg.eigvals(x).dtype, dtype)
608 x = np.array([[1, 0.5], [-1, 1]], dtype=dtype)
609 assert_equal(linalg.eigvals(x).dtype, get_complex_dtype(dtype))
610
611 def test_0_size(self):
612 # Check that all kinds of 0-sized arrays work
613 class ArraySubclass(np.ndarray):
614 pass
615 a = np.zeros((0, 1, 1), dtype=np.int_).view(ArraySubclass)
616 res = linalg.eigvals(a)
617 assert_(res.dtype.type is np.float64)
618 assert_equal((0, 1), res.shape)
619 # This is just for documentation, it might make sense to change:
620 assert_(isinstance(res, np.ndarray))
621
622 a = np.zeros((0, 0), dtype=np.complex64).view(ArraySubclass)
623 res = linalg.eigvals(a)
624 assert_(res.dtype.type is np.complex64)
625 assert_equal((0,), res.shape)
626 # This is just for documentation, it might make sense to change:
627 assert_(isinstance(res, np.ndarray))
628
629
630class EigCases(LinalgSquareTestCase, LinalgGeneralizedSquareTestCase):
631
632 def do(self, a, b, tags):
633 res = linalg.eig(a)
634 eigenvalues, eigenvectors = res.eigenvalues, res.eigenvectors
635 assert_allclose(matmul(a, eigenvectors),
636 np.asarray(eigenvectors) * np.asarray(eigenvalues)[..., None, :],
637 rtol=get_rtol(eigenvalues.dtype))
638 assert_(consistent_subclass(eigenvectors, a))
639
640
641class TestEig(EigCases):
642 @pytest.mark.parametrize('dtype', [single, double, csingle, cdouble])
643 def test_types(self, dtype):
644 x = np.array([[1, 0.5], [0.5, 1]], dtype=dtype)
645 w, v = np.linalg.eig(x)
646 assert_equal(w.dtype, dtype)
647 assert_equal(v.dtype, dtype)
648
649 x = np.array([[1, 0.5], [-1, 1]], dtype=dtype)
650 w, v = np.linalg.eig(x)
651 assert_equal(w.dtype, get_complex_dtype(dtype))
652 assert_equal(v.dtype, get_complex_dtype(dtype))
653
654 def test_0_size(self):
655 # Check that all kinds of 0-sized arrays work
656 class ArraySubclass(np.ndarray):
657 pass
658 a = np.zeros((0, 1, 1), dtype=np.int_).view(ArraySubclass)
659 res, res_v = linalg.eig(a)
660 assert_(res_v.dtype.type is np.float64)
661 assert_(res.dtype.type is np.float64)
662 assert_equal(a.shape, res_v.shape)
663 assert_equal((0, 1), res.shape)
664 # This is just for documentation, it might make sense to change:
665 assert_(isinstance(a, np.ndarray))
666
667 a = np.zeros((0, 0), dtype=np.complex64).view(ArraySubclass)
668 res, res_v = linalg.eig(a)
669 assert_(res_v.dtype.type is np.complex64)
670 assert_(res.dtype.type is np.complex64)
671 assert_equal(a.shape, res_v.shape)
672 assert_equal((0,), res.shape)
673 # This is just for documentation, it might make sense to change:
674 assert_(isinstance(a, np.ndarray))
675
676
677class SVDBaseTests:
678 hermitian = False
679
680 @pytest.mark.parametrize('dtype', [single, double, csingle, cdouble])
681 def test_types(self, dtype):
682 x = np.array([[1, 0.5], [0.5, 1]], dtype=dtype)
683 res = linalg.svd(x)
684 U, S, Vh = res.U, res.S, res.Vh
685 assert_equal(U.dtype, dtype)
686 assert_equal(S.dtype, get_real_dtype(dtype))
687 assert_equal(Vh.dtype, dtype)
688 s = linalg.svd(x, compute_uv=False, hermitian=self.hermitian)
689 assert_equal(s.dtype, get_real_dtype(dtype))
690
691
692class SVDCases(LinalgSquareTestCase, LinalgGeneralizedSquareTestCase):
693
694 def do(self, a, b, tags):
695 u, s, vt = linalg.svd(a, False)
696 assert_allclose(a, matmul(np.asarray(u) * np.asarray(s)[..., None, :],
697 np.asarray(vt)),
698 rtol=get_rtol(u.dtype))
699 assert_(consistent_subclass(u, a))
700 assert_(consistent_subclass(vt, a))
701
702
703class TestSVD(SVDCases, SVDBaseTests):
704 def test_empty_identity(self):
705 """ Empty input should put an identity matrix in u or vh """
706 x = np.empty((4, 0))
707 u, s, vh = linalg.svd(x, compute_uv=True, hermitian=self.hermitian)
708 assert_equal(u.shape, (4, 4))
709 assert_equal(vh.shape, (0, 0))
710 assert_equal(u, np.eye(4))
711
712 x = np.empty((0, 4))
713 u, s, vh = linalg.svd(x, compute_uv=True, hermitian=self.hermitian)
714 assert_equal(u.shape, (0, 0))
715 assert_equal(vh.shape, (4, 4))
716 assert_equal(vh, np.eye(4))
717
718 def test_svdvals(self):
719 x = np.array([[1, 0.5], [0.5, 1]])
720 s_from_svd = linalg.svd(x, compute_uv=False, hermitian=self.hermitian)
721 s_from_svdvals = linalg.svdvals(x)
722 assert_almost_equal(s_from_svd, s_from_svdvals)
723
724
725class SVDHermitianCases(HermitianTestCase, HermitianGeneralizedTestCase):
726
727 def do(self, a, b, tags):
728 u, s, vt = linalg.svd(a, False, hermitian=True)
729 assert_allclose(a, matmul(np.asarray(u) * np.asarray(s)[..., None, :],
730 np.asarray(vt)),
731 rtol=get_rtol(u.dtype))
732
733 def hermitian(mat):
734 axes = list(range(mat.ndim))
735 axes[-1], axes[-2] = axes[-2], axes[-1]
736 return np.conj(np.transpose(mat, axes=axes))
737
738 assert_almost_equal(np.matmul(u, hermitian(u)), np.broadcast_to(np.eye(u.shape[-1]), u.shape))
739 assert_almost_equal(np.matmul(vt, hermitian(vt)), np.broadcast_to(np.eye(vt.shape[-1]), vt.shape))
740 assert_equal(np.sort(s)[..., ::-1], s)
741 assert_(consistent_subclass(u, a))
742 assert_(consistent_subclass(vt, a))
743
744
745class TestSVDHermitian(SVDHermitianCases, SVDBaseTests):
746 hermitian = True
747
748
749class CondCases(LinalgSquareTestCase, LinalgGeneralizedSquareTestCase):
750 # cond(x, p) for p in (None, 2, -2)
751
752 def do(self, a, b, tags):
753 c = asarray(a) # a might be a matrix
754 if 'size-0' in tags:
755 assert_raises(LinAlgError, linalg.cond, c)
756 return
757
758 # +-2 norms
759 s = linalg.svd(c, compute_uv=False)
760 assert_almost_equal(
761 linalg.cond(a), s[..., 0] / s[..., -1],
762 single_decimal=5, double_decimal=11)
763 assert_almost_equal(
764 linalg.cond(a, 2), s[..., 0] / s[..., -1],
765 single_decimal=5, double_decimal=11)
766 assert_almost_equal(
767 linalg.cond(a, -2), s[..., -1] / s[..., 0],
768 single_decimal=5, double_decimal=11)
769
770 # Other norms
771 cinv = np.linalg.inv(c)
772 assert_almost_equal(
773 linalg.cond(a, 1),
774 abs(c).sum(-2).max(-1) * abs(cinv).sum(-2).max(-1),
775 single_decimal=5, double_decimal=11)
776 assert_almost_equal(
777 linalg.cond(a, -1),
778 abs(c).sum(-2).min(-1) * abs(cinv).sum(-2).min(-1),
779 single_decimal=5, double_decimal=11)
780 assert_almost_equal(
781 linalg.cond(a, np.inf),
782 abs(c).sum(-1).max(-1) * abs(cinv).sum(-1).max(-1),
783 single_decimal=5, double_decimal=11)
784 assert_almost_equal(
785 linalg.cond(a, -np.inf),
786 abs(c).sum(-1).min(-1) * abs(cinv).sum(-1).min(-1),
787 single_decimal=5, double_decimal=11)
788 assert_almost_equal(
789 linalg.cond(a, 'fro'),
790 np.sqrt((abs(c)**2).sum(-1).sum(-1)
791 * (abs(cinv)**2).sum(-1).sum(-1)),
792 single_decimal=5, double_decimal=11)
793
794
795class TestCond(CondCases):
796 @pytest.mark.parametrize('is_complex', [False, True])
797 def test_basic_nonsvd(self, is_complex):
798 # Smoketest the non-svd norms
799 A = array([[1., 0, 1], [0, -2., 0], [0, 0, 3.]])
800 if is_complex:
801 # Since A is linearly scaled, the condition number should not change
802 A = A * (1 + 1j)
803 assert_almost_equal(linalg.cond(A, inf), 4)
804 assert_almost_equal(linalg.cond(A, -inf), 2 / 3)
805 assert_almost_equal(linalg.cond(A, 1), 4)
806 assert_almost_equal(linalg.cond(A, -1), 0.5)
807 assert_almost_equal(linalg.cond(A, 'fro'), np.sqrt(265 / 12))
808
809 @pytest.mark.parametrize('dtype', [single, double, csingle, cdouble])
810 @pytest.mark.parametrize('norm_ord', [1, -1, 2, -2, 'fro', np.inf, -np.inf])
811 def test_cond_dtypes(self, dtype, norm_ord):
812 # Check that the condition number is computed in the same dtype
813 # as the input matrix
814 A = array([[1., 0, 1], [0, -2., 0], [0, 0, 3.]], dtype=dtype)
815 out_type = get_real_dtype(dtype)
816 assert_equal(linalg.cond(A, p=norm_ord).dtype, out_type)
817
818 def test_singular(self):
819 # Singular matrices have infinite condition number for
820 # positive norms, and negative norms shouldn't raise
821 # exceptions
822 As = [np.zeros((2, 2)), np.ones((2, 2))]
823 p_pos = [None, 1, 2, 'fro']
824 p_neg = [-1, -2]
825 for A, p in itertools.product(As, p_pos):
826 # Inversion may not hit exact infinity, so just check the
827 # number is large
828 assert_(linalg.cond(A, p) > 1e15)
829 for A, p in itertools.product(As, p_neg):
830 linalg.cond(A, p)
831
832 @pytest.mark.xfail(True, run=False,
833 reason="Platform/LAPACK-dependent failure, "
834 "see gh-18914")
835 def test_nan(self):
836 # nans should be passed through, not converted to infs
837 ps = [None, 1, -1, 2, -2, 'fro']
838 p_pos = [None, 1, 2, 'fro']
839
840 A = np.ones((2, 2))
841 A[0, 1] = np.nan
842 for p in ps:
843 c = linalg.cond(A, p)
844 assert_(isinstance(c, np.float64))
845 assert_(np.isnan(c))
846
847 A = np.ones((3, 2, 2))
848 A[1, 0, 1] = np.nan
849 for p in ps:
850 c = linalg.cond(A, p)
851 assert_(np.isnan(c[1]))
852 if p in p_pos:
853 assert_(c[0] > 1e15)
854 assert_(c[2] > 1e15)
855 else:
856 assert_(not np.isnan(c[0]))
857 assert_(not np.isnan(c[2]))
858
859 def test_stacked_singular(self):
860 # Check behavior when only some of the stacked matrices are
861 # singular
862 np.random.seed(1234)
863 A = np.random.rand(2, 2, 2, 2)
864 A[0, 0] = 0
865 A[1, 1] = 0
866
867 for p in (None, 1, 2, 'fro', -1, -2):
868 c = linalg.cond(A, p)
869 assert_equal(c[0, 0], np.inf)
870 assert_equal(c[1, 1], np.inf)
871 assert_(np.isfinite(c[0, 1]))
872 assert_(np.isfinite(c[1, 0]))
873
874
875class PinvCases(LinalgSquareTestCase,
876 LinalgNonsquareTestCase,
877 LinalgGeneralizedSquareTestCase,
878 LinalgGeneralizedNonsquareTestCase):
879
880 def do(self, a, b, tags):
881 a_ginv = linalg.pinv(a)
882 # `a @ a_ginv == I` does not hold if a is singular
883 dot = matmul
884 assert_almost_equal(dot(dot(a, a_ginv), a), a, single_decimal=5, double_decimal=11)
885 assert_(consistent_subclass(a_ginv, a))
886
887
888class TestPinv(PinvCases):
889 pass
890
891
892class PinvHermitianCases(HermitianTestCase, HermitianGeneralizedTestCase):
893
894 def do(self, a, b, tags):
895 a_ginv = linalg.pinv(a, hermitian=True)
896 # `a @ a_ginv == I` does not hold if a is singular
897 dot = matmul
898 assert_almost_equal(dot(dot(a, a_ginv), a), a, single_decimal=5, double_decimal=11)
899 assert_(consistent_subclass(a_ginv, a))
900
901
902class TestPinvHermitian(PinvHermitianCases):
903 pass
904
905
906def test_pinv_rtol_arg():
907 a = np.array([[1, 2, 3], [4, 1, 1], [2, 3, 1]])
908
909 assert_almost_equal(
910 np.linalg.pinv(a, rcond=0.5),
911 np.linalg.pinv(a, rtol=0.5),
912 )
913
914 with pytest.raises(
915 ValueError, match=r"`rtol` and `rcond` can't be both set."
916 ):
917 np.linalg.pinv(a, rcond=0.5, rtol=0.5)
918
919
920class DetCases(LinalgSquareTestCase, LinalgGeneralizedSquareTestCase):
921
922 def do(self, a, b, tags):
923 d = linalg.det(a)
924 res = linalg.slogdet(a)
925 s, ld = res.sign, res.logabsdet
926 if asarray(a).dtype.type in (single, double):
927 ad = asarray(a).astype(double)
928 else:
929 ad = asarray(a).astype(cdouble)
930 ev = linalg.eigvals(ad)
931 assert_almost_equal(d, multiply.reduce(ev, axis=-1))
932 assert_almost_equal(s * np.exp(ld), multiply.reduce(ev, axis=-1))
933
934 s = np.atleast_1d(s)
935 ld = np.atleast_1d(ld)
936 m = (s != 0)
937 assert_almost_equal(np.abs(s[m]), 1)
938 assert_equal(ld[~m], -inf)
939
940
941class TestDet(DetCases):
942 def test_zero(self):
943 assert_equal(linalg.det([[0.0]]), 0.0)
944 assert_equal(type(linalg.det([[0.0]])), double)
945 assert_equal(linalg.det([[0.0j]]), 0.0)
946 assert_equal(type(linalg.det([[0.0j]])), cdouble)
947
948 assert_equal(linalg.slogdet([[0.0]]), (0.0, -inf))
949 assert_equal(type(linalg.slogdet([[0.0]])[0]), double)
950 assert_equal(type(linalg.slogdet([[0.0]])[1]), double)
951 assert_equal(linalg.slogdet([[0.0j]]), (0.0j, -inf))
952 assert_equal(type(linalg.slogdet([[0.0j]])[0]), cdouble)
953 assert_equal(type(linalg.slogdet([[0.0j]])[1]), double)
954
955 @pytest.mark.parametrize('dtype', [single, double, csingle, cdouble])
956 def test_types(self, dtype):
957 x = np.array([[1, 0.5], [0.5, 1]], dtype=dtype)
958 assert_equal(np.linalg.det(x).dtype, dtype)
959 ph, s = np.linalg.slogdet(x)
960 assert_equal(s.dtype, get_real_dtype(dtype))
961 assert_equal(ph.dtype, dtype)
962
963 def test_0_size(self):
964 a = np.zeros((0, 0), dtype=np.complex64)
965 res = linalg.det(a)
966 assert_equal(res, 1.)
967 assert_(res.dtype.type is np.complex64)
968 res = linalg.slogdet(a)
969 assert_equal(res, (1, 0))
970 assert_(res[0].dtype.type is np.complex64)
971 assert_(res[1].dtype.type is np.float32)
972
973 a = np.zeros((0, 0), dtype=np.float64)
974 res = linalg.det(a)
975 assert_equal(res, 1.)
976 assert_(res.dtype.type is np.float64)
977 res = linalg.slogdet(a)
978 assert_equal(res, (1, 0))
979 assert_(res[0].dtype.type is np.float64)
980 assert_(res[1].dtype.type is np.float64)
981
982
983class LstsqCases(LinalgSquareTestCase, LinalgNonsquareTestCase):
984
985 def do(self, a, b, tags):
986 arr = np.asarray(a)
987 m, n = arr.shape
988 u, s, vt = linalg.svd(a, False)
989 x, residuals, rank, sv = linalg.lstsq(a, b, rcond=-1)
990 if m == 0:
991 assert_((x == 0).all())
992 if m <= n:
993 assert_almost_equal(b, dot(a, x))
994 assert_equal(rank, m)
995 else:
996 assert_equal(rank, n)
997 assert_almost_equal(sv, sv.__array_wrap__(s))
998 if rank == n and m > n:
999 expect_resids = (
1000 np.asarray(abs(np.dot(a, x) - b)) ** 2).sum(axis=0)
1001 expect_resids = np.asarray(expect_resids)
1002 if np.asarray(b).ndim == 1:
1003 expect_resids.shape = (1,)
1004 assert_equal(residuals.shape, expect_resids.shape)
1005 else:
1006 expect_resids = np.array([]).view(type(x))
1007 assert_almost_equal(residuals, expect_resids)
1008 assert_(np.issubdtype(residuals.dtype, np.floating))
1009 assert_(consistent_subclass(x, b))
1010 assert_(consistent_subclass(residuals, b))
1011
1012
1013class TestLstsq(LstsqCases):
1014 def test_rcond(self):
1015 a = np.array([[0., 1., 0., 1., 2., 0.],
1016 [0., 2., 0., 0., 1., 0.],
1017 [1., 0., 1., 0., 0., 4.],
1018 [0., 0., 0., 2., 3., 0.]]).T
1019
1020 b = np.array([1, 0, 0, 0, 0, 0])
1021
1022 x, residuals, rank, s = linalg.lstsq(a, b, rcond=-1)
1023 assert_(rank == 4)
1024 x, residuals, rank, s = linalg.lstsq(a, b)
1025 assert_(rank == 3)
1026 x, residuals, rank, s = linalg.lstsq(a, b, rcond=None)
1027 assert_(rank == 3)
1028
1029 @pytest.mark.parametrize(["m", "n", "n_rhs"], [
1030 (4, 2, 2),
1031 (0, 4, 1),
1032 (0, 4, 2),
1033 (4, 0, 1),
1034 (4, 0, 2),
1035 (4, 2, 0),
1036 (0, 0, 0)
1037 ])
1038 def test_empty_a_b(self, m, n, n_rhs):
1039 a = np.arange(m * n).reshape(m, n)
1040 b = np.ones((m, n_rhs))
1041 x, residuals, rank, s = linalg.lstsq(a, b, rcond=None)
1042 if m == 0:
1043 assert_((x == 0).all())
1044 assert_equal(x.shape, (n, n_rhs))
1045 assert_equal(residuals.shape, ((n_rhs,) if m > n else (0,)))
1046 if m > n and n_rhs > 0:
1047 # residuals are exactly the squared norms of b's columns
1048 r = b - np.dot(a, x)
1049 assert_almost_equal(residuals, (r * r).sum(axis=-2))
1050 assert_equal(rank, min(m, n))
1051 assert_equal(s.shape, (min(m, n),))
1052
1053 def test_incompatible_dims(self):
1054 # use modified version of docstring example
1055 x = np.array([0, 1, 2, 3])
1056 y = np.array([-1, 0.2, 0.9, 2.1, 3.3])
1057 A = np.vstack([x, np.ones(len(x))]).T
1058 with assert_raises_regex(LinAlgError, "Incompatible dimensions"):
1059 linalg.lstsq(A, y, rcond=None)
1060
1061
1062@pytest.mark.parametrize('dt', [np.dtype(c) for c in '?bBhHiIqQefdgFDGO'])
1063class TestMatrixPower:
1064
1065 rshft_0 = np.eye(4)
1066 rshft_1 = rshft_0[[3, 0, 1, 2]]
1067 rshft_2 = rshft_0[[2, 3, 0, 1]]
1068 rshft_3 = rshft_0[[1, 2, 3, 0]]
1069 rshft_all = [rshft_0, rshft_1, rshft_2, rshft_3]
1070 noninv = array([[1, 0], [0, 0]])
1071 stacked = np.block([[[rshft_0]]] * 2)
1072 # FIXME the 'e' dtype might work in future
1073 dtnoinv = [object, np.dtype('e'), np.dtype('g'), np.dtype('G')]
1074
1075 def test_large_power(self, dt):
1076 rshft = self.rshft_1.astype(dt)
1077 assert_equal(
1078 matrix_power(rshft, 2**100 + 2**10 + 2**5 + 0), self.rshft_0)
1079 assert_equal(
1080 matrix_power(rshft, 2**100 + 2**10 + 2**5 + 1), self.rshft_1)
1081 assert_equal(
1082 matrix_power(rshft, 2**100 + 2**10 + 2**5 + 2), self.rshft_2)
1083 assert_equal(
1084 matrix_power(rshft, 2**100 + 2**10 + 2**5 + 3), self.rshft_3)
1085
1086 def test_power_is_zero(self, dt):
1087 def tz(M):
1088 mz = matrix_power(M, 0)
1089 assert_equal(mz, identity_like_generalized(M))
1090 assert_equal(mz.dtype, M.dtype)
1091
1092 for mat in self.rshft_all:
1093 tz(mat.astype(dt))
1094 if dt != object:
1095 tz(self.stacked.astype(dt))
1096
1097 def test_power_is_one(self, dt):
1098 def tz(mat):
1099 mz = matrix_power(mat, 1)
1100 assert_equal(mz, mat)
1101 assert_equal(mz.dtype, mat.dtype)
1102
1103 for mat in self.rshft_all:
1104 tz(mat.astype(dt))
1105 if dt != object:
1106 tz(self.stacked.astype(dt))
1107
1108 def test_power_is_two(self, dt):
1109 def tz(mat):
1110 mz = matrix_power(mat, 2)
1111 mmul = matmul if mat.dtype != object else dot
1112 assert_equal(mz, mmul(mat, mat))
1113 assert_equal(mz.dtype, mat.dtype)
1114
1115 for mat in self.rshft_all:
1116 tz(mat.astype(dt))
1117 if dt != object:
1118 tz(self.stacked.astype(dt))
1119
1120 def test_power_is_minus_one(self, dt):
1121 def tz(mat):
1122 invmat = matrix_power(mat, -1)
1123 mmul = matmul if mat.dtype != object else dot
1124 assert_almost_equal(
1125 mmul(invmat, mat), identity_like_generalized(mat))
1126
1127 for mat in self.rshft_all:
1128 if dt not in self.dtnoinv:
1129 tz(mat.astype(dt))
1130
1131 def test_exceptions_bad_power(self, dt):
1132 mat = self.rshft_0.astype(dt)
1133 assert_raises(TypeError, matrix_power, mat, 1.5)
1134 assert_raises(TypeError, matrix_power, mat, [1])
1135
1136 def test_exceptions_non_square(self, dt):
1137 assert_raises(LinAlgError, matrix_power, np.array([1], dt), 1)
1138 assert_raises(LinAlgError, matrix_power, np.array([[1], [2]], dt), 1)
1139 assert_raises(LinAlgError, matrix_power, np.ones((4, 3, 2), dt), 1)
1140
1141 @pytest.mark.skipif(IS_WASM, reason="fp errors don't work in wasm")
1142 def test_exceptions_not_invertible(self, dt):
1143 if dt in self.dtnoinv:
1144 return
1145 mat = self.noninv.astype(dt)
1146 assert_raises(LinAlgError, matrix_power, mat, -1)
1147
1148
1149class TestEigvalshCases(HermitianTestCase, HermitianGeneralizedTestCase):
1150
1151 def do(self, a, b, tags):
1152 # note that eigenvalue arrays returned by eig must be sorted since
1153 # their order isn't guaranteed.
1154 ev = linalg.eigvalsh(a, 'L')
1155 evalues, evectors = linalg.eig(a)
1156 evalues.sort(axis=-1)
1157 assert_allclose(ev, evalues, rtol=get_rtol(ev.dtype))
1158
1159 ev2 = linalg.eigvalsh(a, 'U')
1160 assert_allclose(ev2, evalues, rtol=get_rtol(ev.dtype))
1161
1162
1163class TestEigvalsh:
1164 @pytest.mark.parametrize('dtype', [single, double, csingle, cdouble])
1165 def test_types(self, dtype):
1166 x = np.array([[1, 0.5], [0.5, 1]], dtype=dtype)
1167 w = np.linalg.eigvalsh(x)
1168 assert_equal(w.dtype, get_real_dtype(dtype))
1169
1170 def test_invalid(self):
1171 x = np.array([[1, 0.5], [0.5, 1]], dtype=np.float32)
1172 assert_raises(ValueError, np.linalg.eigvalsh, x, UPLO="lrong")
1173 assert_raises(ValueError, np.linalg.eigvalsh, x, "lower")
1174 assert_raises(ValueError, np.linalg.eigvalsh, x, "upper")
1175
1176 def test_UPLO(self):
1177 Klo = np.array([[0, 0], [1, 0]], dtype=np.double)
1178 Kup = np.array([[0, 1], [0, 0]], dtype=np.double)
1179 tgt = np.array([-1, 1], dtype=np.double)
1180 rtol = get_rtol(np.double)
1181
1182 # Check default is 'L'
1183 w = np.linalg.eigvalsh(Klo)
1184 assert_allclose(w, tgt, rtol=rtol)
1185 # Check 'L'
1186 w = np.linalg.eigvalsh(Klo, UPLO='L')
1187 assert_allclose(w, tgt, rtol=rtol)
1188 # Check 'l'
1189 w = np.linalg.eigvalsh(Klo, UPLO='l')
1190 assert_allclose(w, tgt, rtol=rtol)
1191 # Check 'U'
1192 w = np.linalg.eigvalsh(Kup, UPLO='U')
1193 assert_allclose(w, tgt, rtol=rtol)
1194 # Check 'u'
1195 w = np.linalg.eigvalsh(Kup, UPLO='u')
1196 assert_allclose(w, tgt, rtol=rtol)
1197
1198 def test_0_size(self):
1199 # Check that all kinds of 0-sized arrays work
1200 class ArraySubclass(np.ndarray):
