evalsuite-python 0.1.0__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,51 @@
1
+ """Classification metrics."""
2
+
3
+ from .metrics import (
4
+ accuracy,
5
+ average_precision,
6
+ balanced_accuracy,
7
+ brier_score,
8
+ calibration_curve,
9
+ cohen_kappa,
10
+ confusion_matrix,
11
+ expected_calibration_error,
12
+ f1,
13
+ fbeta,
14
+ hamming_loss,
15
+ jaccard,
16
+ log_loss,
17
+ mcc,
18
+ npv,
19
+ pr_curve,
20
+ precision,
21
+ recall,
22
+ roc_auc,
23
+ roc_curve,
24
+ specificity,
25
+ top_k_accuracy,
26
+ )
27
+
28
+ __all__ = [
29
+ "accuracy",
30
+ "average_precision",
31
+ "balanced_accuracy",
32
+ "brier_score",
33
+ "calibration_curve",
34
+ "cohen_kappa",
35
+ "confusion_matrix",
36
+ "expected_calibration_error",
37
+ "f1",
38
+ "fbeta",
39
+ "hamming_loss",
40
+ "jaccard",
41
+ "log_loss",
42
+ "mcc",
43
+ "npv",
44
+ "pr_curve",
45
+ "precision",
46
+ "recall",
47
+ "roc_auc",
48
+ "roc_curve",
49
+ "specificity",
50
+ "top_k_accuracy",
51
+ ]
@@ -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)