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(),您只需要执行一次昂贵的部分。