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

40 lines
1.8 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.ROC
=====================
.. py:class:: mindspore.nn.ROC(class_num=None, pos_label=None)
计算ROC曲线。适用于求解二分类和多分类问题。在多分类的情况下将基于one-vs-the-rest的方法进行计算。
**参数:**
- **class_num** (int) - 类别数。对于二分类问题此入参可以不设置。默认值None。
- **pos_label** (int) - 正类的类别值。二进制问题中默认为1多分类问题中不应设置此参数因为它将在[0,num_classes-1]范围内迭代更改。默认值None。
.. py:method:: clear()
内部评估结果清零。
.. py:method:: eval()
计算ROC曲线。
**返回:**
tuple`fpr``tpr``thresholds` 组成。
- **fpr** (np.array) - 假正率。二分类情况下返回不同阈值下的fpr多分类情况下则为fpr的列表列表的每个元素代表一个类别。
- **tps** (np.array) - 真正率。二分类情况下返回不同阈值下的tps多分类情况下则为tps的列表列表的每个元素代表一个类别。
- **thresholds** (np.array) - 用于计算假正率和真正率的阈值。
**异常:**
- **RuntimeError** - 如果没有先调用update方法则会报错。
.. py:method:: update(*inputs)
使用 `y_pred``y` 更新内部评估结果。
**参数:**
- **inputs** - 输入 `y_pred``y``y_pred``y` 是Tensor、list或numpy.ndarray。`y_pred` 一般情况下是范围为 :math:`[0, 1]` 的浮点数列表shape为 :math:`(N, C)`,其中 :math:`N` 是用例数,:math:`C` 是类别数。`y` 为整数值如果为one-hot格式shape为 :math:`(N, C)`如果是类别索引shape为 :math:`(N,)`