>>> a.argmax(axis=0)
array([1, 1, 0])
Answer from eumiro on Stack OverflowHow does NumPy handle multiple maximum values?
Can I use argmax on a multi-dimensional array?
What is the difference between np.argmax() and np.max()?
Use built-in function for it:
prediction.argmax()
output:
9
Also, that index 0 is the row number, so the max is at row 0 and column 9.
As the other answers mentioned, you have a 2D array, so you end up with two indices. Since the array is just a row, the first index is always zero. You can bypass this in a number of ways:
Use
prediction.argmax(). The defaultaxisargument isNone, which means operate on a flattened array. Other options that will get you the same result areprediction.argmax(-1)(last axis) andprediction.argmax(1)(second axis). Keep in mind that you will only ever get the index of the first maximum this way. That's fine if you only ever expect to have one, or only need one.Use
np.flatnonzeroto get the linear indices similarly to the way you were doing:np.flatnonzero(perdiction == prediction.max())Use
np.nonzeroornp.where, but extract the axis you care about:np.nonzero(prediction == prediction.max())[1]ravelthe array on input:np.where(prediction.ravel() == prediction.max())Do the same thing, but with
np.squeeze:np.nonzero(prediction.squeeze() == prediction.max())
Numpy has an argmax function that returns just that, although you will have to deal with the nans manually. nans always get sorted to the end of an array, so with that in mind you can do:
a = np.random.rand(10000)
a[np.random.randint(10000, size=(10,))] = np.nan
a = a.reshape(100, 100)
def nanargmax(a):
idx = np.argmax(a, axis=None)
multi_idx = np.unravel_index(idx, a.shape)
if np.isnan(a[multi_idx]):
nan_count = np.sum(np.isnan(a))
# In numpy < 1.8 use idx = np.argsort(a, axis=None)[-nan_count-1]
idx = np.argpartition(a, -nan_count-1, axis=None)[-nan_count-1]
multi_idx = np.unravel_index(idx, a.shape)
return multi_idx
>>> nanargmax(a)
(20, 93)
You should use np.where
In [17]: a=np.random.uniform(0, 10, size=10)
In [18]: a
Out[18]:
array([ 1.43249468, 4.93950873, 7.22094395, 1.20248629, 4.66783985,
6.17578054, 4.6542771 , 7.09244492, 7.58580515, 5.72501954])
In [20]: np.where(a==a.max())
Out[20]: (array([8]),)
This also works for 2 arrays, the returned value, is the index. Here we create a range from 1 to 9:
x = np.arange(9.).reshape(3, 3)
This returns the index, of the the items that equal 5:
In [34]: np.where(x == 5)
Out[34]: (array([1]), array([2])) # the first one is the row index, the second is the column
You can use this value directly to slice your array:
In [35]: x[np.where(x == 5)]
Out[35]: array([ 5.])