mindspore2022/docs/api/api_python/nn/mindspore.nn.ConfusionMatri...

57 lines
2.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.ConfusionMatrixMetric
==================================
.. py:class:: mindspore.nn.ConfusionMatrixMetric(skip_channel=True, metric_name='sensitivity', calculation_method=False, decrease='mean')
度量分类模型的性能矩阵是输出为二进制或多类的模型。
从满量程Tensor计算混淆矩阵的相关性度量并收集批次、类通道和迭代的平均值。
此函数支持计算以下描述的所有度量参数metric_name中的度量名称。
如果要使用混淆矩阵计算,如"PPV"、"TPR"、"TNR",请使用此类。
如果您只想计算混淆矩阵,请使用'mindspore.nn.ConfusionMatrix'。
**参数:**
- **skip_channel** (bool) - 是否跳过预测输出的第一个通道的度量计算。默认值True。
- **metric_name** (str) - 建议采用如下指标。当然,也可以为这些指标设置通用别名。
取值范围:["sensitivity", "specificity", "precision", "negative predictive value", "miss rate", "fall out", "false discovery rate", "false omission rate", "prevalence threshold", "threat score", "accuracy", "balanced accuracy", "f1 score", "matthews correlation coefficient", "fowlkes mallows index", "informedness", "markedness"]。
默认值:"sensitivity"。
- **calculation_method** (bool) - 如果为True则计算每个样品的度量值。如果为False则累积所有样本的混淆矩阵。
对于分类任务, `calculation_method` 应为False。默认值False。
- **decrease** (str) - 定义减少一批数据计算结果的模式。仅当 `calculation_method` 为True时才生效。
取值范围:["none", "mean", "sum", "mean_batch", "sum_batch", "mean_channel", "sum_channel"]。默认值:"mean"。
.. py:method:: clear()
重置评估结果。
.. py:method:: eval()
计算混淆矩阵度量。
**返回:**
numpy.ndarray计算的结果。
.. py:method:: update(*inputs)
使用预测值和目标值更新状态。
**参数:**
- **inputs** (tuple) - `y_pred``y``y_pred``y``Tensor` 、列表或数组。
- **y_pred** (ndarray)待计算的输入数据。格式必须为one-hot且第一个维度是batch。
`y_pred` 的shape是 :math:`(N, C, ...)`:math:`(N, ...)`
至于分类任务, `y_pred` 的shape应为[BN]其中N大于1。对于分割任务shape应为[BNHW]或[BNHWD]。
- **y** (ndarray)计算度量值的真实值。格式必须为one-hot且第一个维度是batch。`y` 的shape是 :math:`(N, C, ...)`
**返回:**
numpynumpy类型评估结果。
**异常:**
- **ValueError** - 输入参数的数量不等于2。