Spec-Zone.ru › PyTorch 2

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/2.1/generated/torch.argsort.html

Spec-Zone.ru

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