mindspore/docs/api/api_python/ops/mindspore.ops.func_topk.rst

37 lines
2.0 KiB
ReStructuredText
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

mindspore.ops.topk
===================
.. py:function:: mindspore.ops.topk(input_x, k, dim=None, largest=True, sorted=True)
沿给定维度查找 `k` 个最大或最小元素和对应的索引。
.. warning::
- 如果 `sorted` 设置为False它将使用aicpu运算符性能可能会降低另外由于在不同平台上存在内存排布以及遍历方式不同等问题`sorted` 设置为False时计算结果的显示顺序可能会出现不一致的情况。
如果 `input_x` 是一维Tensor则查找Tensor中 `k` 个最大或最小元素并将其值和索引输出为Tensor。`values[k]``input_x``k` 个最大元素,其索引是 `indices[k]`
对于多维矩阵,计算给定维度中最大或最小的 `k` 个元素,因此:
.. math::
values.shape = indices.shape
如果两个比较的元素相同,则优先返回索引值较小的元素。
参数:
- **input_x** (Tensor) - 需计算的输入数据类型必须为float16、float32或int32。
- **k** (int) - 指定计算最大或最小元素的数量,必须为常量。
- **dim** (int, 可选) - 需要排序的维度。默认值None。
- **largest** (bool, 可选) - 如果为False则会返回前k个最小值。默认值True。
- **sorted** (bool, 可选) - 如果为True则获取的元素将按值降序排序。如果为False则不对获取的元素进行排序。默认值True。
返回:
`values``indices` 组成的tuple。
- **values** (Tensor) - 给定维度的每个切片中的 `k` 最大元素或最小元素。
- **indices** (Tensor) - `k` 最大元素的对应索引。
异常:
- **TypeError** - 如果 `sorted` 不是bool。
- **TypeError** - 如果 `input_x` 不是Tensor。
- **TypeError** - 如果 `k` 不是int。
- **TypeError** - 如果 `input_x` 的数据类型不是以下之一float16、float32或int32。