summaryrefslogtreecommitdiff
path: root/numpy/linalg/tests
diff options
context:
space:
mode:
authorToshiki Kataoka <kataoka@preferred.jp>2022-04-08 07:40:12 +0900
committerGitHub <noreply@github.com>2022-04-07 15:40:12 -0700
commit0d13f9f747887b290108a909dd92c3cb47239921 (patch)
tree3d8736b1cf2b8c42b16bb384b9ab70af710eee7c /numpy/linalg/tests
parentb43367e8b59dcaa81dba7cba2df7672c807817b1 (diff)
downloadnumpy-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.py46
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)