Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_linalg.py2443 linesDownload Raw Back to tests
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):

Showing the first 1,200 of 2443 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai