花园宝宝战队 ----- 一阶段代码注释成果 #10

Open
bjutsecurity22 wants to merge 139 commits from bjutsecurity22/mindspore2022:master into master
1 changed files with 58 additions and 2 deletions
Showing only changes of commit 4a05f889c0 - Show all commits

View File

@ -12,16 +12,22 @@
# See the License for the specific language governing permissions and
# limitations under the License.
# ============================================================================
# Recall类用于计算召回率Recall
"""Recall."""
# 用于访问与Python解释器相关的变量和函数
import sys
# 导入了所需的库和模块如numpy和mindspore
import numpy as np
# 从mindspore._checkparam模块中导入了一个名为validator的类用于验证参数
from mindspore._checkparam import Validator as validator
# Recall类继承自EvaluationBase类
from .metric import EvaluationBase, rearrange_inputs, _check_onehot_data
# 从mindspore.metric模块中导入了一个名为EvaluationBase的类和三个函数分别是rearrange_inputs、_check_onehot_data
class Recall(EvaluationBase):
# 用于计算召回率Recall。Recall类继承自EvaluationBase类。在计算召回率时会使用两个本地变量true_positive和false_negative来存储真实阳性和假负例的数量。
# 需要注意的是,在多分类 cases 中元素的y和y_pred必须为0或1
r"""
Calculates recall for classification and multilabel data.
@ -54,25 +60,38 @@ class Recall(EvaluationBase):
>>> print(recall)
[1. 0.5]
"""
# 在定义Recall类时需要初始化一个eval_type参数默认为'classification'
def __init__(self, eval_type='classification'):
# 然后调用父类的构造函数将self传递给父类
super(Recall, self).__init__(eval_type)
# 接着初始化一个eps属性用于存储一个极小值
self.eps = sys.float_info.min
# 最后调用clear方法来清空内部变量
self.clear()
def clear(self):
# 用于清空内部评估结果
"""Clears the internal evaluation result."""
# 首先将_class_num属性设置为0
self._class_num = 0
# 如果_type属性为"multilabel"
if self._type == "multilabel":
# 则将_true_positives和_actual_positives属性设置为空数组
self._true_positives = np.empty(0)
self._actual_positives = np.empty(0)
# 并将_true_positives_average和_actual_positives_average属性设置为0
self._true_positives_average = 0
self._actual_positives_average = 0
else:
# 否则将_true_positives和_actual_positives属性设置为0
self._true_positives = 0
self._actual_positives = 0
@rearrange_inputs
def update(self, *inputs):
# 用于更新内部评估结果。接收一个或多个输入参数分别是y_pred和y。对于'classification'评估类型y_pred通常是一个浮点数列表范围在0到1之间
# 形状为(N, C)其中N是样本数量C是类别数量。对于'multilabel'评估类型y_pred和y必须是one-hot编码的数组值全为0或1。
# indices中值为1的索引表示正类别。y_pred和y的形状都是(N, C)
"""
Updates the internal evaluation result with `y_pred` and `y`.
@ -91,49 +110,78 @@ class Recall(EvaluationBase):
Raises:
ValueError: If the number of inputs is not 2.
"""
# 首先检查输入参数的数量是否为2
if len(inputs) != 2:
# 如果不是则抛出一个ValueError异常
raise ValueError("For 'Recall.update', it needs 2 inputs (predicted value, true value), "
"but got {}.".format(len(inputs)))
# 然后将输入的y_pred和true value转换为适当的数据格式
y_pred = self._convert_data(inputs[0])
y = self._convert_data(inputs[1])
# 如果评估类型为'classification'并且y_pred的形状与y相同并且是one-hot编码
if self._type == 'classification' and y_pred.ndim == y.ndim and _check_onehot_data(y):
# 那么将y转换为argmax轴1的值
y = y.argmax(axis=1)
# 最后检查y_pred和y的形状和值是否符合要求
self._check_shape(y_pred, y)
self._check_value(y_pred, y)
# 检查_class_num属性是否为0
if self._class_num == 0:
# 如果是则将其设置为y_pred的形状[1]
self._class_num = y_pred.shape[1]
# 如果y_pred的形状[1]与_class_num不同
elif y_pred.shape[1] != self._class_num:
# 则抛出一个ValueError异常
raise ValueError("For 'Recall.update', class number not match, last input predicted data contain {} "
"classes, but current predicted data contain {} classes, please check your predicted "
"value(inputs[0]).".format(self._class_num, y_pred.shape[1]))
# 首先获取_class_num属性然后根据评估类型进行相应的处理
class_num = self._class_num
# 如果评估类型为'classification'
if self._type == "classification":
# 首先检查y的最大值是否大于_class_num
if y.max() + 1 > class_num:
# 如果是则抛出一个ValueError异常
raise ValueError("For 'Recall.update', predicted value (input[0]) should have the same classes number "
"as true value (input[1]), but got predicted value classes {}, true value classes {}."
.format(class_num, y.max() + 1))
# 接着将y转换为one-hot编码
y = np.eye(class_num)[y.reshape(-1)]
# 并获取y_pred中每个类别概率最大的索引
indices = y_pred.argmax(axis=1).reshape(-1)
# 最后将y_pred转换回one-hot编码
y_pred = np.eye(class_num)[indices]
# 如果评估类型为'multilabel'
elif self._type == "multilabel":
# 则将y_pred和y交换轴
y_pred = y_pred.swapaxes(1, 0).reshape(class_num, -1)
# 并reshape为(class_num, -1)
y = y.swapaxes(1, 0).reshape(class_num, -1)
# 首先使用sum函数将y沿axis=0轴求和得到actual_positives
actual_positives = y.sum(axis=0)
# 然后使用点乘运算将y_pred和y相乘再沿axis=0轴求和得到true_positives
true_positives = (y * y_pred).sum(axis=0)
# 如果评估类型为'multilabel'
if self._type == "multilabel":
# 那么将true_positives除以(actual_positives + self.eps)后求和然后累加到self._true_positives_average中
self._true_positives_average += np.sum(true_positives / (actual_positives + self.eps))
# 接着将actual_positives累加到self._actual_positives_average中
self._actual_positives_average += len(actual_positives)
# 最后将true_positives和actual_positives拼接在一起存储在self._true_positives和self._actual_positives中
self._true_positives = np.concatenate((self._true_positives, true_positives), axis=0)
self._actual_positives = np.concatenate((self._actual_positives, actual_positives), axis=0)
else:
# 如果评估类型为'classification'那么直接将true_positives和actual_positives累加在一起
self._true_positives += true_positives
# 存储在self._true_positives和self._actual_positives中
self._actual_positives += actual_positives
def eval(self, average=False):
# 用于计算召回率。方法接收一个名为average的布尔参数用于指定是否计算平均召回率
"""
Computes the recall.
@ -143,16 +191,24 @@ class Recall(EvaluationBase):
Returns:
numpy.float64, the computed result.
"""
# 检查输入参数是否为空
if self._class_num == 0:
# 如果是则抛出一个RuntimeError异常
raise RuntimeError("The 'Recall' can not be calculated, because the number of samples is 0, please check "
"whether your inputs (predicted value, true value) are empty, or has called update "
"method before calling eval method.")
# 检查输入参数average是否为布尔值如果不是则抛出一个TypeError异常
validator.check_value_type("average", average, [bool], self.__class__.__name__)
# 然后,计算召回率
result = self._true_positives / (self._actual_positives + self.eps)
# 如果average为True则计算平均结果
if average:
# 如果type为多标签则计算多标签评估下的召回率
if self._type == "multilabel":
result = self._true_positives_average / (self._actual_positives_average + self.eps)
# 否则,计算正确结果的平均值
return result.mean()
# 否则,返回计算得到的召回率
return result