【问题标题】:Sort paired array of 3d array (replace for loop)对 3d 数组的配对数组进行排序(替换循环)
【发布时间】:2020-05-01 08:48:13
【问题描述】:

我有以下 3d 数组:

import numpy as np

z = np.array([[[10,  2],
               [ 5,  3],
               [ 4,  4]],
              [[ 7,  6],
               [ 4,  2],
               [ 5,  8]]])

我想根据 3rd dim & 1st value 对它们进行排序。

目前我正在使用以下代码:

from operator import itemgetter

np.array([sorted(x,key=itemgetter(0)) for x in z])
array([[[ 4,  4],
        [ 5,  3],
        [10,  2]],

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

我想通过删除 for 循环使代码更高效/更快?

【问题讨论】:

    标签: python arrays numpy sorting vectorization


    【解决方案1】:

    对于 numpy one 班轮,您可以使用 numpy.argsort:

    import numpy as np
    
    a = np.array([[[10,  2],
                   [ 5,  3],
                   [ 4,  4]],
                  [[ 7,  6],
                   [ 4,  2],
                   [ 5,  8]]])
    
    a[np.arange(0,2)[:,None], a[:,:,0].argsort()]
    array([[[ 4,  4],
            [ 5,  3],
            [10,  2]],
           [[ 4,  2],
            [ 5,  8],
            [ 7,  6]]])
    

    对于如此小尺寸的数组,这大约需要相同的时间,但扩大尺寸会带来相当大的改进,例如:

    from operator import itemgetter
    
    a = np.random.randint(0,10, (2,100_000,2))
    
    %timeit a[np.arange(0,2)[:,None], a[:,:,0].argsort()]
    26.9 ms ± 351 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)
    
    %timeit [sorted(x,key=itemgetter(0)) for x in a]
    327 ms ± 6.39 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
    

    【讨论】:

      【解决方案2】:

      您可以使用map() 来获得相同的结果,而无需使用for 循环。并且排序函数可以是用户定义的,也可以是 lambda,或者是sorted 的一部分:

      1. 首先创建一个排序函数:

        >>> def mysort(it):
        ...   return sorted(it, key=itemgetter(0))
        ...
        >>> list(map(mysort, z))
        [[[4, 4], [5, 3], [10, 2]], [[4, 2], [5, 8], [7, 6]]]
        
      2. 与上述相同,但使用 lambda:

        >>> list(map(lambda it: sorted(it, key=itemgetter(0)), z))
        [[[4, 4], [5, 3], [10, 2]], [[4, 2], [5, 8], [7, 6]]]
        
      3. partial:

        >>> from functools import partial
        >>> psort = partial(sorted, key=itemgetter(0))
        >>> list(map(psort, z))
        [[[4, 4], [5, 3], [10, 2]], [[4, 2], [5, 8], [7, 6]]]
        

        或部分就地定义:

        >>> list(map(partial(sorted, key=itemgetter(0)), z))
        [[[4, 4], [5, 3], [10, 2]], [[4, 2], [5, 8], [7, 6]]]
        
      4. 您的问题有一个列表列表,而不是 3d numpy 数组。对于面向 numpy 的解决方案,see this answer

      仅供参考,(2)和(3b)是roughly equivalent, but have their differences
      在选项 1-3 中,我更喜欢 (2) 中的 lambda。

      【讨论】:

      • 感谢您的建议。我对它们都进行了测试,似乎 for 循环仍然是最快的。你觉得这听起来对吗?
      • 您是否正在对照原始列表z 的更长版本进行检查?对于如此短的列表,我认为性能差异并不明显。
      【解决方案3】:

      为什么不简单:np.sort(z,axis=1)

      import numpy as np
      
      z = np.array([[[10,  2],
                     [ 5,  3],
                     [ 4,  4]],
                    [[ 7,  6],
                     [ 4,  2],
                     [ 5,  8]]])
      
      print(np.sort(z,axis=1))
      
      [[[ 4  2]
        [ 5  3]
        [10  4]]
      
       [[ 4  2]
        [ 5  6]
        [ 7  8]]]
      

      【讨论】:

        猜你喜欢
        • 1970-01-01
        • 2022-09-23
        • 2021-10-11
        • 2019-03-06
        • 1970-01-01
        • 1970-01-01
        • 1970-01-01
        • 2021-10-15
        • 1970-01-01
        相关资源
        最近更新 更多