使用np.argmax时如何在数组元素全相同或多最大值时返回-1?
自定义argmax实现方案
你可以直接基于numpy现有接口封装函数实现需求,以下是两种适用不同场景的实现:
一维数组场景
直接针对最常用的一维输入场景实现,逻辑简单易读:
import numpy as np def custom_argmax(arr, default=-1): max_val = arr.max() # 统计最大值出现的总次数 max_count = np.sum(arr == max_val) # 最大值出现超过1次就返回自定义值 if max_count > 1: return default # 否则返回原生argmax的结果 return np.argmax(arr)
测试效果:
- 输入
a = np.array([2, 2, 2]),返回-1,符合预期 - 输入
b = np.array([1, 5, 3, 2]),返回1,和原生np.argmax表现一致 - 输入
c = np.array([5, 1, 5, 3]),返回-1,符合预期
支持高维数组、指定axis场景
如果需要和原生np.argmax一样支持高维输入、指定计算维度,可使用以下兼容版本:
def custom_argmax_nd(arr, axis=None, default=-1): max_val = arr.max(axis=axis, keepdims=True) max_count = np.sum(arr == max_val, axis=axis) res = np.argmax(arr, axis=axis) # 把所有存在多个最大值的位置替换为自定义值 res[max_count > 1] = default return res
内容的提问来源于stack exchange,提问作者eng2019
相关产品推荐
相关产品推荐

