Spec-Zone.ru › PyTorch 1

torch.Tensor.index_reduce_

Tensor.index_reduce_(dim, index, source, reduce, *, include_self=True) → Tensor

Накапливать элементы source в тензор self, накапливая по индексам в порядке, указанном в index, используя операцию сокращения, заданную аргументом reduce. Например, если dim == 0, index[i] == j, reduce == prod и include_self == True, то i-я строка source умножается на j-ю строку self. Если include_self="True", значения в тензоре self включаются в сокращение, иначе строки в тензоре self обрабатываются так, как если бы они были заполнены тождественными элементами для выбранного типа сокращения.

Размерность dim тензора source должна совпадать с длиной index (которая должна быть вектором), а все остальные размерности должны совпадать с self, иначе произойдет ошибка.

Для 3-мерного тензора с reduce="prod" и include_self=True результат вычисляется следующим образом:

self[index[i], :, :] *= src[i, :, :]  # if dim == 0
self[:, index[i], :] *= src[:, i, :]  # if dim == 1
self[:, :, index[i]] *= src[:, :, i]  # if dim == 2

Примечание

Данная операция может вести себя недетерминированно при использовании тензоров на устройстве CUDA. Для получения дополнительной информации см. Воспроизводимость.

Примечание

Эта функция поддерживает только тензоры с плавающей точкой.

Предупреждение

Эта функция находится в стадии бета-тестирования и может быть изменена в ближайшем будущем.

Параметры:
  • dim (int) – размерность для индексирования
  • index (Tensor) – индексы source для выбора, тип данных должен быть либо torch.int64, либо torch.int32
  • source (FloatTensor) – тензор, содержащий значения для накопления
  • reduce (str) – операция сокращения ("prod", "mean", "amax", "amin")
Ключевые аргументы:

include_self (bool) – включать ли элементы из тензора self в сокращение

Пример:

>>> x = torch.empty(5, 3).fill_(2)
>>> t = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12]], dtype=torch.float)
>>> index = torch.tensor([0, 4, 2, 0])
>>> x.index_reduce_(0, index, t, 'prod')
tensor([[20., 44., 72.],
        [ 2.,  2.,  2.],
        [14., 16., 18.],
        [ 2.,  2.,  2.],
        [ 8., 10., 12.]])
>>> x = torch.empty(5, 3).fill_(2)
>>> x.index_reduce_(0, index, t, 'prod', include_self=False)
tensor([[10., 22., 36.],
        [ 2.,  2.,  2.],
        [ 7.,  8.,  9.],
        [ 2.,  2.,  2.],
        [ 4.,  5.,  6.]])

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.Tensor.index_reduce_.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API