【问题标题】:numpy - minimum sub-arraynumpy - 最小子数组
【发布时间】:2020-02-26 14:17:53
【问题描述】:

我有一个 3 维 numpy ndarray

例如,一个 4x4x2 数组如下:

array = [
 [ [7 1] [8 0] [2 0] [7 1] ]
 [ [5 4] [1 4] [6 7] [8 1] ]
 [ [3 2] [4 5] [8 6] [6 2] ]
 [ [6 4] [1 2] [5 5] [7 1] ]
]

我需要找到一个最小嵌套的最内层数组(一个有两个数字的数组)和它的索引,就像 Python 通常做的那样:比较第一对元素,如果相等,比较下一对,...

示例数组的预期结果:value=[1 2]indices=(3, 1)

用于查找元素本身的纯 Python 代码如下所示:

min(nested2 for nested1 in array for nested2 in nested1)

不过,我更喜欢 numpy 解决方案,因为数组非常庞大...

【问题讨论】:

  • Python min 正在使用 Python sort,它对于整数列表的列表是“词法”的。您似乎从一个 numpy 数组开始,在这种情况下,您的最后一个表达式可以缩短为 min(array.reshape(-1,2).tolist())。即用数组reshape替换嵌套。
  • @hpaulj 是的,重塑效果很好。但是.tolist() 每次迭代的平均成本约为 5.4 秒。其他一切都在 100 毫秒内运行。所以我希望把繁重的工作交给 numpy...

标签: numpy multidimensional-array numpy-ndarray


【解决方案1】:

您想要的行为取决于 Python 对列表列表的词法排序。即子列表被比较为:

In [275]: [1,2]<[2,0]                                                                          
Out[275]: True

np.sort 只对结构化数组和复数值进行词法排序。

In [288]: alist = [ 
     ...:  [ [7, 1], [8, 0], [2, 0], [7, 1] ], 
     ...:  [ [5, 4], [1, 4], [6, 7], [8, 1] ], 
     ...:  [ [3, 2], [4, 5], [8, 6], [6, 2] ], 
     ...:  [ [6, 4], [1, 2], [5, 5], [7, 1] ] 
     ...: ]                                                                                    
In [289]: arr = np.array(alist)                                                                
In [290]: arr                                                                                  
Out[290]: 
array([[[7, 1],
        [8, 0],
        [2, 0],
        [7, 1]],

       [[5, 4],
        [1, 4],
        [6, 7],
        [8, 1]],

       [[3, 2],
        [4, 5],
        [8, 6],
        [6, 2]],

       [[6, 4],
        [1, 2],
        [5, 5],
        [7, 1]]])

让我们试试复杂的路线:

In [291]: x = np.dot(arr, [1,1j])                                                              
In [292]: x                                                                                    
Out[292]: 
array([[7.+1.j, 8.+0.j, 2.+0.j, 7.+1.j],
       [5.+4.j, 1.+4.j, 6.+7.j, 8.+1.j],
       [3.+2.j, 4.+5.j, 8.+6.j, 6.+2.j],
       [6.+4.j, 1.+2.j, 5.+5.j, 7.+1.j]])
In [293]: np.min(x)                                                                            
Out[293]: (1+2j)
In [294]: np.argmin(x)                                                                         
Out[294]: 13
In [295]: np.unravel_index(13, x.shape)                                                        
Out[295]: (3, 1)

一些时间测试:

In [302]: timeit min(nested2 for nested1 in alist for nested2 in nested1)                      
2.24 µs ± 19.9 ns per loop (mean ± std. dev. of 7 runs, 100000 loops each)
In [303]: timeit min(arr.reshape(-1,2).tolist())                                               
2.53 µs ± 13.5 ns per loop (mean ± std. dev. of 7 runs, 100000 loops each)
In [304]: timeit np.min(arr.dot([1,1j]))                                                       
19.5 µs ± 46.2 ns per loop (mean ± std. dev. of 7 runs, 10000 loops each)

当数组更大时,复杂的路径可能会更快,但对于这个样本,它就差了。

===

结构化数组方法:

In [320]: import numpy.lib.recfunctions as rf                                                  
In [321]: rf.unstructured_to_structured(arr, names=['x','y'])                                  
Out[321]: 
array([[(7, 1), (8, 0), (2, 0), (7, 1)],
       [(5, 4), (1, 4), (6, 7), (8, 1)],
       [(3, 2), (4, 5), (8, 6), (6, 2)],
       [(6, 4), (1, 2), (5, 5), (7, 1)]],
      dtype=[('x', '<i8'), ('y', '<i8')])
In [322]: np.argsort(rf.unstructured_to_structured(arr, names=['x','y']).ravel())              
Out[322]: array([13,  5,  2,  8,  9,  4, 14, 11, 12,  6,  0,  3, 15,  1,  7, 10])
In [323]: timeit np.argsort(rf.unstructured_to_structured(arr, names=['x','y']).ravel())       
41.3 µs ± 192 ns per loop (mean ± std. dev. of 7 runs, 10000 loops each)

【讨论】:

  • 排序效果很好。已经从 5.5s 减少到 0.3s,我仍然可以摆脱创建一些中间数组。结构化数组看起来像是我一直缺少的东西! (我之前没有使用numpy)。 np.amin 不适用于结构化数组,因为 TypeError: cannot perform reduce with flexible type
【解决方案2】:
arr1 = np.array(array) # convert your list to numpy array
arr2 = np.sum(arr1,axis=2) # get the sum of the two elements in each array
ind1 = np.unravel_index(arr2.argmin(), arr2.shape) # the indices of the minimum sum
min1 = arr1[ind1] # the value of the minimum array
print(ind1)
(3, 1)
print(min1)
array([1, 2])

【讨论】:

  • 不幸的是,没有描述它的作用,但它绝对不能解决问题。相反,它只是使用 sum 函数减少初始数组的维度。
  • 我修复了最小值的代码。如果您的标准是数组中两个值的最小总和,这应该有效
  • "Indices of minimum sum" - 你在这里丢失了信息,因为顺序很重要。
猜你喜欢
  • 2016-06-13
  • 1970-01-01
  • 2021-07-31
  • 2018-12-28
  • 1970-01-01
  • 2018-02-13
  • 2018-07-24
  • 1970-01-01
相关资源
最近更新 更多