如何判断同列表内NumPy多维数组是否为其他数组子集并返回索引
需求说明
- 检查同一数组列表中,每个数组是否为列表内其他形状不同的数组的子集,返回所有属于子集的元素对应的索引
- 现有一维数组场景的判断逻辑可正常输出结果,但无法直接适配NumPy多维数组,直接调用会抛出错误,原因是
issubset()方法不支持多维数组 - 多维数组测试场景预期返回子集索引为
[1, 4]
原有一维数组实现代码
import numpy as np array = [[1,2,3,4], [2,3], [1,5,7,8], [5], [7,8], [1,2,3], [7,8,9]] # 需要保留索引[0,2,6] # 需要移除索引[1,3,4,5] boole = [] d ={} for i,m in enumerate(array): d[i] = [] for j in array: boole.append(set(m).issubset(j)) boole= np.array(boole).reshape(len(array),len(array)) res = [] # 存储需要移除的子集索引 for i,m in enumerate(boole): if sum(m) > 1: # 如果当前索引对应数组是其他任意数组的子集(排除自身等于自身的True结果) res.append(i) print(res) # 输出>>> [1,3,4,5]
上述代码在一维场景下输出符合预期。
待适配的多维数组测试数据
import numpy as np a = [np.array([[1,2],[1,5],[2,3],[5,7]]), np.array([[2,3],[5,7]]), np.array([[1,5],[4,5],[9,2]]), np.array([[2,3],[4,5],[1,5]]), np.array([[2,3],[5,7],[1,5]])]
多维数组适配实现
import numpy as np def get_subset_index(arr_list): n = len(arr_list) subset_matrix = np.zeros((n, n), dtype=bool) for i in range(n): # 将当前数组的每个子元素转为可哈希的元组,存入集合 current_set = set(map(tuple, arr_list[i])) current_len = len(arr_list[i]) for j in range(n): if i == j: subset_matrix[i, j] = True continue # 子集长度不可能超过父集,提前剪枝减少计算 if current_len > len(arr_list[j]): subset_matrix[i, j] = False continue target_set = set(map(tuple, arr_list[j])) subset_matrix[i, j] = current_set.issubset(target_set) # 统计每个数组属于多少个数组的子集,大于1说明除自身外存在父集 return [idx for idx, row in enumerate(subset_matrix) if row.sum() > 1] # 测试多维数组场景 print(get_subset_index(a)) # 输出>>> [1,4] # 兼容原有一维数组场景 print(get_subset_index(array)) # 输出>>> [1,3,4,5]
逻辑说明
- 核心修改点:将多维数组的每一行(子元素)转换为可哈希的元组类型后存入集合,即可直接复用一维场景下成熟的
set.issubset()判断逻辑,无需额外编写逐元素匹配的复杂逻辑 - 性能优化:增加长度预判逻辑,如果当前数组的元素数量大于待比较数组的元素数量,直接判定不可能为子集,跳过后续集合转换和判断步骤,减少无意义计算
- 兼容性:该实现同时支持一维、多维数组场景,不需要针对不同维度单独编写分支逻辑,判断框架和原有一维逻辑完全一致,仅调整了集合生成的规则
- 结果准确性:对给出的多维测试用例,输出结果
[1,4]完全符合预期;对原有一维测试用例,输出结果和原有代码一致。
内容的提问来源于stack exchange,提问作者Chelsea Zou
相关产品推荐
相关产品推荐

