evalsuite-python 0.1.0a1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- evalsuite/__init__.py +118 -0
- evalsuite/api.py +249 -0
- evalsuite/classification/__init__.py +47 -0
- evalsuite/classification/_common.py +111 -0
- evalsuite/classification/metrics.py +872 -0
- evalsuite/core/__init__.py +1 -0
- evalsuite/core/context.py +181 -0
- evalsuite/core/exceptions.py +52 -0
- evalsuite/core/registry.py +90 -0
- evalsuite/core/result.py +277 -0
- evalsuite/core/types.py +23 -0
- evalsuite/core/validation.py +202 -0
- evalsuite/py.typed +0 -0
- evalsuite/regression/__init__.py +41 -0
- evalsuite/regression/metrics.py +587 -0
- evalsuite/version.py +3 -0
- evalsuite_python-0.1.0a1.dist-info/METADATA +150 -0
- evalsuite_python-0.1.0a1.dist-info/RECORD +21 -0
- evalsuite_python-0.1.0a1.dist-info/WHEEL +4 -0
- evalsuite_python-0.1.0a1.dist-info/entry_points.txt +2 -0
- evalsuite_python-0.1.0a1.dist-info/licenses/LICENSE +21 -0
evalsuite/__init__.py
ADDED
|
@@ -0,0 +1,118 @@
|
|
|
1
|
+
"""EvalSuite: unified, reproducible evaluation for machine learning and research.
|
|
2
|
+
|
|
3
|
+
>>> import evalsuite as es
|
|
4
|
+
>>> result = es.evaluate([0, 1, 1, 0], [0, 1, 0, 0])
|
|
5
|
+
>>> print(result.summary()) # doctest: +SKIP
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from . import classification, regression
|
|
9
|
+
from .api import evaluate
|
|
10
|
+
from .classification import (
|
|
11
|
+
accuracy,
|
|
12
|
+
average_precision,
|
|
13
|
+
balanced_accuracy,
|
|
14
|
+
brier_score,
|
|
15
|
+
cohen_kappa,
|
|
16
|
+
confusion_matrix,
|
|
17
|
+
f1,
|
|
18
|
+
fbeta,
|
|
19
|
+
hamming_loss,
|
|
20
|
+
jaccard,
|
|
21
|
+
log_loss,
|
|
22
|
+
mcc,
|
|
23
|
+
npv,
|
|
24
|
+
pr_curve,
|
|
25
|
+
precision,
|
|
26
|
+
recall,
|
|
27
|
+
roc_auc,
|
|
28
|
+
roc_curve,
|
|
29
|
+
specificity,
|
|
30
|
+
top_k_accuracy,
|
|
31
|
+
)
|
|
32
|
+
from .core.exceptions import (
|
|
33
|
+
EvalSuiteError,
|
|
34
|
+
InputValidationError,
|
|
35
|
+
MetricInputError,
|
|
36
|
+
OptionalDependencyError,
|
|
37
|
+
StatisticalTestError,
|
|
38
|
+
UndefinedMetricWarning,
|
|
39
|
+
UnsupportedTaskError,
|
|
40
|
+
)
|
|
41
|
+
from .core.registry import MetricInfo, list_metrics, metric_info
|
|
42
|
+
from .core.result import EvaluationResult, MetricResult
|
|
43
|
+
from .regression import (
|
|
44
|
+
adjusted_r2,
|
|
45
|
+
explained_variance,
|
|
46
|
+
huber_loss,
|
|
47
|
+
mae,
|
|
48
|
+
mape,
|
|
49
|
+
max_error,
|
|
50
|
+
mean_bias_error,
|
|
51
|
+
median_absolute_error,
|
|
52
|
+
mse,
|
|
53
|
+
msle,
|
|
54
|
+
quantile_loss,
|
|
55
|
+
r2,
|
|
56
|
+
rae,
|
|
57
|
+
rmse,
|
|
58
|
+
rmsle,
|
|
59
|
+
rse,
|
|
60
|
+
smape,
|
|
61
|
+
)
|
|
62
|
+
from .version import __version__
|
|
63
|
+
|
|
64
|
+
__all__ = [
|
|
65
|
+
"EvalSuiteError",
|
|
66
|
+
"EvaluationResult",
|
|
67
|
+
"InputValidationError",
|
|
68
|
+
"MetricInfo",
|
|
69
|
+
"MetricInputError",
|
|
70
|
+
"MetricResult",
|
|
71
|
+
"OptionalDependencyError",
|
|
72
|
+
"StatisticalTestError",
|
|
73
|
+
"UndefinedMetricWarning",
|
|
74
|
+
"UnsupportedTaskError",
|
|
75
|
+
"__version__",
|
|
76
|
+
"accuracy",
|
|
77
|
+
"adjusted_r2",
|
|
78
|
+
"average_precision",
|
|
79
|
+
"balanced_accuracy",
|
|
80
|
+
"brier_score",
|
|
81
|
+
"classification",
|
|
82
|
+
"cohen_kappa",
|
|
83
|
+
"confusion_matrix",
|
|
84
|
+
"evaluate",
|
|
85
|
+
"explained_variance",
|
|
86
|
+
"f1",
|
|
87
|
+
"fbeta",
|
|
88
|
+
"hamming_loss",
|
|
89
|
+
"huber_loss",
|
|
90
|
+
"jaccard",
|
|
91
|
+
"list_metrics",
|
|
92
|
+
"log_loss",
|
|
93
|
+
"mae",
|
|
94
|
+
"mape",
|
|
95
|
+
"max_error",
|
|
96
|
+
"mcc",
|
|
97
|
+
"mean_bias_error",
|
|
98
|
+
"median_absolute_error",
|
|
99
|
+
"metric_info",
|
|
100
|
+
"mse",
|
|
101
|
+
"msle",
|
|
102
|
+
"npv",
|
|
103
|
+
"pr_curve",
|
|
104
|
+
"precision",
|
|
105
|
+
"quantile_loss",
|
|
106
|
+
"r2",
|
|
107
|
+
"rae",
|
|
108
|
+
"recall",
|
|
109
|
+
"regression",
|
|
110
|
+
"rmse",
|
|
111
|
+
"rmsle",
|
|
112
|
+
"roc_auc",
|
|
113
|
+
"roc_curve",
|
|
114
|
+
"rse",
|
|
115
|
+
"smape",
|
|
116
|
+
"specificity",
|
|
117
|
+
"top_k_accuracy",
|
|
118
|
+
]
|
evalsuite/api.py
ADDED
|
@@ -0,0 +1,249 @@
|
|
|
1
|
+
"""High-level API: ``evaluate()``."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import platform
|
|
6
|
+
from collections.abc import Sequence
|
|
7
|
+
from typing import Any, Callable, Optional
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
|
|
11
|
+
from .classification import metrics as clf
|
|
12
|
+
from .classification._common import averaged
|
|
13
|
+
from .core.context import ClassificationContext
|
|
14
|
+
from .core.exceptions import InputValidationError, UnsupportedTaskError
|
|
15
|
+
from .core.result import EvaluationResult, MetricResult
|
|
16
|
+
from .core.types import ArrayLike, ZeroDivision
|
|
17
|
+
from .core.validation import target_type, to_numpy
|
|
18
|
+
from .regression import metrics as reg
|
|
19
|
+
from .version import __version__
|
|
20
|
+
|
|
21
|
+
__all__ = ["evaluate"]
|
|
22
|
+
|
|
23
|
+
CtxMetric = Callable[[ClassificationContext], MetricResult]
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def _classification_registry(
|
|
27
|
+
average: Optional[str], pos_label: Any, zero_division: ZeroDivision
|
|
28
|
+
) -> dict[str, tuple[CtxMetric, bool]]:
|
|
29
|
+
"""metric id -> (function of a shared context, needs probabilities)."""
|
|
30
|
+
|
|
31
|
+
def avg(fn: Any, metric: str, name: str) -> CtxMetric:
|
|
32
|
+
return lambda ctx: averaged(
|
|
33
|
+
ctx,
|
|
34
|
+
fn,
|
|
35
|
+
metric=metric,
|
|
36
|
+
name=name,
|
|
37
|
+
average=average,
|
|
38
|
+
pos_label=pos_label,
|
|
39
|
+
zero_division=zero_division,
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
return {
|
|
43
|
+
"accuracy": (clf._accuracy, False),
|
|
44
|
+
"balanced_accuracy": (clf._balanced_accuracy, False),
|
|
45
|
+
"precision": (avg(clf._precision, "precision", "Precision"), False),
|
|
46
|
+
"recall": (avg(clf._recall, "recall", "Recall"), False),
|
|
47
|
+
"f1": (avg(clf._fbeta_fn(1.0), "f1", "F1"), False),
|
|
48
|
+
"specificity": (avg(clf._specificity, "specificity", "Specificity"), False),
|
|
49
|
+
"npv": (avg(clf._npv, "npv", "NPV"), False),
|
|
50
|
+
"jaccard": (avg(clf._jaccard, "jaccard", "Jaccard"), False),
|
|
51
|
+
"mcc": (lambda ctx: clf._mcc(ctx, zero_division), False),
|
|
52
|
+
"cohen_kappa": (lambda ctx: clf._cohen_kappa(ctx, None, zero_division), False),
|
|
53
|
+
"hamming_loss": (_hamming, False),
|
|
54
|
+
"roc_auc": (lambda ctx: clf._roc_auc(ctx, average=average, pos_label=pos_label), True),
|
|
55
|
+
"average_precision": (
|
|
56
|
+
lambda ctx: clf._average_precision(ctx, average=average, pos_label=pos_label),
|
|
57
|
+
True,
|
|
58
|
+
),
|
|
59
|
+
"log_loss": (lambda ctx: clf._log_loss(ctx, pos_label), True),
|
|
60
|
+
"brier_score": (lambda ctx: clf._brier(ctx, pos_label), True),
|
|
61
|
+
}
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def _hamming(ctx: ClassificationContext) -> MetricResult:
|
|
65
|
+
if ctx.target_type == "multilabel":
|
|
66
|
+
assert ctx.y_pred is not None # noqa: S101
|
|
67
|
+
wrong = (ctx.y_true != ctx.y_pred).mean(axis=1)
|
|
68
|
+
else:
|
|
69
|
+
wrong = (ctx.true_idx != ctx.pred_idx).astype(float)
|
|
70
|
+
return MetricResult("hamming_loss", "Hamming loss", float(np.average(wrong, weights=ctx.weights)))
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
_DEFAULT_SINGLE = [
|
|
74
|
+
"accuracy",
|
|
75
|
+
"balanced_accuracy",
|
|
76
|
+
"precision",
|
|
77
|
+
"recall",
|
|
78
|
+
"f1",
|
|
79
|
+
"specificity",
|
|
80
|
+
"mcc",
|
|
81
|
+
"cohen_kappa",
|
|
82
|
+
]
|
|
83
|
+
_DEFAULT_MULTILABEL = ["accuracy", "hamming_loss", "precision", "recall", "f1", "jaccard"]
|
|
84
|
+
_DEFAULT_PROB = ["roc_auc", "average_precision", "log_loss", "brier_score"]
|
|
85
|
+
_DEFAULT_PROB_MULTILABEL = ["roc_auc", "average_precision"]
|
|
86
|
+
|
|
87
|
+
_REGRESSION: dict[str, Callable[..., MetricResult]] = {
|
|
88
|
+
"mae": reg.mae,
|
|
89
|
+
"mse": reg.mse,
|
|
90
|
+
"rmse": reg.rmse,
|
|
91
|
+
"r2": reg.r2,
|
|
92
|
+
"explained_variance": reg.explained_variance,
|
|
93
|
+
"median_absolute_error": reg.median_absolute_error,
|
|
94
|
+
"max_error": reg.max_error,
|
|
95
|
+
"mean_bias_error": reg.mean_bias_error,
|
|
96
|
+
"mape": reg.mape,
|
|
97
|
+
"smape": reg.smape,
|
|
98
|
+
"msle": reg.msle,
|
|
99
|
+
"rmsle": reg.rmsle,
|
|
100
|
+
"rae": reg.rae,
|
|
101
|
+
"rse": reg.rse,
|
|
102
|
+
"huber_loss": reg.huber_loss,
|
|
103
|
+
"quantile_loss": reg.quantile_loss,
|
|
104
|
+
}
|
|
105
|
+
_DEFAULT_REGRESSION = [
|
|
106
|
+
"mae",
|
|
107
|
+
"mse",
|
|
108
|
+
"rmse",
|
|
109
|
+
"r2",
|
|
110
|
+
"explained_variance",
|
|
111
|
+
"median_absolute_error",
|
|
112
|
+
"max_error",
|
|
113
|
+
"mean_bias_error",
|
|
114
|
+
]
|
|
115
|
+
_NO_WEIGHTS = {"max_error"}
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def _infer_task(y_true: ArrayLike, y_prob: Optional[ArrayLike]) -> str:
|
|
119
|
+
if y_prob is not None:
|
|
120
|
+
return "classification"
|
|
121
|
+
return (
|
|
122
|
+
"regression"
|
|
123
|
+
if target_type(to_numpy(y_true, "y_true", allow_2d=True)) == "continuous"
|
|
124
|
+
else "classification"
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def evaluate(
|
|
129
|
+
y_true: ArrayLike,
|
|
130
|
+
y_pred: Optional[ArrayLike] = None,
|
|
131
|
+
*,
|
|
132
|
+
y_prob: Optional[ArrayLike] = None,
|
|
133
|
+
task: Optional[str] = None,
|
|
134
|
+
metrics: Optional[Sequence[str]] = None,
|
|
135
|
+
average: Optional[str] = "auto",
|
|
136
|
+
labels: Optional[ArrayLike] = None,
|
|
137
|
+
pos_label: Any = None,
|
|
138
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
139
|
+
zero_division: ZeroDivision = "warn",
|
|
140
|
+
) -> EvaluationResult:
|
|
141
|
+
"""Evaluate predictions with a standard set of metrics (or the ``metrics`` you name).
|
|
142
|
+
|
|
143
|
+
The task is inferred when omitted: probabilities or integer/string labels mean classification, non-integer
|
|
144
|
+
numbers mean regression. Inputs are validated once and shared intermediate results (the confusion matrix)
|
|
145
|
+
are computed once for all metrics.
|
|
146
|
+
|
|
147
|
+
>>> import evalsuite as es
|
|
148
|
+
>>> r = es.evaluate([0, 1, 1, 0], [0, 1, 0, 0])
|
|
149
|
+
>>> round(r["accuracy"], 2)
|
|
150
|
+
0.75
|
|
151
|
+
"""
|
|
152
|
+
task = task or _infer_task(y_true, y_prob)
|
|
153
|
+
if task == "classification":
|
|
154
|
+
return _evaluate_classification(
|
|
155
|
+
y_true, y_pred, y_prob, metrics, average, labels, pos_label, sample_weight, zero_division
|
|
156
|
+
)
|
|
157
|
+
if task == "regression":
|
|
158
|
+
if y_pred is None:
|
|
159
|
+
raise InputValidationError("Regression evaluation needs y_pred.")
|
|
160
|
+
return _evaluate_regression(y_true, y_pred, metrics, sample_weight)
|
|
161
|
+
raise UnsupportedTaskError("task must be 'classification' or 'regression' in EvalSuite 0.1.")
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
def _metadata(**extra: Any) -> dict[str, Any]:
|
|
165
|
+
return {
|
|
166
|
+
"evalsuite_version": __version__,
|
|
167
|
+
"numpy_version": np.__version__,
|
|
168
|
+
"python_version": platform.python_version(),
|
|
169
|
+
**extra,
|
|
170
|
+
}
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def _evaluate_classification(
|
|
174
|
+
y_true: ArrayLike,
|
|
175
|
+
y_pred: Optional[ArrayLike],
|
|
176
|
+
y_prob: Optional[ArrayLike],
|
|
177
|
+
metrics: Optional[Sequence[str]],
|
|
178
|
+
average: Optional[str],
|
|
179
|
+
labels: Optional[ArrayLike],
|
|
180
|
+
pos_label: Any,
|
|
181
|
+
sample_weight: Optional[ArrayLike],
|
|
182
|
+
zero_division: ZeroDivision,
|
|
183
|
+
) -> EvaluationResult:
|
|
184
|
+
if y_pred is None and y_prob is None:
|
|
185
|
+
raise InputValidationError("Classification evaluation needs y_pred, y_prob, or both.")
|
|
186
|
+
ctx = ClassificationContext(y_true, y_pred, y_prob=y_prob, labels=labels, sample_weight=sample_weight)
|
|
187
|
+
registry = _classification_registry(average, pos_label, zero_division)
|
|
188
|
+
if metrics is None:
|
|
189
|
+
multilabel = ctx.target_type == "multilabel"
|
|
190
|
+
names = [] if y_pred is None else list(_DEFAULT_MULTILABEL if multilabel else _DEFAULT_SINGLE)
|
|
191
|
+
if y_prob is not None:
|
|
192
|
+
names += _DEFAULT_PROB_MULTILABEL if multilabel else _DEFAULT_PROB
|
|
193
|
+
else:
|
|
194
|
+
names = list(metrics)
|
|
195
|
+
unknown = [m for m in names if m not in registry]
|
|
196
|
+
if unknown:
|
|
197
|
+
raise InputValidationError(
|
|
198
|
+
f"Unknown classification metric(s): {', '.join(unknown)}. Available: {', '.join(registry)}."
|
|
199
|
+
)
|
|
200
|
+
out: dict[str, MetricResult] = {}
|
|
201
|
+
for name in names:
|
|
202
|
+
fn, needs_prob = registry[name]
|
|
203
|
+
if needs_prob and y_prob is None:
|
|
204
|
+
raise InputValidationError(f"Metric '{name}' needs y_prob (predicted probabilities).")
|
|
205
|
+
if not needs_prob and y_pred is None:
|
|
206
|
+
raise InputValidationError(f"Metric '{name}' needs y_pred (predicted labels).")
|
|
207
|
+
out[name] = fn(ctx)
|
|
208
|
+
cm = None if ctx.target_type == "multilabel" or y_pred is None else ctx.confusion_matrix
|
|
209
|
+
return EvaluationResult(
|
|
210
|
+
task="classification",
|
|
211
|
+
metrics=out,
|
|
212
|
+
n_samples=ctx.n,
|
|
213
|
+
target_type=ctx.target_type,
|
|
214
|
+
labels=tuple(ctx.labels.tolist()),
|
|
215
|
+
confusion_matrix=cm,
|
|
216
|
+
metadata=_metadata(average=average, zero_division=zero_division, weighted=sample_weight is not None),
|
|
217
|
+
)
|
|
218
|
+
|
|
219
|
+
|
|
220
|
+
def _evaluate_regression(
|
|
221
|
+
y_true: ArrayLike, y_pred: ArrayLike, metrics: Optional[Sequence[str]], sample_weight: Optional[ArrayLike]
|
|
222
|
+
) -> EvaluationResult:
|
|
223
|
+
yt = to_numpy(y_true, "y_true", allow_2d=True)
|
|
224
|
+
multi = yt.ndim == 2
|
|
225
|
+
if metrics is None:
|
|
226
|
+
names = [m for m in _DEFAULT_REGRESSION if not (multi and m == "max_error")]
|
|
227
|
+
else:
|
|
228
|
+
names = list(metrics)
|
|
229
|
+
unknown = [m for m in names if m not in _REGRESSION]
|
|
230
|
+
if unknown:
|
|
231
|
+
raise InputValidationError(
|
|
232
|
+
f"Unknown regression metric(s): {', '.join(unknown)}. Available: {', '.join(_REGRESSION)}."
|
|
233
|
+
)
|
|
234
|
+
out: dict[str, MetricResult] = {}
|
|
235
|
+
for name in names:
|
|
236
|
+
fn = _REGRESSION[name]
|
|
237
|
+
if name in _NO_WEIGHTS:
|
|
238
|
+
if sample_weight is not None:
|
|
239
|
+
raise InputValidationError(f"Metric '{name}' does not support sample_weight.")
|
|
240
|
+
out[name] = fn(y_true, y_pred)
|
|
241
|
+
else:
|
|
242
|
+
out[name] = fn(y_true, y_pred, sample_weight=sample_weight)
|
|
243
|
+
return EvaluationResult(
|
|
244
|
+
task="regression",
|
|
245
|
+
metrics=out,
|
|
246
|
+
n_samples=yt.shape[0],
|
|
247
|
+
target_type="continuous",
|
|
248
|
+
metadata=_metadata(weighted=sample_weight is not None, outputs=yt.shape[1] if multi else 1),
|
|
249
|
+
)
|
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
"""Classification metrics."""
|
|
2
|
+
|
|
3
|
+
from .metrics import (
|
|
4
|
+
accuracy,
|
|
5
|
+
average_precision,
|
|
6
|
+
balanced_accuracy,
|
|
7
|
+
brier_score,
|
|
8
|
+
cohen_kappa,
|
|
9
|
+
confusion_matrix,
|
|
10
|
+
f1,
|
|
11
|
+
fbeta,
|
|
12
|
+
hamming_loss,
|
|
13
|
+
jaccard,
|
|
14
|
+
log_loss,
|
|
15
|
+
mcc,
|
|
16
|
+
npv,
|
|
17
|
+
pr_curve,
|
|
18
|
+
precision,
|
|
19
|
+
recall,
|
|
20
|
+
roc_auc,
|
|
21
|
+
roc_curve,
|
|
22
|
+
specificity,
|
|
23
|
+
top_k_accuracy,
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
__all__ = [
|
|
27
|
+
"accuracy",
|
|
28
|
+
"average_precision",
|
|
29
|
+
"balanced_accuracy",
|
|
30
|
+
"brier_score",
|
|
31
|
+
"cohen_kappa",
|
|
32
|
+
"confusion_matrix",
|
|
33
|
+
"f1",
|
|
34
|
+
"fbeta",
|
|
35
|
+
"hamming_loss",
|
|
36
|
+
"jaccard",
|
|
37
|
+
"log_loss",
|
|
38
|
+
"mcc",
|
|
39
|
+
"npv",
|
|
40
|
+
"pr_curve",
|
|
41
|
+
"precision",
|
|
42
|
+
"recall",
|
|
43
|
+
"roc_auc",
|
|
44
|
+
"roc_curve",
|
|
45
|
+
"specificity",
|
|
46
|
+
"top_k_accuracy",
|
|
47
|
+
]
|
|
@@ -0,0 +1,111 @@
|
|
|
1
|
+
"""Shared averaging logic for count-based classification metrics."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Callable
|
|
6
|
+
from typing import Any, Optional
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
|
|
10
|
+
from ..core.context import ClassificationContext
|
|
11
|
+
from ..core.exceptions import InputValidationError, UnsupportedTaskError
|
|
12
|
+
from ..core.result import MetricResult
|
|
13
|
+
from ..core.types import FloatArray, ZeroDivision
|
|
14
|
+
from ..core.validation import safe_divide, validate_zero_division
|
|
15
|
+
|
|
16
|
+
CountFn = Callable[[FloatArray, FloatArray, FloatArray, FloatArray, ZeroDivision], FloatArray]
|
|
17
|
+
|
|
18
|
+
AVERAGES = ("auto", "binary", "micro", "macro", "weighted", "samples", None)
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def resolve_average(ctx: ClassificationContext, average: Optional[str]) -> Optional[str]:
|
|
22
|
+
if average not in AVERAGES:
|
|
23
|
+
raise UnsupportedTaskError(
|
|
24
|
+
f"average={average!r} is not supported. Use one of: 'binary', 'micro', 'macro', 'weighted', "
|
|
25
|
+
"'samples' (multilabel) or None (per-class values)."
|
|
26
|
+
)
|
|
27
|
+
if average == "auto":
|
|
28
|
+
return "binary" if ctx.target_type == "binary" else "macro"
|
|
29
|
+
if average == "binary" and ctx.target_type != "binary":
|
|
30
|
+
raise UnsupportedTaskError(
|
|
31
|
+
f"average='binary' needs a binary target, but this target is {ctx.target_type} "
|
|
32
|
+
f"({len(ctx.labels)} labels). Choose average='macro', 'micro' or 'weighted', "
|
|
33
|
+
"or None for per-class values."
|
|
34
|
+
)
|
|
35
|
+
if average == "samples" and ctx.target_type != "multilabel":
|
|
36
|
+
raise UnsupportedTaskError("average='samples' is only defined for multilabel targets.")
|
|
37
|
+
return average
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def positive_index(ctx: ClassificationContext, pos_label: Any) -> Optional[int]:
|
|
41
|
+
"""Index of the positive class in ctx.labels, or None if it never occurs (all-negative data)."""
|
|
42
|
+
labels = ctx.labels
|
|
43
|
+
if pos_label is None:
|
|
44
|
+
if set(labels.tolist()) <= {0, 1, -1}:
|
|
45
|
+
pos_label = 1
|
|
46
|
+
else:
|
|
47
|
+
raise InputValidationError(
|
|
48
|
+
f"Labels are {labels.tolist()}; specify which one is positive with pos_label=..."
|
|
49
|
+
)
|
|
50
|
+
hits = np.flatnonzero(labels == pos_label)
|
|
51
|
+
if hits.size:
|
|
52
|
+
return int(hits[0])
|
|
53
|
+
if labels.shape[0] <= 1:
|
|
54
|
+
return None
|
|
55
|
+
raise InputValidationError(f"pos_label={pos_label!r} is not one of the labels {labels.tolist()}.")
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def averaged(
|
|
59
|
+
ctx: ClassificationContext,
|
|
60
|
+
fn: CountFn,
|
|
61
|
+
*,
|
|
62
|
+
metric: str,
|
|
63
|
+
name: str,
|
|
64
|
+
average: Optional[str],
|
|
65
|
+
pos_label: Any = None,
|
|
66
|
+
zero_division: ZeroDivision = "warn",
|
|
67
|
+
extra_params: Optional[dict[str, Any]] = None,
|
|
68
|
+
) -> MetricResult:
|
|
69
|
+
"""Apply a vectorised count metric fn(tp, fp, fn, tn) with the requested averaging."""
|
|
70
|
+
zd = validate_zero_division(zero_division)
|
|
71
|
+
avg = resolve_average(ctx, average)
|
|
72
|
+
c = ctx.counts
|
|
73
|
+
params: dict[str, Any] = {"average": avg, "zero_division": zero_division, **(extra_params or {})}
|
|
74
|
+
|
|
75
|
+
if avg == "binary":
|
|
76
|
+
idx = positive_index(ctx, pos_label)
|
|
77
|
+
if idx is None:
|
|
78
|
+
tp = fp = fnn = np.zeros(1)
|
|
79
|
+
tn = np.array([ctx.total_weight])
|
|
80
|
+
else:
|
|
81
|
+
tp, fp, fnn, tn = (c[k][idx : idx + 1] for k in ("tp", "fp", "fn", "tn"))
|
|
82
|
+
params["pos_label"] = 1 if pos_label is None else pos_label
|
|
83
|
+
return MetricResult(metric, name, float(fn(tp, fp, fnn, tn, zd)[0]), params)
|
|
84
|
+
|
|
85
|
+
if avg == "micro":
|
|
86
|
+
tp_s, fp_s, fn_s, tn_s = (np.array([c[k].sum()]) for k in ("tp", "fp", "fn", "tn"))
|
|
87
|
+
return MetricResult(metric, name, float(fn(tp_s, fp_s, fn_s, tn_s, zd)[0]), params)
|
|
88
|
+
|
|
89
|
+
if avg == "samples":
|
|
90
|
+
w = ctx.weights
|
|
91
|
+
t = ctx.y_true.astype(bool)
|
|
92
|
+
if ctx.y_pred is None:
|
|
93
|
+
raise InputValidationError("This metric needs y_pred (predicted labels).")
|
|
94
|
+
p = ctx.y_pred.astype(bool)
|
|
95
|
+
per_sample = fn(
|
|
96
|
+
(t & p).sum(1).astype(float),
|
|
97
|
+
(~t & p).sum(1).astype(float),
|
|
98
|
+
(t & ~p).sum(1).astype(float),
|
|
99
|
+
(~t & ~p).sum(1).astype(float),
|
|
100
|
+
zd,
|
|
101
|
+
)
|
|
102
|
+
return MetricResult(metric, name, float(np.average(per_sample, weights=w)), params)
|
|
103
|
+
|
|
104
|
+
per_class = fn(c["tp"], c["fp"], c["fn"], c["tn"], zd)
|
|
105
|
+
if avg is None:
|
|
106
|
+
return MetricResult(metric, name, per_class, params, labels=tuple(ctx.labels.tolist()))
|
|
107
|
+
if avg == "macro":
|
|
108
|
+
return MetricResult(metric, name, float(np.mean(per_class)), params)
|
|
109
|
+
support = c["support"]
|
|
110
|
+
value = safe_divide((per_class * support).sum(), support.sum(), zero_division=zd, metric=name)
|
|
111
|
+
return MetricResult(metric, name, float(value), params)
|