【问题标题】:numpy multidimensional indexing and the function 'take'numpy 多维索引和函数“take”
【发布时间】:2017-07-23 06:00:05
【问题描述】:

在一周中的奇数天,我几乎了解 numpy 中的多维索引。 Numpy 有一个函数'take',它似乎可以做我想做的事,但额外的好处是,如果索引超出范围,我可以控制会发生什么 具体来说,我有一个 3 维数组作为查找表来询问

lut = np.ones([13,13,13],np.bool)

和一个 2x2 的 3 长向量数组作为表的索引

arr = np.arange(12).reshape([2,2,3]) % 13 

IIUC,如果我要写 lut[arr],那么 arr 被视为 2x2x3 数字数组,当这些数字用作 lut 的索引时,它们每个都返回一个 13x13 数组。这就解释了为什么lut[arr].shape is (2, 2, 3, 13, 13)

我可以通过写作让它做我想做的事

lut[ arr[:,:,0],arr[:,:,1],arr[:,:,2] ] #(is there a better way to write this?)

现在这三个术语的行为就好像它们已被压缩以生成一个 2x2 元组数组,而 lut[<tuple>]lut 生成单个元素。最终结果是来自lut 的 2x2 条目数组,正是我想要的。

我已经阅读了“take”功能的文档...

此功能与“花式”索引功能相同 (使用数组索引数组);但是,它可以更容易 如果您需要沿给定轴的元素,请使用。

轴:int,可选
选择值的轴。

也许天真地,我认为设置 axis=2 会得到三个值用作 3 元组来执行查找,但实际上

np.take(lut,arr).shape =  (2, 2, 3)
np.take(lut,arr,axis=0).shape =  (2, 2, 3, 13, 13)
np.take(lut,arr,axis=1).shape =  (13, 2, 2, 3, 13)
np.take(lut,arr,axis=2).shape =  (13, 13, 2, 2, 3)

所以很明显我不明白发生了什么。谁能告诉我如何实现我想要的?

【问题讨论】:

    标签: python numpy multidimensional-array indexing


    【解决方案1】:

    我们可以计算线性索引,然后使用np.take -

    np.take(lut, np.ravel_multi_index(arr.T, lut.shape)).T
    

    如果您对替代方案持开放态度,我们可以将索引数组重新整形为2D,转换为元组,用它索引到数据数组中,给我们1D,它可以重新整形回2D -

    lut[tuple(arr.reshape(-1,arr.shape[-1]).T)].reshape(arr.shape[:2])
    

    示例运行 -

    In [49]: lut = np.random.randint(11,99,(13,13,13))
    
    In [50]: arr = np.arange(12).reshape([2,2,3])
    
    In [51]: lut[ arr[:,:,0],arr[:,:,1],arr[:,:,2] ] # Original approach
    Out[51]: 
    array([[41, 21],
           [94, 22]])
    
    In [52]: np.take(lut, np.ravel_multi_index(arr.T, lut.shape)).T
    Out[52]: 
    array([[41, 21],
           [94, 22]])
    
    In [53]: lut[tuple(arr.reshape(-1,arr.shape[-1]).T)].reshape(arr.shape[:2])
    Out[53]: 
    array([[41, 21],
           [94, 22]])
    

    我们可以避免 np.take 方法的双重转置,就像这样 -

    In [55]: np.take(lut, np.ravel_multi_index(arr.transpose(2,0,1), lut.shape))
    Out[55]: 
    array([[41, 21],
           [94, 22]])
    

    泛化为通用维度的多维数组

    这可以推广到通用编号的 ndarrays。昏暗的,像这样 -

    np.take(lut, np.ravel_multi_index(np.rollaxis(arr,-1,0), lut.shape))
    

    tuple-based 方法无需任何更改即可工作。

    这是相同的运行示例 -

    In [95]: lut = np.random.randint(11,99,(13,13,13,13))
    
    In [96]: arr = np.random.randint(0,13,(2,3,4,4))
    
    In [97]: lut[ arr[:,:,:,0] , arr[:,:,:,1],arr[:,:,:,2],arr[:,:,:,3] ]
    Out[97]: 
    array([[[95, 11, 40, 75],
            [38, 82, 11, 38],
            [30, 53, 69, 21]],
    
           [[61, 74, 33, 94],
            [90, 35, 89, 72],
            [52, 64, 85, 22]]])
    
    In [98]: np.take(lut, np.ravel_multi_index(np.rollaxis(arr,-1,0), lut.shape))
    Out[98]: 
    array([[[95, 11, 40, 75],
            [38, 82, 11, 38],
            [30, 53, 69, 21]],
    
           [[61, 74, 33, 94],
            [90, 35, 89, 72],
            [52, 64, 85, 22]]])
    

    【讨论】:

    • 不管怎样,np.take 等同于索引lut.flat: lut.flat[np.ravel_multi_index(arr.T, lut.shape)].T
    • 感谢两位的解决方案。由于 ravel_multi_index 允许剪裁,我根本不需要“拍摄”,所以我会选择 @hjpauli 对 Divakar 解决方案的修正
    【解决方案2】:

    我没有尝试 3 维。但是在 2-dimensions 中,我使用 numpy.take 得到了我想要的结果:

    np.take(np.take(T,ix,axis=0), iy,axis=1 )
    

    也许您可以将其扩展到 3 维。

    作为一个示例,我可以为索引 ix 和 iy 使用两个 1-dim 数组来寻址离散拉普拉斯方程的二维模板,

    ΔT = T[ix-1,iy] + T[ix+1, iy] + T[ix,iy-1] + T[ix,iy+1] - 4*T[ix,iy]
    

    介绍更精简的写作:

    def q(Φ,kx,ky):
        return np.take(np.take(Φ,kx,axis=0), ky,axis=1 )
    

    然后我可以使用 numpy.take 运行以下 python 代码

    nx = 6; ny= 10
    T  = np.arange(nx*ny).reshape(nx, ny)
    
    ix = np.linspace(1,nx-2,nx-2,dtype=int) 
    iy = np.linspace(1,ny-2,ny-2,dtype=int)
    
    ΔT = q(T,ix-1,iy)  + q(T,ix+1,iy)  + q(T,ix,iy-1)  + q(T,ix,iy+1)  - 4.0 * q(T,ix,iy)
    

    【讨论】:

      【解决方案3】:

      最初的问题是尝试在表中进行查找,但 一些索引超出范围,我想控制 发生这种情况时的行为。

      import numpy as np
      lut = np.ones((5,7,11),np.int) # a 3-dimensional lookup table
      print("lut.shape = ",lut.shape ) # (5,7,11)
      
      # valid points are in the interior with value 99,
      # invalid points are on the faces with value 0
      lut[:,:,:] = 0
      lut[1:-1,1:-1,1:-1] = 99
      
      # set up an array of indexes with many of them too large or too small
      start = -35
      arr = np.arange(start,2*11*3+start,1).reshape(2,11,3)
      
      # This solution has the advantage that I can understand what is going on
      # and so I can amend it if I need to
      
      # split arr into tuples along axis=2
      arrchannels = arr[:,:,0],arr[:,:,1],arr[:,:,2]
      
      # convert into a flat array but clip the values
      ravelledarr = np.ravel_multi_index(arrchannels, lut.shape, mode='clip')
      
      # and now turn back into a list of numpy arrays
      # (not an array of the original shape )
      clippedarr = np.unravel_index( ravelledarr, lut.shape)
      print(clippedarr[0].shape,"*",len(clippedarr)) # produces (2, 11) * 3
      
      # and now I can do the lookup with the indexes clipped to fit
      print(lut[clippedarr])
      
      # these are more succinct but opaque ways of doing the same
      # due to @Divakar and @hjpauli respectively
      print( np.take(lut, np.ravel_multi_index(arr.T, lut.shape, mode='clip')).T )
      print( lut.flat[np.ravel_multi_index(arr.T, lut.shape, mode='clip')].T )
      

      实际的应用是我有一个 rgb 图像,其中包含一些带有一些标记的纹理木材,并且我已经确定了它的一部分。我想获取这个补丁中的一组像素,并标记整个图像中与其中一个匹配的所有点。 256x256x256 存在表太大,所以我对补丁中的像素运行聚类算法并为每个集群设置存在表(补丁中的颜色通过 rgb 或 hsv 空间形成细长的线,因此集群周围的框很小)。

      我使存在表比需要的稍大,并用 False 填充每个面。

      一旦我设置了这些小的存在表,我现在可以通过查找表中的每个像素并使用剪裁来制作通常不会映射到表中的像素来测试图像的其余部分是否匹配补丁实际上映射到表的一个面(并得到值'False')

      【讨论】:

        猜你喜欢
        • 1970-01-01
        • 2015-04-19
        • 1970-01-01
        • 1970-01-01
        • 2020-03-25
        • 1970-01-01
        • 2020-11-18
        • 2021-07-27
        • 2018-01-22
        相关资源
        最近更新 更多