torch.gather
-
torch.gather(input, dim, index, *, sparse_grad=False, out=None) → Tensor -
Сбор значений по оси, указанной параметром
dim.Для 3-мерного тензора вывод определяется следующим образом:
out[i][j][k] = input[index[i][j][k]][j][k] # if dim == 0 out[i][j][k] = input[i][index[i][j][k]][k] # if dim == 1 out[i][j][k] = input[i][j][index[i][j][k]] # if dim == 2
inputиindexдолжны иметь одинаковое количество измерений. Также необходимо, чтобыindex.size(d) <= input.size(d)для всех измеренийd != dim.outбудет иметь такую же форму, какindex. Обратите внимание, чтоinputиindexне производят распределение друг против друга.- Параметры:
- Ключевые аргументы:
Пример:
>>> t = torch.tensor([[1, 2], [3, 4]]) >>> torch.gather(t, 1, torch.tensor([[0, 0], [1, 0]])) tensor([[ 1, 1], [ 4, 3]])
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.gather.html