mindspore2022/docs/api/api_python/nn/mindspore.nn.OneHot.rst

54 lines
2.1 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.nn.OneHot
====================
.. py:class:: mindspore.nn.OneHot(axis=-1, depth=1, on_value=1.0, off_value=0.0, dtype=mstype.float32)
返回一个one-hot类型的Tensor。
参数 `indices` 表示的位置取值为on_value其他所有位置取值为off_value。
.. note::
如果indices是n阶Tensor那么返回的one-hot Tensor则为n+1阶Tensor。
如果 `indices` 是Scalar则输出shape将是长度为 `depth` 的向量。
如果 `indices` 是长度为 `features` 的向量则输出shape为
.. code-block::
features * depth if axis == -1
depth * features if axis == 0
如果 `indices` 是shape为 `[batch, features]` 的矩阵则输出shape为
.. code-block::
batch * features * depth if axis == -1
batch * depth * features if axis == 1
depth * batch * features if axis == 0
**参数:**
- **axis** (int) - 指定第几阶为depth维one-hot向量如果轴为-1则 features x depth如果轴为0则 depth x features。默认值-1。
- **depth** (int) - 定义one-hot向量的维度深度。默认值1。
- **on_value** (float) - one-hot值当indices[j] = i时填充output[i][j]的取值。默认值1.0。
- **off_value** (float) - 非one-hot值当indices[j] != i时填充output[i][j]的取值。默认值0.0。
- **dtype** (:class:`mindspore.dtype`) - 是'on_value'和'off_value'的数据类型而不是索引的数据类型。默认值mindspore.float32。
**输入:**
**indices** (Tensor) - 输入索引任意维度的Tensor数据类型为int32或int64。
**输出:**
Tensor数据类型 `dtype` 的独热Tensor维度为 `axis` 扩展到 `depth`并填充on_value和off_value。`Outputs` 的维度等于 `indices` 的维度加1。
**异常:**
- **TypeError** - `axis``depth` 不是整数。
- **TypeError** - `indices` 的dtype既不是int32也不是int64。
- **ValueError** - 如果 `axis` 不在范围[-1, len(indices_shape)]内。
- **ValueError** - `depth` 小于0。