【问题标题】:Move position of row in Numpy 2d array在 Numpy 二维数组中移动行的位置
【发布时间】:2020-05-18 17:34:54
【问题描述】:

给定数组:

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

如果我想将索引 1 处的行移动到索引 3。 输出应该是:

[[1, 2],
 [5, 6],
 [7, 8],
 [3, 4],
 [9, 10]]

如果我想将索引 4 处的行移动到索引 1。 输出应该是:

[[1, 2],
 [9, 10],
 [3, 4],
 [5, 6],
 [7, 8]]

执行此移动操作的最快方法是什么?

【问题讨论】:

  • 这能回答你的问题吗? Rearrange columns of numpy 2D array
  • 速度可能取决于必须移动(或向上/向下移动)的行数。从概念上讲,最简单的方法是使用高级索引创建一个新数组(如@norok2 所示)。如果只切换更大数组的几行(但您的示例移动了 5 个中的 3 个和 5 个中的 4 个),则就地更改可能会更快。
  • 是的,速度肯定与受影响的行数成正比。高级索引会访问所有数组元素,这就是为什么当行增加时它的伸缩性很差。

标签: numpy


【解决方案1】:

如果你仔细看,如果你想把行i放在j的位置,那么只有ij之间的行会受到影响;外面的行不需要改变。而这个改动基本上就是roll操作。对于a,b,c,d,e 的项目,将项目放在i=1j=3 意味着b,c,d 将变为c,d,b,给我们a,c,d,b,e。班次是-1+1,具体取决于i<j

i, j = 1,3
i, j, s = (i, j, -1) if i<j else (j, i, 1)
arr[i:j+1] = np.roll(arr[i:j+1],shift=s,axis=0)

【讨论】:

  • 当然。虽然简单地插入正确排列的索引很容易,但您正在手动计算这些索引;这很差。想象一下,您有一个大小为 (5000,5000) 的数组,不可能像您写的 (0,2,3,1,4) 那样手动编写索引。相反,您需要编写某种逻辑来创建正确排列的索引,这将消耗相当于长度为 5000 的数组的内存开销。
  • 我认为您的评论应该出现在 norok 的回答中。但你是对的,我的数组很大,因此我相信你的解决方案是最优的。
  • @Mercury 这与速度无关。查看我的答案的修改。
【解决方案2】:

tuple() 第一个轴上的索引怎么样?

例如:

arr[(0, 2, 3, 1, 4), :]

和:

arr[(0, 4, 1, 2, 3), :]

分别用于您的预期输出。


对于从两个索引开始生成索引的方法,您可以使用以下内容:

def inner_roll(arr, first, last, axis):
    stop = last + 1
    indices = list(range(arr.shape[axis]))
    indices.insert(first, last)
    indices.pop(last + 1)
    slicing = tuple(
        slice(None) if i != axis else indices
        for i, d in enumerate(arr.shape))
    return arr[slicing]

对于沿您正在操作的轴相对较小的输入(例如问题中的输入),这非常快。

将其与@Mercury's answer 的稍微完善的版本进行比较,以将其包装在一个函数中并使其对任意axis 正常工作:

import numpy as np


def inner_roll2(arr, first, last, axis):
    if first > last:
        first, last = last, first
        shift = 1
    else:
        shift = -1
    slicing = tuple(
        slice(None) if i != axis else slice(first, last + 1)
        for i, d in enumerate(arr.shape))
    arr[slicing] = np.roll(arr[slicing], shift=shift, axis=axis)
    return arr

并获得一些时间安排:

funcs = inner_roll, inner_roll2
for n in (5, 50, 500):
    for m in (2, 20, 200):
        arr = np.arange(n * m).reshape((n, m))
        print(f'({n:<3d}, {m:<3d})', end='    ')
        for func in funcs:
            results = %timeit -o -q func(arr, 1, 2, 0)
            print(f'{func.__name__:>12s}  {results.best* 1e6:>7.3f} µs', end='    ')
        print()
# (5  , 2  )      inner_roll    5.613 µs     inner_roll2   15.393 µs    
# (5  , 20 )      inner_roll    5.592 µs     inner_roll2   15.468 µs    
# (5  , 200)      inner_roll    5.916 µs     inner_roll2   15.815 µs    
# (50 , 2  )      inner_roll   10.117 µs     inner_roll2   15.517 µs    
# (50 , 20 )      inner_roll   10.360 µs     inner_roll2   15.505 µs    
# (50 , 200)      inner_roll   12.067 µs     inner_roll2   15.886 µs    
# (500, 2  )      inner_roll   55.833 µs     inner_roll2   15.409 µs    
# (500, 20 )      inner_roll   57.364 µs     inner_roll2   15.319 µs    
# (500, 200)      inner_roll  194.408 µs     inner_roll2   15.731 µs    

这表明inner_roll() 是您输入的最快方法。 然而,inner_roll2() 似乎可以更好地适应输入大小,即使是适度的输入大小,这已经比 inner_roll() 快。

请注意,虽然inner_roll() 创建了一个副本,但inner_roll2() 在原地工作(修改输入arr)。可以通过在 inner_roll2() 的主体的开头添加 arr = arr.copy() 来修改此行为,这会使该函数变慢(当然),并且其时间将受到 m 值的更大影响(大小非滚动轴)。

另一方面,如果您要进行多次连续滚动操作,inner_roll2() 的时间只会叠加,而对于inner_roll(),您只需要执行一次昂贵的部分。

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2021-06-27
    • 2018-04-10
    • 2011-12-31
    • 2017-04-11
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2020-07-30
    相关资源
    最近更新 更多