如何为NumPy数组中的每个子数组设置不同起止索引进行切片?
为NumPy数组的每个子数组应用不同起止索引切片的高效方法
先明确你的场景:
import numpy as np X = np.array([[1,2,3,4],[5,6,7,8]])
你需要实现类似自定义eg_slice函数的效果,传入每一行的起止索引对,输出对应切片后的二维数组,同时替代linspace生成全量索引的高内存方案。
以下是几种高效低内存的实现方式:
方法1:广播生成紧凑索引提取
这种方式避免生成全量冗余索引,仅按需创建必要的行、列索引数组:
def slice_per_row(X, slices): # 检查所有切片长度是否一致(否则无法组成二维数组) slice_lengths = [end - begin for begin, end in slices] if len(set(slice_lengths)) != 1: raise ValueError("所有切片的长度必须相同") length = slice_lengths[0] # 生成重复的行索引,对应每个切片元素 row_idx = np.repeat(np.arange(X.shape[0]), length) # 拼接每行的连续列索引 col_idx = np.concatenate([np.arange(begin, begin+length) for begin, _ in slices]) return X[row_idx, col_idx].reshape(X.shape[0], length) # 测试调用 result = slice_per_row(X, [[1,3],[0,2]]) print(result) # 输出: # [[2 3] # [5 6]]
方法2:利用广播直接索引(零额外内存开销)
通过广播机制生成索引视图,完全避免存储额外的索引数组,内存效率最高:
def slice_per_row_broadcast(X, slices): slice_lengths = [end - begin for begin, end in slices] if len(set(slice_lengths)) != 1: raise ValueError("所有切片的长度必须相同") length = slice_lengths[0] # 提取每行的起始索引,通过广播生成对应列索引矩阵 starts = np.array([b for b, _ in slices]) col_idx = starts[:, None] + np.arange(length) # 结合行索引广播完成切片 return X[np.arange(X.shape[0])[:, None], col_idx] # 测试调用 result = slice_per_row_broadcast(X, [[1,3],[0,2]]) print(result) # 输出: # [[2 3] # [5 6]]
方法3:列表推导式(简洁高效,适合小规模数组)
如果数组规模不大,直接循环处理每行,代码最简洁且无额外内存开销:
def slice_per_row_simple(X, slices): return np.array([X[i, begin:end] for i, (begin, end) in enumerate(slices)]) # 测试调用 result = slice_per_row_simple(X, [[1,3],[0,2]]) print(result) # 输出: # [[2 3] # [5 6]]
方案对比
上述所有方法都比linspace生成全量索引的方式更节省内存:linspace会生成完整的索引数组,而这些方法要么仅生成紧凑的必要索引,要么利用广播/视图机制避免存储冗余数据。
内容的提问来源于stack exchange,提问作者ROgercad
相关产品推荐
相关产品推荐

