diff options
Diffstat (limited to 'numpy/linalg/tests')
| -rw-r--r-- | numpy/linalg/tests/test_linalg.py | 46 |
1 files changed, 27 insertions, 19 deletions
diff --git a/numpy/linalg/tests/test_linalg.py b/numpy/linalg/tests/test_linalg.py index 5f9f3b920..f27b4be7f 100644 --- a/numpy/linalg/tests/test_linalg.py +++ b/numpy/linalg/tests/test_linalg.py @@ -1233,6 +1233,14 @@ class _TestNormBase: dt = None dec = None + @staticmethod + def check_dtype(x, res): + if issubclass(x.dtype.type, np.inexact): + assert_equal(res.dtype, x.real.dtype) + else: + # For integer input, don't have to test float precision of output. + assert_(issubclass(res.dtype.type, np.floating)) + class _TestNormGeneral(_TestNormBase): @@ -1249,37 +1257,37 @@ class _TestNormGeneral(_TestNormBase): all_types = exact_types + inexact_types - for each_inexact_types in all_types: - at = a.astype(each_inexact_types) + for each_type in all_types: + at = a.astype(each_type) an = norm(at, -np.inf) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 0.0) with suppress_warnings() as sup: sup.filter(RuntimeWarning, "divide by zero encountered") an = norm(at, -1) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 0.0) an = norm(at, 0) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 2) an = norm(at, 1) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 2.0) an = norm(at, 2) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, an.dtype.type(2.0)**an.dtype.type(1.0/2.0)) an = norm(at, 4) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, an.dtype.type(2.0)**an.dtype.type(1.0/4.0)) an = norm(at, np.inf) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 1.0) def test_vector(self): @@ -1412,41 +1420,41 @@ class _TestNorm2D(_TestNormBase): all_types = exact_types + inexact_types - for each_inexact_types in all_types: - at = a.astype(each_inexact_types) + for each_type in all_types: + at = a.astype(each_type) an = norm(at, -np.inf) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 2.0) with suppress_warnings() as sup: sup.filter(RuntimeWarning, "divide by zero encountered") an = norm(at, -1) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 1.0) an = norm(at, 1) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 2.0) an = norm(at, 2) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 3.0**(1.0/2.0)) an = norm(at, -2) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 1.0) an = norm(at, np.inf) - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 2.0) an = norm(at, 'fro') - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) assert_almost_equal(an, 2.0) an = norm(at, 'nuc') - assert_(issubclass(an.dtype.type, np.floating)) + self.check_dtype(at, an) # Lower bar needed to support low precision floats. # They end up being off by 1 in the 7th place. np.testing.assert_almost_equal(an, 2.7320508075688772, decimal=6) |
