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

如何向量化实现从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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 07:24:13