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.
@@ -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
+ )