torch.sort
-
torch.sort(input, dim=- 1, descending=False, stable=False, *, out=None) -
Сортирует элементы
inputтензора по заданной размерности в порядке возрастания значений.Если
dimне указано, выбирается последняя размерностьinput.Если
descendingравноTrue, то элементы сортируются в порядке убывания значений.Если
stableравноTrue, то процедура сортировки становится устойчивой, сохраняя порядок эквивалентных элементов.Возвращается кортеж из namedtuple (значения, индексы), где
values— отсортированные значения, аindices— индексы элементов в исходномinputтензоре.- Параметры:
-
- input (Tensor) – входной тензор.
- dim (int, необязательно) – размерность для сортировки.
- descending (bool, необязательно) – определяет порядок сортировки (возрастающий или убывающий).
- stable (bool, необязательно) – делает процедуру сортировки устойчивой, гарантируя, что порядок эквивалентных элементов сохраняется.
- Ключевые аргументы:
-
out (tuple, необязательно) – выходной кортеж (
Tensor,LongTensor) который можно необязательно указать для использования в качестве буферов вывода.
Пример:
>>> x = torch.randn(3, 4) >>> sorted, indices = torch.sort(x) >>> sorted tensor([[-0.2162, 0.0608, 0.6719, 2.3332], [-0.5793, 0.0061, 0.6058, 0.9497], [-0.5071, 0.3343, 0.9553, 1.0960]]) >>> indices tensor([[ 1, 0, 2, 3], [ 3, 1, 0, 2], [ 0, 3, 1, 2]]) >>> sorted, indices = torch.sort(x, 0) >>> sorted tensor([[-0.5071, -0.2162, 0.6719, -0.5793], [ 0.0608, 0.0061, 0.9497, 0.3343], [ 0.6058, 0.9553, 1.0960, 2.3332]]) >>> indices tensor([[ 2, 0, 0, 1], [ 0, 1, 1, 2], [ 1, 2, 2, 0]]) >>> x = torch.tensor([0, 1] * 9) >>> x.sort() torch.return_types.sort( values=tensor([0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1]), indices=tensor([ 2, 16, 4, 6, 14, 8, 0, 10, 12, 9, 17, 15, 13, 11, 7, 5, 3, 1])) >>> x.sort(stable=True) torch.return_types.sort( values=tensor([0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1]), indices=tensor([ 0, 2, 4, 6, 8, 10, 12, 14, 16, 1, 3, 5, 7, 9, 11, 13, 15, 17]))
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.sort.html