如何获取多维数组每行最大值的索引?numpy代码问题求助
问题分析与解决方案
你的代码当前存在两个核心问题:
- 错误计算全局最大值而非每行最大值:
np.max(input_array)会提取整个数组的全局最大值,再用np.where只能找到该全局值的位置,而非每行各自最大值的索引。 - 输出格式不符合需求:
np.where返回的是多维索引的元组,无法直接得到像示例那样简洁的列表格式。
修正后的代码
使用np.argmax并指定维度参数,可以直接获取每行最大值的索引,再转换为普通列表即可得到简洁输出:
import numpy as np def display_max(input_array): arr = np.array(input_array) # 沿最后一个维度(每行)计算最大值索引 max_indices = np.argmax(arr, axis=-1) # 转换为Python列表格式 return max_indices.tolist() # 测试示例输入 print(display_max([[1,2,3,4],[5,10,2,3]])) # 输出: [3, 1] # 测试你的三维输入 print(display_max([[[7,15,3,10],[2,6,9,0],[20,45,71,500]]])) # 输出: [[1, 2, 3]]
如果希望三维输入也返回一维列表,可以添加flatten()处理:
def display_max(input_array): arr = np.array(input_array) max_indices = np.argmax(arr, axis=-1) return max_indices.flatten().tolist() print(display_max([[[7,15,3,10],[2,6,9,0],[20,45,71,500]]])) # 输出: [1, 2, 3]
关键说明
np.argmax(arr, axis=-1):axis=-1表示沿着数组的最后一个维度(即每行的元素方向)查找最大值的索引,完美匹配“每行最大值索引”的需求。对于你的三维输入(1,3,4),最后一个维度是4(每行的元素数),所以会得到每行的索引。tolist():将numpy数组转换为普通Python列表,输出格式和示例一致。
内容的提问来源于stack exchange,提问作者Amelia Putri Damayanti
相关产品推荐
相关产品推荐

