如何无循环实现3D NumPy数组N次滚动并保留所有版本
NumPy数组滚动并保留每次结果的实现需求
我正在使用NumPy处理一个二进制数组G,结构如下:
import numpy as np G = np.array([ [ [1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.], [1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.], [1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.] ], [ [0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0., 0.], [0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0., 0.], [0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0., 0.] ] ])
我需要实现一种方法,将该数组沿最后一个维度滚动N次,并保留每次滚动后的版本。例如当N=3时,输出结构如下:
output = np.array([ [ [ [1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.], [1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.], [1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.] ], [ [0., 1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.], [0., 1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.], [0., 1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.] ], [ [0., 0., 1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.], [0., 0., 1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.], [0., 0., 1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.] ] ], [ [ [0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0., 0.], [0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0., 0.], [0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0., 0.] ], [ [0., 0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0.], [0., 0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0.], [0., 0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0.] ], [ [0., 0., 0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0.], [0., 0., 0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0.], [0., 0., 0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0.] ] ] ])
解决方案
可以结合NumPy的np.roll和np.stack函数高效实现需求,避免循环带来的性能损耗,具体代码如下:
import numpy as np def roll_and_stack(arr, N): # 生成0到N-1的滚动偏移量 shifts = np.arange(N) # 对每个偏移量执行滚动,再沿轴1堆叠结果 rolled_arrays = [np.roll(arr, shift, axis=-1) for shift in shifts] return np.stack(rolled_arrays, axis=1) # 测试示例 G = np.array([ [ [1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.], [1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.], [1., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.] ], [ [0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0., 0.], [0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0., 0.], [0., 0., 0., 1., 0., 0., 1., 0., 0., 0., 0., 0., 0., 0.] ] ]) output = roll_and_stack(G, 3) print(output.shape) # 输出 (2, 3, 3, 14),与示例结构一致
关键说明
np.roll(arr, shift, axis=-1):将数组沿最后一个维度循环滚动shift个位置,正数表示向右滚动。np.stack(..., axis=1):把N次滚动后的数组在轴1位置堆叠,完美匹配需求的输出结构。- 该方法比循环逐个处理更高效,尤其适合数组规模大或N值较大的场景。
内容的提问来源于stack exchange,提问作者OIBOI
相关产品推荐
相关产品推荐

