codekingpro/portable-devtools
114k
1""" Test functions for linalg module using the matrix class."""
2import pytest
3
4import numpy as np
5from numpy.linalg.tests.test_linalg import (
6 CondCases,
7 DetCases,
8 EigCases,
9 EigvalsCases,
10 InvCases,
11 LinalgCase,
12 LinalgTestCase,
13 LstsqCases,
14 PinvCases,
15 SolveCases,
16 SVDCases,
17 TestQR as _TestQR,
18 _TestNorm2D,
19 _TestNormDoubleBase,
20 _TestNormInt64Base,
21 _TestNormSingleBase,
22 apply_tag,
23)
24
25CASES = []
26
27# square test cases
28CASES += apply_tag('square', [
29 LinalgCase("0x0_matrix",
30 np.empty((0, 0), dtype=np.double).view(np.matrix),
31 np.empty((0, 1), dtype=np.double).view(np.matrix),
32 tags={'size-0'}),
33 LinalgCase("matrix_b_only",
34 np.array([[1., 2.], [3., 4.]]),
35 np.matrix([2., 1.]).T),
36 LinalgCase("matrix_a_and_b",
37 np.matrix([[1., 2.], [3., 4.]]),
38 np.matrix([2., 1.]).T),
39])
40
41# hermitian test-cases
42CASES += apply_tag('hermitian', [
43 LinalgCase("hmatrix_a_and_b",
44 np.matrix([[1., 2.], [2., 1.]]),
45 None),
46])
47# No need to make generalized or strided cases for matrices.
48
49
50class MatrixTestCase(LinalgTestCase):
51 TEST_CASES = CASES
52
53
54class TestSolveMatrix(SolveCases, MatrixTestCase):
55 pass
56
57
58class TestInvMatrix(InvCases, MatrixTestCase):
59 pass
60
61
62class TestEigvalsMatrix(EigvalsCases, MatrixTestCase):
63 pass
64
65
66class TestEigMatrix(EigCases, MatrixTestCase):
67 pass
68
69
70class TestSVDMatrix(SVDCases, MatrixTestCase):
71 pass
72
73
74class TestCondMatrix(CondCases, MatrixTestCase):
75 pass
76
77
78class TestPinvMatrix(PinvCases, MatrixTestCase):
79 pass
80
81
82class TestDetMatrix(DetCases, MatrixTestCase):
83 pass
84
85
86@pytest.mark.thread_unsafe(
87 reason="residuals not calculated properly for square tests (gh-29851)"
88)
89class TestLstsqMatrix(LstsqCases, MatrixTestCase):
90 pass
91
92
93class _TestNorm2DMatrix(_TestNorm2D):
94 array = np.matrix
95
96
97class TestNormDoubleMatrix(_TestNorm2DMatrix, _TestNormDoubleBase):
98 pass
99
100
101class TestNormSingleMatrix(_TestNorm2DMatrix, _TestNormSingleBase):
102 pass
103
104
105class TestNormInt64Matrix(_TestNorm2DMatrix, _TestNormInt64Base):
106 pass
107
108
109class TestQRMatrix(_TestQR):
110 array = np.matrix
111 