diff options
Diffstat (limited to 'numpy/linalg')
| -rw-r--r-- | numpy/linalg/linalg.py | 5 | ||||
| -rw-r--r-- | numpy/linalg/tests/test_linalg.py | 46 |
2 files changed, 30 insertions, 21 deletions
diff --git a/numpy/linalg/linalg.py b/numpy/linalg/linalg.py index 36c5eb85c..7ce854197 100644 --- a/numpy/linalg/linalg.py +++ b/numpy/linalg/linalg.py @@ -24,7 +24,8 @@ from numpy.core import ( add, multiply, sqrt, fastCopyAndTranspose, sum, isfinite, finfo, errstate, geterrobj, moveaxis, amin, amax, product, abs, atleast_2d, intp, asanyarray, object_, matmul, - swapaxes, divide, count_nonzero, isnan, sign, argsort, sort + swapaxes, divide, count_nonzero, isnan, sign, argsort, sort, + reciprocal ) from numpy.core.multiarray import normalize_axis_index from numpy.core.overrides import set_module @@ -2561,7 +2562,7 @@ def norm(x, ord=None, axis=None, keepdims=False): absx = abs(x) absx **= ord ret = add.reduce(absx, axis=axis, keepdims=keepdims) - ret **= (1 / ord) + ret **= reciprocal(ord, dtype=ret.dtype) return ret elif len(axis) == 2: row_axis, col_axis = axis 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) |
