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
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)."""
|