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

40 lines
2.0 KiB
ReStructuredText
Raw Normal View History

2022-02-14 11:35:00 +08:00
mindspore.nn.ROC
=====================
.. py:class:: mindspore.nn.ROC(class_num=None, pos_label=None)
计算ROC曲线。适用于求解二分类和多分类问题。在多分类的情况下将基于one-vs-the-rest的方法进行计算。
**参数:**
- **class_num** (int) - 类别数。对于二分类问题此入参可以不设置。默认值None。
2022-03-25 16:24:33 +08:00
- **pos_label** (int) - 正类的类别值。二分类问题中,不设置此入参,即 `pos_label` 为None时正类类别值默认为1用户可以自行设置正类类别值为其他值。多分类问题中用户不应设置此参数因为它将在[0,num_classes-1]范围内迭代更改。默认值None。
2022-02-14 11:35:00 +08:00
.. py:method:: clear()
内部评估结果清零。
.. py:method:: eval()
计算ROC曲线。
**返回:**
tuple`fpr``tpr``thresholds` 组成。
2022-03-25 16:24:33 +08:00
- **fpr** (np.array) - 假正率。二分类情况下返回不同阈值下的fpr多分类情况下则为fpr(false positive rate)的列表,列表的每个元素代表一个类别。
- **tps** (np.array) - 真正率。二分类情况下返回不同阈值下的tps多分类情况下则为tps(true positive rate)的列表,列表的每个元素代表一个类别。
2022-02-14 11:35:00 +08:00
- **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,)`