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
|
@@ -0,0 +1,872 @@
|
|
|
1
|
+
"""Classification metrics for binary, multiclass and multilabel targets.
|
|
2
|
+
|
|
3
|
+
Conventions (stated once, applied everywhere):
|
|
4
|
+
|
|
5
|
+
* ``labels`` order defines per-class outputs and the columns of 2-D ``y_prob``. By default it is the sorted
|
|
6
|
+
set of labels present in ``y_true`` and ``y_pred``.
|
|
7
|
+
* ``average="auto"`` means ``"binary"`` for binary targets and ``"macro"`` otherwise; the resolved value is
|
|
8
|
+
recorded in ``result.params["average"]``.
|
|
9
|
+
* Undefined ratios (zero denominators) follow ``zero_division``: ``"warn"`` returns 0 with an
|
|
10
|
+
:class:`~evalsuite.core.exceptions.UndefinedMetricWarning`; ``np.nan`` propagates NaN.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
from typing import Any, Literal, Optional
|
|
16
|
+
|
|
17
|
+
import numpy as np
|
|
18
|
+
|
|
19
|
+
from ..core.context import ClassificationContext
|
|
20
|
+
from ..core.exceptions import InputValidationError, MetricInputError, UnsupportedTaskError
|
|
21
|
+
from ..core.registry import register
|
|
22
|
+
from ..core.result import MetricResult
|
|
23
|
+
from ..core.types import ArrayLike, FloatArray, ZeroDivision
|
|
24
|
+
from ..core.validation import safe_divide, validate_zero_division
|
|
25
|
+
from ._common import averaged, positive_index, resolve_average
|
|
26
|
+
|
|
27
|
+
__all__ = [
|
|
28
|
+
"accuracy",
|
|
29
|
+
"average_precision",
|
|
30
|
+
"balanced_accuracy",
|
|
31
|
+
"brier_score",
|
|
32
|
+
"cohen_kappa",
|
|
33
|
+
"confusion_matrix",
|
|
34
|
+
"f1",
|
|
35
|
+
"fbeta",
|
|
36
|
+
"hamming_loss",
|
|
37
|
+
"jaccard",
|
|
38
|
+
"log_loss",
|
|
39
|
+
"mcc",
|
|
40
|
+
"npv",
|
|
41
|
+
"pr_curve",
|
|
42
|
+
"precision",
|
|
43
|
+
"recall",
|
|
44
|
+
"roc_auc",
|
|
45
|
+
"roc_curve",
|
|
46
|
+
"specificity",
|
|
47
|
+
"top_k_accuracy",
|
|
48
|
+
]
|
|
49
|
+
|
|
50
|
+
_C = "classification"
|
|
51
|
+
_REF_SOKOLOVA = "Sokolova M, Lapalme G. A systematic analysis of performance measures for classification tasks. Information Processing & Management. 2009;45(4):427-437."
|
|
52
|
+
_REF_POWERS = "Powers DMW. Evaluation: from precision, recall and F-measure to ROC, informedness, markedness and correlation. Journal of Machine Learning Technologies. 2011;2(1):37-63."
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def _ctx(y_true: ArrayLike, y_pred: Optional[ArrayLike] = None, **kw: Any) -> ClassificationContext:
|
|
56
|
+
return ClassificationContext(y_true, y_pred, **kw)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
# ---- count formulas (vectorised; shared by every averaging mode) ----------------------------------
|
|
60
|
+
def _precision(tp: FloatArray, fp: FloatArray, fn: FloatArray, tn: FloatArray, zd: ZeroDivision) -> FloatArray:
|
|
61
|
+
return safe_divide(tp, tp + fp, zero_division=zd, metric="Precision")
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def _recall(tp: FloatArray, fp: FloatArray, fn: FloatArray, tn: FloatArray, zd: ZeroDivision) -> FloatArray:
|
|
65
|
+
return safe_divide(tp, tp + fn, zero_division=zd, metric="Recall")
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _specificity(tp: FloatArray, fp: FloatArray, fn: FloatArray, tn: FloatArray, zd: ZeroDivision) -> FloatArray:
|
|
69
|
+
return safe_divide(tn, tn + fp, zero_division=zd, metric="Specificity")
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def _npv(tp: FloatArray, fp: FloatArray, fn: FloatArray, tn: FloatArray, zd: ZeroDivision) -> FloatArray:
|
|
73
|
+
return safe_divide(tn, tn + fn, zero_division=zd, metric="NPV")
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def _jaccard(tp: FloatArray, fp: FloatArray, fn: FloatArray, tn: FloatArray, zd: ZeroDivision) -> FloatArray:
|
|
77
|
+
return safe_divide(tp, tp + fp + fn, zero_division=zd, metric="Jaccard index")
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def _fbeta_fn(beta: float) -> Any:
|
|
81
|
+
b2 = beta * beta
|
|
82
|
+
|
|
83
|
+
def fn(tp: FloatArray, fp: FloatArray, fn_: FloatArray, tn: FloatArray, zd: ZeroDivision) -> FloatArray:
|
|
84
|
+
return safe_divide((1 + b2) * tp, (1 + b2) * tp + b2 * fn_ + fp, zero_division=zd, metric=f"F{beta:g}")
|
|
85
|
+
|
|
86
|
+
return fn
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
_AVG_DOC = "average, labels, pos_label, sample_weight, zero_division"
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
# ---- label-based metrics --------------------------------------------------------------------------
|
|
93
|
+
@register(
|
|
94
|
+
category=_C,
|
|
95
|
+
task="binary, multiclass, multilabel",
|
|
96
|
+
name="Accuracy",
|
|
97
|
+
definition="Proportion of observations whose predicted label equals the true label (exact match for multilabel).",
|
|
98
|
+
formula="(1/n) Σ 1[ŷᵢ = yᵢ]",
|
|
99
|
+
range="[0, 1]",
|
|
100
|
+
input_requirements=("y_true", "y_pred"),
|
|
101
|
+
references=(_REF_SOKOLOVA,),
|
|
102
|
+
)
|
|
103
|
+
def accuracy(y_true: ArrayLike, y_pred: ArrayLike, *, sample_weight: Optional[ArrayLike] = None) -> MetricResult:
|
|
104
|
+
"""Accuracy. For multilabel targets this is subset accuracy (all labels of a row must match)."""
|
|
105
|
+
return _accuracy(_ctx(y_true, y_pred, sample_weight=sample_weight))
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def _accuracy(ctx: ClassificationContext) -> MetricResult:
|
|
109
|
+
if ctx.target_type == "multilabel":
|
|
110
|
+
assert ctx.y_pred is not None # noqa: S101 - ensured by context for label metrics
|
|
111
|
+
correct = (ctx.y_true == ctx.y_pred).all(axis=1)
|
|
112
|
+
else:
|
|
113
|
+
correct = ctx.true_idx == ctx.pred_idx
|
|
114
|
+
return MetricResult("accuracy", "Accuracy", float(np.average(correct, weights=ctx.weights)))
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
@register(
|
|
118
|
+
category=_C,
|
|
119
|
+
task="binary, multiclass, multilabel",
|
|
120
|
+
name="Precision",
|
|
121
|
+
definition="Of the observations predicted positive, the proportion that are truly positive.",
|
|
122
|
+
formula="TP / (TP + FP)",
|
|
123
|
+
range="[0, 1]",
|
|
124
|
+
input_requirements=("y_true", "y_pred"),
|
|
125
|
+
references=(_REF_SOKOLOVA,),
|
|
126
|
+
)
|
|
127
|
+
def precision(
|
|
128
|
+
y_true: ArrayLike,
|
|
129
|
+
y_pred: ArrayLike,
|
|
130
|
+
*,
|
|
131
|
+
average: Optional[str] = "auto",
|
|
132
|
+
labels: Optional[ArrayLike] = None,
|
|
133
|
+
pos_label: Any = None,
|
|
134
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
135
|
+
zero_division: ZeroDivision = "warn",
|
|
136
|
+
) -> MetricResult:
|
|
137
|
+
"""Precision (positive predictive value)."""
|
|
138
|
+
ctx = _ctx(y_true, y_pred, labels=labels, sample_weight=sample_weight)
|
|
139
|
+
return averaged(
|
|
140
|
+
ctx,
|
|
141
|
+
_precision,
|
|
142
|
+
metric="precision",
|
|
143
|
+
name="Precision",
|
|
144
|
+
average=average,
|
|
145
|
+
pos_label=pos_label,
|
|
146
|
+
zero_division=zero_division,
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
@register(
|
|
151
|
+
category=_C,
|
|
152
|
+
task="binary, multiclass, multilabel",
|
|
153
|
+
name="Recall (sensitivity)",
|
|
154
|
+
definition="Of the truly positive observations, the proportion predicted positive.",
|
|
155
|
+
formula="TP / (TP + FN)",
|
|
156
|
+
range="[0, 1]",
|
|
157
|
+
input_requirements=("y_true", "y_pred"),
|
|
158
|
+
references=(_REF_SOKOLOVA,),
|
|
159
|
+
)
|
|
160
|
+
def recall(
|
|
161
|
+
y_true: ArrayLike,
|
|
162
|
+
y_pred: ArrayLike,
|
|
163
|
+
*,
|
|
164
|
+
average: Optional[str] = "auto",
|
|
165
|
+
labels: Optional[ArrayLike] = None,
|
|
166
|
+
pos_label: Any = None,
|
|
167
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
168
|
+
zero_division: ZeroDivision = "warn",
|
|
169
|
+
) -> MetricResult:
|
|
170
|
+
"""Recall (sensitivity, true positive rate)."""
|
|
171
|
+
ctx = _ctx(y_true, y_pred, labels=labels, sample_weight=sample_weight)
|
|
172
|
+
return averaged(
|
|
173
|
+
ctx,
|
|
174
|
+
_recall,
|
|
175
|
+
metric="recall",
|
|
176
|
+
name="Recall",
|
|
177
|
+
average=average,
|
|
178
|
+
pos_label=pos_label,
|
|
179
|
+
zero_division=zero_division,
|
|
180
|
+
)
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
@register(
|
|
184
|
+
category=_C,
|
|
185
|
+
task="binary, multiclass, multilabel",
|
|
186
|
+
name="Specificity",
|
|
187
|
+
definition="Of the truly negative observations, the proportion predicted negative.",
|
|
188
|
+
formula="TN / (TN + FP)",
|
|
189
|
+
range="[0, 1]",
|
|
190
|
+
input_requirements=("y_true", "y_pred"),
|
|
191
|
+
references=(_REF_SOKOLOVA,),
|
|
192
|
+
)
|
|
193
|
+
def specificity(
|
|
194
|
+
y_true: ArrayLike,
|
|
195
|
+
y_pred: ArrayLike,
|
|
196
|
+
*,
|
|
197
|
+
average: Optional[str] = "auto",
|
|
198
|
+
labels: Optional[ArrayLike] = None,
|
|
199
|
+
pos_label: Any = None,
|
|
200
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
201
|
+
zero_division: ZeroDivision = "warn",
|
|
202
|
+
) -> MetricResult:
|
|
203
|
+
"""Specificity (true negative rate). Per class, the class is treated as positive versus the rest."""
|
|
204
|
+
ctx = _ctx(y_true, y_pred, labels=labels, sample_weight=sample_weight)
|
|
205
|
+
return averaged(
|
|
206
|
+
ctx,
|
|
207
|
+
_specificity,
|
|
208
|
+
metric="specificity",
|
|
209
|
+
name="Specificity",
|
|
210
|
+
average=average,
|
|
211
|
+
pos_label=pos_label,
|
|
212
|
+
zero_division=zero_division,
|
|
213
|
+
)
|
|
214
|
+
|
|
215
|
+
|
|
216
|
+
@register(
|
|
217
|
+
category=_C,
|
|
218
|
+
task="binary, multiclass, multilabel",
|
|
219
|
+
name="Negative predictive value",
|
|
220
|
+
definition="Of the observations predicted negative, the proportion that are truly negative.",
|
|
221
|
+
formula="TN / (TN + FN)",
|
|
222
|
+
range="[0, 1]",
|
|
223
|
+
input_requirements=("y_true", "y_pred"),
|
|
224
|
+
references=(_REF_POWERS,),
|
|
225
|
+
)
|
|
226
|
+
def npv(
|
|
227
|
+
y_true: ArrayLike,
|
|
228
|
+
y_pred: ArrayLike,
|
|
229
|
+
*,
|
|
230
|
+
average: Optional[str] = "auto",
|
|
231
|
+
labels: Optional[ArrayLike] = None,
|
|
232
|
+
pos_label: Any = None,
|
|
233
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
234
|
+
zero_division: ZeroDivision = "warn",
|
|
235
|
+
) -> MetricResult:
|
|
236
|
+
"""Negative predictive value."""
|
|
237
|
+
ctx = _ctx(y_true, y_pred, labels=labels, sample_weight=sample_weight)
|
|
238
|
+
return averaged(
|
|
239
|
+
ctx, _npv, metric="npv", name="NPV", average=average, pos_label=pos_label, zero_division=zero_division
|
|
240
|
+
)
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
@register(
|
|
244
|
+
category=_C,
|
|
245
|
+
task="binary, multiclass, multilabel",
|
|
246
|
+
name="F1 score",
|
|
247
|
+
definition="Harmonic mean of precision and recall.",
|
|
248
|
+
formula="2·TP / (2·TP + FP + FN)",
|
|
249
|
+
range="[0, 1]",
|
|
250
|
+
input_requirements=("y_true", "y_pred"),
|
|
251
|
+
references=("van Rijsbergen CJ. Information Retrieval. 2nd ed. Butterworths; 1979.", _REF_SOKOLOVA),
|
|
252
|
+
)
|
|
253
|
+
def f1(
|
|
254
|
+
y_true: ArrayLike,
|
|
255
|
+
y_pred: ArrayLike,
|
|
256
|
+
*,
|
|
257
|
+
average: Optional[str] = "auto",
|
|
258
|
+
labels: Optional[ArrayLike] = None,
|
|
259
|
+
pos_label: Any = None,
|
|
260
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
261
|
+
zero_division: ZeroDivision = "warn",
|
|
262
|
+
) -> MetricResult:
|
|
263
|
+
"""F1 score."""
|
|
264
|
+
ctx = _ctx(y_true, y_pred, labels=labels, sample_weight=sample_weight)
|
|
265
|
+
return averaged(
|
|
266
|
+
ctx,
|
|
267
|
+
_fbeta_fn(1.0),
|
|
268
|
+
metric="f1",
|
|
269
|
+
name="F1",
|
|
270
|
+
average=average,
|
|
271
|
+
pos_label=pos_label,
|
|
272
|
+
zero_division=zero_division,
|
|
273
|
+
)
|
|
274
|
+
|
|
275
|
+
|
|
276
|
+
@register(
|
|
277
|
+
category=_C,
|
|
278
|
+
task="binary, multiclass, multilabel",
|
|
279
|
+
name="F-beta score",
|
|
280
|
+
definition="Weighted harmonic mean of precision and recall; beta > 1 favours recall, beta < 1 precision.",
|
|
281
|
+
formula="(1+β²)·TP / ((1+β²)·TP + β²·FN + FP)",
|
|
282
|
+
range="[0, 1]",
|
|
283
|
+
input_requirements=("y_true", "y_pred"),
|
|
284
|
+
references=("van Rijsbergen CJ. Information Retrieval. 2nd ed. Butterworths; 1979.",),
|
|
285
|
+
)
|
|
286
|
+
def fbeta(
|
|
287
|
+
y_true: ArrayLike,
|
|
288
|
+
y_pred: ArrayLike,
|
|
289
|
+
*,
|
|
290
|
+
beta: float,
|
|
291
|
+
average: Optional[str] = "auto",
|
|
292
|
+
labels: Optional[ArrayLike] = None,
|
|
293
|
+
pos_label: Any = None,
|
|
294
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
295
|
+
zero_division: ZeroDivision = "warn",
|
|
296
|
+
) -> MetricResult:
|
|
297
|
+
"""F-beta score."""
|
|
298
|
+
if not (isinstance(beta, (int, float)) and beta > 0 and np.isfinite(beta)):
|
|
299
|
+
raise InputValidationError("beta must be a positive finite number.")
|
|
300
|
+
ctx = _ctx(y_true, y_pred, labels=labels, sample_weight=sample_weight)
|
|
301
|
+
return averaged(
|
|
302
|
+
ctx,
|
|
303
|
+
_fbeta_fn(float(beta)),
|
|
304
|
+
metric="fbeta",
|
|
305
|
+
name=f"F{beta:g}",
|
|
306
|
+
average=average,
|
|
307
|
+
pos_label=pos_label,
|
|
308
|
+
zero_division=zero_division,
|
|
309
|
+
extra_params={"beta": float(beta)},
|
|
310
|
+
)
|
|
311
|
+
|
|
312
|
+
|
|
313
|
+
@register(
|
|
314
|
+
category=_C,
|
|
315
|
+
task="binary, multiclass, multilabel",
|
|
316
|
+
name="Jaccard index",
|
|
317
|
+
definition="Size of the intersection over the size of the union of predicted and true positives.",
|
|
318
|
+
formula="TP / (TP + FP + FN)",
|
|
319
|
+
range="[0, 1]",
|
|
320
|
+
input_requirements=("y_true", "y_pred"),
|
|
321
|
+
references=(
|
|
322
|
+
"Jaccard P. The distribution of the flora in the alpine zone. New Phytologist. 1912;11(2):37-50.",
|
|
323
|
+
),
|
|
324
|
+
)
|
|
325
|
+
def jaccard(
|
|
326
|
+
y_true: ArrayLike,
|
|
327
|
+
y_pred: ArrayLike,
|
|
328
|
+
*,
|
|
329
|
+
average: Optional[str] = "auto",
|
|
330
|
+
labels: Optional[ArrayLike] = None,
|
|
331
|
+
pos_label: Any = None,
|
|
332
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
333
|
+
zero_division: ZeroDivision = "warn",
|
|
334
|
+
) -> MetricResult:
|
|
335
|
+
"""Jaccard index (intersection over union)."""
|
|
336
|
+
ctx = _ctx(y_true, y_pred, labels=labels, sample_weight=sample_weight)
|
|
337
|
+
return averaged(
|
|
338
|
+
ctx,
|
|
339
|
+
_jaccard,
|
|
340
|
+
metric="jaccard",
|
|
341
|
+
name="Jaccard",
|
|
342
|
+
average=average,
|
|
343
|
+
pos_label=pos_label,
|
|
344
|
+
zero_division=zero_division,
|
|
345
|
+
)
|
|
346
|
+
|
|
347
|
+
|
|
348
|
+
@register(
|
|
349
|
+
category=_C,
|
|
350
|
+
task="binary, multiclass",
|
|
351
|
+
name="Balanced accuracy",
|
|
352
|
+
definition="Mean recall over classes present in y_true; robust to class imbalance.",
|
|
353
|
+
formula="(1/K) Σₖ TPₖ / (TPₖ + FNₖ)",
|
|
354
|
+
range="[0, 1] (adjusted: chance = 0)",
|
|
355
|
+
input_requirements=("y_true", "y_pred"),
|
|
356
|
+
references=(
|
|
357
|
+
"Brodersen KH, Ong CS, Stephan KE, Buhmann JM. The balanced accuracy and its posterior distribution. ICPR 2010:3121-3124.",
|
|
358
|
+
),
|
|
359
|
+
)
|
|
360
|
+
def balanced_accuracy(
|
|
361
|
+
y_true: ArrayLike, y_pred: ArrayLike, *, adjusted: bool = False, sample_weight: Optional[ArrayLike] = None
|
|
362
|
+
) -> MetricResult:
|
|
363
|
+
"""Balanced accuracy. With ``adjusted=True`` chance performance scores 0 and perfect performance 1."""
|
|
364
|
+
return _balanced_accuracy(_ctx(y_true, y_pred, sample_weight=sample_weight), adjusted=adjusted)
|
|
365
|
+
|
|
366
|
+
|
|
367
|
+
def _balanced_accuracy(ctx: ClassificationContext, *, adjusted: bool = False) -> MetricResult:
|
|
368
|
+
if ctx.target_type == "multilabel":
|
|
369
|
+
raise UnsupportedTaskError("Balanced accuracy is defined for binary and multiclass targets only.")
|
|
370
|
+
c = ctx.counts
|
|
371
|
+
present = c["support"] > 0
|
|
372
|
+
score = float(np.mean(c["tp"][present] / c["support"][present]))
|
|
373
|
+
if adjusted:
|
|
374
|
+
k = int(present.sum())
|
|
375
|
+
chance = 1.0 / k
|
|
376
|
+
score = (score - chance) / (1 - chance) if k > 1 else 0.0
|
|
377
|
+
return MetricResult("balanced_accuracy", "Balanced accuracy", score, {"adjusted": adjusted})
|
|
378
|
+
|
|
379
|
+
|
|
380
|
+
@register(
|
|
381
|
+
category=_C,
|
|
382
|
+
task="binary, multiclass",
|
|
383
|
+
name="Matthews correlation coefficient",
|
|
384
|
+
definition="Correlation between true and predicted labels using all confusion-matrix cells; 0 is chance level.",
|
|
385
|
+
formula="(c·s − Σₖ pₖtₖ) / √((s² − Σₖ pₖ²)(s² − Σₖ tₖ²))",
|
|
386
|
+
range="[-1, 1]",
|
|
387
|
+
input_requirements=("y_true", "y_pred"),
|
|
388
|
+
references=(
|
|
389
|
+
"Matthews BW. Comparison of the predicted and observed secondary structure of T4 phage lysozyme. Biochim Biophys Acta. 1975;405(2):442-451.",
|
|
390
|
+
"Gorodkin J. Comparing two K-category assignments by a K-category correlation coefficient. Comput Biol Chem. 2004;28(5-6):367-374.",
|
|
391
|
+
"Chicco D, Jurman G. The advantages of the Matthews correlation coefficient (MCC) over F1 score and accuracy in binary classification evaluation. BMC Genomics. 2020;21:6.",
|
|
392
|
+
),
|
|
393
|
+
)
|
|
394
|
+
def mcc(
|
|
395
|
+
y_true: ArrayLike,
|
|
396
|
+
y_pred: ArrayLike,
|
|
397
|
+
*,
|
|
398
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
399
|
+
zero_division: ZeroDivision = "warn",
|
|
400
|
+
) -> MetricResult:
|
|
401
|
+
"""Matthews correlation coefficient (Gorodkin's multiclass generalisation)."""
|
|
402
|
+
return _mcc(_ctx(y_true, y_pred, sample_weight=sample_weight), zero_division)
|
|
403
|
+
|
|
404
|
+
|
|
405
|
+
def _mcc(ctx: ClassificationContext, zero_division: ZeroDivision = "warn") -> MetricResult:
|
|
406
|
+
if ctx.target_type == "multilabel":
|
|
407
|
+
raise UnsupportedTaskError("MCC is defined for binary and multiclass targets; use it per label instead.")
|
|
408
|
+
cm = ctx.confusion_matrix
|
|
409
|
+
t, p = cm.sum(1), cm.sum(0)
|
|
410
|
+
c, s = np.trace(cm), cm.sum()
|
|
411
|
+
cov_ytyp = c * s - (t * p).sum()
|
|
412
|
+
cov_ypyp = s * s - (p * p).sum()
|
|
413
|
+
cov_ytyt = s * s - (t * t).sum()
|
|
414
|
+
value = safe_divide(
|
|
415
|
+
cov_ytyp,
|
|
416
|
+
np.sqrt(cov_ytyt * cov_ypyp),
|
|
417
|
+
zero_division=validate_zero_division(zero_division),
|
|
418
|
+
metric="MCC",
|
|
419
|
+
)
|
|
420
|
+
return MetricResult("mcc", "MCC", float(value), {"zero_division": zero_division})
|
|
421
|
+
|
|
422
|
+
|
|
423
|
+
@register(
|
|
424
|
+
category=_C,
|
|
425
|
+
task="binary, multiclass",
|
|
426
|
+
name="Cohen's kappa",
|
|
427
|
+
definition="Agreement between true and predicted labels corrected for agreement expected by chance.",
|
|
428
|
+
formula="κ = 1 − Σ wᵢⱼ Oᵢⱼ / Σ wᵢⱼ Eᵢⱼ",
|
|
429
|
+
range="[-1, 1]",
|
|
430
|
+
input_requirements=("y_true", "y_pred"),
|
|
431
|
+
references=(
|
|
432
|
+
"Cohen J. A coefficient of agreement for nominal scales. Educ Psychol Meas. 1960;20(1):37-46.",
|
|
433
|
+
"Cohen J. Weighted kappa. Psychol Bull. 1968;70(4):213-220.",
|
|
434
|
+
),
|
|
435
|
+
)
|
|
436
|
+
def cohen_kappa(
|
|
437
|
+
y_true: ArrayLike,
|
|
438
|
+
y_pred: ArrayLike,
|
|
439
|
+
*,
|
|
440
|
+
weights: Optional[Literal["linear", "quadratic"]] = None,
|
|
441
|
+
labels: Optional[ArrayLike] = None,
|
|
442
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
443
|
+
zero_division: ZeroDivision = "warn",
|
|
444
|
+
) -> MetricResult:
|
|
445
|
+
"""Cohen's kappa, optionally linearly or quadratically weighted (for ordinal labels, in label order)."""
|
|
446
|
+
return _cohen_kappa(_ctx(y_true, y_pred, labels=labels, sample_weight=sample_weight), weights, zero_division)
|
|
447
|
+
|
|
448
|
+
|
|
449
|
+
def _cohen_kappa(
|
|
450
|
+
ctx: ClassificationContext, weights: Optional[str] = None, zero_division: ZeroDivision = "warn"
|
|
451
|
+
) -> MetricResult:
|
|
452
|
+
if ctx.target_type == "multilabel":
|
|
453
|
+
raise UnsupportedTaskError("Cohen's kappa is defined for binary and multiclass targets.")
|
|
454
|
+
cm = ctx.confusion_matrix
|
|
455
|
+
k = cm.shape[0]
|
|
456
|
+
expected = np.outer(cm.sum(1), cm.sum(0)) / cm.sum()
|
|
457
|
+
if weights is None:
|
|
458
|
+
w = 1.0 - np.eye(k)
|
|
459
|
+
elif weights in ("linear", "quadratic"):
|
|
460
|
+
grid = np.abs(np.subtract.outer(np.arange(k), np.arange(k))).astype(float)
|
|
461
|
+
w = grid if weights == "linear" else grid**2
|
|
462
|
+
else:
|
|
463
|
+
raise InputValidationError("weights must be None, 'linear' or 'quadratic'.")
|
|
464
|
+
ratio = safe_divide(
|
|
465
|
+
(w * cm).sum(),
|
|
466
|
+
(w * expected).sum(),
|
|
467
|
+
zero_division=validate_zero_division(zero_division),
|
|
468
|
+
metric="Cohen's kappa",
|
|
469
|
+
)
|
|
470
|
+
return MetricResult("cohen_kappa", "Cohen's kappa", float(1 - ratio), {"weights": weights})
|
|
471
|
+
|
|
472
|
+
|
|
473
|
+
@register(
|
|
474
|
+
category=_C,
|
|
475
|
+
task="binary, multiclass, multilabel",
|
|
476
|
+
name="Hamming loss",
|
|
477
|
+
definition="Fraction of labels predicted incorrectly (for single-label targets equal to 1 − accuracy).",
|
|
478
|
+
formula="(1/(n·L)) Σᵢ Σₗ 1[ŷᵢₗ ≠ yᵢₗ]",
|
|
479
|
+
range="[0, 1]",
|
|
480
|
+
input_requirements=("y_true", "y_pred"),
|
|
481
|
+
references=(
|
|
482
|
+
"Tsoumakas G, Katakis I. Multi-label classification: an overview. Int J Data Warehousing and Mining. 2007;3(3):1-13.",
|
|
483
|
+
),
|
|
484
|
+
higher_is_better=False,
|
|
485
|
+
)
|
|
486
|
+
def hamming_loss(
|
|
487
|
+
y_true: ArrayLike, y_pred: ArrayLike, *, sample_weight: Optional[ArrayLike] = None
|
|
488
|
+
) -> MetricResult:
|
|
489
|
+
"""Hamming loss."""
|
|
490
|
+
ctx = _ctx(y_true, y_pred, sample_weight=sample_weight)
|
|
491
|
+
if ctx.target_type == "multilabel":
|
|
492
|
+
assert ctx.y_pred is not None # noqa: S101
|
|
493
|
+
wrong = (ctx.y_true != ctx.y_pred).mean(axis=1)
|
|
494
|
+
else:
|
|
495
|
+
wrong = (ctx.true_idx != ctx.pred_idx).astype(float)
|
|
496
|
+
return MetricResult("hamming_loss", "Hamming loss", float(np.average(wrong, weights=ctx.weights)))
|
|
497
|
+
|
|
498
|
+
|
|
499
|
+
def confusion_matrix(
|
|
500
|
+
y_true: ArrayLike,
|
|
501
|
+
y_pred: ArrayLike,
|
|
502
|
+
*,
|
|
503
|
+
labels: Optional[ArrayLike] = None,
|
|
504
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
505
|
+
normalize: Optional[Literal["true", "pred", "all"]] = None,
|
|
506
|
+
) -> np.ndarray:
|
|
507
|
+
"""Confusion matrix: rows are true labels, columns predicted labels, both in ``labels`` order
|
|
508
|
+
(sorted by default). Counts are integers unless ``sample_weight`` or ``normalize`` is given."""
|
|
509
|
+
ctx = _ctx(y_true, y_pred, labels=labels, sample_weight=sample_weight)
|
|
510
|
+
cm = ctx.confusion_matrix
|
|
511
|
+
if normalize is not None:
|
|
512
|
+
if normalize not in ("true", "pred", "all"):
|
|
513
|
+
raise InputValidationError("normalize must be None, 'true', 'pred' or 'all'.")
|
|
514
|
+
den = (
|
|
515
|
+
cm.sum(1, keepdims=True)
|
|
516
|
+
if normalize == "true"
|
|
517
|
+
else cm.sum(0, keepdims=True)
|
|
518
|
+
if normalize == "pred"
|
|
519
|
+
else cm.sum()
|
|
520
|
+
)
|
|
521
|
+
with np.errstate(invalid="ignore", divide="ignore"):
|
|
522
|
+
return np.nan_to_num(cm / den)
|
|
523
|
+
return cm if sample_weight is not None else cm.astype(np.int64)
|
|
524
|
+
|
|
525
|
+
|
|
526
|
+
# ---- probability-based metrics --------------------------------------------------------------------
|
|
527
|
+
def _binary_curve(y: FloatArray, score: FloatArray, w: FloatArray) -> tuple[FloatArray, FloatArray, FloatArray]:
|
|
528
|
+
"""Cumulative weighted false/true positives at each distinct score threshold (descending)."""
|
|
529
|
+
order = np.argsort(-score, kind="mergesort")
|
|
530
|
+
score, y, w = score[order], y[order], w[order]
|
|
531
|
+
distinct = np.flatnonzero(np.diff(score))
|
|
532
|
+
thr_idx = np.r_[distinct, y.size - 1]
|
|
533
|
+
tps = np.cumsum(y * w)[thr_idx]
|
|
534
|
+
fps = np.cumsum((1 - y) * w)[thr_idx]
|
|
535
|
+
return fps, tps, score[thr_idx]
|
|
536
|
+
|
|
537
|
+
|
|
538
|
+
def _require_both_classes(tps: FloatArray, fps: FloatArray, what: str) -> None:
|
|
539
|
+
if tps[-1] == 0 or fps[-1] == 0:
|
|
540
|
+
missing = "positive" if tps[-1] == 0 else "negative"
|
|
541
|
+
raise MetricInputError(
|
|
542
|
+
f"{what} is undefined because y_true contains no {missing} observations. "
|
|
543
|
+
"It needs at least one positive and one negative example."
|
|
544
|
+
)
|
|
545
|
+
|
|
546
|
+
|
|
547
|
+
def _binary_auc(y: FloatArray, score: FloatArray, w: FloatArray) -> float:
|
|
548
|
+
fps, tps, _ = _binary_curve(y, score, w)
|
|
549
|
+
_require_both_classes(tps, fps, "ROC AUC")
|
|
550
|
+
fpr = np.r_[0.0, fps / fps[-1]]
|
|
551
|
+
tpr = np.r_[0.0, tps / tps[-1]]
|
|
552
|
+
integrate = getattr(np, "trapezoid", None) or getattr(np, "trapz") # noqa: B009 - NumPy 1.x/2.x
|
|
553
|
+
return float(integrate(tpr, fpr))
|
|
554
|
+
|
|
555
|
+
|
|
556
|
+
def _binary_ap(y: FloatArray, score: FloatArray, w: FloatArray) -> float:
|
|
557
|
+
fps, tps, _ = _binary_curve(y, score, w)
|
|
558
|
+
if tps[-1] == 0:
|
|
559
|
+
raise MetricInputError("Average precision is undefined because y_true contains no positive observations.")
|
|
560
|
+
prec = tps / (tps + fps)
|
|
561
|
+
rec = tps / tps[-1]
|
|
562
|
+
return float(np.sum(np.diff(np.r_[0.0, rec]) * prec))
|
|
563
|
+
|
|
564
|
+
|
|
565
|
+
def _binary_target(ctx: ClassificationContext, pos_label: Any) -> FloatArray:
|
|
566
|
+
idx = positive_index(ctx, pos_label)
|
|
567
|
+
if idx is None:
|
|
568
|
+
return np.zeros(ctx.n)
|
|
569
|
+
target: FloatArray = (ctx.true_idx == idx).astype(np.float64)
|
|
570
|
+
return target
|
|
571
|
+
|
|
572
|
+
|
|
573
|
+
@register(
|
|
574
|
+
category=_C,
|
|
575
|
+
task="binary, multiclass, multilabel",
|
|
576
|
+
name="ROC AUC",
|
|
577
|
+
definition="Area under the receiver operating characteristic curve: the probability that a random positive "
|
|
578
|
+
"is scored above a random negative (ties count half).",
|
|
579
|
+
formula="∫₀¹ TPR d(FPR)",
|
|
580
|
+
range="[0, 1] (0.5 = chance)",
|
|
581
|
+
input_requirements=("y_true", "y_prob"),
|
|
582
|
+
references=(
|
|
583
|
+
"Hanley JA, McNeil BJ. The meaning and use of the area under a receiver operating characteristic (ROC) curve. Radiology. 1982;143(1):29-36.",
|
|
584
|
+
"Fawcett T. An introduction to ROC analysis. Pattern Recognit Lett. 2006;27(8):861-874.",
|
|
585
|
+
"Hand DJ, Till RJ. A simple generalisation of the area under the ROC curve for multiple class classification problems. Mach Learn. 2001;45:171-186.",
|
|
586
|
+
),
|
|
587
|
+
)
|
|
588
|
+
def roc_auc(
|
|
589
|
+
y_true: ArrayLike,
|
|
590
|
+
y_prob: ArrayLike,
|
|
591
|
+
*,
|
|
592
|
+
average: Optional[str] = "auto",
|
|
593
|
+
multi_class: Literal["ovr", "ovo"] = "ovr",
|
|
594
|
+
labels: Optional[ArrayLike] = None,
|
|
595
|
+
pos_label: Any = None,
|
|
596
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
597
|
+
) -> MetricResult:
|
|
598
|
+
"""ROC AUC. Binary: ``y_prob`` is P(positive). Multiclass: one column per class (label order), averaged
|
|
599
|
+
one-vs-rest (``"ovr"``, macro or weighted) or one-vs-one (``"ovo"``, Hand & Till, macro).
|
|
600
|
+
Multilabel: one column per label, macro/micro/weighted or per-label (``average=None``)."""
|
|
601
|
+
ctx = _ctx(y_true, None, y_prob=y_prob, labels=labels, sample_weight=sample_weight)
|
|
602
|
+
return _roc_auc(ctx, average=average, multi_class=multi_class, pos_label=pos_label)
|
|
603
|
+
|
|
604
|
+
|
|
605
|
+
def _per_column(
|
|
606
|
+
ctx: ClassificationContext,
|
|
607
|
+
fn: Any,
|
|
608
|
+
average: Optional[str],
|
|
609
|
+
metric: str,
|
|
610
|
+
name: str,
|
|
611
|
+
params: dict[str, Any],
|
|
612
|
+
) -> MetricResult:
|
|
613
|
+
"""Apply a binary score metric per class/label column with macro, weighted, micro or no averaging."""
|
|
614
|
+
p = ctx.y_prob
|
|
615
|
+
w = ctx.weights
|
|
616
|
+
y = (
|
|
617
|
+
ctx.y_true.astype(float)
|
|
618
|
+
if ctx.target_type == "multilabel"
|
|
619
|
+
else (ctx.true_idx[:, None] == np.arange(ctx.labels.shape[0])[None, :]).astype(float)
|
|
620
|
+
)
|
|
621
|
+
if average == "micro":
|
|
622
|
+
return MetricResult(metric, name, fn(y.ravel(), p.ravel(), np.repeat(w, y.shape[1])), params)
|
|
623
|
+
scores = np.array([fn(y[:, j], p[:, j], w) for j in range(y.shape[1])])
|
|
624
|
+
if average is None:
|
|
625
|
+
return MetricResult(metric, name, scores, params, labels=tuple(ctx.labels.tolist()))
|
|
626
|
+
if average == "macro":
|
|
627
|
+
return MetricResult(metric, name, float(scores.mean()), params)
|
|
628
|
+
support = (y * w[:, None]).sum(0)
|
|
629
|
+
return MetricResult(metric, name, float(np.average(scores, weights=support)), params)
|
|
630
|
+
|
|
631
|
+
|
|
632
|
+
def _roc_auc(
|
|
633
|
+
ctx: ClassificationContext,
|
|
634
|
+
*,
|
|
635
|
+
average: Optional[str] = "auto",
|
|
636
|
+
multi_class: str = "ovr",
|
|
637
|
+
pos_label: Any = None,
|
|
638
|
+
) -> MetricResult:
|
|
639
|
+
avg = resolve_average(ctx, average)
|
|
640
|
+
params: dict[str, Any] = {"average": avg}
|
|
641
|
+
if ctx.target_type == "binary":
|
|
642
|
+
params["pos_label"] = 1 if pos_label is None else pos_label
|
|
643
|
+
return MetricResult(
|
|
644
|
+
"roc_auc", "ROC AUC", _binary_auc(_binary_target(ctx, pos_label), ctx.y_prob, ctx.weights), params
|
|
645
|
+
)
|
|
646
|
+
if ctx.target_type == "multiclass":
|
|
647
|
+
params["multi_class"] = multi_class
|
|
648
|
+
if multi_class == "ovo":
|
|
649
|
+
if avg != "macro":
|
|
650
|
+
raise UnsupportedTaskError("One-vs-one ROC AUC supports average='macro' only.")
|
|
651
|
+
return MetricResult("roc_auc", "ROC AUC", _ovo_auc(ctx), params)
|
|
652
|
+
if multi_class != "ovr":
|
|
653
|
+
raise UnsupportedTaskError("multi_class must be 'ovr' or 'ovo'.")
|
|
654
|
+
if avg == "micro":
|
|
655
|
+
raise UnsupportedTaskError(
|
|
656
|
+
"average='micro' is not defined for multiclass ROC AUC; use 'macro', 'weighted' or None."
|
|
657
|
+
)
|
|
658
|
+
return _per_column(ctx, _binary_auc, avg, "roc_auc", "ROC AUC", params)
|
|
659
|
+
|
|
660
|
+
|
|
661
|
+
def _ovo_auc(ctx: ClassificationContext) -> float:
|
|
662
|
+
"""Hand & Till (2001) multiclass AUC: mean over class pairs of the two directional AUCs."""
|
|
663
|
+
k = ctx.labels.shape[0]
|
|
664
|
+
p, t, w = ctx.y_prob, ctx.true_idx, ctx.weights
|
|
665
|
+
pair_scores = []
|
|
666
|
+
for a in range(k):
|
|
667
|
+
for b in range(a + 1, k):
|
|
668
|
+
mask = (t == a) | (t == b)
|
|
669
|
+
ya = (t[mask] == a).astype(float)
|
|
670
|
+
a_vs_b = _binary_auc(ya, p[mask, a], w[mask])
|
|
671
|
+
b_vs_a = _binary_auc(1 - ya, p[mask, b], w[mask])
|
|
672
|
+
pair_scores.append((a_vs_b + b_vs_a) / 2)
|
|
673
|
+
return float(np.mean(pair_scores))
|
|
674
|
+
|
|
675
|
+
|
|
676
|
+
@register(
|
|
677
|
+
category=_C,
|
|
678
|
+
task="binary, multiclass, multilabel",
|
|
679
|
+
name="Average precision (PR AUC)",
|
|
680
|
+
definition="Precision averaged over recall levels, weighting each precision by the increase in recall "
|
|
681
|
+
"(step interpolation, no trapezoidal optimism).",
|
|
682
|
+
formula="AP = Σₙ (Rₙ − Rₙ₋₁) Pₙ",
|
|
683
|
+
range="[0, 1] (chance = prevalence)",
|
|
684
|
+
input_requirements=("y_true", "y_prob"),
|
|
685
|
+
references=(
|
|
686
|
+
"Davis J, Goadrich M. The relationship between precision-recall and ROC curves. ICML 2006:233-240.",
|
|
687
|
+
"Saito T, Rehmsmeier M. The precision-recall plot is more informative than the ROC plot when evaluating binary classifiers on imbalanced datasets. PLoS ONE. 2015;10(3):e0118432.",
|
|
688
|
+
),
|
|
689
|
+
)
|
|
690
|
+
def average_precision(
|
|
691
|
+
y_true: ArrayLike,
|
|
692
|
+
y_prob: ArrayLike,
|
|
693
|
+
*,
|
|
694
|
+
average: Optional[str] = "auto",
|
|
695
|
+
labels: Optional[ArrayLike] = None,
|
|
696
|
+
pos_label: Any = None,
|
|
697
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
698
|
+
) -> MetricResult:
|
|
699
|
+
"""Average precision, the recommended summary of the precision-recall curve."""
|
|
700
|
+
ctx = _ctx(y_true, None, y_prob=y_prob, labels=labels, sample_weight=sample_weight)
|
|
701
|
+
return _average_precision(ctx, average=average, pos_label=pos_label)
|
|
702
|
+
|
|
703
|
+
|
|
704
|
+
def _average_precision(
|
|
705
|
+
ctx: ClassificationContext, *, average: Optional[str] = "auto", pos_label: Any = None
|
|
706
|
+
) -> MetricResult:
|
|
707
|
+
avg = resolve_average(ctx, average)
|
|
708
|
+
params: dict[str, Any] = {"average": avg}
|
|
709
|
+
if ctx.target_type == "binary":
|
|
710
|
+
params["pos_label"] = 1 if pos_label is None else pos_label
|
|
711
|
+
return MetricResult(
|
|
712
|
+
"average_precision",
|
|
713
|
+
"Average precision",
|
|
714
|
+
_binary_ap(_binary_target(ctx, pos_label), ctx.y_prob, ctx.weights),
|
|
715
|
+
params,
|
|
716
|
+
)
|
|
717
|
+
return _per_column(ctx, _binary_ap, avg, "average_precision", "Average precision", params)
|
|
718
|
+
|
|
719
|
+
|
|
720
|
+
def roc_curve(
|
|
721
|
+
y_true: ArrayLike,
|
|
722
|
+
y_prob: ArrayLike,
|
|
723
|
+
*,
|
|
724
|
+
pos_label: Any = None,
|
|
725
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
726
|
+
) -> tuple[FloatArray, FloatArray, FloatArray]:
|
|
727
|
+
"""Binary ROC curve: ``(fpr, tpr, thresholds)``, starting at (0, 0) with threshold +inf."""
|
|
728
|
+
ctx = _ctx(y_true, None, y_prob=y_prob, sample_weight=sample_weight)
|
|
729
|
+
if ctx.target_type != "binary":
|
|
730
|
+
raise UnsupportedTaskError(
|
|
731
|
+
"roc_curve is binary; for multiclass, compute one curve per class (one-vs-rest)."
|
|
732
|
+
)
|
|
733
|
+
fps, tps, thr = _binary_curve(_binary_target(ctx, pos_label), ctx.y_prob, ctx.weights)
|
|
734
|
+
_require_both_classes(tps, fps, "The ROC curve")
|
|
735
|
+
return np.r_[0.0, fps / fps[-1]], np.r_[0.0, tps / tps[-1]], np.r_[np.inf, thr]
|
|
736
|
+
|
|
737
|
+
|
|
738
|
+
def pr_curve(
|
|
739
|
+
y_true: ArrayLike,
|
|
740
|
+
y_prob: ArrayLike,
|
|
741
|
+
*,
|
|
742
|
+
pos_label: Any = None,
|
|
743
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
744
|
+
) -> tuple[FloatArray, FloatArray, FloatArray]:
|
|
745
|
+
"""Binary precision-recall curve: ``(precision, recall, thresholds)`` ordered by increasing threshold
|
|
746
|
+
(one point per distinct score) and ending at (precision=1, recall=0)."""
|
|
747
|
+
ctx = _ctx(y_true, None, y_prob=y_prob, sample_weight=sample_weight)
|
|
748
|
+
if ctx.target_type != "binary":
|
|
749
|
+
raise UnsupportedTaskError("pr_curve is binary; for multiclass, compute one curve per class.")
|
|
750
|
+
fps, tps, thr = _binary_curve(_binary_target(ctx, pos_label), ctx.y_prob, ctx.weights)
|
|
751
|
+
if tps[-1] == 0:
|
|
752
|
+
raise MetricInputError("The PR curve is undefined because y_true contains no positive observations.")
|
|
753
|
+
prec = tps / (tps + fps)
|
|
754
|
+
rec = tps / tps[-1]
|
|
755
|
+
rev = slice(None, None, -1) # increasing thresholds; every threshold is kept
|
|
756
|
+
return np.r_[prec[rev], 1.0], np.r_[rec[rev], 0.0], thr[rev]
|
|
757
|
+
|
|
758
|
+
|
|
759
|
+
@register(
|
|
760
|
+
category=_C,
|
|
761
|
+
task="binary, multiclass",
|
|
762
|
+
name="Log loss (cross-entropy)",
|
|
763
|
+
definition="Negative mean log-likelihood of the true labels under the predicted probabilities.",
|
|
764
|
+
formula="−(1/n) Σᵢ log p̂ᵢ,yᵢ",
|
|
765
|
+
range="[0, ∞)",
|
|
766
|
+
input_requirements=("y_true", "y_prob"),
|
|
767
|
+
references=("Good IJ. Rational decisions. J R Stat Soc B. 1952;14(1):107-114.",),
|
|
768
|
+
higher_is_better=False,
|
|
769
|
+
)
|
|
770
|
+
def log_loss(
|
|
771
|
+
y_true: ArrayLike,
|
|
772
|
+
y_prob: ArrayLike,
|
|
773
|
+
*,
|
|
774
|
+
labels: Optional[ArrayLike] = None,
|
|
775
|
+
pos_label: Any = None,
|
|
776
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
777
|
+
) -> MetricResult:
|
|
778
|
+
"""Log loss. Probabilities are clipped to [ε, 1 − ε] (ε = float64 machine epsilon) to keep it finite."""
|
|
779
|
+
ctx = _ctx(y_true, None, y_prob=y_prob, labels=labels, sample_weight=sample_weight)
|
|
780
|
+
return _log_loss(ctx, pos_label)
|
|
781
|
+
|
|
782
|
+
|
|
783
|
+
def _log_loss(ctx: ClassificationContext, pos_label: Any = None) -> MetricResult:
|
|
784
|
+
if ctx.target_type == "multilabel":
|
|
785
|
+
raise UnsupportedTaskError("Log loss here is for binary and multiclass targets.")
|
|
786
|
+
eps = np.finfo(np.float64).eps
|
|
787
|
+
p = np.clip(ctx.y_prob, eps, 1 - eps)
|
|
788
|
+
if ctx.target_type == "binary":
|
|
789
|
+
y = _binary_target(ctx, pos_label)
|
|
790
|
+
losses = -(y * np.log(p) + (1 - y) * np.log(1 - p))
|
|
791
|
+
else:
|
|
792
|
+
losses = -np.log(p[np.arange(ctx.n), ctx.true_idx])
|
|
793
|
+
return MetricResult(
|
|
794
|
+
"log_loss", "Log loss", float(np.average(losses, weights=ctx.weights)), {"eps": float(eps)}
|
|
795
|
+
)
|
|
796
|
+
|
|
797
|
+
|
|
798
|
+
@register(
|
|
799
|
+
category=_C,
|
|
800
|
+
task="binary, multiclass",
|
|
801
|
+
name="Brier score",
|
|
802
|
+
definition="Mean squared difference between predicted probabilities and the outcome; for multiclass, the "
|
|
803
|
+
"squared error summed over classes (Brier's original definition).",
|
|
804
|
+
formula="(1/n) Σᵢ Σₖ (p̂ᵢₖ − yᵢₖ)² (binary: (1/n) Σᵢ (p̂ᵢ − yᵢ)²)",
|
|
805
|
+
range="[0, 1] binary; [0, 2] multiclass",
|
|
806
|
+
input_requirements=("y_true", "y_prob"),
|
|
807
|
+
references=(
|
|
808
|
+
"Brier GW. Verification of forecasts expressed in terms of probability. Mon Weather Rev. 1950;78(1):1-3.",
|
|
809
|
+
),
|
|
810
|
+
higher_is_better=False,
|
|
811
|
+
)
|
|
812
|
+
def brier_score(
|
|
813
|
+
y_true: ArrayLike,
|
|
814
|
+
y_prob: ArrayLike,
|
|
815
|
+
*,
|
|
816
|
+
labels: Optional[ArrayLike] = None,
|
|
817
|
+
pos_label: Any = None,
|
|
818
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
819
|
+
) -> MetricResult:
|
|
820
|
+
"""Brier score."""
|
|
821
|
+
ctx = _ctx(y_true, None, y_prob=y_prob, labels=labels, sample_weight=sample_weight)
|
|
822
|
+
return _brier(ctx, pos_label)
|
|
823
|
+
|
|
824
|
+
|
|
825
|
+
def _brier(ctx: ClassificationContext, pos_label: Any = None) -> MetricResult:
|
|
826
|
+
if ctx.target_type == "multilabel":
|
|
827
|
+
raise UnsupportedTaskError("Brier score here is for binary and multiclass targets.")
|
|
828
|
+
if ctx.target_type == "binary":
|
|
829
|
+
err = (ctx.y_prob - _binary_target(ctx, pos_label)) ** 2
|
|
830
|
+
else:
|
|
831
|
+
onehot = np.eye(ctx.labels.shape[0])[ctx.true_idx]
|
|
832
|
+
err = ((ctx.y_prob - onehot) ** 2).sum(axis=1)
|
|
833
|
+
return MetricResult("brier_score", "Brier score", float(np.average(err, weights=ctx.weights)))
|
|
834
|
+
|
|
835
|
+
|
|
836
|
+
@register(
|
|
837
|
+
category=_C,
|
|
838
|
+
task="multiclass",
|
|
839
|
+
name="Top-k accuracy",
|
|
840
|
+
definition="Proportion of observations whose true class is among the k classes with the highest probability.",
|
|
841
|
+
formula="(1/n) Σᵢ 1[yᵢ ∈ top-k(p̂ᵢ)]",
|
|
842
|
+
range="[0, 1]",
|
|
843
|
+
input_requirements=("y_true", "y_prob"),
|
|
844
|
+
references=(
|
|
845
|
+
"Russakovsky O, et al. ImageNet Large Scale Visual Recognition Challenge. IJCV. 2015;115:211-252.",
|
|
846
|
+
),
|
|
847
|
+
)
|
|
848
|
+
def top_k_accuracy(
|
|
849
|
+
y_true: ArrayLike,
|
|
850
|
+
y_prob: ArrayLike,
|
|
851
|
+
*,
|
|
852
|
+
k: int = 2,
|
|
853
|
+
labels: Optional[ArrayLike] = None,
|
|
854
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
855
|
+
) -> MetricResult:
|
|
856
|
+
"""Top-k accuracy for multiclass probabilities. Ties at the k-th place count as a hit only if the true
|
|
857
|
+
class's probability is strictly greater than the (k+1)-th largest (no credit for unbroken ties)."""
|
|
858
|
+
ctx = _ctx(y_true, None, y_prob=y_prob, labels=labels, sample_weight=sample_weight)
|
|
859
|
+
if ctx.target_type != "multiclass":
|
|
860
|
+
raise UnsupportedTaskError(
|
|
861
|
+
"top_k_accuracy needs a multiclass target with one probability column per class."
|
|
862
|
+
)
|
|
863
|
+
n_classes = ctx.labels.shape[0]
|
|
864
|
+
if not (isinstance(k, (int, np.integer)) and 1 <= k <= n_classes):
|
|
865
|
+
raise InputValidationError(f"k must be an integer between 1 and the number of classes ({n_classes}).")
|
|
866
|
+
p = ctx.y_prob
|
|
867
|
+
true_p = p[np.arange(ctx.n), ctx.true_idx]
|
|
868
|
+
n_higher = (p > true_p[:, None]).sum(axis=1)
|
|
869
|
+
hit = n_higher < k
|
|
870
|
+
return MetricResult(
|
|
871
|
+
"top_k_accuracy", f"Top-{k} accuracy", float(np.average(hit, weights=ctx.weights)), {"k": int(k)}
|
|
872
|
+
)
|