diff options
| author | Toshiki Kataoka <kataoka@preferred.jp> | 2022-04-08 07:40:12 +0900 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2022-04-07 15:40:12 -0700 |
| commit | 0d13f9f747887b290108a909dd92c3cb47239921 (patch) | |
| tree | 3d8736b1cf2b8c42b16bb384b9ab70af710eee7c /numpy/linalg/tests | |
| parent | b43367e8b59dcaa81dba7cba2df7672c807817b1 (diff) | |
| download | numpy-0d13f9f747887b290108a909dd92c3cb47239921.tar.gz | |
BUG: Consistent promotion for norm for all values of ord (#17709)
Previously, numpy.linalg.norm would return values with the same floating-point
type as input arrays for most values of the ``ord`` parameter, but not all.
This PR fixes this so that the output dtype matches the input for all (valid) values
of ``ord``.
Co-authored-by: Kenichi Maehashi <webmaster@kenichimaehashi.com>
Co-authored-by: Ross Barnowski <rossbar@berkeley.edu>
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) |
