如何向量化实现从3D数组提取2D补丁并计算每个补丁均值?
向量化实现从三维数组的每个内层2D数组提取补丁并计算均值
要高效解决这个问题,我们可以直接利用NumPy的as_strided(也就是extract_patches_2d底层依赖的机制)来实现全向量化操作,避免循环调用extract_patches_2d,大幅提升性能。以下是具体步骤和代码示例:
1. 明确输入数组的形状
假设你的三维数组a的形状是(M, H, W),其中:
M是最外层的数组数量(每个元素都是一个2D数组)H和W是每个内层2D数组的高和宽
我们要提取的补丁大小是(p, p) = (3, 3),步长默认设为1(如果需要调整步长,后面也可以修改参数)。
2. 用as_strided创建补丁的跨步视图
as_strided可以在不复制数据的前提下,创建一个新的数组视图,对应所有需要的补丁。关键是正确计算新数组的形状和跨步:
- 新形状:
(M, H - p + 1, W - p + 1, p, p),其中前三个维度分别对应外层数组索引、补丁的行位置、补丁的列位置,最后两个维度是补丁本身的3x3结构。 - 新跨步:利用原数组的跨步信息推导,原数组
a的跨步是a.strides,假设为(s0, s1, s2),那么新跨步应该是(s0, s1, s2, s1, s2)——这样每个补丁的元素都对应原数组中连续的3x3区域。
3. 计算每个补丁的均值
创建跨步视图后,直接对最后两个维度(补丁的行和列)计算均值即可,最后可以根据需求调整输出形状。
完整代码示例
import numpy as np from sklearn.feature_extraction.image import extract_patches_2d # 构造示例三维数组:5个5x5的2D数组 a = np.random.rand(5, 5, 5) patch_size = (3, 3) p = patch_size[0] # 计算补丁的输出维度 M, H, W = a.shape num_patches_h = H - p + 1 num_patches_w = W - p + 1 # 用as_strided创建全向量化的补丁视图 patches = np.lib.stride_tricks.as_strided( a, shape=(M, num_patches_h, num_patches_w, p, p), strides=(a.strides[0], a.strides[1], a.strides[2], a.strides[1], a.strides[2]) ) # 计算每个补丁的均值:对最后两个轴取均值 patch_means = patches.mean(axis=(-2, -1)) # 如果需要把每个外层数组的补丁均值展平成一维,可以reshape patch_means_flat = patch_means.reshape(M, -1) # 对比循环调用extract_patches_2d的结果(验证正确性) def loop_version(arr): means = [] for img in arr: patches = extract_patches_2d(img, patch_size) means.append(patches.mean(axis=(1,2))) return np.array(means) # 验证结果一致(浮点误差范围内) assert np.allclose(patch_means_flat, loop_version(a))
关键说明
- 效率优势:
as_strided创建的是视图而非副本,比循环调用extract_patches_2d(每次调用都会生成新数组)快得多,尤其是当M很大时。 - 步长调整:如果需要非1的步长,比如步长为2,只需要修改形状和跨步:
stride = 2 num_patches_h = (H - p) // stride + 1 num_patches_w = (W - p) // stride + 1 patches = np.lib.stride_tricks.as_strided( a, shape=(M, num_patches_h, num_patches_w, p, p), strides=(a.strides[0], a.strides[1]*stride, a.strides[2]*stride, a.strides[1], a.strides[2]) ) - 注意事项:使用
as_strided必须确保形状和跨步的计算准确,否则会导致内存越界访问。务必验证结果的正确性(比如和循环版本对比)。
内容的提问来源于stack exchange,提问作者Brad Solomon
相关产品推荐
相关产品推荐

