如何解读np.argmax()的输出结果?相关案例疑问咨询
Hey there! Let's clear up these two confusion points about np.argmax() with your examples.
np.argmax(a, axis=0) return array([1, 1, 1])? First, let's recap your 2D array:
import numpy as np a = np.arange(6).reshape(2,3) # a looks like: # array([[0, 1, 2], # [3, 4, 5]])
When you specify axis=0, np.argmax() looks for the index of the maximum value along the vertical axis (columns). Let's break down each column:
- First column values:
[0, 3]→ max is 3, which is at index 1 (the second row) - Second column values:
[1, 4]→ max is 4, at index 1 - Third column values:
[2, 5]→ max is 5, at index 1
That's why the result is array([1, 1, 1]) — it's not pointing to the global maximum's position, but rather giving the row index of the maximum in each individual column.
np.argmax(b) return 11 when the max value is 9? Your 3D array b has shape (2, 2, 3):
b = np.array([[[2,3,4],[4,5,6]],[[3,7,1],[2,5,9]]])
When you don't specify an axis parameter, np.argmax() automatically flattens the entire array into a 1D sequence first, then finds the index of the maximum value in this flattened array.
Let's flatten b manually:
Flattened b: [2, 3, 4, 4, 5, 6, 3, 7, 1, 2, 5, 9] Indices (0-based): 0 1 2 3 4 5 6 7 8 9 10 11
The maximum value 9 is at index 11 in this flattened list — that's exactly what np.argmax(b) returns. If you want to map this 1D index back to the original 3D coordinates, you can use np.unravel_index():
np.unravel_index(11, b.shape) # returns (1, 1, 2)
And indeed, b[1][1][2] gives you the value 9, which is the global maximum.
内容的提问来源于stack exchange,提问作者FaCoffee

