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

32 lines
1.9 KiB
ReStructuredText
Raw Permalink 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.OneHot
====================
.. py:class:: mindspore.ops.OneHot(axis=-1)
返回一个one-hot类型的Tensor。
生成一个新的Tensor由索引 `indices` 表示的位置取值为 `on_value` ,而在其他所有位置取值为 `off_value`
.. note::
如果输入索引为秩 `N` ,则输出为秩 `N+1` 。新轴在 `axis` 处创建。当执行设备是 Ascend 时,如果 `on_value` 为int64类型`indices` 也必须为int64类型`on_value``off_value` 的取值只能是1和0。
参数:
- **axis** (int) - 指定one-hot的计算维度。例如如果 `indices` 的shape为 :math:`(N, C)` `axis` 为-1则输出shape为 :math:`(N, C, D)` ,如果 `axis` 为0则输出shape为 :math:`(D, N, C)` 。默认值: ``-1``
输入:
- **indices** (Tensor) - 输入索引shape为 :math:`(X_0, \ldots, X_n)` 的Tensor。数据类型必须为int32或int64。
- **depth** (Union[int, Tensor]) - 输入的Scalar定义one-hot的深度。
- **on_value** (Tensor) - 当 `indices[j] = i`用来填充输出的值。数据类型必须为int32、int64、float16或float32。
- **off_value** (Tensor) - 当 `indices[j] != i` 时,用来填充输出的值。数据类型与 `on_value` 的相同。
输出:
Tensorone-hot类型的Tensor。shape为 :math:`(X_0, \ldots, X_{axis}, \text{depth} ,X_{axis+1}, \ldots, X_n)` ,输出数据类型与 `on_value` 的相同。
异常:
- **TypeError** - `axis``depth` 不是int。
- **TypeError** - `indices` 的数据类型不是int32或者int64。
- **TypeError** - `on_value` 的数据类型不是int32、int64、float16或者float32。
- **TypeError** - `indices``on_value``off_value` 不是Tensor。
- **ValueError** - `axis` 不在[-1, ndim]范围内。
- **ValueError** - `depth` 小于0。