如何确定子numpy数组源自原数组的列索引?
如何识别NumPy子数组对应的原数组列索引?
假设你有一个大的NumPy数组large,以及一个从其中提取若干列得到的子数组small,需要找出small的每一列在large中对应的原始列索引(0起始),这里提供两种通用的解决方法,适配不同场景。
首先先明确问题中的示例数据:
import numpy as np large = np.array([ [-0.047391, -0.10926778, -0.00899118, 0.07461428, -0.07667476, 0.06961918, 0.09440736, 0.01648382, -0.04102225, -0.05038805, -0.00930337, 0.3667651, -0.02803499, 0.02597451, -0.1218804, 0.00561949], [-0.00253788, -0.08670117, -0.00466262, 0.07330351, -0.06403728, 0.00301005, 0.12807456, 0.01198117, -0.04290793, -0.06138136, -0.01369276, 0.37094407, -0.03747804, 0.04444246, -0.01162705, 0.00554793] ]) small = np.array([ [-0.10926778, -0.07667476, 0.09440736], [-0.08670117, -0.06403728, 0.12807456] ])
方法1:遍历匹配(直观易懂)
这种方法逻辑清晰,适合新手理解,通过遍历small的每一列,再在large中查找完全匹配的列:
def find_column_indices(large_arr, small_arr): # 先校验两个数组的行数是否一致(毕竟是从原数组提取列,行数必须相同) if large_arr.shape[0] != small_arr.shape[0]: raise ValueError("两个数组的行数必须相同") indices = [] # 转置数组,方便按列遍历 for small_col in small_arr.T: # 遍历原数组的每一列 for idx, large_col in enumerate(large_arr.T): # 精确匹配列(如果有浮点精度问题,换成np.allclose) if np.all(large_col == small_col): indices.append(idx) break else: # 遍历完所有列都没找到匹配,抛出错误 raise ValueError(f"small中的列{small_col}在large中未找到匹配项") return indices # 调用函数 result = find_column_indices(large, small) print(result) # 输出: [1, 4, 6]
注意点:
- 如果你的数组涉及浮点数运算,可能存在微小精度误差,这时候建议用
np.allclose(large_col, small_col)代替==,它会在允许的误差范围内判断两数组是否相等。
方法2:NumPy广播(高效简洁)
利用NumPy的广播机制,可以一次性完成所有列的比较,代码更简洁,效率也更高,适合处理较大的数组:
# 转置数组后,利用广播比较所有列对的匹配情况 # large.T[:, None, :] 形状是(大数组列数, 1, 行数) # small.T[None, :, :] 形状是(1, 小数组列数, 行数) # 广播后逐元素比较,再按行(原数组的行)取all,得到(大数组列数, 小数组列数)的布尔矩阵 matches = np.all(large.T[:, None, :] == small.T[None, :, :], axis=2) # 找到每个small列对应的large列索引 indices = np.where(matches.any(axis=0))[0] print(indices) # 输出: [1, 4, 6]
解释:
matches.any(axis=0)会检查small的每一列是否在large中存在匹配的列,返回一个长度等于small列数的布尔数组。np.where则会提取出这些布尔值为True对应的索引,也就是我们需要的原数组列索引。- 同样,浮点数场景下把
==换成np.allclose即可。
内容的提问来源于stack exchange,提问作者Jack Arnestad
相关产品推荐
相关产品推荐

