mindspore/docs/api/api_python/train/mindspore.train.Recall.rst

46 lines
1.8 KiB
ReStructuredText
Raw Normal View History

mindspore.train.Recall
=======================
2021-12-04 13:55:42 +08:00
.. py:class:: mindspore.train.Recall(eval_type='classification')
2021-12-04 13:55:42 +08:00
2021-12-31 16:45:06 +08:00
计算数据分类的召回率,包括单标签场景和多标签场景。
2021-12-04 13:55:42 +08:00
2021-12-31 16:45:06 +08:00
Recall类创建两个局部变量 :math:`\text{true_positive}`:math:`\text{false_negative}` 用于计算召回率。计算方式为:
2021-12-04 13:55:42 +08:00
.. math::
\text{recall} = \frac{\text{true_positive}}{\text{true_positive} + \text{false_negative}}
2021-12-04 20:36:47 +08:00
.. note::
2021-12-04 13:55:42 +08:00
在多标签情况下, :math:`y`:math:`y_{pred}` 的元素必须为0或1。
参数:
2022-07-26 11:11:41 +08:00
- **eval_type** (str) - 支持'classification'和'multilabel'。默认值:'classification'。
2021-12-04 13:55:42 +08:00
.. py:method:: clear()
内部评估结果清零。
.. py:method:: eval(average=False)
计算召回率。
参数:
- **average** (bool) - 指定是否计算平均召回率。默认值False。
2021-12-04 20:36:47 +08:00
返回:
numpy.float64计算结果。
2021-12-04 13:55:42 +08:00
.. py:method:: update(*inputs)
使用预测值 `y_pred` 和真实标签 `y` 更新局部变量。
参数:
- **inputs** - 输入 `y_pred``y``y_pred``y` 支持Tensor、list或numpy.ndarray类型。
2021-12-04 13:55:42 +08:00
对于'classification'情况,`y_pred` 在大多数情况下由范围 :math:`[0, 1]` 中的浮点数组成shape为 :math:`(N, C)` ,其中 :math:`N` 是样本数, :math:`C` 是类别数。`y` 由整数值组成如果是one_hot编码格式shape是 :math:`(N,C)` 如果是类别索引shape是 :math:`(N,)`
2021-12-04 13:55:42 +08:00
对于'multilabel'情况,`y_pred``y` 只能是值为0或1的one-hot编码格式其中值为1的索引表示正类别。`y_pred``y` 的shape都是 :math:`(N,C)`
2021-12-04 20:36:47 +08:00
异常:
- **ValueError** - inputs数量不是2。