如何从指定结构的NumPy数组中匹配输入对并获取对应输出值?
解决方案
问题原因
你用np.where(arr == [2,8])返回空数组,是因为这个NumPy数组是object类型的不规则数组(子元素长度不一致),直接用数组相等比较无法正确匹配嵌套的输入列表[2,8]。
方法一:修复NumPy的查找逻辑
遍历数组中的每个条目,检查输入部分是否匹配[a,b]:
import numpy as np arr = np.array([[[0,0],[93]],[[2,8],[94]]], dtype=object) a = 2 b = 8 # 遍历查找匹配项 for entry in arr: input_pair, output_val = entry if np.array_equal(input_pair, [a, b]): print(output_val) # 输出94 break
方法二:使用字典实现高效查找(推荐)
这种键值对查找场景,字典是最优选择,查找时间复杂度为O(1)。如果要保留([inp,inp],[out])的条目格式,可以先定义条目列表再转换为字典:
a = 2 b = 8 # 保留原始条目格式 entries = [([0,0],[93]), ([2,8],[94])] # 转换为字典(列表不能做键,转成元组) lookup_dict = {tuple(inp): out for inp, out in entries} # 直接查找 result = lookup_dict[(a, b)] print(result) # 输出94
如果不需要保留原始条目格式,直接定义字典更简洁:
lookup_dict = {(0, 0): 93, (2, 8): 94} result = lookup_dict[(a, b)]
内容的提问来源于stack exchange,提问作者user24644752
相关产品推荐
相关产品推荐

