Looks like you want the row that contains the maximum value, right?

max(axis=0) returns the maximum of [1,0] and [2,4] independently.

argmax without axis parameter finds the maximum over the whole array - in flattened form. To turn that index into row number we have to use unravel_index:

In [464]: a.argmax()
Out[464]: 3
In [465]: np.unravel_index(3,(2,2))
Out[465]: (1, 1)
In [466]: a[1,:]
Out[466]: array([0, 4])

or in one expression:

In [467]: a[np.unravel_index(a.argmax(), a.shape)[0], :]
Out[467]: array([0, 4])

As you can see from the length of the answer it's not the usual definition of maximum along/over an axis.

Sum along axis in numpy array may give more insight into the meaning of 'along axis'. The same definitions apply to the sum, mean and max operations.

===================

To pick row with the largest norm, first calculate the norm. norm uses the axis parameter in the same way.

In [537]: np.linalg.norm(a,axis=1)
Out[537]: array([ 2.23606798,  4.        ])
In [538]: np.argmax(_)
Out[538]: 1
In [539]: a[_,:]
Out[539]: array([0, 4])
Answer from hpaulj on Stack Overflow
🌐
NumPy
numpy.org › doc › 2.2 › reference › generated › numpy.max.html
numpy.max — NumPy v2.2 Manual
numpy.max(a, axis=None, out=None, keepdims=<no value>, initial=<no value>, where=<no value>)[source]#
🌐
NumPy
numpy.org › devdocs › reference › generated › numpy.max.html
numpy.max — NumPy v2.6.dev0 Manual
numpy.max(a, axis=None, out=None, keepdims=<no value>, initial=<no value>, where=<no value>)[source]#
🌐
DataCamp
datacamp.com › doc › numpy › max
NumPy max()
In this syntax, `array` is the input array, `axis` specifies the axis along which to find the maximum, and `out` is an optional parameter to store the result. The `keepdims` parameter, when set to `True`, retains reduced dimensions as dimensions with size one.
🌐
Sharp Sight
sharpsight.ai › blog › numpy-max
How to use the NumPy max function - Sharp Sight
February 6, 2024 - Keep in mind that the axis parameter is optional. If you don’t specify an axis, NumPy max will find the maximum value in the whole NumPy array. The out parameter allows you to specify a special output array where you can store the output of np.max.
🌐
Programiz
programiz.com › python-programming › numpy › methods › max
NumPy max()
If axis = 1, the maximum of the largest element in each row is returned. import numpy as np array = np.array([[10, 17, 25], [15, 11, 22]])
🌐
NumPy
numpy.org › doc › 2.3 › reference › generated › numpy.max.html
numpy.max — NumPy v2.3 Manual
numpy.max(a, axis=None, out=None, keepdims=<no value>, initial=<no value>, where=<no value>)[source]#
🌐
NumPy
numpy.org › doc › stable › reference › generated › numpy.amax.html
numpy.amax — NumPy v2.5 Manual
numpy.amax(a, axis=None, out=None, keepdims=<no value>, initial=<no value>, where=<no value>)[source]# Return the maximum of an array or maximum along an axis. amax is an alias of max. See also · max · alias of this function · ndarray.max · equivalent method ·
Find elsewhere
🌐
Codecademy
codecademy.com › docs › python:numpy › built-in functions › .max()
Python:NumPy | Built-in Functions | .max() | Codecademy
July 2, 2025 - Returns the maximum value of an array or maximum values along a specified axis.
🌐
TutorialsPoint
tutorialspoint.com › numpy › numpy_max.htm
NumPy - Max
The axis parameter refers to the direction along which the maximum value should be calculated. For example, in a 2D array − · axis=0: Calculate the maximum value along the columns (vertical axis).
🌐
NumPy
numpy.org › doc › 2.1 › reference › generated › numpy.max.html
numpy.max — NumPy v2.1 Manual
The minimum value of an array along a given axis, propagating any NaNs. ... The maximum value of an array along a given axis, ignoring any NaNs.
🌐
GeeksforGeeks
geeksforgeeks.org › numpy-amax-python
numpy.amax() in Python | GeeksforGeeks
April 28, 2022 - The numpy.amax() method returns the maximum of an array or maximum along the axis(if mentioned).
🌐
Python Examples
pythonexamples.org › python-numpy-get-maximum-value-of-array-along-axis
Find Maximum Value in NumPy Array along an Axis
np.amax(arr, axis=(0, 2)) calculates the maximum values along both axis 0 (depth) and axis 2 (columns).
🌐
Codecademy
codecademy.com › docs › python:numpy › built-in functions › .amax()
Python:NumPy | Built-in Functions | .amax() | Codecademy
May 15, 2024 - The a parameter is required and represents the array of elements to choose the maximum from. All other parameters are optional. ... axis: (Default = None) An integer or a tuple of integers specifying the axis/axes along which to operate.
🌐
Real Python
realpython.com › numpy-max-maximum
NumPy's max() and maximum(): Find Extreme Values in Arrays – Real Python
October 22, 2025 - The .max() method has scanned the whole array and returned the largest element. Using this method is exactly equivalent to calling np.max(n_scores). But perhaps you want some more detailed information. What was the top score for each test? Here you can use the axis parameter:
🌐
Vultr Docs
docs.vultr.com › python › third party › numpy › max()
Python Numpy max() - Find Maximum Value
November 18, 2024 - arr = np.array([[1, 5, 6], [9, 0, 2], [4, 8, 3]]) max_value_rows = arr.max(axis=1) print("Maximum values per row:", max_value_rows) Explain Code · Powered by Vultr AgentBeta · By setting axis=1, the maximum value from each row is returned: [6, ...
🌐
NumPy
numpy.org › doc › stable › reference › generated › numpy.max.html
numpy.max — NumPy v2.5 Manual
numpy.max(a, axis=None, out=None, keepdims=<no value>, initial=<no value>, where=<no value>)[source]#
🌐
Reddit
reddit.com › r/learnprogramming › numpy: trying to efficiently select maximum value of an axis for each cell of the other axes
r/learnprogramming on Reddit: Numpy: trying to efficiently select maximum value of an axis for each cell of the other axes
August 19, 2022 -

I have an (N, N, 3) numpy array of floats as input.

I want to return an (N, N, 3) array which has, along the third axis, a 1.0 in the position of the highest value for that vector, and 0.0s in the other two values.

That is to say, for N=2 with input:

in = array([[[0.3767, 0.5967, 0.4188],  
             [0.3749, 0.5432, 0.8066]],
            [[0.3265, 0.9366, 0.6033],
             [0.2315, 0.9459, 0.2973]]])

I want to return output:

out = array([[[0.0000, 1.0000, 0.0000],  
              [0.0000, 0.0000, 1.0000]],
             [[0.0000, 1.0000, 0.0000],
              [0.0000, 1.0000, 0.0000]]])

Right now I can get this output using:

import numpy as np
np.set_printoptions(precision=4, floatmode='fixed')
N = 2
inp = np.random.random((N, N, 3))
out = np.empty_like(inp)
for i in range(N):
    for j in range(N):
        out[i, j, :] = 1.0 * (inp[i, j, :] == np.max(inp[i, j, :]))
print(f'in = {inp}')
print(f'out = {out}')

It seems like there should be a way to (at least explicitly) avoid those nested for loops, but I don't know how.

Is there a neater/more efficient way to do this?

🌐
GitHub
github.com › numba › numba › issues › 4414
TypingError: `np.max(..., axis=1)` · Issue #4414 · numba/numba
August 6, 2019 - def ts_max(x,n,d): n = len(x) x_ = np.lib.stride_tricks.as_strided(x,shape=(n-d+1,d),strides=(x.strides[0],x.strides[0])) ret = np.max(x_,axis=1) return ret @jit(nopython=True) def ts_max_mat(m,d): if d == 1: return m rowNum = m.shape[0] colNum ...
Author: numba