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.
- evalsuite/__init__.py +159 -0
- evalsuite/__main__.py +5 -0
- evalsuite/api.py +264 -0
- evalsuite/benchmarks.py +272 -0
- evalsuite/classification/__init__.py +51 -0
- evalsuite/classification/_common.py +111 -0
- evalsuite/classification/metrics.py +943 -0
- evalsuite/cli/__init__.py +5 -0
- evalsuite/cli/main.py +389 -0
- evalsuite/core/__init__.py +1 -0
- evalsuite/core/context.py +181 -0
- evalsuite/core/exceptions.py +52 -0
- evalsuite/core/export.py +100 -0
- evalsuite/core/registry.py +90 -0
- evalsuite/core/result.py +379 -0
- evalsuite/core/types.py +23 -0
- evalsuite/core/validation.py +202 -0
- evalsuite/plot.py +296 -0
- evalsuite/py.typed +0 -0
- evalsuite/regression/__init__.py +41 -0
- evalsuite/regression/metrics.py +604 -0
- evalsuite/reporting.py +241 -0
- evalsuite/stats/__init__.py +25 -0
- evalsuite/stats/_resolve.py +90 -0
- evalsuite/stats/compare.py +414 -0
- evalsuite/stats/effect.py +113 -0
- evalsuite/stats/intervals.py +320 -0
- evalsuite/stats/paired.py +207 -0
- evalsuite/stats/results.py +129 -0
- evalsuite/version.py +3 -0
- evalsuite_python-0.1.0.dist-info/METADATA +247 -0
- evalsuite_python-0.1.0.dist-info/RECORD +35 -0
- evalsuite_python-0.1.0.dist-info/WHEEL +4 -0
- evalsuite_python-0.1.0.dist-info/entry_points.txt +2 -0
- evalsuite_python-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,202 @@
|
|
|
1
|
+
"""Centralised input validation.
|
|
2
|
+
|
|
3
|
+
Inputs are converted to NumPy once, checked once, and never modified in place.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from __future__ import annotations
|
|
7
|
+
|
|
8
|
+
import warnings
|
|
9
|
+
from typing import Any, Literal, Optional
|
|
10
|
+
|
|
11
|
+
import numpy as np
|
|
12
|
+
from numpy.typing import NDArray
|
|
13
|
+
|
|
14
|
+
from .exceptions import InputValidationError, UndefinedMetricWarning
|
|
15
|
+
from .types import ArrayLike, FloatArray, ZeroDivision
|
|
16
|
+
|
|
17
|
+
__all__ = [
|
|
18
|
+
"TargetType",
|
|
19
|
+
"check_consistent_length",
|
|
20
|
+
"resolve_labels",
|
|
21
|
+
"safe_divide",
|
|
22
|
+
"target_type",
|
|
23
|
+
"to_numpy",
|
|
24
|
+
"validate_probabilities",
|
|
25
|
+
"validate_sample_weight",
|
|
26
|
+
"validate_zero_division",
|
|
27
|
+
]
|
|
28
|
+
|
|
29
|
+
TargetType = Literal["binary", "multiclass", "multilabel", "continuous"]
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def to_numpy(x: ArrayLike, name: str, *, allow_2d: bool = False) -> NDArray[Any]:
|
|
33
|
+
"""Convert array-likes (lists, NumPy arrays, pandas objects) to a NumPy array without copying
|
|
34
|
+
when possible. Raises :class:`InputValidationError` for None, scalars, empty input or bad shapes."""
|
|
35
|
+
if x is None:
|
|
36
|
+
raise InputValidationError(f"{name} is required but was None.")
|
|
37
|
+
values = getattr(x, "to_numpy", None)
|
|
38
|
+
arr = np.asarray(values() if callable(values) else x)
|
|
39
|
+
if arr.ndim == 0:
|
|
40
|
+
raise InputValidationError(f"{name} must be one-dimensional array-like; received a scalar.")
|
|
41
|
+
if arr.ndim == 2 and arr.shape[1] == 1 and not allow_2d:
|
|
42
|
+
arr = arr.ravel()
|
|
43
|
+
if arr.ndim > (2 if allow_2d else 1):
|
|
44
|
+
expected = "1-D or 2-D" if allow_2d else "1-D"
|
|
45
|
+
raise InputValidationError(f"{name} must be {expected}; received an array with shape {arr.shape}.")
|
|
46
|
+
if arr.shape[0] == 0:
|
|
47
|
+
raise InputValidationError(f"{name} is empty. Provide at least one observation.")
|
|
48
|
+
return arr
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def check_finite(arr: NDArray[Any], name: str) -> None:
|
|
52
|
+
if np.issubdtype(arr.dtype, np.number) and not np.all(np.isfinite(arr)):
|
|
53
|
+
n_nan = int(np.isnan(arr).sum())
|
|
54
|
+
n_inf = int(np.isinf(arr).sum())
|
|
55
|
+
raise InputValidationError(
|
|
56
|
+
f"{name} contains {n_nan} NaN and {n_inf} infinite value(s). Remove or impute them before evaluating."
|
|
57
|
+
)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def check_consistent_length(**arrays: Optional[NDArray[Any]]) -> int:
|
|
61
|
+
"""All given arrays must have the same number of observations; returns that number."""
|
|
62
|
+
lengths = {name: a.shape[0] for name, a in arrays.items() if a is not None}
|
|
63
|
+
if len(set(lengths.values())) > 1:
|
|
64
|
+
detail = " and ".join(f"{k}={v}" for k, v in lengths.items())
|
|
65
|
+
names = " and ".join(lengths)
|
|
66
|
+
raise InputValidationError(f"{names} must contain the same number of observations. Received {detail}.")
|
|
67
|
+
return int(next(iter(lengths.values())))
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def validate_sample_weight(sample_weight: Optional[ArrayLike], n: int) -> Optional[FloatArray]:
|
|
71
|
+
if sample_weight is None:
|
|
72
|
+
return None
|
|
73
|
+
w = to_numpy(sample_weight, "sample_weight").astype(np.float64, copy=False)
|
|
74
|
+
check_finite(w, "sample_weight")
|
|
75
|
+
if w.shape[0] != n:
|
|
76
|
+
raise InputValidationError(
|
|
77
|
+
f"sample_weight must contain one weight per observation. Received {w.shape[0]} for {n} observations."
|
|
78
|
+
)
|
|
79
|
+
if np.any(w < 0):
|
|
80
|
+
raise InputValidationError("sample_weight must be non-negative.")
|
|
81
|
+
if w.sum() == 0:
|
|
82
|
+
raise InputValidationError("sample_weight must not sum to zero.")
|
|
83
|
+
return w
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def target_type(y: NDArray[Any]) -> TargetType:
|
|
87
|
+
"""Infer the type of a label array."""
|
|
88
|
+
if y.ndim == 2:
|
|
89
|
+
if y.shape[1] > 1 and np.isin(y, (0, 1)).all():
|
|
90
|
+
return "multilabel"
|
|
91
|
+
if np.issubdtype(y.dtype, np.floating) and not np.all(np.mod(y, 1) == 0):
|
|
92
|
+
return "continuous" # multi-output regression
|
|
93
|
+
raise InputValidationError(
|
|
94
|
+
"2-D targets must be a binary indicator matrix (0/1, one column per label) for multilabel tasks; "
|
|
95
|
+
f"received shape {y.shape} with values other than 0 and 1."
|
|
96
|
+
)
|
|
97
|
+
if np.issubdtype(y.dtype, np.floating) and not np.all(np.mod(y, 1) == 0):
|
|
98
|
+
return "continuous"
|
|
99
|
+
try:
|
|
100
|
+
n_unique = np.unique(y).shape[0]
|
|
101
|
+
except TypeError as exc:
|
|
102
|
+
raise InputValidationError(
|
|
103
|
+
"Labels must be mutually comparable (for example all integers or all strings); found a mix of types."
|
|
104
|
+
) from exc
|
|
105
|
+
return "binary" if n_unique <= 2 else "multiclass"
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def resolve_labels(
|
|
109
|
+
y_true: NDArray[Any], y_pred: Optional[NDArray[Any]], labels: Optional[ArrayLike]
|
|
110
|
+
) -> NDArray[Any]:
|
|
111
|
+
"""Sorted labels present in y_true or y_pred, or the user's explicit list (order kept)."""
|
|
112
|
+
if labels is not None:
|
|
113
|
+
lab = to_numpy(labels, "labels")
|
|
114
|
+
if np.unique(lab).shape[0] != lab.shape[0]:
|
|
115
|
+
raise InputValidationError("labels contains duplicates.")
|
|
116
|
+
return lab
|
|
117
|
+
present = y_true if y_pred is None else np.concatenate([y_true, y_pred])
|
|
118
|
+
try:
|
|
119
|
+
return np.unique(present)
|
|
120
|
+
except TypeError as exc:
|
|
121
|
+
raise InputValidationError(
|
|
122
|
+
"Labels in y_true and y_pred must be mutually comparable (for example all integers or all strings)."
|
|
123
|
+
) from exc
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
def validate_probabilities(
|
|
127
|
+
y_prob: ArrayLike,
|
|
128
|
+
n: int,
|
|
129
|
+
*,
|
|
130
|
+
n_classes: Optional[int] = None,
|
|
131
|
+
name: str = "y_prob",
|
|
132
|
+
rows_sum_to_one: bool = True,
|
|
133
|
+
unit: str = "class",
|
|
134
|
+
) -> FloatArray:
|
|
135
|
+
"""Probabilities in [0, 1]. For multiclass (2-D), one column per class and rows summing to 1.
|
|
136
|
+
Multilabel probabilities are independent per label (``rows_sum_to_one=False``)."""
|
|
137
|
+
p = to_numpy(y_prob, name, allow_2d=True).astype(np.float64, copy=False)
|
|
138
|
+
check_finite(p, name)
|
|
139
|
+
if p.shape[0] != n:
|
|
140
|
+
raise InputValidationError(
|
|
141
|
+
f"{name} must contain one row per observation. Received {p.shape[0]} for {n} observations."
|
|
142
|
+
)
|
|
143
|
+
if np.any(p < 0) or np.any(p > 1):
|
|
144
|
+
raise InputValidationError(
|
|
145
|
+
f"{name} must contain probabilities in [0, 1]; found values from {float(np.min(p)):.4g} to "
|
|
146
|
+
f"{float(np.max(p)):.4g}. "
|
|
147
|
+
"If these are scores or logits, convert them to probabilities first."
|
|
148
|
+
)
|
|
149
|
+
if p.ndim == 2:
|
|
150
|
+
if n_classes is not None and p.shape[1] != n_classes:
|
|
151
|
+
raise InputValidationError(
|
|
152
|
+
f"{name} has {p.shape[1]} columns but there are {n_classes} {unit}es. "
|
|
153
|
+
f"Provide one probability column per {unit}, in label order."
|
|
154
|
+
if unit == "class"
|
|
155
|
+
else f"{name} has {p.shape[1]} columns but there are {n_classes} {unit}s. "
|
|
156
|
+
f"Provide one probability column per {unit}."
|
|
157
|
+
)
|
|
158
|
+
sums = p.sum(axis=1)
|
|
159
|
+
if rows_sum_to_one and not np.allclose(sums, 1.0, atol=1e-6):
|
|
160
|
+
worst = float(np.abs(sums - 1).max())
|
|
161
|
+
raise InputValidationError(
|
|
162
|
+
f"Each row of {name} must sum to 1 for multiclass probabilities (largest deviation {worst:.3g})."
|
|
163
|
+
)
|
|
164
|
+
return p
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def validate_zero_division(zero_division: ZeroDivision) -> ZeroDivision:
|
|
168
|
+
if zero_division == "warn":
|
|
169
|
+
return zero_division
|
|
170
|
+
if isinstance(zero_division, (int, float)) and (np.isnan(zero_division) or 0 <= zero_division <= 1):
|
|
171
|
+
return float(zero_division)
|
|
172
|
+
raise InputValidationError('zero_division must be "warn", 0, 1 or np.nan.')
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
def safe_divide(
|
|
176
|
+
num: NDArray[np.float64] | float,
|
|
177
|
+
den: NDArray[np.float64] | float,
|
|
178
|
+
*,
|
|
179
|
+
zero_division: ZeroDivision = "warn",
|
|
180
|
+
metric: str = "metric",
|
|
181
|
+
) -> NDArray[np.float64]:
|
|
182
|
+
"""Element-wise num/den. Where den == 0 the result is ``zero_division``.
|
|
183
|
+
|
|
184
|
+
``"warn"`` (default) returns 0 and emits :class:`UndefinedMetricWarning`, matching scikit-learn's
|
|
185
|
+
convention; pass ``zero_division=np.nan`` to propagate undefined values instead. Never silent.
|
|
186
|
+
"""
|
|
187
|
+
num_a = np.asarray(num, dtype=np.float64)
|
|
188
|
+
den_a = np.asarray(den, dtype=np.float64)
|
|
189
|
+
zero = den_a == 0
|
|
190
|
+
out = np.divide(num_a, den_a, out=np.zeros(np.broadcast(num_a, den_a).shape), where=~zero)
|
|
191
|
+
if np.any(zero):
|
|
192
|
+
if zero_division == "warn":
|
|
193
|
+
warnings.warn(
|
|
194
|
+
f"{metric} is undefined for {int(np.sum(zero))} case(s) because the denominator is zero; "
|
|
195
|
+
"using 0. Set zero_division=0, 1 or np.nan to choose the value explicitly.",
|
|
196
|
+
UndefinedMetricWarning,
|
|
197
|
+
stacklevel=4,
|
|
198
|
+
)
|
|
199
|
+
out[zero] = 0.0
|
|
200
|
+
else:
|
|
201
|
+
out[zero] = float(zero_division)
|
|
202
|
+
return out
|
evalsuite/plot.py
ADDED
|
@@ -0,0 +1,296 @@
|
|
|
1
|
+
"""Publication-ready plots (optional dependency: ``pip install "evalsuite-python[plot]"``).
|
|
2
|
+
|
|
3
|
+
Every function draws on ``ax`` (or a new figure), returns the matplotlib ``Axes`` and computes its numbers
|
|
4
|
+
with EvalSuite's own metrics, so the plot and the reported values always agree. Several models can be passed
|
|
5
|
+
as a dict ``{name: values}``; they get distinct colours *and* line styles, so plots stay readable in
|
|
6
|
+
greyscale. Matplotlib is imported only when a plot function is called.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
from collections.abc import Mapping
|
|
12
|
+
from typing import TYPE_CHECKING, Any, Optional, Union, cast
|
|
13
|
+
|
|
14
|
+
import numpy as np
|
|
15
|
+
|
|
16
|
+
from .core.exceptions import InputValidationError, OptionalDependencyError
|
|
17
|
+
from .core.types import ArrayLike
|
|
18
|
+
|
|
19
|
+
if TYPE_CHECKING:
|
|
20
|
+
from matplotlib.axes import Axes
|
|
21
|
+
|
|
22
|
+
from .stats.compare import ComparisonResult
|
|
23
|
+
|
|
24
|
+
__all__ = ["calibration", "comparison", "confusion_matrix", "pr", "residuals", "roc"]
|
|
25
|
+
|
|
26
|
+
Scores = Union[ArrayLike, Mapping[str, ArrayLike]]
|
|
27
|
+
_STYLES = ("-", "--", "-.", ":")
|
|
28
|
+
_MARKERS = ("o", "s", "^", "D", "v", "P")
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def _plt() -> Any:
|
|
32
|
+
try:
|
|
33
|
+
import matplotlib.pyplot as plt
|
|
34
|
+
except ImportError as exc:
|
|
35
|
+
raise OptionalDependencyError("matplotlib", "plot", "Plotting") from exc
|
|
36
|
+
return plt
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def _axes(ax: Optional[Axes], figsize: tuple[float, float] = (5.0, 4.2)) -> Axes:
|
|
40
|
+
if ax is not None:
|
|
41
|
+
return ax
|
|
42
|
+
_, new_ax = _plt().subplots(figsize=figsize, layout="constrained")
|
|
43
|
+
return cast("Axes", new_ax)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def _models(values: Scores, label: Optional[str]) -> list[tuple[Optional[str], Any]]:
|
|
47
|
+
if isinstance(values, Mapping):
|
|
48
|
+
if not values:
|
|
49
|
+
raise InputValidationError("Pass at least one model.")
|
|
50
|
+
return [(str(k), v) for k, v in values.items()]
|
|
51
|
+
return [(label, values)]
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _style(i: int) -> dict[str, Any]:
|
|
55
|
+
return {"linestyle": _STYLES[i % len(_STYLES)], "linewidth": 1.8}
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def roc(
|
|
59
|
+
y_true: ArrayLike,
|
|
60
|
+
y_prob: Scores,
|
|
61
|
+
*,
|
|
62
|
+
ax: Optional[Axes] = None,
|
|
63
|
+
label: Optional[str] = None,
|
|
64
|
+
pos_label: Any = None,
|
|
65
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
66
|
+
chance: bool = True,
|
|
67
|
+
) -> Axes:
|
|
68
|
+
"""ROC curve(s) for binary probabilities, with the AUC in the legend.
|
|
69
|
+
|
|
70
|
+
``y_prob``: P(positive) for one model, or ``{name: probabilities}`` for several.
|
|
71
|
+
"""
|
|
72
|
+
from .classification.metrics import roc_auc, roc_curve
|
|
73
|
+
|
|
74
|
+
ax = _axes(ax)
|
|
75
|
+
for i, (name, prob) in enumerate(_models(y_prob, label)):
|
|
76
|
+
fpr, tpr, _ = roc_curve(y_true, prob, pos_label=pos_label, sample_weight=sample_weight)
|
|
77
|
+
auc = float(roc_auc(y_true, prob, pos_label=pos_label, sample_weight=sample_weight))
|
|
78
|
+
ax.plot(
|
|
79
|
+
fpr, tpr, drawstyle="steps-post", label=f"{name + ': ' if name else ''}AUC = {auc:.3f}", **_style(i)
|
|
80
|
+
)
|
|
81
|
+
if chance:
|
|
82
|
+
ax.plot([0, 1], [0, 1], color="0.6", linewidth=1, linestyle=":", label="chance")
|
|
83
|
+
ax.set(
|
|
84
|
+
xlabel="False positive rate (1 − specificity)",
|
|
85
|
+
ylabel="True positive rate (sensitivity)",
|
|
86
|
+
title="ROC curve",
|
|
87
|
+
xlim=(-0.01, 1.01),
|
|
88
|
+
ylim=(-0.01, 1.01),
|
|
89
|
+
)
|
|
90
|
+
ax.set_aspect("equal")
|
|
91
|
+
ax.legend(loc="lower right", frameon=False)
|
|
92
|
+
return ax
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def pr(
|
|
96
|
+
y_true: ArrayLike,
|
|
97
|
+
y_prob: Scores,
|
|
98
|
+
*,
|
|
99
|
+
ax: Optional[Axes] = None,
|
|
100
|
+
label: Optional[str] = None,
|
|
101
|
+
pos_label: Any = None,
|
|
102
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
103
|
+
chance: bool = True,
|
|
104
|
+
) -> Axes:
|
|
105
|
+
"""Precision-recall curve(s) with average precision in the legend; the chance line is the prevalence."""
|
|
106
|
+
from .classification.metrics import _binary_target, _ctx, average_precision, pr_curve
|
|
107
|
+
|
|
108
|
+
ax = _axes(ax)
|
|
109
|
+
for i, (name, prob) in enumerate(_models(y_prob, label)):
|
|
110
|
+
prec, rec, _ = pr_curve(y_true, prob, pos_label=pos_label, sample_weight=sample_weight)
|
|
111
|
+
ap = float(average_precision(y_true, prob, pos_label=pos_label, sample_weight=sample_weight))
|
|
112
|
+
ax.plot(
|
|
113
|
+
rec, prec, drawstyle="steps-post", label=f"{name + ': ' if name else ''}AP = {ap:.3f}", **_style(i)
|
|
114
|
+
)
|
|
115
|
+
if chance:
|
|
116
|
+
first = _models(y_prob, label)[0][1]
|
|
117
|
+
ctx = _ctx(y_true, None, y_prob=first, sample_weight=sample_weight)
|
|
118
|
+
prevalence = float(np.average(_binary_target(ctx, pos_label), weights=ctx.weights))
|
|
119
|
+
ax.axhline(prevalence, color="0.6", linewidth=1, linestyle=":", label=f"chance ({prevalence:.2f})")
|
|
120
|
+
ax.set(
|
|
121
|
+
xlabel="Recall (sensitivity)",
|
|
122
|
+
ylabel="Precision (PPV)",
|
|
123
|
+
title="Precision-recall curve",
|
|
124
|
+
xlim=(-0.01, 1.01),
|
|
125
|
+
ylim=(-0.01, 1.03),
|
|
126
|
+
)
|
|
127
|
+
ax.legend(loc="lower left", frameon=False)
|
|
128
|
+
return ax
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
def confusion_matrix(
|
|
132
|
+
y_true: ArrayLike,
|
|
133
|
+
y_pred: ArrayLike,
|
|
134
|
+
*,
|
|
135
|
+
ax: Optional[Axes] = None,
|
|
136
|
+
labels: Optional[ArrayLike] = None,
|
|
137
|
+
normalize: Optional[str] = None,
|
|
138
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
139
|
+
cmap: str = "Blues",
|
|
140
|
+
colorbar: bool = True,
|
|
141
|
+
) -> Axes:
|
|
142
|
+
"""Annotated confusion matrix (rows: true, columns: predicted). ``normalize``: None, "true", "pred", "all"."""
|
|
143
|
+
from .classification.metrics import confusion_matrix as cm_fn
|
|
144
|
+
from .core.validation import resolve_labels, to_numpy
|
|
145
|
+
|
|
146
|
+
cm = np.asarray(cm_fn(y_true, y_pred, labels=labels, sample_weight=sample_weight, normalize=normalize)) # type: ignore[arg-type]
|
|
147
|
+
names = (
|
|
148
|
+
to_numpy(labels, "labels")
|
|
149
|
+
if labels is not None
|
|
150
|
+
else resolve_labels(to_numpy(y_true, "y_true"), to_numpy(y_pred, "y_pred"), None)
|
|
151
|
+
).tolist()
|
|
152
|
+
k = cm.shape[0]
|
|
153
|
+
ax = _axes(ax, (max(3.6, 0.7 * k + 2.4), max(3.2, 0.7 * k + 1.8)))
|
|
154
|
+
image = ax.imshow(cm, cmap=cmap, vmin=0, vmax=1 if normalize else None)
|
|
155
|
+
threshold = (cm.max() + cm.min()) / 2
|
|
156
|
+
integral = normalize is None and np.all(np.mod(cm, 1) == 0)
|
|
157
|
+
for i in range(k):
|
|
158
|
+
for j in range(k):
|
|
159
|
+
text = f"{int(cm[i, j])}" if integral else f"{cm[i, j]:.2f}"
|
|
160
|
+
ax.text(
|
|
161
|
+
j,
|
|
162
|
+
i,
|
|
163
|
+
text,
|
|
164
|
+
ha="center",
|
|
165
|
+
va="center",
|
|
166
|
+
fontsize=9,
|
|
167
|
+
color="white" if cm[i, j] > threshold else "black",
|
|
168
|
+
)
|
|
169
|
+
ax.set(
|
|
170
|
+
xticks=range(k),
|
|
171
|
+
yticks=range(k),
|
|
172
|
+
xticklabels=names,
|
|
173
|
+
yticklabels=names,
|
|
174
|
+
xlabel="Predicted label",
|
|
175
|
+
ylabel="True label",
|
|
176
|
+
title="Confusion matrix" + (f" (normalised by {normalize})" if normalize else ""),
|
|
177
|
+
)
|
|
178
|
+
if colorbar:
|
|
179
|
+
ax.figure.colorbar(image, ax=ax, fraction=0.046, pad=0.04)
|
|
180
|
+
return ax
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
def calibration(
|
|
184
|
+
y_true: ArrayLike,
|
|
185
|
+
y_prob: Scores,
|
|
186
|
+
*,
|
|
187
|
+
ax: Optional[Axes] = None,
|
|
188
|
+
label: Optional[str] = None,
|
|
189
|
+
n_bins: int = 10,
|
|
190
|
+
strategy: str = "uniform",
|
|
191
|
+
pos_label: Any = None,
|
|
192
|
+
sample_weight: Optional[ArrayLike] = None,
|
|
193
|
+
legend_loc: str = "below",
|
|
194
|
+
) -> Axes:
|
|
195
|
+
"""Reliability diagram: observed frequency against mean predicted probability per bin, with ECE and Brier
|
|
196
|
+
score in the legend; the diagonal is perfect calibration.
|
|
197
|
+
|
|
198
|
+
``legend_loc="below"`` (default) puts the legend under the axes so it never covers the curves; any
|
|
199
|
+
matplotlib location (e.g. ``"upper left"``) places it inside instead."""
|
|
200
|
+
from .classification.metrics import brier_score, calibration_curve, expected_calibration_error
|
|
201
|
+
|
|
202
|
+
ax = _axes(ax)
|
|
203
|
+
ax.plot([0, 1], [0, 1], color="0.6", linewidth=1, linestyle=":", label="perfect calibration")
|
|
204
|
+
kw = {"n_bins": n_bins, "strategy": strategy, "pos_label": pos_label, "sample_weight": sample_weight}
|
|
205
|
+
for i, (name, prob) in enumerate(_models(y_prob, label)):
|
|
206
|
+
frac, mean_p, _ = calibration_curve(y_true, prob, **kw)
|
|
207
|
+
ece = float(expected_calibration_error(y_true, prob, **kw))
|
|
208
|
+
brier = float(brier_score(y_true, prob, pos_label=pos_label, sample_weight=sample_weight))
|
|
209
|
+
ax.plot(
|
|
210
|
+
mean_p,
|
|
211
|
+
frac,
|
|
212
|
+
marker=_MARKERS[i % len(_MARKERS)],
|
|
213
|
+
markersize=4,
|
|
214
|
+
label=f"{name + ': ' if name else ''}ECE = {ece:.3f}, Brier = {brier:.3f}",
|
|
215
|
+
**_style(i),
|
|
216
|
+
)
|
|
217
|
+
ax.set(
|
|
218
|
+
xlabel="Mean predicted probability",
|
|
219
|
+
ylabel="Observed frequency of positives",
|
|
220
|
+
title="Calibration",
|
|
221
|
+
xlim=(-0.01, 1.01),
|
|
222
|
+
ylim=(-0.01, 1.01),
|
|
223
|
+
)
|
|
224
|
+
ax.set_aspect("equal")
|
|
225
|
+
if legend_loc == "below":
|
|
226
|
+
ax.legend(loc="upper center", bbox_to_anchor=(0.5, -0.14), frameon=False, fontsize="small")
|
|
227
|
+
else:
|
|
228
|
+
ax.legend(loc=legend_loc, frameon=False) # type: ignore[call-overload]
|
|
229
|
+
return ax
|
|
230
|
+
|
|
231
|
+
|
|
232
|
+
def residuals(
|
|
233
|
+
y_true: ArrayLike,
|
|
234
|
+
y_pred: ArrayLike,
|
|
235
|
+
*,
|
|
236
|
+
ax: Optional[Axes] = None,
|
|
237
|
+
kind: str = "residuals",
|
|
238
|
+
) -> Axes:
|
|
239
|
+
"""Regression diagnostics. ``kind="residuals"``: residual (predicted − true) against predicted value;
|
|
240
|
+
``kind="predicted"``: predicted against true with the identity line. RMSE and R² in the title."""
|
|
241
|
+
from .regression.metrics import _Inputs, r2, rmse
|
|
242
|
+
|
|
243
|
+
inp = _Inputs(y_true, y_pred, None)
|
|
244
|
+
if inp.multi:
|
|
245
|
+
raise InputValidationError("residuals() plots single-output targets; pass one output column at a time.")
|
|
246
|
+
yt, yp = inp.y_true[:, 0], inp.y_pred[:, 0]
|
|
247
|
+
ax = _axes(ax)
|
|
248
|
+
stats = f"RMSE = {float(rmse(yt, yp)):.3g}, R² = {float(r2(yt, yp)):.3f}"
|
|
249
|
+
if kind == "residuals":
|
|
250
|
+
ax.scatter(yp, yp - yt, s=14, alpha=0.7, edgecolors="none")
|
|
251
|
+
ax.axhline(0, color="0.4", linewidth=1)
|
|
252
|
+
ax.set(xlabel="Predicted value", ylabel="Residual (predicted − true)", title=f"Residuals ({stats})")
|
|
253
|
+
elif kind == "predicted":
|
|
254
|
+
lo, hi = float(min(np.min(yt), np.min(yp))), float(max(np.max(yt), np.max(yp)))
|
|
255
|
+
ax.scatter(yt, yp, s=14, alpha=0.7, edgecolors="none")
|
|
256
|
+
ax.plot([lo, hi], [lo, hi], color="0.4", linewidth=1, linestyle=":", label="y = x")
|
|
257
|
+
ax.set(xlabel="True value", ylabel="Predicted value", title=f"Predicted vs true ({stats})")
|
|
258
|
+
ax.legend(loc="upper left", frameon=False)
|
|
259
|
+
else:
|
|
260
|
+
raise InputValidationError("kind must be 'residuals' or 'predicted'.")
|
|
261
|
+
return ax
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
def comparison(
|
|
265
|
+
result: ComparisonResult,
|
|
266
|
+
*,
|
|
267
|
+
metrics: Optional[list[str]] = None,
|
|
268
|
+
ax: Optional[Axes] = None,
|
|
269
|
+
) -> Axes:
|
|
270
|
+
"""Forest plot of an :func:`evalsuite.compare` result: each model's estimate with its confidence interval,
|
|
271
|
+
grouped by metric; the best model per metric is drawn filled."""
|
|
272
|
+
names = list(metrics or result.metrics)
|
|
273
|
+
unknown = [m for m in names if m not in result.metrics]
|
|
274
|
+
if unknown:
|
|
275
|
+
raise InputValidationError(f"Metric(s) {unknown} were not compared. Compared: {list(result.metrics)}.")
|
|
276
|
+
rows = [(m, model) for m in names for model in result.models]
|
|
277
|
+
ax = _axes(ax, (6.0, max(2.5, 0.38 * len(rows) + 1.2)))
|
|
278
|
+
for y_pos, (metric, model) in enumerate(rows):
|
|
279
|
+
est = result.estimate(model, metric)
|
|
280
|
+
best = result.best(metric) == model
|
|
281
|
+
ax.errorbar(
|
|
282
|
+
est["estimate"],
|
|
283
|
+
y_pos,
|
|
284
|
+
xerr=[[est["estimate"] - est["low"]], [est["high"] - est["estimate"]]],
|
|
285
|
+
fmt="o",
|
|
286
|
+
capsize=3,
|
|
287
|
+
color="C0",
|
|
288
|
+
markerfacecolor="C0" if best else "white",
|
|
289
|
+
markersize=6,
|
|
290
|
+
)
|
|
291
|
+
ax.set_yticks(range(len(rows)), [f"{metric} · {model}" for metric, model in rows])
|
|
292
|
+
ax.invert_yaxis()
|
|
293
|
+
level = round(result.settings["level"] * 100, 6)
|
|
294
|
+
ax.set(xlabel=f"Estimate with {level:g}% CI (filled: best per metric)", title="Model comparison")
|
|
295
|
+
ax.grid(axis="x", color="0.9")
|
|
296
|
+
return ax
|
evalsuite/py.typed
ADDED
|
File without changes
|
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
"""Regression metrics."""
|
|
2
|
+
|
|
3
|
+
from .metrics import (
|
|
4
|
+
adjusted_r2,
|
|
5
|
+
explained_variance,
|
|
6
|
+
huber_loss,
|
|
7
|
+
mae,
|
|
8
|
+
mape,
|
|
9
|
+
max_error,
|
|
10
|
+
mean_bias_error,
|
|
11
|
+
median_absolute_error,
|
|
12
|
+
mse,
|
|
13
|
+
msle,
|
|
14
|
+
quantile_loss,
|
|
15
|
+
r2,
|
|
16
|
+
rae,
|
|
17
|
+
rmse,
|
|
18
|
+
rmsle,
|
|
19
|
+
rse,
|
|
20
|
+
smape,
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
__all__ = [
|
|
24
|
+
"adjusted_r2",
|
|
25
|
+
"explained_variance",
|
|
26
|
+
"huber_loss",
|
|
27
|
+
"mae",
|
|
28
|
+
"mape",
|
|
29
|
+
"max_error",
|
|
30
|
+
"mean_bias_error",
|
|
31
|
+
"median_absolute_error",
|
|
32
|
+
"mse",
|
|
33
|
+
"msle",
|
|
34
|
+
"quantile_loss",
|
|
35
|
+
"r2",
|
|
36
|
+
"rae",
|
|
37
|
+
"rmse",
|
|
38
|
+
"rmsle",
|
|
39
|
+
"rse",
|
|
40
|
+
"smape",
|
|
41
|
+
]
|