Spec-Zone.ru › PyTorch 1

torch.argsort

torch.argsort(input, dim=- 1, descending=False, stable=False) → Tensor

Возвращает индексы, которые сортируют тензор вдоль заданного измерения в порядке возрастания значений.

Это второе значение, возвращаемое torch.sort(). Смотрите его документацию для точного определения семантики этого метода.

Если stable равно True, то сортировка становится стабильной, сохраняя порядок эквивалентных элементов. Если False, то относительный порядок значений, которые сравниваются как равные, не гарантируется. True медленнее.

Параметры:
  • input (Тензор) – входной тензор.
  • dim (int, опционально) – измерение для сортировки.
  • descending (bool, опционально) – управляет порядком сортировки (возрастание или убывание).
  • stable (bool, опционально) – управляет относительным порядком эквивалентных элементов.

Пример:

>>> a = torch.randn(4, 4)
>>> a
tensor([[ 0.0785,  1.5267, -0.8521,  0.4065],
        [ 0.1598,  0.0788, -0.0745, -1.2700],
        [ 1.2208,  1.0722, -0.7064,  1.2564],
        [ 0.0669, -0.2318, -0.8229, -0.9280]])


>>> torch.argsort(a, dim=1)
tensor([[2, 0, 3, 1],
        [3, 2, 1, 0],
        [2, 1, 0, 3],
        [3, 2, 1, 0]])

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

Spec-Zone.ru

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