【问题标题】:Numpy: Reshape/horizontally split 3D array into 4D arrayNumpy:重塑/水平分割 3D 数组为 4D 数组
【发布时间】:2021-08-31 11:06:42
【问题描述】:

我确实有一个像这样的 3D np.array:

arr3d = np.arange(36).reshape(3, 2, 6)

array([[[ 0,  1,  2,  3,  4,  5],
        [ 6,  7,  8,  9, 10, 11]],

       [[12, 13, 14, 15, 16, 17],
        [18, 19, 20, 21, 22, 23]],

       [[24, 25, 26, 27, 28, 29],
        [30, 31, 32, 33, 34, 35]]])

我需要将arr3d的每个窗格水平拆分成3个块,例如:

np.array(np.hsplit(arr3d[0, :, :], 3))

array([[[ 0,  1],
        [ 6,  7]],

       [[ 2,  3],
        [ 8,  9]],

       [[ 4,  5],
        [10, 11]]])

这应该会导致一个 4D 数组。

arr4d[0, :, :, :] 应包含原始 3D 数组的第一个窗格的新拆分 3D 数组 (np.array(np.hsplit(arr3d[0, :, :], 3)))

最终的结果应该是这样的:

result = np.array(
    [
        [[[0, 1], [6, 7]], [[2, 3], [8, 9]], [[4, 5], [10, 11]]],
        [[[12, 13], [18, 19]], [[14, 15], [20, 21]], [[16, 17], [22, 23]]],
        [[[24, 25], [30, 31]], [[26, 27], [32, 33]], [[28, 29], [34, 35]]],
    ]
)

result.shape
(3, 3, 2, 2)

array([[[[ 0,  1],
         [ 6,  7]],

        [[ 2,  3],
         [ 8,  9]],

        [[ 4,  5],
         [10, 11]]],


       [[[12, 13],
         [18, 19]],

        [[14, 15],
         [20, 21]],

        [[16, 17],
         [22, 23]]],


       [[[24, 25],
         [30, 31]],

        [[26, 27],
         [32, 33]],

        [[28, 29],
         [34, 35]]]])

我正在寻找一种 Pythonic 方式来执行此重塑/拆分。

【问题讨论】:

    标签: python numpy reshape


    【解决方案1】:

    试试:

    sh = arr3d.shape[:-1] + (3, -1)
    arr4d = arr3d.reshape(*sh).swapaxes(1, 2)
    
    >>> arr4d
    array([[[[ 0,  1],
             [ 6,  7]],
    
            [[ 2,  3],
             [ 8,  9]],
    
            [[ 4,  5],
             [10, 11]]],
    
    
           [[[12, 13],
             [18, 19]],
    
            [[14, 15],
             [20, 21]],
    
            [[16, 17],
             [22, 23]]],
    
    
           [[[24, 25],
             [30, 31]],
    
            [[26, 27],
             [32, 33]],
    
            [[28, 29],
             [34, 35]]]])
    

    说明

    这是您要拆分为(3, -1) 的最后一个维度(在您的示例中,大小为 6)。这就是为什么我们首先重塑为(a, b, 3, -1)(其中(a, b, _) 是arr3d 的形状)。但是因为你对每一行做了一个hsplit(),那么你想要的实际形状是(a, 3, b, -1),所以我们需要交换轴1和2(更准确地说:滚动它们,我们将看到下面是更高的尺寸)。

    另一个例子

    shape = 7, 2, 3*3
    arr3d = np.arange(np.prod(shape)).reshape(*shape)
    check = np.array([np.array(np.hsplit(arr3d[k], 3)) for k in range(shape[0)])
    
    sh = arr3d.shape[:-1] + (3, -1)
    arr4d = arr3d.reshape(*sh).swapaxes(1, 2)
    >>> np.equal(arr4d, check).all()
    True
    

    泛化到更高维度

    shape = 4, 5, 2, 3*3
    ar = np.arange(np.prod(shape)).reshape(*shape)
    check = np.array([np.array(np.split(ar[k], 3, axis=-1)) for k in range(shape[0])])
    
    # any dimension
    sh = ar.shape[:-1] + (3, -1)
    out = np.rollaxis(ar.reshape(*sh), -2, 1)
    >>> np.equal(out, check).all()
    True
    

    【讨论】:

    • 不错的答案。这不是微不足道的,必须考虑一段时间;-) 老实说,您的 check = ... 比使用 reshape/swapaxes 更直观,但可能效率不高。
    • 对n D --> (n+1) D 的概括特别有趣...
    猜你喜欢
    • 2018-03-17
    • 2021-08-06
    • 2021-10-14
    • 2016-02-10
    • 2017-09-18
    • 2020-04-22
    • 2016-05-22
    • 1970-01-01
    • 2015-10-19
    相关资源
    最近更新 更多