如何使用二维索引数组对NumPy二维数组进行切片?
如何使用二维索引数组对NumPy二维数组进行切片?
你需要对二维数组的每一行,根据对应的start和end索引截取子数组,再补零到原数组的列数以得到固定形状的输出。下面我会给出两种常见场景的解决方案,分别对应你提到的两种预期输出格式。
先明确需求细节
观察你的示例可以发现:每一行的start和end是闭区间(即包含start和end索引的元素),所以在Python切片中需要用start:end+1来获取对应的子数组(因为Python切片是左闭右开的)。
方法1:实现「截取子数组前置,后面补零」的格式
这种格式是把每一行截取的子数组放在行的开头,剩余位置填充0。
方案A:循环实现(简单直观)
适合小规模数据,代码易读、易调试:
import numpy as np np.random.seed(0) a = np.random.randint(0,999,(4,5)) idx = np.array([[2,4], [0,3], [2,3], [1,3]]) # 初始化和原数组形状一致的全零数组 output = np.zeros_like(a) for i in range(a.shape[0]): start, end = idx[i] # 截取当前行的目标子数组(闭区间转切片) sub_arr = a[i, start:end+1] # 将子数组填充到当前行的开头,剩余位置保持0 output[i, :len(sub_arr)] = sub_arr print(output)
运行结果:
[[629 192 835 0 0] [763 707 359 9 0] [804 599 0 0 0] [600 396 314 0 0]]
方案B:向量化实现(高效大数据)
避免循环,利用NumPy的向量化操作提升速度,适合大规模数组:
import numpy as np np.random.seed(0) a = np.random.randint(0,999,(4,5)) idx = np.array([[2,4], [0,3], [2,3], [1,3]]) output = np.zeros_like(a) # 计算每一行截取的子数组长度 row_lengths = idx[:, 1] - idx[:, 0] + 1 # 生成行索引,对应每个要提取的元素 row_indices = np.repeat(np.arange(a.shape[0]), row_lengths) # 生成原数组中要提取的元素的列索引 col_a = np.concatenate([np.arange(s, e+1) for s, e in idx]) # 生成输出数组中要填充的列索引(从0开始连续) col_out = np.concatenate([np.arange(l) for l in row_lengths]) # 扁平化赋值 output[row_indices, col_out] = a[row_indices, col_a] print(output)
运行结果和方案A完全一致。
方法2:实现「截取子数组保留原位置,其他补零」的格式
这种格式是把每一行的目标子数组保留在原索引位置,其余位置填充0。
方案A:循环实现
import numpy as np np.random.seed(0) a = np.random.randint(0,999,(4,5)) idx = np.array([[2,4], [0,3], [2,3], [1,3]]) output = np.zeros_like(a) for i in range(a.shape[0]): start, end = idx[i] # 直接将原数组的目标切片填充到输出的对应位置 output[i, start:end+1] = a[i, start:end+1] print(output)
运行结果:
[[ 0 0 629 192 835] [763 707 359 9 0] [ 0 0 804 599 0] [ 0 600 396 314 0]]
方案B:向量化实现(高效大数据)
用掩码方式一次性完成赋值,无需循环:
import numpy as np np.random.seed(0) a = np.random.randint(0,999,(4,5)) idx = np.array([[2,4], [0,3], [2,3], [1,3]]) # 生成列索引数组,形状为(5,) cols = np.arange(a.shape[1]) # 生成掩码:判断每个位置是否属于对应行的[start, end]区间 # 利用广播,将cols扩展为(4,5),和idx的(4,2)比较 mask = (cols >= idx[:, 0, None]) & (cols <= idx[:, 1, None]) # 用掩码选择原数组元素,否则填充0 output = np.where(mask, a, 0) print(output)
运行结果和方案A完全一致。
备注:内容来源于stack exchange,提问作者Abhishek Jain
相关产品推荐
相关产品推荐

