Функция take_along_axis
Функция take_along_axis выбирает элементы
из входного массива по индексам, переданным
во втором аргументе, вдоль указанной оси.
Первый параметр - исходный массив, второй
- массив индексов, третий - ось (axis),
вдоль которой производится выборка.
Возвращает новый массив с выбранными
элементами, имеющий ту же форму, что
и массив индексов.
Синтаксис
np.take_along_axis(arr, indices, axis)
Пример
Давайте создадим одномерный массив и выберем из него элементы по индексам с помощью аргумента axis=0 (для одномерного массива это единственная ось):
import numpy as np
arr = np.array([10, 20, 30, 40, 50])
indices = np.array([0, 3, 1])
res = np.take_along_axis(arr, indices, axis=0)
print(res)
Результат выполнения кода:
[10 40 20]
Пример
Теперь создадим двумерный массив и выберем элементы вдоль оси 0 (по строкам). Индексы указывают, из какой строки брать элемент для каждого столбца:
import numpy as np
arr = np.array([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
indices = np.array([[0, 1, 2],
[2, 0, 1]])
res = np.take_along_axis(arr, indices, axis=0)
print(res)
Результат выполнения кода:
[[1 5 9]
[7 2 6]]
Пример
А теперь выберем элементы вдоль оси 1 (по столбцам). Индексы указывают, из какого столбца брать элемент для каждой строки:
import numpy as np
arr = np.array([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
indices = np.array([[2, 0],
[1, 2],
[0, 1]])
res = np.take_along_axis(arr, indices, axis=1)
print(res)
Результат выполнения кода:
[[3 1]
[5 6]
[7 8]]
Пример
Часто take_along_axis используется вместе с argsort для сортировки массива по значениям из другого массива. Получим индексы сортировки с помощью argsort и применим их для выборки:
import numpy as np
arr = np.array([[3, 1, 2],
[6, 4, 5]])
indices = np.argsort(arr, axis=1)
res = np.take_along_axis(arr, indices, axis=1)
print(res)
Результат выполнения кода:
[[1 2 3]
[4 5 6]]
Пример
Индексы должны иметь ту же размерность, что и исходный массив, за исключением оси, вдоль которой производится выборка. Вот пример с трехмерным массивом:
import numpy as np
arr = np.array([[[1, 2],
[3, 4]],
[[5, 6],
[7, 8]]])
indices = np.array([[[0, 1],
[0, 0]],
[[1, 0],
[1, 1]]])
res = np.take_along_axis(arr, indices, axis=2)
print(res)
Результат выполнения кода:
[[[1 2]
[3 3]]
[[6 5]
[8 8]]]
Смотрите также
-
функцию
take,
которая выбирает элементы из массива по плоским индексам -
функцию
put_along_axis,
которая вставляет значения в массив по индексам вдоль оси -
функцию
compress,
которая выбирает элементы по условию -
функцию
choose,
которая выбирает элементы из нескольких массивов по индексам