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,100 @@
1
+ """Shared HTML, CSV and file-saving helpers for every result type."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import csv
6
+ import html
7
+ import io
8
+ import os
9
+ from collections.abc import Sequence
10
+ from typing import Any, Callable, Optional, Union
11
+
12
+ from .exceptions import InputValidationError
13
+
14
+ __all__ = ["csv_text", "html_document", "html_table", "save_as"]
15
+
16
+ PathLike = Union[str, "os.PathLike[str]"]
17
+
18
+ _CSS = """
19
+ body{font-family:system-ui,-apple-system,"Segoe UI",Roboto,sans-serif;margin:2rem;color:#1f2328}
20
+ h1{font-size:1.25rem;margin:0 0 .25rem}p.meta{color:#59636e;margin:0 0 1rem;font-size:.875rem}
21
+ table.evalsuite{border-collapse:collapse;margin:0 0 1.5rem;font-variant-numeric:tabular-nums}
22
+ table.evalsuite caption{text-align:left;font-weight:600;padding:0 0 .5rem}
23
+ table.evalsuite th,table.evalsuite td{padding:.35rem .75rem;border-bottom:1px solid #d1d9e0}
24
+ table.evalsuite thead th{border-bottom:2px solid #1f2328;text-align:left}
25
+ table.evalsuite td.num,table.evalsuite th.num{text-align:right}
26
+ table.evalsuite tr.avg td{background:#f6f8fa}table.evalsuite .best{font-weight:700}
27
+ """
28
+
29
+
30
+ def html_table(
31
+ header: Sequence[str],
32
+ rows: Sequence[Sequence[Any]],
33
+ *,
34
+ caption: Optional[str] = None,
35
+ numeric: Optional[Sequence[bool]] = None,
36
+ row_classes: Optional[Sequence[str]] = None,
37
+ bold: Optional[set[tuple[int, int]]] = None,
38
+ ) -> str:
39
+ """An escaped HTML table with class ``evalsuite``. ``bold`` marks (row, column) cells."""
40
+ numeric = list(numeric) if numeric is not None else [i > 0 for i in range(len(header))]
41
+ out = ['<table class="evalsuite">']
42
+ if caption:
43
+ out.append(f"<caption>{html.escape(caption)}</caption>")
44
+ out.append(
45
+ "<thead><tr>"
46
+ + "".join(
47
+ f'<th{" class=num" if numeric[i] else ""} scope="col">{html.escape(str(h))}</th>'.replace(
48
+ " class=num", ' class="num"'
49
+ )
50
+ for i, h in enumerate(header)
51
+ )
52
+ + "</tr></thead><tbody>"
53
+ )
54
+ for r, row in enumerate(rows):
55
+ cls = f' class="{html.escape(row_classes[r])}"' if row_classes and row_classes[r] else ""
56
+ cells = []
57
+ for c, value in enumerate(row):
58
+ classes = [n for n, on in (("num", numeric[c]), ("best", bool(bold and (r, c) in bold))) if on]
59
+ attr = f' class="{" ".join(classes)}"' if classes else ""
60
+ tag = "th" if c == 0 else "td"
61
+ scope = ' scope="row"' if c == 0 else ""
62
+ cells.append(f"<{tag}{attr}{scope}>{html.escape(str(value))}</{tag}>")
63
+ out.append(f"<tr{cls}>" + "".join(cells) + "</tr>")
64
+ out.append("</tbody></table>")
65
+ return "\n".join(out)
66
+
67
+
68
+ def html_document(title: str, body: str, meta: Optional[str] = None) -> str:
69
+ """A standalone, dependency-free HTML page (inline CSS, no scripts)."""
70
+ meta_html = f'<p class="meta">{html.escape(meta)}</p>' if meta else ""
71
+ return (
72
+ '<!doctype html>\n<html lang="en">\n<head>\n<meta charset="utf-8">\n'
73
+ '<meta name="viewport" content="width=device-width, initial-scale=1">\n'
74
+ f"<title>{html.escape(title)}</title>\n<style>{_CSS}</style>\n</head>\n<body>\n"
75
+ f"<h1>{html.escape(title)}</h1>\n{meta_html}\n{body}\n</body>\n</html>\n"
76
+ )
77
+
78
+
79
+ def csv_text(header: Sequence[str], rows: Sequence[Sequence[Any]]) -> str:
80
+ """RFC 4180 CSV. Floats keep full precision; NaN is written as an empty field."""
81
+ buf = io.StringIO()
82
+ writer = csv.writer(buf, lineterminator="\n")
83
+ writer.writerow(header)
84
+ for row in rows:
85
+ writer.writerow(["" if isinstance(v, float) and v != v else v for v in row])
86
+ return buf.getvalue()
87
+
88
+
89
+ def save_as(path: PathLike, writers: dict[str, Callable[[], str]]) -> str:
90
+ """Write the format chosen by the file extension; returns the path written."""
91
+ p = os.fspath(path)
92
+ ext = os.path.splitext(p)[1].lower().lstrip(".")
93
+ aliases = {"tex": "latex", "md": "markdown", "htm": "html", "txt": "text"}
94
+ fmt = aliases.get(ext, ext)
95
+ if fmt not in writers:
96
+ exts = ", ".join(sorted({"." + e for e in (*writers, *aliases) if aliases.get(e, e) in writers}))
97
+ raise InputValidationError(f"Cannot save as '.{ext}'. Supported extensions: {exts}.")
98
+ with open(p, "w", encoding="utf-8", newline="") as fh:
99
+ fh.write(writers[fmt]())
100
+ return p
@@ -0,0 +1,90 @@
1
+ """Metric registry: programmatic discovery of every metric and its documentation."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Callable
6
+ from dataclasses import asdict, dataclass, field
7
+ from typing import Any, Optional, TypeVar
8
+
9
+ __all__ = ["MetricInfo", "list_metrics", "metric_info", "register"]
10
+
11
+ F = TypeVar("F", bound=Callable[..., Any])
12
+
13
+
14
+ @dataclass(frozen=True)
15
+ class MetricInfo:
16
+ """Documentation for one metric. ``id`` is ``"<category>.<name>"``, e.g. ``"classification.f1"``."""
17
+
18
+ id: str
19
+ name: str
20
+ category: str
21
+ task: str
22
+ definition: str
23
+ formula: str
24
+ range: str
25
+ input_requirements: tuple[str, ...] = ()
26
+ references: tuple[str, ...] = ()
27
+ higher_is_better: Optional[bool] = True
28
+ function: Optional[Callable[..., Any]] = field(default=None, repr=False, compare=False)
29
+
30
+ def to_dict(self) -> dict[str, Any]:
31
+ out = asdict(self)
32
+ out.pop("function", None)
33
+ out["input_requirements"] = list(self.input_requirements)
34
+ out["references"] = list(self.references)
35
+ out["api"] = f"evalsuite.{self.id}"
36
+ return out
37
+
38
+
39
+ _REGISTRY: dict[str, MetricInfo] = {}
40
+
41
+
42
+ def register(
43
+ *,
44
+ category: str,
45
+ task: str,
46
+ name: str,
47
+ definition: str,
48
+ formula: str,
49
+ range: str,
50
+ input_requirements: tuple[str, ...] = (),
51
+ references: tuple[str, ...] = (),
52
+ higher_is_better: Optional[bool] = True,
53
+ ) -> Callable[[F], F]:
54
+ """Decorator registering a public metric function under ``<category>.<function name>``."""
55
+
56
+ def deco(fn: F) -> F:
57
+ metric_id = f"{category}.{fn.__name__}"
58
+ if metric_id in _REGISTRY:
59
+ raise RuntimeError(f"Metric {metric_id} is registered twice.")
60
+ _REGISTRY[metric_id] = MetricInfo(
61
+ id=metric_id,
62
+ name=name,
63
+ category=category,
64
+ task=task,
65
+ definition=definition,
66
+ formula=formula,
67
+ range=range,
68
+ input_requirements=input_requirements,
69
+ references=references,
70
+ higher_is_better=higher_is_better,
71
+ function=fn,
72
+ )
73
+ return fn
74
+
75
+ return deco
76
+
77
+
78
+ def list_metrics(category: Optional[str] = None) -> list[str]:
79
+ """Sorted metric ids, optionally filtered by category (``"classification"``, ``"regression"``)."""
80
+ return sorted(k for k, v in _REGISTRY.items() if category is None or v.category == category)
81
+
82
+
83
+ def metric_info(metric_id: str) -> MetricInfo:
84
+ """Documentation for one metric, e.g. ``metric_info("classification.mcc")``."""
85
+ try:
86
+ return _REGISTRY[metric_id]
87
+ except KeyError:
88
+ close = [k for k in _REGISTRY if metric_id.split(".")[-1] in k]
89
+ hint = f" Did you mean: {', '.join(close)}?" if close else " Use evalsuite.list_metrics() to see all ids."
90
+ raise KeyError(f"Unknown metric '{metric_id}'.{hint}") from None
@@ -0,0 +1,379 @@
1
+ """Immutable result objects with portable exports (dict, JSON, DataFrame, Markdown, LaTeX).
2
+
3
+ JSON uses the standard library only; NaN and infinity become ``null`` so the output is valid JSON.
4
+ Pickle is never used.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import json
10
+ import math
11
+ from collections.abc import Iterator, Mapping
12
+ from dataclasses import dataclass, field
13
+ from types import MappingProxyType
14
+ from typing import TYPE_CHECKING, Any, Optional, Union
15
+
16
+ import numpy as np
17
+
18
+ if TYPE_CHECKING:
19
+ import pandas as pd
20
+
21
+ __all__ = ["EvaluationResult", "MetricResult"]
22
+
23
+ Number = Union[float, int]
24
+
25
+
26
+ def _json_safe(value: Any) -> Any:
27
+ if isinstance(value, (np.floating, float)):
28
+ f = float(value)
29
+ return None if math.isnan(f) or math.isinf(f) else f
30
+ if isinstance(value, (np.integer,)):
31
+ return int(value)
32
+ if isinstance(value, np.ndarray):
33
+ return [_json_safe(v) for v in value.tolist()]
34
+ if isinstance(value, (list, tuple)):
35
+ return [_json_safe(v) for v in value]
36
+ if isinstance(value, Mapping):
37
+ return {str(k): _json_safe(v) for k, v in value.items()}
38
+ if isinstance(value, (np.bool_,)):
39
+ return bool(value)
40
+ return value
41
+
42
+
43
+ def _fmt(value: Any, digits: int) -> str:
44
+ if value is None:
45
+ return "–"
46
+ if isinstance(value, (float, np.floating)):
47
+ return "NaN" if math.isnan(value) else f"{float(value):.{digits}f}"
48
+ return str(value)
49
+
50
+
51
+ def _latex_escape(text: str) -> str:
52
+ repl = {
53
+ "\\": r"\textbackslash{}",
54
+ "&": r"\&",
55
+ "%": r"\%",
56
+ "$": r"\$",
57
+ "#": r"\#",
58
+ "_": r"\_",
59
+ "{": r"\{",
60
+ "}": r"\}",
61
+ "~": r"\textasciitilde{}",
62
+ "^": r"\textasciicircum{}",
63
+ }
64
+ return "".join(repl.get(c, c) for c in text)
65
+
66
+
67
+ @dataclass(frozen=True, eq=False)
68
+ class MetricResult:
69
+ """One metric's value.
70
+
71
+ ``value`` is a float, or a read-only array of per-class values when ``average=None``
72
+ (aligned with ``labels``). ``params`` records how it was computed (average, zero_division, ...).
73
+ """
74
+
75
+ metric: str
76
+ name: str
77
+ value: Union[float, np.ndarray]
78
+ params: Mapping[str, Any] = field(default_factory=dict)
79
+ labels: Optional[tuple[Any, ...]] = None
80
+
81
+ def __post_init__(self) -> None:
82
+ if isinstance(self.value, np.ndarray):
83
+ arr = np.array(self.value, dtype=np.float64)
84
+ arr.setflags(write=False)
85
+ object.__setattr__(self, "value", arr)
86
+ else:
87
+ object.__setattr__(self, "value", float(self.value))
88
+ object.__setattr__(self, "params", MappingProxyType(dict(self.params)))
89
+
90
+ def __float__(self) -> float:
91
+ if isinstance(self.value, np.ndarray):
92
+ raise TypeError(f"{self.name} has per-class values; use .value, .per_class() or .to_dict().")
93
+ return float(self.value)
94
+
95
+ # Scalar results behave like numbers: f"{r:.3f}", r > 0.8, round(r, 3), r == 0.75.
96
+ def __format__(self, spec: str) -> str:
97
+ return format(float(self), spec) if spec else repr(self)
98
+
99
+ def __round__(self, ndigits: Optional[int] = None) -> float:
100
+ return round(float(self), ndigits)
101
+
102
+ def __lt__(self, other: Number) -> bool:
103
+ return float(self) < float(other)
104
+
105
+ def __le__(self, other: Number) -> bool:
106
+ return float(self) <= float(other)
107
+
108
+ def __gt__(self, other: Number) -> bool:
109
+ return float(self) > float(other)
110
+
111
+ def __ge__(self, other: Number) -> bool:
112
+ return float(self) >= float(other)
113
+
114
+ def __eq__(self, other: object) -> bool:
115
+ if isinstance(other, (int, float, np.floating, np.integer)) and not isinstance(self.value, np.ndarray):
116
+ return float(self.value) == float(other)
117
+ return self is other
118
+
119
+ def __hash__(self) -> int:
120
+ return id(self)
121
+
122
+ def __repr__(self) -> str:
123
+ if isinstance(self.value, np.ndarray):
124
+ pairs = ", ".join(f"{lab}: {_fmt(v, 4)}" for lab, v in zip(self.labels or (), self.value))
125
+ return f"MetricResult({self.metric}: {{{pairs}}})"
126
+ return f"MetricResult({self.metric}={_fmt(self.value, 6)})"
127
+
128
+ def per_class(self) -> dict[Any, float]:
129
+ if not isinstance(self.value, np.ndarray):
130
+ raise TypeError(f"{self.name} is a single value; compute it with average=None for per-class values.")
131
+ return {lab: float(v) for lab, v in zip(self.labels or (), self.value)}
132
+
133
+ def to_dict(self) -> dict[str, Any]:
134
+ out: dict[str, Any] = {"metric": self.metric, "name": self.name, "value": _json_safe(self.value)}
135
+ if self.labels is not None:
136
+ out["labels"] = _json_safe(list(self.labels))
137
+ if self.params:
138
+ out["params"] = _json_safe(dict(self.params))
139
+ return out
140
+
141
+ def to_json(self, *, indent: Optional[int] = 2) -> str:
142
+ return json.dumps(self.to_dict(), indent=indent, allow_nan=False)
143
+
144
+ def to_dataframe(self) -> pd.DataFrame:
145
+ import pandas as pd
146
+
147
+ frame: pd.DataFrame
148
+ if isinstance(self.value, np.ndarray):
149
+ frame = pd.DataFrame({"label": list(self.labels or ()), self.metric: self.value})
150
+ else:
151
+ frame = pd.DataFrame({"metric": [self.metric], "name": [self.name], "value": [self.value]})
152
+ return frame
153
+
154
+ def to_markdown(self, *, digits: int = 4) -> str:
155
+ if isinstance(self.value, np.ndarray):
156
+ rows = [f"| {lab} | {_fmt(v, digits)} |" for lab, v in zip(self.labels or (), self.value)]
157
+ return "\n".join([f"| Label | {self.name} |", "| --- | ---: |", *rows])
158
+ return "\n".join(["| Metric | Value |", "| --- | ---: |", f"| {self.name} | {_fmt(self.value, digits)} |"])
159
+
160
+ def to_latex(self, *, digits: int = 4) -> str:
161
+ return _latex_table(
162
+ ["Metric", "Value"] if not isinstance(self.value, np.ndarray) else ["Label", _latex_escape(self.name)],
163
+ [[_latex_escape(self.name), _fmt(self.value, digits)]]
164
+ if not isinstance(self.value, np.ndarray)
165
+ else [[_latex_escape(str(lab)), _fmt(v, digits)] for lab, v in zip(self.labels or (), self.value)],
166
+ )
167
+
168
+ def _rows(self) -> list[list[Any]]:
169
+ if isinstance(self.value, np.ndarray):
170
+ return [[self.metric, self.name, lab, float(v)] for lab, v in zip(self.labels or (), self.value)]
171
+ return [[self.metric, self.name, "", float(self.value)]]
172
+
173
+ def to_csv(self, path: Optional[str] = None) -> str:
174
+ """CSV with columns metric, name, label, value (one row per class for per-class results)."""
175
+ from .export import csv_text
176
+
177
+ text = csv_text(["metric", "name", "label", "value"], self._rows())
178
+ if path is not None:
179
+ with open(path, "w", encoding="utf-8", newline="") as fh:
180
+ fh.write(text)
181
+ return text
182
+
183
+ def to_html(self, *, digits: int = 4, full: bool = False) -> str:
184
+ """HTML table (``full=True``: a standalone page)."""
185
+ from .export import html_document, html_table
186
+
187
+ if isinstance(self.value, np.ndarray):
188
+ table = html_table(
189
+ ["Label", self.name], [[lab, _fmt(v, digits)] for lab, v in zip(self.labels or (), self.value)]
190
+ )
191
+ else:
192
+ table = html_table(["Metric", "Value"], [[self.name, _fmt(self.value, digits)]])
193
+ return html_document(self.name, table) if full else table
194
+
195
+
196
+ def _latex_table(
197
+ header: list[str], rows: list[list[str]], caption: Optional[str] = None, label: Optional[str] = None
198
+ ) -> str:
199
+ cols = "l" + "r" * (len(header) - 1)
200
+ lines = [r"\begin{table}[ht]", r"\centering"]
201
+ if caption:
202
+ lines.append(rf"\caption{{{_latex_escape(caption)}}}")
203
+ if label:
204
+ lines.append(rf"\label{{{label}}}")
205
+ lines += [rf"\begin{{tabular}}{{{cols}}}", r"\toprule", " & ".join(header) + r" \\", r"\midrule"]
206
+ lines += [" & ".join(r) + r" \\" for r in rows]
207
+ lines += [r"\bottomrule", r"\end{tabular}", r"\end{table}"]
208
+ return "\n".join(lines)
209
+
210
+
211
+ @dataclass(frozen=True, eq=False)
212
+ class EvaluationResult(Mapping[str, MetricResult]):
213
+ """All metrics from one evaluation, keyed by metric id (e.g. ``"accuracy"``)."""
214
+
215
+ task: str
216
+ metrics: Mapping[str, MetricResult]
217
+ n_samples: int
218
+ target_type: str
219
+ labels: Optional[tuple[Any, ...]] = None
220
+ confusion_matrix: Optional[np.ndarray] = None
221
+ metadata: Mapping[str, Any] = field(default_factory=dict)
222
+
223
+ def __post_init__(self) -> None:
224
+ object.__setattr__(self, "metrics", MappingProxyType(dict(self.metrics)))
225
+ object.__setattr__(self, "metadata", MappingProxyType(dict(self.metadata)))
226
+ if self.confusion_matrix is not None:
227
+ cm = np.array(self.confusion_matrix)
228
+ cm.setflags(write=False)
229
+ object.__setattr__(self, "confusion_matrix", cm)
230
+
231
+ # Mapping interface: result["f1"]
232
+ def __getitem__(self, key: str) -> MetricResult:
233
+ try:
234
+ return self.metrics[key]
235
+ except KeyError:
236
+ raise KeyError(f"No metric '{key}' in this result. Available: {', '.join(self.metrics)}") from None
237
+
238
+ def __iter__(self) -> Iterator[str]:
239
+ return iter(self.metrics)
240
+
241
+ def __len__(self) -> int:
242
+ return len(self.metrics)
243
+
244
+ @property
245
+ def value(self) -> dict[str, Any]:
246
+ return {k: m.value for k, m in self.metrics.items()}
247
+
248
+ def summary(self, *, digits: int = 4) -> str:
249
+ width = max((len(m.name) for m in self.metrics.values()), default=6)
250
+ lines = [f"EvalSuite {self.task} evaluation ({self.target_type}, n={self.n_samples})"]
251
+ for m in self.metrics.values():
252
+ if isinstance(m.value, np.ndarray):
253
+ lines.append(
254
+ f" {m.name:<{width}} "
255
+ + ", ".join(f"{lab}={_fmt(v, digits)}" for lab, v in zip(m.labels or (), m.value))
256
+ )
257
+ else:
258
+ lines.append(f" {m.name:<{width}} {_fmt(m.value, digits)}")
259
+ return "\n".join(lines)
260
+
261
+ def __repr__(self) -> str:
262
+ return self.summary()
263
+
264
+ def to_dict(self) -> dict[str, Any]:
265
+ out: dict[str, Any] = {
266
+ "task": self.task,
267
+ "target_type": self.target_type,
268
+ "n_samples": self.n_samples,
269
+ "metrics": {k: m.to_dict() for k, m in self.metrics.items()},
270
+ "metadata": _json_safe(dict(self.metadata)),
271
+ }
272
+ if self.labels is not None:
273
+ out["labels"] = _json_safe(list(self.labels))
274
+ if self.confusion_matrix is not None:
275
+ out["confusion_matrix"] = _json_safe(self.confusion_matrix)
276
+ return out
277
+
278
+ def to_json(self, *, indent: Optional[int] = 2) -> str:
279
+ return json.dumps(self.to_dict(), indent=indent, allow_nan=False)
280
+
281
+ def _scalar_rows(self) -> list[MetricResult]:
282
+ return [m for m in self.metrics.values() if not isinstance(m.value, np.ndarray)]
283
+
284
+ def to_dataframe(self) -> pd.DataFrame:
285
+ """One row per scalar metric (per-class metrics: use ``result[id].to_dataframe()``)."""
286
+ import pandas as pd
287
+
288
+ rows = self._scalar_rows()
289
+ frame: pd.DataFrame = pd.DataFrame(
290
+ {
291
+ "metric": [m.metric for m in rows],
292
+ "name": [m.name for m in rows],
293
+ "value": [m.value for m in rows],
294
+ }
295
+ )
296
+ return frame
297
+
298
+ def to_markdown(self, *, digits: int = 4) -> str:
299
+ rows = [f"| {m.name} | {_fmt(m.value, digits)} |" for m in self._scalar_rows()]
300
+ return "\n".join(["| Metric | Value |", "| --- | ---: |", *rows])
301
+
302
+ def to_latex(self, *, digits: int = 4, caption: Optional[str] = None, label: Optional[str] = None) -> str:
303
+ rows = [[_latex_escape(m.name), _fmt(m.value, digits)] for m in self._scalar_rows()]
304
+ return _latex_table(["Metric", "Value"], rows, caption=caption, label=label)
305
+
306
+ def _csv_rows(self) -> list[list[Any]]:
307
+ rows: list[list[Any]] = []
308
+ for m in self.metrics.values():
309
+ rows.extend(m._rows())
310
+ return rows
311
+
312
+ def to_csv(self, path: Optional[str] = None) -> str:
313
+ """CSV with columns metric, name, label, value; per-class metrics get one row per label."""
314
+ from .export import csv_text
315
+
316
+ text = csv_text(["metric", "name", "label", "value"], self._csv_rows())
317
+ if path is not None:
318
+ with open(path, "w", encoding="utf-8", newline="") as fh:
319
+ fh.write(text)
320
+ return text
321
+
322
+ def _meta_line(self) -> str:
323
+ return (
324
+ f"{self.task}, {self.target_type}, n = {self.n_samples}; EvalSuite "
325
+ f"{self.metadata.get('evalsuite_version', '')}"
326
+ )
327
+
328
+ def to_html(self, *, digits: int = 4, caption: Optional[str] = None, full: bool = False) -> str:
329
+ """HTML report: the metric table, per-class tables and the confusion matrix (``full=True``: standalone
330
+ page with inline CSS and no scripts)."""
331
+ from .export import html_document, html_table
332
+
333
+ parts = [
334
+ html_table(
335
+ ["Metric", "Value"],
336
+ [[m.name, _fmt(m.value, digits)] for m in self._scalar_rows()],
337
+ caption=caption or "Metrics",
338
+ )
339
+ ]
340
+ for m in self.metrics.values():
341
+ if isinstance(m.value, np.ndarray):
342
+ parts.append(
343
+ html_table(
344
+ ["Label", m.name],
345
+ [[lab, _fmt(v, digits)] for lab, v in zip(m.labels or (), m.value)],
346
+ caption=f"{m.name} per class",
347
+ )
348
+ )
349
+ if self.confusion_matrix is not None and self.labels is not None:
350
+ cm = self.confusion_matrix
351
+ integral = bool(np.all(np.mod(cm, 1) == 0))
352
+ parts.append(
353
+ html_table(
354
+ ["True \\ predicted", *(str(lab) for lab in self.labels)],
355
+ [
356
+ [lab, *(int(v) if integral else _fmt(v, digits) for v in row)]
357
+ for lab, row in zip(self.labels, cm)
358
+ ],
359
+ caption="Confusion matrix (rows: true, columns: predicted)",
360
+ )
361
+ )
362
+ body = "\n".join(parts)
363
+ return html_document("EvalSuite evaluation report", body, self._meta_line()) if full else body
364
+
365
+ def save(self, path: str, *, digits: int = 4) -> str:
366
+ """Save in the format given by the extension: .json .csv .md .tex .html .txt."""
367
+ from .export import save_as
368
+
369
+ return save_as(
370
+ path,
371
+ {
372
+ "json": self.to_json,
373
+ "csv": self.to_csv,
374
+ "markdown": lambda: self.to_markdown(digits=digits) + "\n",
375
+ "latex": lambda: self.to_latex(digits=digits) + "\n",
376
+ "html": lambda: self.to_html(digits=digits, full=True),
377
+ "text": lambda: self.summary(digits=digits) + "\n",
378
+ },
379
+ )
@@ -0,0 +1,23 @@
1
+ """Shared type aliases."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import Any, Literal, Union
6
+
7
+ import numpy as np
8
+ from numpy.typing import ArrayLike, NDArray
9
+
10
+ __all__ = ["Average", "ArrayLike", "FloatArray", "IntArray", "Task", "ZeroDivision"]
11
+
12
+ FloatArray = NDArray[np.float64]
13
+ IntArray = NDArray[np.int64]
14
+
15
+ #: How per-class results are combined. ``None`` returns per-class values.
16
+ Average = Literal["binary", "micro", "macro", "weighted", "samples", None]
17
+
18
+ #: Value used when a metric's denominator is zero. ``"warn"`` returns NaN and warns.
19
+ ZeroDivision = Union[Literal["warn"], float]
20
+
21
+ Task = Literal["classification", "regression"]
22
+
23
+ JSONValue = Union[None, bool, int, float, str, list[Any], dict[str, Any]]