如何遍历3D数组中所有相邻2x2子矩阵?有无Numpy/PyTorch内置方法?
完全不需要自己手动实现遍历,Numpy和PyTorch都有专门的内置工具可以高效完成这个需求,具体方案如下:
Numpy 实现
使用numpy.lib.stride_tricks.sliding_window_view(Numpy 1.20及以上版本支持),它能直接生成滑动窗口视图,不会额外复制数据,效率极高。
import numpy as np # 原数组 array = np.array([[[ 72, 11, 158, 252], [160, 50, 131, 174], [245, 127, 99, 6], [152, 25, 58, 96]], [[ 29, 37, 211, 215], [195, 72, 186, 33], [ 12, 68, 44, 241], [ 95, 184, 188, 176]], [[238, 90, 177, 15], [ 48, 221, 41, 236], [ 86, 14, 130, 192], [ 64, 17, 44, 251]]]) # 生成所有2x2滑动窗口,针对每个矩阵的行、列维度滑动 windows = np.lib.stride_tricks.sliding_window_view(array, window_shape=(2, 2), axis=(-2, -1)) # 调整形状,把每个原矩阵的子矩阵整理成一维列表(可选,按需调整) windows_reshaped = windows.reshape(3, -1, 2, 2) # 遍历第一个矩阵的所有子矩阵 print("第一个矩阵的2x2子矩阵:") for sub_mat in windows_reshaped[0]: print(sub_mat) print()
生成的windows_reshaped形状为(3, 9, 2, 2),对应:3个原矩阵,每个原矩阵包含9个2x2子矩阵。
PyTorch 实现
使用PyTorch的torch.Tensor.unfold方法,通过两次unfold分别处理行、列维度,提取滑动窗口。
import torch # 转成PyTorch张量 tensor = torch.tensor(array) # 先在列维度提取窗口大小为2、步长为1的片段 unfolded_col = tensor.unfold(dim=2, size=2, step=1) # 再在行维度提取窗口大小为2、步长为1的片段 unfolded = unfolded_col.unfold(dim=1, size=2, step=1) # 调整维度顺序并重塑,得到(3, 9, 2, 2)的结果 unfolded_reshaped = unfolded.permute(0, 1, 2, 4, 3).reshape(3, -1, 2, 2) # 遍历第一个矩阵的所有子矩阵 print("第一个矩阵的2x2子矩阵:") for sub_mat in unfolded_reshaped[0]: print(sub_mat.numpy()) print()
内容的提问来源于stack exchange,提问作者Tamir
相关产品推荐
相关产品推荐

