Метод index_select
Метод index_select выполняет выборку элементов из тензора вдоль указанной размерности. Первым параметром передаётся размерность dim, вдоль которой происходит выборка, вторым параметром - одномерный тензор index с индексами элементов. Метод возвращает новый тензор, содержащий выбранные элементы. Размерность результата совпадает с исходной, но по выбранной оси заменяется на длину тензора индексов.
Синтаксис
tensor.index_select(dim, index)
Пример
Выполним выборку элементов из одномерного тензора по заданным индексам:
import torch
t = torch.tensor([10, 20, 30, 40, 50])
indices = torch.tensor([0, 2, 4])
res = t.index_select(dim=0, index=indices)
print(res)
Результат выполнения кода:
tensor([10, 30, 50])
Пример
Выполним выборку строк из двумерного тензора по указанным индексам:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
[10, 11, 12],
])
indices = torch.tensor([0, 2])
res = t.index_select(dim=0, index=indices)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[7, 8, 9],
])
Пример
Выполним выборку столбцов из двумерного тензора по указанным индексам:
import torch
t = torch.tensor([
[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12],
])
indices = torch.tensor([1, 3])
res = t.index_select(dim=1, index=indices)
print(res)
Результат выполнения кода:
tensor([
[2, 4],
[6, 8],
[10, 12],
])
Пример
Используем метод index_select для выборки элементов из трёхмерного тензора вдоль оси глубины:
import torch
t = torch.tensor([
[
[1, 2],
[3, 4],
],
[
[5, 6],
[7, 8],
],
[
[9, 10],
[11, 12],
],
])
indices = torch.tensor([0, 2])
res = t.index_select(dim=0, index=indices)
print(res)
Результат выполнения кода:
tensor([
[
[1, 2],
[3, 4],
],
[
[9, 10],
[11, 12],
],
])
Смотрите также
-
метод
masked_select,
который выбирает элементы по маске -
метод
gather,
который собирает элементы по индексам с поддержкой многомерной выборки -
метод
scatter,
который записывает значения по индексам -
метод
index_fill_,
который заполняет элементы по индексам указанным значением