You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用起止索引数组高效切片多维NumPy数组?

问题

如何使用包含起止索引、维度相近的第二个np.array,高效切片多维np.array?

最小可运行示例

我有一个需要切片的4维数组,示例如下:

import numpy as np
shape = np.array([[3,3], [5,5]])
arr = np.arange(np.prod(shape), dtype=np.uint8).reshape(shape.flatten())

此示例中,数组是包含0到225所有整数的(3,3,5,5)矩阵。

此外,我还有一组索引:

idxs = np.broadcast_to(np.array([[1,1], [4,3]]), (*shape[0],2,2))
rand = np.random.randint(-1, 2, shape[0])
idxs = idxs + np.broadcast_to(rand[:,:,None,None], (*shape[0],2,2))

这是一个存储二维坐标的4维数组,结构如下:

np.array([[r1, c1], [r2,c2]])

其中[r1, r2]和[c1,c2]分别表示要切片的行范围和列范围。尽管加入了随机值,但可保证所有idxs元素对应的切片维度相同(此示例中为3x2矩阵)。切片维度也可通过以下代码计算:

extent = np.min(idxs[:,:,1] - idxs[:,:,0], axis=(0,1))

简单Python实现

以下函数可实现需求,但使用了较慢的常规Python for循环:

def sliceMultiDim(arr: np.ndarray, idxs: np.ndarray, extent: np.ndarray) -> np.ndarray:
    sliced = np.zeros((*arr.shape[:2], *extent))        
    for row in range(arr.shape[0]):
        for col in range(arr.shape[1]):
            sub = arr[row,col]
            idx = idxs[row,col]
            sliced[row,col] = sub[idx[0,0] : idx[1,0], idx[0,1] : idx[1,1]]
    return sliced

基准测试

定义向量化解决方案:

def sliceMultiDimVect(arr: np.ndarray, idxs: np.ndarray) -> np.ndarray:
    # 计算滑动窗口的最小尺寸
    extent = np.min(idxs[:,:,1] - idxs[:,:,0], axis=(0,1))
    # 创建包含所有`(3, 2)`窗口的六维视图
    windowed = swv(arr, tuple(extent), axis=(2, 3))
    # 每个切片的左上角坐标,可直接索引windowed中的目标切片
    r1, c1 = idxs[:, :, 0, 0], idxs[:, :, 0, 1]
    # 前两维的高级索引网格
    h, w = np.indices(arr.shape[:2], sparse=True)
    return windowed[h, w, r1, c1]

当扩大问题规模时,根据以下基准测试脚本,向量化解决方案比朴素实现快约8倍:

import timeit
# 初始化包含1亿个元素的100个数组
shape = np.array([[10,10], [10000,10000]])
arr = np.arange(np.prod(shape), dtype=np.uint8).reshape(shape.flatten())
# 准备切片索引,目标切片包含100万个元素
idxs = np.broadcast_to(np.array([[0,0], shape[1]//10]), (*shape[0],2,2))
rand = np.random.randint(0, shape[1,0]//10, shape[0])
idxs = idxs + np.broadcast_to(rand[:,:,None,None], (*shape[0],2,2))
def non_vect():
    sliceMultiDim(arr, idxs)
def vect():
    sliceMultiDimVect(arr, idxs)
print("non_vect:", timeit.timeit("non_vect()", setup="from __main__ import non_vect", number=100))
print("vect:", timeit.timeit("vect()", setup="from __main__ import vect", number=100))
解决方案耗时 [秒]相对耗时
vect8.4304994000121951
non_vect64.53972959998646~7.7

测试环境:Intel Core i7-1185G7 @ 3.00 GHz 处理器。

内容的提问来源于Stack Exchange,提问作者Ben Zeen

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 14:38:12