如何高效将二维NumPy数组按窗口大小分割为三维数组?
高效分割二维NumPy数组为滑动窗口三维数组的方法
问题描述
需要将形状为(M, N)的二维NumPy数组按指定窗口大小分割为三维数组,当前通过循环实现的代码执行速度极慢:
X = list() for j in range(size): end_idx = j + seq if end_idx >= size: break seq_x = data[j:end_idx, :] X.append(seq_x) final_data = np.array(X)
数据示例与预期输出
示例输入data:
import numpy as np data = np.array([ [0, 1], [2, 3], [3, 4], [4, 5], [5, 6], [6, 7], [7, 8], [8, 9], [9, 7] ])
当窗口大小w=2时,预期输出是形状为(8, 2, 2)的三维数组,每个子数组对应原始数据中连续的2行:
res = np.array([ [[0, 1], [2, 3]], [[2, 3], [3, 4]], [[3, 4], [4, 5]], # ... 中间省略 ... [[8, 9], [9, 7]] ])
高效实现方案
1. 使用np.lib.stride_tricks.sliding_window_view(NumPy 1.20+ 推荐)
这是NumPy官方提供的滑动窗口工具,基于数组stride机制创建视图,无需复制数据,性能远超循环:
import numpy as np w = 2 # 窗口大小 # 生成滑动窗口视图,自动处理边界 windowed_data = np.lib.stride_tricks.sliding_window_view(data, window_shape=(w, data.shape[1])) # 去掉多余的中间维度,得到目标形状 res = windowed_data.squeeze(axis=2)
- 输出形状:
(M - w + 1, w, N),完全符合预期。
2. 使用np.lib.stride_tricks.as_strided(兼容旧版本NumPy)
如果你的NumPy版本低于1.20,可以手动通过as_strided实现,同样基于stride机制:
import numpy as np w = 2 M, N = data.shape # 定义新数组的形状 new_shape = (M - w + 1, w, N) # 计算stride(字节为单位,复用原始数组的内存步长) row_stride = data.strides[0] elem_stride = data.strides[1] new_strides = (row_stride, row_stride, elem_stride) # 创建视图,无内存复制 res = np.lib.stride_tricks.as_strided(data, shape=new_shape, strides=new_strides)
- 注意:使用
as_strided时需确保计算的new_shape和new_strides正确,避免越界访问内存。
性能优势
循环实现每次切片都会生成新数组并复制数据,而上述两种方法都是创建原始数据的视图,没有内存复制操作,速度可提升数倍至数十倍,尤其当数组规模较大时差异更明显。
内容的提问来源于stack exchange,提问作者narutoArea51
相关产品推荐
相关产品推荐

