torch.take_along_dim
-
torch.take_along_dim(input, indices, dim, *, out=None) → Tensor -
Выбирает значения из
inputпо одномерным индексам изindicesвдоль указаннойdim.Функции, возвращающие индексы вдоль размерности, такие как
torch.argmax()иtorch.argsort(), предназначены для работы с этой функцией. См. примеры ниже.Примечание
Эта функция аналогична функции NumPy’s
take_along_axis. См. такжеtorch.gather().- Параметры:
- Ключевые аргументы:
-
out (Тензор, необязательно) – выходной тензор.
Пример:
>>> t = torch.tensor([[10, 30, 20], [60, 40, 50]]) >>> max_idx = torch.argmax(t) >>> torch.take_along_dim(t, max_idx) tensor([60]) >>> sorted_idx = torch.argsort(t, dim=1) >>> torch.take_along_dim(t, sorted_idx, dim=1) tensor([[10, 20, 30], [40, 50, 60]])
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.take_along_dim.html