如何确定Numpy索引表达式使用的索引数量,并提前获取索引后数组的形状?
在NumPy中预计算索引后的数组形状与索引数量
一、通用方法获取索引后的数组形状
如果你想避免实际对大数组执行索引操作(节省内存和时间),最简单高效的方式是创建一个和原数组形状相同的空数组,通过索引这个空数组来获取目标形状。空数组的索引操作几乎不消耗内存,且能准确复用NumPy原生的索引规则:
import numpy as np def get_indexed_shape(arr, indexer): # 创建与原数组形状一致的空布尔数组(内存占用极小) dummy = np.empty(arr.shape, dtype=np.bool_) return dummy[:, indexer].shape # 这里可根据你的索引维度调整逻辑 # 测试你的示例 f.M = np.repeat([[1,2,3,4,5]],3,axis=0) print(get_indexed_shape(f.M, [1,4])) # 输出 (3, 2) print(get_indexed_shape(f.M, slice(0,None,2))) # 输出 (3, 3) print(get_indexed_shape(f.M, [True, False, True, False, False])) # 输出 (3, 2)
这个方法完全兼容你提到的整数列表、切片、布尔表达式三种索引器,甚至能处理更复杂的索引类型(比如整数数组、np.newaxis等),因为它直接利用NumPy本身的索引逻辑来计算形状,不需要你手动处理每种索引器的规则。
二、针对布尔/整数列表的简化方法
如果你只需要处理布尔表达式和整数列表索引器,也可以直接计算索引器对应的长度,再结合原数组的形状:
def get_index_length(indexer, axis_length): if isinstance(indexer, (list, np.ndarray)): if np.asarray(indexer).dtype == bool: # 布尔索引:统计True的数量 return np.count_nonzero(indexer) else: # 整数列表:直接取长度 return len(indexer) else: raise ValueError("仅支持布尔表达式或整数列表索引器") # 计算形状 def get_indexed_shape_simple(arr, indexer): return arr.shape[:-1] + (get_index_length(indexer, arr.shape[-1]),) # 测试 print(get_indexed_shape_simple(f.M, [1,4])) # (3, 2) print(get_indexed_shape_simple(f.M, [True, False, True, False, False])) # (3, 2)
三、确定索引表达式中使用的索引数量
这里分两种场景:
用户传入的索引器个数:不管NumPy是否自动补全维度,直接统计用户提供的索引器元素数量:
def get_user_index_count(indexer): # 如果不是元组,转换成单元素元组 if not isinstance(indexer, tuple): indexer = (indexer,) return len(indexer) # 测试 print(get_user_index_count([1,4])) # 1 print(get_user_index_count((slice(None), [1,4]))) # 2NumPy实际使用的索引维度数:这其实等于原数组的维度数,因为NumPy会自动用
slice(None)(即:)补全缺失的维度。比如你传入[1,4]对二维数组索引,NumPy会自动处理成[:, [1,4]],实际使用2个索引维度:def get_actual_index_count(arr): return arr.ndim
如果你指的是索引器本身的维度(比如整数数组索引器的形状),可以直接用np.asarray(indexer).ndim来获取。
内容的提问来源于stack exchange,提问作者JPhibs
相关产品推荐
相关产品推荐

