Функция index_select
Функция index_select выбирает элементы из тензора вдоль указанной оси на основе переданных индексов. Первым параметром функция принимает исходный тензор, вторым параметром - ось для выбора, третьим параметром - тензор с индексами. Индексы должны быть целочисленными и не выходить за пределы размера оси.
Синтаксис
torch.index_select(input, dim, index)
Пример
Давайте выберем элементы из одномерного тензора по индексам:
import torch
t = torch.tensor([10, 20, 30, 40, 50])
indices = torch.tensor([0, 2, 4])
res = torch.index_select(t, 0, indices)
print(res)
Результат выполнения кода:
tensor([10, 30, 50])
Пример
Давайте выберем строки из двумерного тензора по индексам:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
])
indices = torch.tensor([0, 2])
res = torch.index_select(t, 0, indices)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[7, 8, 9],
])
Пример
Давайте выберем столбцы из двумерного тензора по индексам:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
indices = torch.tensor([1, 2])
res = torch.index_select(t, 1, indices)
print(res)
Результат выполнения кода:
tensor([
[2, 3],
[5, 6],
])
Пример
Давайте используем неупорядоченные индексы для выбора элементов:
import torch
t = torch.tensor([100, 200, 300, 400, 500])
indices = torch.tensor([3, 1, 4, 0])
res = torch.index_select(t, 0, indices)
print(res)
Результат выполнения кода:
tensor([400, 200, 500, 100])
Смотрите также
-
функцию
take,
которая выбирает элементы по индексам из плоского представления тензора -
функцию
gather,
которая собирает значения по индексам с учетом формы тензора -
функцию
masked_select,
которая выбирает элементы по маске -
функцию
nonzero,
которая возвращает индексы ненулевых элементов