Python检索评估函数中numpy.ndarray转集合遇不可哈希类型错误求助
解决Python检索评估函数中numpy.ndarray不可哈希的问题
问题场景
编写Python检索评估函数时,执行retrieved_indices_set = set(retrieved_indices)触发TypeError: unhashable type: 'numpy.ndarray'错误。其中retrieved_indices是128×128×3的多维结构,尝试过将元素转为元组的方法但未解决。
用户尝试的代码如下:
def evaluate_retrieval(query_idx, retrieved_indices, relevant_indices): # Convert each list element to a tuple # Flatten the two-layer list and convert elements to tuples arr = np.array(retrieved_indices) #retrieved_indices => 128 * 128 * 3 # Transpose the array and convert it to a list of tuples retrieved_indices = tuple(list(map(tuple, np.vstack(arr.T)))) print(type(retrieved_indices)) # Create a set from the tuples retrieved_indices_set = set(retrieved_indices) relevant_retrieved = len(retrieved_indices_set.intersection(relevant_indices_set)) precision = relevant_retrieved / len(retrieved_indices_set) if len(retrieved_indices_set) > 0 else 0 return precision # 另一种尝试也失败 retrieved_indices_tuples = tuple(tuple(tuple(pixel) for pixel in row) for row in retrieved_indices)
错误原因
核心问题是numpy数组属于不可哈希类型,无法直接存入集合。你之前的转换逻辑要么没彻底把所有嵌套的numpy数组转成元组,要么操作(比如np.vstack(arr.T))打乱了结构,导致最终元素仍包含numpy数组。
解决方案
需要把128×128×3结构里的每一个最内层3维元素都转成元组,同时确保整个结构中没有残留的numpy数组。以下两种方案任选其一:
方案1:基于Python列表的逐层转换
先将numpy数组转为纯Python列表,再逐层把每个像素转成元组并扁平化:
import numpy as np def evaluate_retrieval(query_idx, retrieved_indices, relevant_indices): # 先转成纯Python列表(兼容输入是numpy数组或嵌套列表的情况) retrieved_list = retrieved_indices.tolist() if isinstance(retrieved_indices, np.ndarray) else retrieved_indices # 扁平化所有像素并转成元组 flattened_pixels = tuple(tuple(pixel) for row in retrieved_list for pixel in row) retrieved_indices_set = set(flattened_pixels) # 确保相关索引也是集合类型 relevant_indices_set = set(relevant_indices) relevant_retrieved = len(retrieved_indices_set & relevant_indices_set) precision = relevant_retrieved / len(retrieved_indices_set) if len(retrieved_indices_set) > 0 else 0 return precision
方案2:直接遍历numpy数组转换
如果输入确定是numpy数组,可直接遍历每个像素位置转换:
import numpy as np def evaluate_retrieval(query_idx, retrieved_indices, relevant_indices): flattened_pixels = [] # 遍历128×128的每个位置,把3维numpy元素转成元组 for i in range(retrieved_indices.shape[0]): for j in range(retrieved_indices.shape[1]): flattened_pixels.append(tuple(retrieved_indices[i,j])) retrieved_indices_set = set(flattened_pixels) relevant_indices_set = set(relevant_indices) relevant_retrieved = len(retrieved_indices_set.intersection(relevant_indices_set)) precision = relevant_retrieved / len(retrieved_indices_set) if len(retrieved_indices_set) > 0 else 0 return precision
关键注意点
- 必须确保最内层的每个3维元素都被转成元组,不能有残留的numpy数组;
- 如果你的需求是统计所有唯一像素的交集,一定要先扁平化二维结构,否则集合里会是整行的元组,而非单个像素,这会导致交集计算完全错误。
内容的提问来源于stack exchange,提问作者Erfan Hamidi
相关产品推荐
相关产品推荐

