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

31 lines
1.6 KiB
ReStructuredText
Raw Permalink Normal View History

mindspore.ops.OneHot
====================
2022-03-10 15:13:21 +08:00
.. py:class:: mindspore.ops.OneHot(axis=-1)
返回一个one-hot类型的Tensor。
生成一个新的Tensor由索引 `indices` 表示的位置取值为 `on_value` ,而在其他所有位置取值为 `off_value`
.. note::
如果输入索引为秩 `N` ,则输出为秩 `N+1` 。新轴在 `axis` 处创建。
2022-07-26 16:39:37 +08:00
参数:
- **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** (int) - 输入的Scalar定义one-hot的深度。
- **on_value** (Tensor) - 当 `indices[j] = i`用来填充输出的值。数据类型为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)`
异常:
- **TypeError** - `axis``depth` 不是int。
- **TypeError** - `indices` 的数据类型既不是uint8也不是int32或者int64。
- **TypeError** - `indices``on_value``off_value` 不是Tensor。
- **ValueError** - `axis` 不在[-1, ndim]范围内。
- **ValueError** - `depth` 小于0。