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,5 @@
1
+ """Command-line interface: ``evalsuite --help``."""
2
+
3
+ from .main import main
4
+
5
+ __all__ = ["main"]
evalsuite/cli/main.py ADDED
@@ -0,0 +1,389 @@
1
+ """``evalsuite``: evaluate, compare, report, plot and benchmark from the command line.
2
+
3
+ Examples::
4
+
5
+ evalsuite evaluate predictions.csv --y-true label --y-pred pred --y-prob prob
6
+ evalsuite report predictions.csv --y-true label --y-pred pred --format markdown
7
+ evalsuite compare predictions.csv --y-true label --pred lr=pred_lr --pred rf=pred_rf \\
8
+ --prob lr=p_lr --prob rf=p_rf
9
+ evalsuite plot roc predictions.csv --y-true label --y-prob prob --output roc.png
10
+ evalsuite metrics --category regression
11
+ evalsuite info classification.mcc
12
+ evalsuite benchmark --quick
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import argparse
18
+ import os
19
+ import sys
20
+ from collections.abc import Sequence
21
+ from typing import Any, Optional
22
+
23
+ from ..core.exceptions import EvalSuiteError
24
+ from ..version import __version__
25
+
26
+ FORMATS = ("text", "json", "csv", "markdown", "latex", "html")
27
+ _EXT = {
28
+ ".json": "json",
29
+ ".csv": "csv",
30
+ ".md": "markdown",
31
+ ".tex": "latex",
32
+ ".html": "html",
33
+ ".htm": "html",
34
+ ".txt": "text",
35
+ }
36
+
37
+
38
+ class CLIError(Exception):
39
+ """A user-facing error: printed without a traceback, exit code 2."""
40
+
41
+
42
+ # ---- input ---------------------------------------------------------------------------------------
43
+ def read_table(path: str) -> Any:
44
+ import pandas as pd
45
+
46
+ if path == "-":
47
+ return pd.read_csv(sys.stdin)
48
+ if not os.path.exists(path):
49
+ raise CLIError(f"File not found: {path}")
50
+ ext = os.path.splitext(path)[1].lower()
51
+ try:
52
+ if ext in (".tsv", ".tab"):
53
+ return pd.read_csv(path, sep="\t")
54
+ if ext == ".parquet":
55
+ return pd.read_parquet(path)
56
+ if ext in (".json", ".jsonl"):
57
+ return pd.read_json(path, lines=ext == ".jsonl")
58
+ return pd.read_csv(path)
59
+ except ImportError as exc:
60
+ raise CLIError(f"Reading {ext} files needs an extra package: {exc}") from exc
61
+ except Exception as exc:
62
+ raise CLIError(f"Could not read {path}: {exc}") from exc
63
+
64
+
65
+ def column(df: Any, name: str, role: str) -> Any:
66
+ if name not in df.columns:
67
+ raise CLIError(
68
+ f"Column '{name}' ({role}) not found. Columns in the file: {', '.join(map(str, df.columns))}."
69
+ )
70
+ return df[name].to_numpy()
71
+
72
+
73
+ def columns(df: Any, names: Sequence[str], role: str) -> Any:
74
+ import numpy as np
75
+
76
+ cols = [column(df, n, role) for n in names]
77
+ return cols[0] if len(cols) == 1 else np.column_stack(cols)
78
+
79
+
80
+ def pairs(items: Optional[Sequence[str]], flag: str) -> dict[str, list[str]]:
81
+ """``name=col`` or ``name=col1,col2,col3`` (multiclass probabilities) -> {name: [cols]}."""
82
+ out: dict[str, list[str]] = {}
83
+ for item in items or []:
84
+ name, sep, cols = item.partition("=")
85
+ if not sep or not name or not cols:
86
+ raise CLIError(f"{flag} expects NAME=COLUMN (or NAME=COL1,COL2,... for multiclass); got '{item}'.")
87
+ out[name] = [c for c in cols.split(",") if c]
88
+ return out
89
+
90
+
91
+ def metric_list(text: Optional[str]) -> Optional[list[str]]:
92
+ return [m.strip() for m in text.split(",") if m.strip()] if text else None
93
+
94
+
95
+ def label_value(text: Optional[str]) -> Any:
96
+ if text is None:
97
+ return None
98
+ try:
99
+ return int(text)
100
+ except ValueError:
101
+ return text
102
+
103
+
104
+ # ---- output --------------------------------------------------------------------------------------
105
+ def render(obj: Any, fmt: str, digits: int) -> str:
106
+ if fmt == "text":
107
+ return obj.summary(digits=digits) if hasattr(obj, "summary") else str(obj)
108
+ if fmt == "json":
109
+ return str(obj.to_json())
110
+ if fmt == "csv":
111
+ return str(obj.to_csv())
112
+ if fmt == "markdown":
113
+ return str(obj.to_markdown(digits=digits))
114
+ if fmt == "latex":
115
+ return str(obj.to_latex(digits=digits))
116
+ return str(obj.to_html(digits=digits, full=True))
117
+
118
+
119
+ def emit(obj: Any, args: argparse.Namespace) -> None:
120
+ fmt = args.format
121
+ if fmt is None:
122
+ fmt = _EXT.get(os.path.splitext(args.output)[1].lower(), "text") if args.output else "text"
123
+ text = render(obj, fmt, args.digits)
124
+ if not text.endswith("\n"):
125
+ text += "\n"
126
+ if args.output:
127
+ with open(args.output, "w", encoding="utf-8", newline="") as fh:
128
+ fh.write(text)
129
+ print(f"Wrote {fmt} to {args.output}", file=sys.stderr)
130
+ else:
131
+ sys.stdout.write(text)
132
+
133
+
134
+ def add_output(p: argparse.ArgumentParser, digits: int = 4) -> None:
135
+ p.add_argument(
136
+ "--format",
137
+ "-f",
138
+ choices=FORMATS,
139
+ default=None,
140
+ help="output format (default: from --output extension, else text)",
141
+ )
142
+ p.add_argument("--output", "-o", help="write to this file instead of standard output")
143
+ p.add_argument("--digits", type=int, default=digits, help=f"decimal places (default {digits})")
144
+
145
+
146
+ def add_targets(p: argparse.ArgumentParser, pred_required: bool = False) -> None:
147
+ p.add_argument(
148
+ "file", help="CSV/TSV/Parquet/JSON file with one row per observation ('-' reads CSV from stdin)"
149
+ )
150
+ p.add_argument("--y-true", required=True, metavar="COL", help="column with the true labels or values")
151
+ p.add_argument(
152
+ "--y-pred", required=pred_required, metavar="COL", help="column with predicted labels or values"
153
+ )
154
+ p.add_argument(
155
+ "--y-prob",
156
+ nargs="+",
157
+ metavar="COL",
158
+ help="probability column(s): one for P(positive), or one per class in label order",
159
+ )
160
+ p.add_argument("--weight", metavar="COL", help="column with sample weights")
161
+
162
+
163
+ # ---- commands ------------------------------------------------------------------------------------
164
+ def cmd_evaluate(args: argparse.Namespace) -> int:
165
+ import evalsuite as es
166
+
167
+ df = read_table(args.file)
168
+ if not args.y_pred and not args.y_prob:
169
+ raise CLIError("Give --y-pred, --y-prob, or both.")
170
+ result = es.evaluate(
171
+ column(df, args.y_true, "--y-true"),
172
+ column(df, args.y_pred, "--y-pred") if args.y_pred else None,
173
+ y_prob=columns(df, args.y_prob, "--y-prob") if args.y_prob else None,
174
+ task=args.task,
175
+ metrics=metric_list(args.metrics),
176
+ average=args.average,
177
+ pos_label=label_value(args.pos_label),
178
+ sample_weight=column(df, args.weight, "--weight") if args.weight else None,
179
+ )
180
+ emit(result, args)
181
+ return 0
182
+
183
+
184
+ def cmd_report(args: argparse.Namespace) -> int:
185
+ import evalsuite as es
186
+
187
+ df = read_table(args.file)
188
+ report = es.classification_report(
189
+ column(df, args.y_true, "--y-true"),
190
+ column(df, args.y_pred, "--y-pred"),
191
+ sample_weight=column(df, args.weight, "--weight") if args.weight else None,
192
+ )
193
+ emit(report, args)
194
+ return 0
195
+
196
+
197
+ def cmd_compare(args: argparse.Namespace) -> int:
198
+ import evalsuite as es
199
+
200
+ df = read_table(args.file)
201
+ preds = {k: columns(df, v, "--pred") for k, v in pairs(args.pred, "--pred").items()}
202
+ probs = {k: columns(df, v, "--prob") for k, v in pairs(args.prob, "--prob").items()}
203
+ if len(set(preds) | set(probs)) < 2:
204
+ raise CLIError(
205
+ "compare needs at least two models: --pred NAME=COL (and/or --prob NAME=COL), twice or more."
206
+ )
207
+ result = es.compare(
208
+ column(df, args.y_true, "--y-true"),
209
+ preds or None,
210
+ probabilities=probs or None,
211
+ metrics=metric_list(args.metrics),
212
+ baseline=args.baseline,
213
+ level=args.level,
214
+ alpha=args.alpha,
215
+ n_resamples=args.resamples,
216
+ correction=args.correction,
217
+ random_state=args.seed,
218
+ )
219
+ emit(result, args)
220
+ if args.plot:
221
+ _save_figure(lambda ax: es.plot.comparison(result, ax=ax), args.plot)
222
+ return 0
223
+
224
+
225
+ def _save_figure(draw: Any, path: str) -> None:
226
+ from ..plot import _plt
227
+
228
+ plt = _plt()
229
+ fig, ax = plt.subplots(figsize=(5.5, 4.4), layout="constrained")
230
+ draw(ax)
231
+ fig.savefig(path, dpi=200)
232
+ plt.close(fig)
233
+ print(f"Wrote figure to {path}", file=sys.stderr)
234
+
235
+
236
+ def cmd_plot(args: argparse.Namespace) -> int:
237
+ import evalsuite as es
238
+
239
+ df = read_table(args.file)
240
+ y = column(df, args.y_true, "--y-true")
241
+ kind = args.kind
242
+ if kind in ("roc", "pr", "calibration"):
243
+ if not args.y_prob:
244
+ raise CLIError(f"'{kind}' needs --y-prob.")
245
+ prob = columns(df, args.y_prob, "--y-prob")
246
+ fn: Any = {"roc": es.plot.roc, "pr": es.plot.pr, "calibration": es.plot.calibration}[kind]
247
+ _save_figure(lambda ax: fn(y, prob, ax=ax), args.output)
248
+ else:
249
+ if not args.y_pred:
250
+ raise CLIError(f"'{kind}' needs --y-pred.")
251
+ pred = column(df, args.y_pred, "--y-pred")
252
+ if kind == "confusion":
253
+ _save_figure(
254
+ lambda ax: es.plot.confusion_matrix(y, pred, ax=ax, normalize=args.normalize), args.output
255
+ )
256
+ else:
257
+ _save_figure(
258
+ lambda ax: es.plot.residuals(
259
+ y, pred, ax=ax, kind="predicted" if kind == "predicted" else "residuals"
260
+ ),
261
+ args.output,
262
+ )
263
+ return 0
264
+
265
+
266
+ def cmd_metrics(args: argparse.Namespace) -> int:
267
+ import evalsuite as es
268
+
269
+ for mid in es.list_metrics(args.category):
270
+ info = es.metric_info(mid)
271
+ print(f"{mid:<45} {info.name}")
272
+ return 0
273
+
274
+
275
+ def cmd_info(args: argparse.Namespace) -> int:
276
+ import evalsuite as es
277
+
278
+ try:
279
+ info = es.metric_info(args.metric)
280
+ except KeyError as exc:
281
+ raise CLIError(str(exc).strip("'\"")) from exc
282
+ better = {True: "higher is better", False: "lower is better", None: "closer to 0 is better"}
283
+ print(f"{info.name} ({info.id})\n")
284
+ print(f"Definition: {info.definition}")
285
+ print(f"Formula: {info.formula}")
286
+ print(f"Range: {info.range} ({better[info.higher_is_better]})")
287
+ print(f"Task: {info.task}")
288
+ print(f"Needs: {', '.join(info.input_requirements)}")
289
+ print(f"Python: es.{info.id.split('.')[-1]}(...)")
290
+ print("References:")
291
+ for ref in info.references:
292
+ print(f" - {ref}")
293
+ return 0
294
+
295
+
296
+ def cmd_benchmark(args: argparse.Namespace) -> int:
297
+ from ..benchmarks import run_benchmarks
298
+
299
+ sizes = (1_000, 10_000) if args.quick else tuple(args.sizes)
300
+ result = run_benchmarks(
301
+ sizes=sizes, repeat=args.repeat, compare_sklearn=not args.no_sklearn, random_state=args.seed
302
+ )
303
+ emit(result, args)
304
+ return 0
305
+
306
+
307
+ # ---- parser --------------------------------------------------------------------------------------
308
+ def build_parser() -> argparse.ArgumentParser:
309
+ parser = argparse.ArgumentParser(
310
+ prog="evalsuite",
311
+ description="EvalSuite: unified, reproducible evaluation for machine learning and research.",
312
+ epilog="Run 'evalsuite COMMAND --help' for the options of each command.",
313
+ )
314
+ parser.add_argument("--version", action="version", version=f"evalsuite-python {__version__}")
315
+ sub = parser.add_subparsers(dest="command", metavar="COMMAND")
316
+
317
+ p = sub.add_parser("evaluate", help="evaluate predictions in a file with a standard or chosen set of metrics")
318
+ add_targets(p)
319
+ p.add_argument("--task", choices=("classification", "regression"), help="default: inferred from the data")
320
+ p.add_argument("--metrics", help="comma-separated metric names, e.g. accuracy,f1,mcc")
321
+ p.add_argument("--average", default="auto", help="binary, micro, macro, weighted, samples (default: auto)")
322
+ p.add_argument("--pos-label", help="positive class for binary metrics (default: 1)")
323
+ add_output(p)
324
+ p.set_defaults(func=cmd_evaluate)
325
+
326
+ p = sub.add_parser(
327
+ "report", help="per-class classification report (precision, recall, F1, specificity, support)"
328
+ )
329
+ add_targets(p, pred_required=True)
330
+ add_output(p)
331
+ p.set_defaults(func=cmd_report)
332
+
333
+ p = sub.add_parser("compare", help="compare models on the same test set with CIs and paired tests")
334
+ p.add_argument("file", help="CSV/TSV/Parquet/JSON file with one row per observation")
335
+ p.add_argument("--y-true", required=True, metavar="COL")
336
+ p.add_argument("--pred", action="append", metavar="NAME=COL", help="a model's predictions (repeat per model)")
337
+ p.add_argument(
338
+ "--prob", action="append", metavar="NAME=COL[,COL...]", help="a model's probabilities (repeat per model)"
339
+ )
340
+ p.add_argument("--metrics", help="comma-separated metric names")
341
+ p.add_argument("--baseline", help="compare every model with this one instead of all pairs")
342
+ p.add_argument("--level", type=float, default=0.95, help="confidence level (default 0.95)")
343
+ p.add_argument("--alpha", type=float, default=0.05, help="significance level (default 0.05)")
344
+ p.add_argument("--resamples", type=int, default=1000, help="bootstrap resamples (default 1000)")
345
+ p.add_argument("--correction", default="holm", choices=("holm", "bonferroni", "bh", "by"))
346
+ p.add_argument("--seed", type=int, default=0, help="random seed (default 0, for reproducibility)")
347
+ p.add_argument("--plot", metavar="PATH", help="also save a forest plot (needs matplotlib)")
348
+ add_output(p, digits=3)
349
+ p.set_defaults(func=cmd_compare)
350
+
351
+ p = sub.add_parser("plot", help="save a ROC, PR, calibration, confusion-matrix or residual plot")
352
+ p.add_argument("kind", choices=("roc", "pr", "calibration", "confusion", "residuals", "predicted"))
353
+ add_targets(p)
354
+ p.add_argument("--normalize", choices=("true", "pred", "all"), help="confusion matrix normalisation")
355
+ p.add_argument("--output", "-o", required=True, help="image file (.png, .pdf, .svg)")
356
+ p.set_defaults(func=cmd_plot)
357
+
358
+ p = sub.add_parser("metrics", help="list available metrics")
359
+ p.add_argument("--category", choices=("classification", "regression"))
360
+ p.set_defaults(func=cmd_metrics)
361
+
362
+ p = sub.add_parser("info", help="show a metric's definition, formula, range and references")
363
+ p.add_argument("metric", help="metric id, e.g. classification.mcc")
364
+ p.set_defaults(func=cmd_info)
365
+
366
+ p = sub.add_parser("benchmark", help="time and memory benchmarks (against scikit-learn if installed)")
367
+ p.add_argument("--sizes", type=int, nargs="+", default=[1_000, 100_000, 1_000_000])
368
+ p.add_argument("--quick", action="store_true", help="small sizes only (1k and 10k)")
369
+ p.add_argument("--repeat", type=int, default=5, help="timed repetitions per case; the fastest is reported")
370
+ p.add_argument("--no-sklearn", action="store_true", help="skip the scikit-learn comparison")
371
+ p.add_argument("--seed", type=int, default=0)
372
+ add_output(p, digits=3)
373
+ p.set_defaults(func=cmd_benchmark)
374
+ return parser
375
+
376
+
377
+ def main(argv: Optional[Sequence[str]] = None) -> int:
378
+ parser = build_parser()
379
+ args = parser.parse_args(argv)
380
+ if not getattr(args, "command", None):
381
+ parser.print_help()
382
+ return 0
383
+ try:
384
+ return int(args.func(args))
385
+ except (CLIError, EvalSuiteError, ValueError) as exc:
386
+ print(f"evalsuite {args.command}: error: {exc}", file=sys.stderr)
387
+ return 2
388
+ except BrokenPipeError: # e.g. piped into `head`
389
+ return 0
@@ -0,0 +1 @@
1
+ """Core infrastructure: exceptions, validation, results, registry and evaluation context."""
@@ -0,0 +1,181 @@
1
+ """EvaluationContext: validate inputs once and cache shared intermediate results.
2
+
3
+ Every classification metric derives from per-class counts (TP, FP, FN, TN). The context computes the
4
+ confusion matrix once; ``evaluate()`` passes one context to all requested metrics, so ten metrics cost one
5
+ pass over the data.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from functools import cached_property
11
+ from typing import Any, Optional
12
+
13
+ import numpy as np
14
+ from numpy.typing import NDArray
15
+
16
+ from .exceptions import InputValidationError, UnsupportedTaskError
17
+ from .types import ArrayLike, FloatArray
18
+ from .validation import (
19
+ TargetType,
20
+ check_consistent_length,
21
+ check_finite,
22
+ resolve_labels,
23
+ target_type,
24
+ to_numpy,
25
+ validate_probabilities,
26
+ validate_sample_weight,
27
+ )
28
+
29
+ __all__ = ["ClassificationContext"]
30
+
31
+
32
+ class ClassificationContext:
33
+ """Validated classification inputs plus cached counts.
34
+
35
+ For binary and multiclass targets, ``labels`` defines the class order used by every per-class output and
36
+ by the columns of 2-D ``y_prob``. For multilabel targets, labels are column indices.
37
+ """
38
+
39
+ def __init__(
40
+ self,
41
+ y_true: ArrayLike,
42
+ y_pred: Optional[ArrayLike] = None,
43
+ *,
44
+ y_prob: Optional[ArrayLike] = None,
45
+ labels: Optional[ArrayLike] = None,
46
+ sample_weight: Optional[ArrayLike] = None,
47
+ ) -> None:
48
+ yt = to_numpy(y_true, "y_true", allow_2d=True)
49
+ yp = None if y_pred is None else to_numpy(y_pred, "y_pred", allow_2d=True)
50
+ check_finite(yt, "y_true")
51
+ if yp is not None:
52
+ check_finite(yp, "y_pred")
53
+ self.n = check_consistent_length(y_true=yt, y_pred=yp)
54
+ kind = target_type(yt)
55
+ if kind == "continuous":
56
+ raise UnsupportedTaskError(
57
+ "y_true contains non-integer numbers, which looks like a regression target. "
58
+ 'Use evalsuite.regression metrics or evaluate(..., task="regression").'
59
+ )
60
+ if yp is not None and yt.ndim != yp.ndim:
61
+ raise InputValidationError(
62
+ f"y_true has {yt.ndim} dimension(s) but y_pred has {yp.ndim}. Multilabel targets need both as "
63
+ "indicator matrices; single-label targets need both as 1-D label arrays."
64
+ )
65
+ if kind == "multilabel" and yp is not None:
66
+ if yp.shape[1] != yt.shape[1]:
67
+ raise InputValidationError(f"y_true has {yt.shape[1]} label columns but y_pred has {yp.shape[1]}.")
68
+ if not np.isin(yp, (0, 1)).all():
69
+ raise InputValidationError("Multilabel y_pred must be a 0/1 indicator matrix.")
70
+ self.target_type: TargetType = kind
71
+ self.y_true = yt
72
+ self.y_pred = yp
73
+ self.sample_weight = validate_sample_weight(sample_weight, self.n)
74
+ if kind == "multilabel":
75
+ self.labels: NDArray[Any] = np.arange(yt.shape[1])
76
+ else:
77
+ self.labels = resolve_labels(yt, yp, labels)
78
+ if labels is not None:
79
+ observed = np.unique(yt if yp is None else np.concatenate([yt, yp]))
80
+ missing = np.setdiff1d(observed, self.labels)
81
+ if missing.size:
82
+ raise InputValidationError(
83
+ f"y_true/y_pred contain label(s) {missing.tolist()} that are not in labels. "
84
+ "Include every label that occurs, or filter the data first."
85
+ )
86
+ if kind == "binary" and self.labels.shape[0] > 2:
87
+ self.target_type = "multiclass"
88
+ self._y_prob_raw = y_prob
89
+
90
+ # ---- encodings -------------------------------------------------------------------------------
91
+ @cached_property
92
+ def _index(self) -> dict[Any, int]:
93
+ return {lab.item() if hasattr(lab, "item") else lab: i for i, lab in enumerate(self.labels)}
94
+
95
+ def encode(self, y: NDArray[Any]) -> NDArray[np.int64]:
96
+ sorted_labels = np.all(self.labels[:-1] <= self.labels[1:]) if self.labels.dtype != object else False
97
+ if sorted_labels:
98
+ return np.searchsorted(self.labels, y).astype(np.int64)
99
+ idx = self._index
100
+ return np.fromiter((idx[v.item() if hasattr(v, "item") else v] for v in y), dtype=np.int64, count=len(y))
101
+
102
+ @cached_property
103
+ def true_idx(self) -> NDArray[np.int64]:
104
+ return self.encode(self.y_true)
105
+
106
+ @cached_property
107
+ def pred_idx(self) -> NDArray[np.int64]:
108
+ if self.y_pred is None:
109
+ raise InputValidationError("This metric needs y_pred (predicted labels).")
110
+ return self.encode(self.y_pred)
111
+
112
+ @cached_property
113
+ def weights(self) -> FloatArray:
114
+ return np.ones(self.n) if self.sample_weight is None else self.sample_weight
115
+
116
+ # ---- cached counts ---------------------------------------------------------------------------
117
+ @cached_property
118
+ def confusion_matrix(self) -> FloatArray:
119
+ """Rows: true labels; columns: predicted labels (weighted counts). Single-label targets only."""
120
+ if self.target_type == "multilabel":
121
+ raise UnsupportedTaskError(
122
+ "A single confusion matrix is not defined for multilabel targets; use multilabel_counts."
123
+ )
124
+ k = self.labels.shape[0]
125
+ flat = self.true_idx * k + self.pred_idx
126
+ return np.bincount(flat, weights=self.weights, minlength=k * k).reshape(k, k).astype(np.float64)
127
+
128
+ @cached_property
129
+ def counts(self) -> dict[str, FloatArray]:
130
+ """Per-class (or per-label) weighted tp, fp, fn, tn and support."""
131
+ if self.target_type == "multilabel":
132
+ if self.y_pred is None:
133
+ raise InputValidationError("This metric needs y_pred (predicted labels).")
134
+ w = self.weights[:, None]
135
+ t = self.y_true.astype(bool)
136
+ p = self.y_pred.astype(bool)
137
+ tp = ((t & p) * w).sum(0)
138
+ fp = ((~t & p) * w).sum(0)
139
+ fn = ((t & ~p) * w).sum(0)
140
+ tn = ((~t & ~p) * w).sum(0)
141
+ else:
142
+ cm = self.confusion_matrix
143
+ tp = np.diag(cm).copy()
144
+ fp = cm.sum(0) - tp
145
+ fn = cm.sum(1) - tp
146
+ tn = cm.sum() - tp - fp - fn
147
+ return {"tp": tp, "fp": fp, "fn": fn, "tn": tn, "support": tp + fn}
148
+
149
+ @cached_property
150
+ def total_weight(self) -> float:
151
+ return float(self.weights.sum())
152
+
153
+ # ---- probabilities ---------------------------------------------------------------------------
154
+ @cached_property
155
+ def y_prob(self) -> FloatArray:
156
+ if self._y_prob_raw is None:
157
+ raise InputValidationError(
158
+ "This metric needs predicted probabilities (y_prob): P(positive class) for binary tasks, "
159
+ "or one column per class in label order for multiclass tasks."
160
+ )
161
+ k = self.labels.shape[0]
162
+ p = validate_probabilities(
163
+ self._y_prob_raw,
164
+ self.n,
165
+ n_classes=None if self.target_type == "binary" else k,
166
+ rows_sum_to_one=self.target_type == "multiclass",
167
+ unit="label" if self.target_type == "multilabel" else "class",
168
+ )
169
+ if self.target_type == "binary" and p.ndim == 2:
170
+ if p.shape[1] != 2:
171
+ raise InputValidationError(
172
+ f"For binary tasks y_prob must be P(positive) or two columns; received {p.shape[1]} columns."
173
+ )
174
+ p = p[:, 1]
175
+ if self.target_type == "multiclass" and p.ndim == 1:
176
+ raise InputValidationError(
177
+ f"Multiclass tasks need y_prob with one column per class ({k} columns, in label order)."
178
+ )
179
+ if self.target_type == "multilabel" and (p.ndim != 2 or p.shape[1] != self.y_true.shape[1]):
180
+ raise InputValidationError("Multilabel y_prob must have one probability column per label.")
181
+ return p
@@ -0,0 +1,52 @@
1
+ """Exception hierarchy.
2
+
3
+ Every message says what failed, why, and how to fix it.
4
+ """
5
+
6
+ from __future__ import annotations
7
+
8
+ __all__ = [
9
+ "EvalSuiteError",
10
+ "InputValidationError",
11
+ "MetricInputError",
12
+ "OptionalDependencyError",
13
+ "StatisticalTestError",
14
+ "UndefinedMetricWarning",
15
+ "UnsupportedTaskError",
16
+ ]
17
+
18
+
19
+ class EvalSuiteError(Exception):
20
+ """Base class for all EvalSuite errors."""
21
+
22
+
23
+ class InputValidationError(EvalSuiteError, ValueError):
24
+ """Inputs have the wrong shape, length, type or values."""
25
+
26
+
27
+ class MetricInputError(EvalSuiteError, ValueError):
28
+ """Inputs are valid arrays but outside the domain a metric is defined on."""
29
+
30
+
31
+ class UnsupportedTaskError(EvalSuiteError, ValueError):
32
+ """The requested task or averaging mode is not supported for this metric."""
33
+
34
+
35
+ class OptionalDependencyError(EvalSuiteError, ImportError):
36
+ """A feature needs an optional dependency that is not installed."""
37
+
38
+ def __init__(self, package: str, extra: str, feature: str) -> None:
39
+ super().__init__(
40
+ f"{feature} requires the optional dependency '{package}'. "
41
+ f'Install it with: pip install "evalsuite-python[{extra}]"'
42
+ )
43
+ self.package = package
44
+ self.extra = extra
45
+
46
+
47
+ class StatisticalTestError(EvalSuiteError, ValueError):
48
+ """A statistical procedure cannot be applied to the given data."""
49
+
50
+
51
+ class UndefinedMetricWarning(UserWarning):
52
+ """A metric is undefined for the input (for example a zero denominator)."""