easyclassifier 0.8.1__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,510 @@
1
+ """All figures, drawn in one of four colour themes.
2
+
3
+ * Every class keeps the same colour (and, in greyscale, the same hatch or
4
+ line style) in every figure.
5
+ * Figures are saved as PNG at 300 dpi, the usual journal requirement, and
6
+ optionally also as PDF or SVG (vector files that stay sharp at any size).
7
+ * A non-interactive matplotlib backend is used, so it works on servers.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import os
13
+ from contextlib import contextmanager
14
+ from dataclasses import dataclass
15
+ from typing import Dict, List, Optional, Sequence
16
+
17
+ import matplotlib
18
+ matplotlib.use("Agg") # headless
19
+ import matplotlib.pyplot as plt # noqa: E402
20
+ import numpy as np # noqa: E402
21
+ import pandas as pd # noqa: E402
22
+ from sklearn.metrics import ( # noqa: E402
23
+ auc,
24
+ precision_recall_curve,
25
+ roc_curve,
26
+ )
27
+
28
+
29
+ # --------------------------------------------------------------------------- #
30
+ # Themes
31
+ # --------------------------------------------------------------------------- #
32
+
33
+ @dataclass(frozen=True)
34
+ class Theme:
35
+ key: str
36
+ name: str
37
+ description: str
38
+ palette: Sequence[str]
39
+ cmap: str # sequential (confusion matrices)
40
+ cmap_div: str # diverging (correlations, -1 .. +1)
41
+ font_scale: float = 1.0
42
+ linewidth: float = 1.8
43
+ hatches: Sequence[str] = ("",)
44
+ linestyles: Sequence[str] = ("-",)
45
+ edge: str = "none"
46
+
47
+
48
+ THEMES: Dict[str, Theme] = {
49
+ "colorblind": Theme(
50
+ "colorblind", "Colour-blind safe",
51
+ "Readable with the common forms of colour blindness (Okabe-Ito "
52
+ "colours). Recommended for papers.",
53
+ # Okabe & Ito (2008) palette
54
+ ["#0072B2", "#E69F00", "#009E73", "#CC79A7", "#56B4E9",
55
+ "#D55E00", "#F0E442", "#000000"],
56
+ cmap="cividis", cmap_div="RdBu_r"),
57
+ "greyscale": Theme(
58
+ "greyscale", "Greyscale",
59
+ "Shades of grey with patterns and line styles, for printed journals "
60
+ "and theses.",
61
+ ["#1a1a1a", "#7f7f7f", "#bdbdbd", "#4d4d4d", "#a6a6a6", "#e0e0e0"],
62
+ cmap="Greys", cmap_div="Greys",
63
+ hatches=("", "//", "..", "xx", "\\\\", "oo"),
64
+ linestyles=("-", "--", "-.", ":"), edge="black"),
65
+ "high_contrast": Theme(
66
+ "high_contrast", "High contrast",
67
+ "Strong colours, thick lines and larger text, for slides and "
68
+ "posters.",
69
+ # Blue and orange first: the pair that stays distinct for readers
70
+ # with red-green colour blindness.
71
+ ["#0033CC", "#FF8C00", "#008000", "#8B00FF", "#E60000", "#00A6A6",
72
+ "#000000"],
73
+ cmap="viridis", cmap_div="coolwarm", font_scale=1.25,
74
+ linewidth=3.0, edge="black"),
75
+ "soft": Theme(
76
+ "soft", "Soft",
77
+ "Gentle pastel colours for reports and teaching material.",
78
+ ["#66C2A5", "#FC8D62", "#8DA0CB", "#E78AC3", "#A6D854", "#FFD92F",
79
+ "#E5C494", "#B3B3B3"],
80
+ cmap="PuBu", cmap_div="PiYG"),
81
+ }
82
+
83
+ DEFAULT_THEME = "colorblind"
84
+
85
+ FORMATS = {
86
+ "png": "PNG images at 300 dpi",
87
+ "png+pdf": "PNG + PDF (vector files, best for journals)",
88
+ "png+svg": "PNG + SVG (vector files, for editing in Inkscape or "
89
+ "Illustrator)",
90
+ }
91
+
92
+ # Figure key -> menu label (order = order in the report)
93
+ FIGURES = {
94
+ "class_distribution": "Class distribution",
95
+ "missing_values": "Missing-value map",
96
+ "correlation": "Correlation matrix of numeric columns",
97
+ "comparison": "Classifier comparison (with fold-to-fold spread)",
98
+ "confusion_matrix": "Confusion matrix (counts and percentages)",
99
+ "roc_curve": "ROC curves",
100
+ "pr_curve": "Precision-recall curves",
101
+ "feature_importance": "Feature importance (which columns matter)",
102
+ "columns_by_class": "Most important columns, by class",
103
+ "learning_curve": "Learning curve (would more data help?)",
104
+ }
105
+
106
+ # The figures most used in classification papers: always made by default.
107
+ STANDARD_FIGURES = ["class_distribution", "comparison", "confusion_matrix",
108
+ "roc_curve", "feature_importance"]
109
+ # Added by default only when the data call for them.
110
+ WHEN_NEEDED = {
111
+ "pr_curve": "classes are imbalanced",
112
+ "missing_values": "the data have empty cells",
113
+ }
114
+ # Optional extras, chosen in Advanced/Research mode.
115
+ OPTIONAL_FIGURES = ["correlation", "columns_by_class", "learning_curve"]
116
+
117
+
118
+ def default_figures(imbalanced: bool, has_missing: bool) -> List[str]:
119
+ figs = list(STANDARD_FIGURES)
120
+ if imbalanced:
121
+ figs.append("pr_curve")
122
+ if has_missing:
123
+ figs.append("missing_values")
124
+ return [k for k in FIGURES if k in figs] # report order
125
+
126
+
127
+ # --------------------------------------------------------------------------- #
128
+ # Figure maker
129
+ # --------------------------------------------------------------------------- #
130
+
131
+ class FigureMaker:
132
+ """Draws figures in a theme and saves them in the chosen formats.
133
+
134
+ ``files`` maps each figure key to {"png": path, "pdf": path, ...}.
135
+ """
136
+
137
+ def __init__(self, out_dir: str, theme: str = DEFAULT_THEME,
138
+ formats: str = "png", dpi: int = 300,
139
+ class_names: Optional[List[str]] = None):
140
+ self.out_dir = out_dir
141
+ os.makedirs(out_dir, exist_ok=True)
142
+ self.theme = THEMES[theme]
143
+ self.exts = formats.split("+")
144
+ self.dpi = dpi
145
+ self.class_names = list(class_names or [])
146
+ self.files: Dict[str, Dict[str, str]] = {}
147
+
148
+ # ---- styling ------------------------------------------------------ #
149
+
150
+ def color(self, i: int) -> str:
151
+ return self.theme.palette[i % len(self.theme.palette)]
152
+
153
+ def hatch(self, i: int) -> str:
154
+ return self.theme.hatches[i % len(self.theme.hatches)]
155
+
156
+ def linestyle(self, i: int) -> str:
157
+ return self.theme.linestyles[i % len(self.theme.linestyles)]
158
+
159
+ @contextmanager
160
+ def _style(self):
161
+ s = self.theme.font_scale
162
+ rc = {
163
+ "font.size": 10 * s, "axes.titlesize": 11.5 * s,
164
+ "axes.labelsize": 10.5 * s, "xtick.labelsize": 9 * s,
165
+ "ytick.labelsize": 9 * s, "legend.fontsize": 9 * s,
166
+ "axes.spines.top": False, "axes.spines.right": False,
167
+ "axes.grid": True, "grid.alpha": 0.3, "grid.linewidth": 0.6,
168
+ "lines.linewidth": self.theme.linewidth,
169
+ "hatch.linewidth": 0.8, "svg.fonttype": "none",
170
+ "pdf.fonttype": 42,
171
+ }
172
+ with plt.rc_context(rc):
173
+ yield
174
+
175
+ def _save(self, fig, key: str) -> Dict[str, str]:
176
+ paths = {}
177
+ for ext in self.exts:
178
+ p = os.path.join(self.out_dir, f"{key}.{ext}")
179
+ fig.savefig(p, dpi=self.dpi, bbox_inches="tight")
180
+ paths[ext] = p
181
+ plt.close(fig)
182
+ self.files[key] = paths
183
+ return paths
184
+
185
+ def _bar_style(self, i: int) -> dict:
186
+ return dict(color=self.color(i), hatch=self.hatch(i),
187
+ edgecolor=(self.theme.edge if self.theme.edge != "none"
188
+ else self.color(i)), linewidth=0.8)
189
+
190
+ # ---- data description --------------------------------------------- #
191
+
192
+ def class_distribution(self, y) -> Dict[str, str]:
193
+ y = np.asarray(y)
194
+ counts = [int((y == i).sum()) for i in range(len(self.class_names))]
195
+ with self._style():
196
+ fig, ax = plt.subplots(figsize=(5.5, 3.8))
197
+ for i, (name, n) in enumerate(zip(self.class_names, counts)):
198
+ ax.bar(i, n, **self._bar_style(i))
199
+ ax.text(i, n, str(n), ha="center", va="bottom")
200
+ ax.set_xticks(range(len(counts)))
201
+ ax.set_xticklabels(self.class_names, rotation=30, ha="right")
202
+ ax.set_ylabel("Rows")
203
+ ax.set_title("Rows in each class")
204
+ ax.grid(axis="x", visible=False)
205
+ return self._save(fig, "class_distribution")
206
+
207
+ def missing_values(self, df: pd.DataFrame,
208
+ max_cols: int = 60) -> Optional[Dict[str, str]]:
209
+ """Map of empty cells (rows x columns) and share missing per
210
+ column. Returns None if nothing is missing."""
211
+ miss = df.isna()
212
+ if not miss.values.any():
213
+ return None
214
+ cols = list(miss.columns)
215
+ if len(cols) > max_cols: # prefer columns that have gaps
216
+ with_gaps = [c for c in cols if miss[c].any()]
217
+ cols = (with_gaps + [c for c in cols if c not in with_gaps])[
218
+ :max_cols]
219
+ m = miss[cols].to_numpy()
220
+ share = m.mean(axis=0) * 100
221
+ with self._style():
222
+ fig, (ax1, ax2) = plt.subplots(
223
+ 1, 2, figsize=(9, max(3.5, 0.22 * len(cols) + 1.5)),
224
+ gridspec_kw={"width_ratios": [3, 1.3]}, sharey=True)
225
+ ax1.imshow(m.T, aspect="auto", interpolation="nearest",
226
+ cmap=matplotlib.colors.ListedColormap(
227
+ ["#f2f2f2", self.color(0)]))
228
+ ax1.set_yticks(range(len(cols)))
229
+ ax1.set_yticklabels([str(c) for c in cols])
230
+ ax1.set_xlabel("Row number")
231
+ ax1.set_title("Empty cells (dark)")
232
+ ax1.grid(False)
233
+ ax2.barh(range(len(cols)), share, **self._bar_style(0))
234
+ ax2.set_xlabel("% empty")
235
+ ax2.set_xlim(0, max(5, share.max() * 1.15))
236
+ ax2.set_title("Share empty")
237
+ ax2.invert_yaxis()
238
+ fig.tight_layout()
239
+ return self._save(fig, "missing_values")
240
+
241
+ def correlation(self, X: pd.DataFrame, max_cols: int = 25,
242
+ prefer: Optional[List[str]] = None
243
+ ) -> Optional[Dict[str, str]]:
244
+ """Pearson correlations between numeric columns."""
245
+ num = [c for c in X.columns if pd.api.types.is_numeric_dtype(X[c])
246
+ and X[c].nunique() > 1]
247
+ if len(num) < 2:
248
+ return None
249
+ if len(num) > max_cols:
250
+ order = [c for c in (prefer or []) if c in num]
251
+ num = (order + [c for c in num if c not in order])[:max_cols]
252
+ corr = X[num].corr().to_numpy()
253
+ n = len(num)
254
+ with self._style():
255
+ size = max(4.5, 0.42 * n + 2)
256
+ fig, ax = plt.subplots(figsize=(size + 1, size))
257
+ im = ax.imshow(corr, cmap=self.theme.cmap_div, vmin=-1, vmax=1)
258
+ ax.set_xticks(range(n))
259
+ ax.set_yticks(range(n))
260
+ ax.set_xticklabels([str(c) for c in num], rotation=45,
261
+ ha="right")
262
+ ax.set_yticklabels([str(c) for c in num])
263
+ ax.grid(False)
264
+ if n <= 12:
265
+ for i in range(n):
266
+ for j in range(n):
267
+ v = corr[i, j]
268
+ ax.text(j, i, f"{v:.2f}", ha="center", va="center",
269
+ fontsize=8 * self.theme.font_scale,
270
+ color="white" if abs(v) > 0.6 else "black")
271
+ cb = fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
272
+ cb.set_label("Correlation")
273
+ ax.set_title("Correlation between numeric columns")
274
+ fig.tight_layout()
275
+ return self._save(fig, "correlation")
276
+
277
+ # ---- results -------------------------------------------------------- #
278
+
279
+ def comparison(self, results: List, selected: str,
280
+ final_score: Optional[float] = None,
281
+ final_label: str = "") -> Dict[str, str]:
282
+ """Balanced accuracy of each classifier: one dot per test fold and
283
+ the overall score; the selected classifier is highlighted."""
284
+ names = [r.classifier_name for r in results][::-1]
285
+ with self._style():
286
+ fig, ax = plt.subplots(
287
+ figsize=(7, max(2.6, 0.5 * len(names) + 1.3)))
288
+ for row, r in enumerate(results[::-1]):
289
+ chosen = r.classifier_name == selected
290
+ col = self.color(0) if chosen else "#8c8c8c"
291
+ if r.fold_scores:
292
+ jitter = np.linspace(-0.12, 0.12, len(r.fold_scores))
293
+ ax.scatter(np.array(r.fold_scores) * 100, row + jitter,
294
+ s=18, color=col, alpha=0.55, zorder=2,
295
+ edgecolors="none")
296
+ ax.scatter(r.selection_score * 100, row, marker="|", s=380,
297
+ linewidths=2.6, color=col, zorder=3)
298
+ if final_score is not None:
299
+ row = names.index(selected)
300
+ ax.scatter(final_score * 100, row, marker="D", s=70,
301
+ color=self.color(1), edgecolors="black",
302
+ zorder=4, label=final_label or "Final score")
303
+ ax.set_yticks(range(len(names)))
304
+ ax.set_yticklabels(names)
305
+ for lab in ax.get_yticklabels():
306
+ if lab.get_text() == selected:
307
+ lab.set_fontweight("bold")
308
+ ax.set_xlabel("Balanced accuracy (%)")
309
+ ax.set_title("Comparison of classifiers")
310
+ from matplotlib.lines import Line2D
311
+ handles = [
312
+ Line2D([], [], marker="o", ls="", color="#8c8c8c",
313
+ alpha=0.6, label="one test fold"),
314
+ Line2D([], [], marker="|", ls="", color="#8c8c8c",
315
+ markersize=14, markeredgewidth=2.6,
316
+ label="overall"),
317
+ ]
318
+ if final_score is not None:
319
+ handles.append(Line2D([], [], marker="D", ls="",
320
+ color=self.color(1),
321
+ markeredgecolor="black",
322
+ label=final_label or "Final score"))
323
+ fig.tight_layout(rect=(0, 0.1, 1, 1))
324
+ fig.legend(handles=handles, loc="lower center", ncol=3,
325
+ frameon=False)
326
+ return self._save(fig, "comparison")
327
+
328
+ def confusion(self, result) -> Dict[str, str]:
329
+ """Counts and row percentages in one matrix. Colours follow the
330
+ percentages, so small classes are as readable as large ones."""
331
+ counts = result.confusion.astype(int)
332
+ names = result.class_names
333
+ rows = counts.sum(axis=1, keepdims=True)
334
+ pct = np.divide(counts * 100.0, rows,
335
+ out=np.zeros(counts.shape, dtype=float),
336
+ where=rows > 0)
337
+ n = len(names)
338
+ with self._style():
339
+ size = max(4.2, 0.75 * n + 2.4)
340
+ fig, ax = plt.subplots(figsize=(size + 0.8, size))
341
+ im = ax.imshow(pct, cmap=self.theme.cmap, vmin=0, vmax=100)
342
+ for i in range(n):
343
+ for j in range(n):
344
+ v = pct[i, j]
345
+ txt = f"{counts[i, j]}\n({v:.0f}%)"
346
+ rgba = im.cmap(im.norm(v))
347
+ light = (0.299 * rgba[0] + 0.587 * rgba[1]
348
+ + 0.114 * rgba[2]) > 0.55
349
+ ax.text(j, i, txt, ha="center", va="center",
350
+ color="black" if light else "white",
351
+ fontweight="bold" if i == j else "normal")
352
+ ax.set_xticks(range(n))
353
+ ax.set_yticks(range(n))
354
+ ax.set_xticklabels(names, rotation=30, ha="right")
355
+ ax.set_yticklabels(names)
356
+ ax.set_xlabel("Predicted class")
357
+ ax.set_ylabel("True class")
358
+ ax.grid(False)
359
+ cb = fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
360
+ cb.set_label("% of the true class")
361
+ ax.set_title("Confusion matrix: rows (and % of each true class)")
362
+ fig.tight_layout()
363
+ return self._save(fig, "confusion_matrix")
364
+
365
+ def _curves(self, result, kind: str) -> Optional[Dict[str, str]]:
366
+ if result.y_proba is None:
367
+ return None
368
+ y, P, names = result.y_true, result.y_proba, result.class_names
369
+ binary = len(names) == 2
370
+ classes = [1] if binary else list(range(len(names)))
371
+ with self._style():
372
+ fig, ax = plt.subplots(figsize=(5.8, 4.8))
373
+ for c in classes:
374
+ yt = (y == c).astype(int)
375
+ if yt.min() == yt.max():
376
+ continue
377
+ label = names[c] if not binary else f"{names[1]} (positive)"
378
+ # One curve (two classes): always a solid line, so it cannot
379
+ # be mistaken for the dotted random-guessing reference.
380
+ style = dict(color=self.color(c),
381
+ ls="-" if binary else self.linestyle(c))
382
+ if kind == "roc":
383
+ fpr, tpr, _ = roc_curve(yt, P[:, c])
384
+ ax.plot(fpr, tpr, label=f"{label} AUC={auc(fpr, tpr):.3f}",
385
+ **style)
386
+ else:
387
+ prec, rec, _ = precision_recall_curve(yt, P[:, c])
388
+ ax.plot(rec, prec,
389
+ label=f"{label} AUC={auc(rec, prec):.3f}",
390
+ **style)
391
+ ax.axhline(yt.mean(), color=self.color(c), lw=0.8,
392
+ ls=":", alpha=0.8)
393
+ if kind == "roc":
394
+ ax.plot([0, 1], [0, 1], ls=":", color="#9e9e9e", lw=1.2,
395
+ label="random guessing")
396
+ ax.set_xlabel("False positive rate (1 - specificity)")
397
+ ax.set_ylabel("True positive rate (sensitivity)")
398
+ ax.set_title("ROC curve" + ("" if binary else
399
+ "s (each class vs the rest)"))
400
+ else:
401
+ ax.set_xlabel("Recall")
402
+ ax.set_ylabel("Precision")
403
+ ax.set_title("Precision-recall curve" + (
404
+ "" if binary else "s (each class vs the rest)"))
405
+ ax.set_xlim(-0.01, 1.01)
406
+ ax.set_ylim(-0.01, 1.02)
407
+ ax.legend(loc="lower right" if kind == "roc" else "lower left")
408
+ fig.tight_layout()
409
+ return self._save(fig, "roc_curve" if kind == "roc"
410
+ else "pr_curve")
411
+
412
+ def roc(self, result):
413
+ return self._curves(result, "roc")
414
+
415
+ def pr(self, result):
416
+ return self._curves(result, "pr")
417
+
418
+ # ---- which columns matter ----------------------------------------- #
419
+
420
+ def importance(self, table: pd.DataFrame, top: int = 20
421
+ ) -> Dict[str, str]:
422
+ t = table.head(top).iloc[::-1]
423
+ with self._style():
424
+ fig, ax = plt.subplots(figsize=(6.5, max(2.6, 0.36 * len(t)
425
+ + 1.2)))
426
+ ax.barh([str(c) for c in t["column"]], t["importance"],
427
+ xerr=t["std"], ecolor="#555555", capsize=2,
428
+ **self._bar_style(0))
429
+ ax.axvline(0, color="black", linewidth=0.8)
430
+ ax.set_xlabel("Drop in balanced accuracy when the column is "
431
+ "shuffled")
432
+ ax.set_title("Which columns matter")
433
+ ax.grid(axis="y", visible=False)
434
+ fig.tight_layout()
435
+ return self._save(fig, "feature_importance")
436
+
437
+ def columns_by_class(self, X: pd.DataFrame, y, columns: List[str]
438
+ ) -> Optional[Dict[str, str]]:
439
+ """For the most important columns: box plots of numeric columns by
440
+ class; for text columns, the share of each category by class."""
441
+ cols = [c for c in columns if c in X.columns][:6]
442
+ if not cols:
443
+ return None
444
+ y = np.asarray(y)
445
+ k = len(self.class_names)
446
+ ncols = min(3, len(cols))
447
+ nrows = int(np.ceil(len(cols) / ncols))
448
+ with self._style():
449
+ fig, axes = plt.subplots(nrows, ncols,
450
+ figsize=(4.2 * ncols, 3.6 * nrows),
451
+ squeeze=False)
452
+ for ax, col in zip(axes.flat, cols):
453
+ s = X[col]
454
+ if pd.api.types.is_numeric_dtype(s):
455
+ data = [s[y == i].dropna().to_numpy() for i in range(k)]
456
+ bp = ax.boxplot(data, patch_artist=True, widths=0.6,
457
+ medianprops=dict(color="black", lw=1.4),
458
+ flierprops=dict(markersize=3,
459
+ alpha=0.5))
460
+ for i, box in enumerate(bp["boxes"]):
461
+ box.set_facecolor(self.color(i))
462
+ box.set_alpha(0.85)
463
+ box.set_hatch(self.hatch(i))
464
+ box.set_edgecolor("black")
465
+ ax.set_xticks(range(1, k + 1))
466
+ ax.set_xticklabels(self.class_names, rotation=30,
467
+ ha="right")
468
+ ax.set_ylabel(str(col))
469
+ else:
470
+ cats = s.astype(str).value_counts().index[:8]
471
+ shares = np.array([[((s[y == i].astype(str) == c).mean()
472
+ if (y == i).any() else 0) * 100
473
+ for c in cats] for i in range(k)])
474
+ width = 0.8 / k
475
+ for i in range(k):
476
+ ax.bar(np.arange(len(cats)) + i * width
477
+ - 0.4 + width / 2, shares[i], width,
478
+ label=self.class_names[i],
479
+ **self._bar_style(i))
480
+ ax.set_xticks(range(len(cats)))
481
+ ax.set_xticklabels([str(c) for c in cats], rotation=30,
482
+ ha="right")
483
+ ax.set_ylabel(f"% of class ({col})")
484
+ ax.legend(fontsize=7 * self.theme.font_scale)
485
+ ax.set_title(str(col))
486
+ ax.grid(axis="x", visible=False)
487
+ for ax in list(axes.flat)[len(cols):]:
488
+ ax.set_visible(False)
489
+ fig.suptitle("Most important columns, by class")
490
+ fig.tight_layout()
491
+ return self._save(fig, "columns_by_class")
492
+
493
+ # ---- more data? ----------------------------------------------------- #
494
+
495
+ def learning_curve(self, d: Dict) -> Dict[str, str]:
496
+ with self._style():
497
+ fig, ax = plt.subplots(figsize=(6, 4.3))
498
+ for key, label, i in (("train", "rows used for training", 1),
499
+ ("val", "new rows (validation)", 0)):
500
+ m, s = d[f"{key}_mean"] * 100, d[f"{key}_std"] * 100
501
+ ax.plot(d["sizes"], m, marker="o", color=self.color(i),
502
+ ls=self.linestyle(i), label=label)
503
+ ax.fill_between(d["sizes"], m - s, m + s,
504
+ color=self.color(i), alpha=0.15)
505
+ ax.set_xlabel("Number of training rows")
506
+ ax.set_ylabel("Balanced accuracy (%)")
507
+ ax.set_title("Learning curve: would more data help?")
508
+ ax.legend(loc="lower right")
509
+ fig.tight_layout()
510
+ return self._save(fig, "learning_curve")
@@ -0,0 +1,113 @@
1
+ """Plain-language explanations for metrics, methods, and algorithms.
2
+
3
+ Every concept the wizard exposes has a short, jargon-free explanation here so
4
+ users can ask "what is this?" at any decision point.
5
+ """
6
+
7
+ GLOSSARY = {
8
+ # Metrics ------------------------------------------------------------- #
9
+ "accuracy": "Accuracy is the share of predictions the model got right, "
10
+ "out of all predictions.",
11
+ "precision": "Precision measures how many of the items the model labeled "
12
+ "positive are actually positive.",
13
+ "recall": "Recall (sensitivity) measures how many of the real positive "
14
+ "items the model successfully found.",
15
+ "f1": "F1 is a single score that balances precision and recall. It is "
16
+ "high only when both are high.",
17
+ "specificity": "Specificity measures how many of the real negative items "
18
+ "the model correctly labeled negative.",
19
+ "sensitivity": "Sensitivity is another name for recall: how many real "
20
+ "positives the model found.",
21
+ "mcc": "Matthews Correlation Coefficient summarises the whole confusion "
22
+ "matrix in one number from -1 (worst) to +1 (perfect). It stays "
23
+ "reliable even when classes are imbalanced.",
24
+ "cohen_kappa": "Cohen's Kappa measures agreement between predictions and "
25
+ "truth, corrected for the agreement you'd expect by chance.",
26
+ "balanced_accuracy": "Balanced accuracy averages the recall of each "
27
+ "class, so every class counts equally regardless of "
28
+ "size.",
29
+ "roc_auc": "ROC AUC measures how well the model separates classes across "
30
+ "all thresholds. 1.0 is perfect, 0.5 is random guessing.",
31
+ "pr_auc": "PR AUC summarises the precision-recall curve. It is especially "
32
+ "informative when the positive class is rare.",
33
+ "log_loss": "Log loss penalises confident wrong answers. Lower is better.",
34
+
35
+ # Preprocessing ------------------------------------------------------- #
36
+ "missing_values": "Missing values are empty cells. Models cannot train on "
37
+ "blanks, so we either remove those rows or fill the gaps "
38
+ "with a sensible value (mean, median, or most common).",
39
+ "duplicates": "Duplicate rows are identical records. They can bias a model "
40
+ "toward repeated examples, so they are often removed.",
41
+ "label_encoding": "Label encoding turns text categories into whole numbers "
42
+ "(e.g. red=0, green=1, blue=2). Best for categories with "
43
+ "a natural order, or for tree-based models.",
44
+ "one_hot_encoding": "One-hot encoding creates a separate yes/no column for "
45
+ "each category. Best when categories have no natural "
46
+ "order.",
47
+ "scaling": "Scaling (normalisation) puts all numeric columns on a similar "
48
+ "range. Distance-based models like KNN and SVM need it; "
49
+ "tree-based models do not.",
50
+ "feature_selection": "Feature selection keeps only the most useful columns. "
51
+ "It can speed up training and reduce noise.",
52
+
53
+ # Validation ---------------------------------------------------------- #
54
+ "hold_out": "Hold-out splits the data once into a training part and a "
55
+ "testing part. Fast, but the score depends on the split.",
56
+ "cross_validation": "Cross-validation splits the data into several folds, "
57
+ "trains and tests multiple times, and averages the "
58
+ "scores for a more reliable estimate.",
59
+ "stratified": "Stratified splitting keeps the class proportions the same in "
60
+ "every fold. Recommended when classes are imbalanced.",
61
+ "leave_one_out": "Leave-One-Out tests on one row at a time and trains on "
62
+ "the rest. Very thorough but slow on large datasets.",
63
+
64
+ "selection": "Automatic uses nested cross-validation for datasets of up "
65
+ "to 2,000 rows - in EasyClassifier's benchmarks it gave the "
66
+ "most accurate final scores - and a final test set for "
67
+ "larger datasets, where it is accurate and much faster.",
68
+ "final_test": "20% of the rows are put aside before anything else and "
69
+ "never looked at while comparing classifiers. The winner is "
70
+ "then scored once on those rows, like a final exam it has "
71
+ "never seen. This is the honest number to report.",
72
+ "nested_cv": "Nested cross-validation repeats the whole comparison "
73
+ "inside each of 5 folds, using only that fold's training "
74
+ "rows, and scores the winner on the fold's test rows. It "
75
+ "uses every row for testing once, so it suits small "
76
+ "datasets. Slower, but honest.",
77
+
78
+ # Classifiers --------------------------------------------------------- #
79
+ "decision_tree": "A Decision Tree asks a series of yes/no questions about "
80
+ "the data to reach a decision. Easy to interpret.",
81
+ "random_forest": "A Random Forest combines many decision trees and lets "
82
+ "them vote. Accurate and robust with little tuning.",
83
+ "svm": "A Support Vector Machine finds the boundary that best separates "
84
+ "the classes. Works well on clean, scaled data.",
85
+ "logistic_regression": "Logistic Regression estimates the probability of "
86
+ "each class using a weighted sum of the features. "
87
+ "Simple and fast.",
88
+ "knn": "K-Nearest Neighbours classifies a row by looking at the most "
89
+ "similar rows around it. You will be asked which distance to use; "
90
+ "EasyClassifier scales the columns automatically for every distance.",
91
+ "naive_bayes": "Naive Bayes uses probability and assumes features are "
92
+ "independent. Very fast, good for a baseline.",
93
+ "xgboost": "XGBoost builds trees one after another, each fixing the errors "
94
+ "of the last. Often a top performer on tabular data.",
95
+ "lightgbm": "LightGBM is a fast gradient-boosting method similar to "
96
+ "XGBoost, designed to be quick on large datasets.",
97
+ "neural_network": "A Neural Network learns patterns through layers of "
98
+ "connected 'neurons'. Flexible but needs more data.",
99
+
100
+ # Tuning -------------------------------------------------------------- #
101
+ "grid_search": "Grid Search tries every combination of the settings you "
102
+ "give it and keeps the best. Thorough but slow.",
103
+ "random_search": "Random Search tries random combinations of settings. "
104
+ "Often finds a good result much faster than grid search.",
105
+ }
106
+
107
+
108
+ def explain(key: str) -> str:
109
+ """Return the plain-language explanation for a concept key."""
110
+ return GLOSSARY.get(
111
+ key.lower().replace(" ", "_"),
112
+ "No explanation is available for this item yet.",
113
+ )