summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorEric Wieser <wieser.eric@gmail.com>2017-02-09 18:51:08 +0000
committerEric Wieser <wieser.eric@gmail.com>2017-02-20 22:03:10 +0000
commit8d6ec65c925ebef5e0567708de1d16df39077c9d (patch)
tree89eb436cf6d453f807becfceef773f053b0ba146
parenteb642f1b6533fcd92e366377f5859e4ea56d5eed (diff)
downloadnumpy-8d6ec65c925ebef5e0567708de1d16df39077c9d.tar.gz
MAINT: Be specific about where AxisError is raised
These were tested by temporarily removing the base classes from AxisError
-rw-r--r--numpy/core/numeric.py4
-rw-r--r--numpy/core/tests/test_multiarray.py16
-rw-r--r--numpy/core/tests/test_numeric.py16
-rw-r--r--numpy/core/tests/test_shape_base.py8
-rw-r--r--numpy/core/tests/test_ufunc.py8
-rw-r--r--numpy/lib/tests/test_function_base.py4
-rw-r--r--numpy/linalg/tests/test_linalg.py4
-rw-r--r--numpy/ma/core.py6
-rw-r--r--numpy/ma/tests/test_core.py8
9 files changed, 37 insertions, 37 deletions
diff --git a/numpy/core/numeric.py b/numpy/core/numeric.py
index e7307a870..066697f3e 100644
--- a/numpy/core/numeric.py
+++ b/numpy/core/numeric.py
@@ -1532,7 +1532,7 @@ def rollaxis(a, axis, start=0):
start += n
msg = "'%s' arg requires %d <= %s < %d, but %d was passed in"
if not (0 <= start < n + 1):
- raise IndexError(msg % ('start', -n, 'start', n + 1, start))
+ raise AxisError(msg % ('start', -n, 'start', n + 1, start))
if axis < start:
# it's been removed
start -= 1
@@ -1551,7 +1551,7 @@ def _validate_axis(axis, ndim, argname):
axis = list(axis)
axis = [a + ndim if a < 0 else a for a in axis]
if not builtins.all(0 <= a < ndim for a in axis):
- raise IndexError('invalid axis for this array in `%s` argument' %
+ raise AxisError('invalid axis for this array in `%s` argument' %
argname)
if len(set(axis)) != len(axis):
raise ValueError('repeated axis in `%s` argument' % argname)
diff --git a/numpy/core/tests/test_multiarray.py b/numpy/core/tests/test_multiarray.py
index 48d532ab9..fa5051ba7 100644
--- a/numpy/core/tests/test_multiarray.py
+++ b/numpy/core/tests/test_multiarray.py
@@ -2013,13 +2013,13 @@ class TestMethods(TestCase):
d = np.array([2, 1])
d.partition(0, kind=k)
assert_raises(ValueError, d.partition, 2)
- assert_raises(IndexError, d.partition, 3, axis=1)
+ assert_raises(np.AxisError, d.partition, 3, axis=1)
assert_raises(ValueError, np.partition, d, 2)
- assert_raises(IndexError, np.partition, d, 2, axis=1)
+ assert_raises(np.AxisError, np.partition, d, 2, axis=1)
assert_raises(ValueError, d.argpartition, 2)
- assert_raises(IndexError, d.argpartition, 3, axis=1)
+ assert_raises(np.AxisError, d.argpartition, 3, axis=1)
assert_raises(ValueError, np.argpartition, d, 2)
- assert_raises(IndexError, np.argpartition, d, 2, axis=1)
+ assert_raises(np.AxisError, np.argpartition, d, 2, axis=1)
d = np.arange(10).reshape((2, 5))
d.partition(1, axis=0, kind=k)
d.partition(4, axis=1, kind=k)
@@ -3522,8 +3522,8 @@ class TestArgmin(TestCase):
class TestMinMax(TestCase):
def test_scalar(self):
- assert_raises(IndexError, np.amax, 1, 1)
- assert_raises(IndexError, np.amin, 1, 1)
+ assert_raises(np.AxisError, np.amax, 1, 1)
+ assert_raises(np.AxisError, np.amin, 1, 1)
assert_equal(np.amax(1, axis=0), 1)
assert_equal(np.amin(1, axis=0), 1)
@@ -3531,7 +3531,7 @@ class TestMinMax(TestCase):
assert_equal(np.amin(1, axis=None), 1)
def test_axis(self):
- assert_raises(IndexError, np.amax, [1, 2, 3], 1000)
+ assert_raises(np.AxisError, np.amax, [1, 2, 3], 1000)
assert_equal(np.amax([[1, 2, 3]], axis=1), 3)
def test_datetime(self):
@@ -3793,7 +3793,7 @@ class TestLexsort(TestCase):
def test_invalid_axis(self): # gh-7528
x = np.linspace(0., 1., 42*3).reshape(42, 3)
- assert_raises(IndexError, np.lexsort, x, axis=2)
+ assert_raises(np.AxisError, np.lexsort, x, axis=2)
class TestIO(object):
"""Test tofile, fromfile, tobytes, and fromstring"""
diff --git a/numpy/core/tests/test_numeric.py b/numpy/core/tests/test_numeric.py
index d7b3d82e8..906280e15 100644
--- a/numpy/core/tests/test_numeric.py
+++ b/numpy/core/tests/test_numeric.py
@@ -1010,7 +1010,7 @@ class TestNonzero(TestCase):
assert_raises(ValueError, np.count_nonzero, m, axis=(1, 1))
assert_raises(TypeError, np.count_nonzero, m, axis='foo')
- assert_raises(IndexError, np.count_nonzero, m, axis=3)
+ assert_raises(np.AxisError, np.count_nonzero, m, axis=3)
assert_raises(TypeError, np.count_nonzero,
m, axis=np.array([[1], [2]]))
@@ -2323,10 +2323,10 @@ class TestRollaxis(TestCase):
def test_exceptions(self):
a = np.arange(1*2*3*4).reshape(1, 2, 3, 4)
- assert_raises(IndexError, np.rollaxis, a, -5, 0)
- assert_raises(IndexError, np.rollaxis, a, 0, -5)
- assert_raises(IndexError, np.rollaxis, a, 4, 0)
- assert_raises(IndexError, np.rollaxis, a, 0, 5)
+ assert_raises(np.AxisError, np.rollaxis, a, -5, 0)
+ assert_raises(np.AxisError, np.rollaxis, a, 0, -5)
+ assert_raises(np.AxisError, np.rollaxis, a, 4, 0)
+ assert_raises(np.AxisError, np.rollaxis, a, 0, 5)
def test_results(self):
a = np.arange(1*2*3*4).reshape(1, 2, 3, 4).copy()
@@ -2413,11 +2413,11 @@ class TestMoveaxis(TestCase):
def test_errors(self):
x = np.random.randn(1, 2, 3)
- assert_raises_regex(IndexError, 'invalid axis .* `source`',
+ assert_raises_regex(np.AxisError, 'invalid axis .* `source`',
np.moveaxis, x, 3, 0)
- assert_raises_regex(IndexError, 'invalid axis .* `source`',
+ assert_raises_regex(np.AxisError, 'invalid axis .* `source`',
np.moveaxis, x, -4, 0)
- assert_raises_regex(IndexError, 'invalid axis .* `destination`',
+ assert_raises_regex(np.AxisError, 'invalid axis .* `destination`',
np.moveaxis, x, 0, 5)
assert_raises_regex(ValueError, 'repeated axis in `source`',
np.moveaxis, x, [0, 0], [0, 1])
diff --git a/numpy/core/tests/test_shape_base.py b/numpy/core/tests/test_shape_base.py
index ac8dc1eea..727608a17 100644
--- a/numpy/core/tests/test_shape_base.py
+++ b/numpy/core/tests/test_shape_base.py
@@ -184,8 +184,8 @@ class TestConcatenate(TestCase):
for ndim in [1, 2, 3]:
a = np.ones((1,)*ndim)
np.concatenate((a, a), axis=0) # OK
- assert_raises(IndexError, np.concatenate, (a, a), axis=ndim)
- assert_raises(IndexError, np.concatenate, (a, a), axis=-(ndim + 1))
+ assert_raises(np.AxisError, np.concatenate, (a, a), axis=ndim)
+ assert_raises(np.AxisError, np.concatenate, (a, a), axis=-(ndim + 1))
# Scalars cannot be concatenated
assert_raises(ValueError, concatenate, (0,))
@@ -294,8 +294,8 @@ def test_stack():
expected_shapes = [(10, 3), (3, 10), (3, 10), (10, 3)]
for axis, expected_shape in zip(axes, expected_shapes):
assert_equal(np.stack(arrays, axis).shape, expected_shape)
- assert_raises_regex(IndexError, 'out of bounds', stack, arrays, axis=2)
- assert_raises_regex(IndexError, 'out of bounds', stack, arrays, axis=-3)
+ assert_raises_regex(np.AxisError, 'out of bounds', stack, arrays, axis=2)
+ assert_raises_regex(np.AxisError, 'out of bounds', stack, arrays, axis=-3)
# all shapes for 2d input
arrays = [np.random.randn(3, 4) for _ in range(10)]
axes = [0, 1, 2, -1, -2, -3]
diff --git a/numpy/core/tests/test_ufunc.py b/numpy/core/tests/test_ufunc.py
index 4fe5a5ce2..f7b66f90c 100644
--- a/numpy/core/tests/test_ufunc.py
+++ b/numpy/core/tests/test_ufunc.py
@@ -703,14 +703,14 @@ class TestUfunc(TestCase):
def test_axis_out_of_bounds(self):
a = np.array([False, False])
- assert_raises(IndexError, a.all, axis=1)
+ assert_raises(np.AxisError, a.all, axis=1)
a = np.array([False, False])
- assert_raises(IndexError, a.all, axis=-2)
+ assert_raises(np.AxisError, a.all, axis=-2)
a = np.array([False, False])
- assert_raises(IndexError, a.any, axis=1)
+ assert_raises(np.AxisError, a.any, axis=1)
a = np.array([False, False])
- assert_raises(IndexError, a.any, axis=-2)
+ assert_raises(np.AxisError, a.any, axis=-2)
def test_scalar_reduction(self):
# The functions 'sum', 'prod', etc allow specifying axis=0
diff --git a/numpy/lib/tests/test_function_base.py b/numpy/lib/tests/test_function_base.py
index f69c24d59..d914260ad 100644
--- a/numpy/lib/tests/test_function_base.py
+++ b/numpy/lib/tests/test_function_base.py
@@ -466,8 +466,8 @@ class TestInsert(TestCase):
insert(a, 1, a[:, 2,:], axis=1))
# invalid axis value
- assert_raises(IndexError, insert, a, 1, a[:, 2, :], axis=3)
- assert_raises(IndexError, insert, a, 1, a[:, 2, :], axis=-4)
+ assert_raises(np.AxisError, insert, a, 1, a[:, 2, :], axis=3)
+ assert_raises(np.AxisError, insert, a, 1, a[:, 2, :], axis=-4)
# negative axis value
a = np.arange(24).reshape((2, 3, 4))
diff --git a/numpy/linalg/tests/test_linalg.py b/numpy/linalg/tests/test_linalg.py
index 2f8058ae6..b0a1f04d0 100644
--- a/numpy/linalg/tests/test_linalg.py
+++ b/numpy/linalg/tests/test_linalg.py
@@ -1102,8 +1102,8 @@ class _TestNorm(object):
assert_raises(ValueError, norm, B, order, (1, 2))
# Invalid axis
- assert_raises(IndexError, norm, B, None, 3)
- assert_raises(IndexError, norm, B, None, (2, 3))
+ assert_raises(np.AxisError, norm, B, None, 3)
+ assert_raises(np.AxisError, norm, B, None, (2, 3))
assert_raises(ValueError, norm, B, None, (0, 1, 2))
diff --git a/numpy/ma/core.py b/numpy/ma/core.py
index c32db0b49..1b25725d1 100644
--- a/numpy/ma/core.py
+++ b/numpy/ma/core.py
@@ -3903,7 +3903,7 @@ class MaskedArray(ndarray):
axis = None
try:
mask = mask.view((bool_, len(self.dtype))).all(axis)
- except (ValueError, IndexError):
+ except (ValueError, np.AxisError):
# TODO: what error are we trying to catch here?
# invalid axis, or invalid view?
mask = np.all([[f[n].all() for n in mask.dtype.names]
@@ -3941,7 +3941,7 @@ class MaskedArray(ndarray):
axis = None
try:
mask = mask.view((bool_, len(self.dtype))).all(axis)
- except (ValueError, IndexError):
+ except (ValueError, np.AxisError):
# TODO: what error are we trying to catch here?
# invalid axis, or invalid view?
mask = np.all([[f[n].all() for n in mask.dtype.names]
@@ -4345,7 +4345,7 @@ class MaskedArray(ndarray):
if self.shape is ():
if axis not in (None, 0):
- raise IndexError("'axis' entry is out of bounds")
+ raise np.AxisError("'axis' entry is out of bounds")
return 1
elif axis is None:
if kwargs.get('keepdims', False):
diff --git a/numpy/ma/tests/test_core.py b/numpy/ma/tests/test_core.py
index 45a6f4e86..9d8002ed0 100644
--- a/numpy/ma/tests/test_core.py
+++ b/numpy/ma/tests/test_core.py
@@ -1030,7 +1030,7 @@ class TestMaskedArrayArithmetic(TestCase):
res = count(ott, 0)
assert_(isinstance(res, ndarray))
assert_(res.dtype.type is np.intp)
- assert_raises(IndexError, ott.count, axis=1)
+ assert_raises(np.AxisError, ott.count, axis=1)
def test_minmax_func(self):
# Tests minimum and maximum.
@@ -4409,7 +4409,7 @@ class TestOptionalArgs(TestCase):
assert_equal(count(a, axis=(0,1), keepdims=True), 4*ones((1,1,4)))
assert_equal(count(a, axis=-2), 2*ones((2,4)))
assert_raises(ValueError, count, a, axis=(1,1))
- assert_raises(IndexError, count, a, axis=3)
+ assert_raises(np.AxisError, count, a, axis=3)
# check the 'nomask' path
a = np.ma.array(d, mask=nomask)
@@ -4423,13 +4423,13 @@ class TestOptionalArgs(TestCase):
assert_equal(count(a, axis=(0,1), keepdims=True), 6*ones((1,1,4)))
assert_equal(count(a, axis=-2), 3*ones((2,4)))
assert_raises(ValueError, count, a, axis=(1,1))
- assert_raises(IndexError, count, a, axis=3)
+ assert_raises(np.AxisError, count, a, axis=3)
# check the 'masked' singleton
assert_equal(count(np.ma.masked), 0)
# check 0-d arrays do not allow axis > 0
- assert_raises(IndexError, count, np.ma.array(1), axis=1)
+ assert_raises(np.AxisError, count, np.ma.array(1), axis=1)
class TestMaskedConstant(TestCase):