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,1266 @@
1
+ """The interactive wizard that ties every phase together.
2
+
3
+ This is the only stateful orchestrator. Each step calls into a focused module
4
+ (ui, dataset, preprocessing, models, evaluation, reporting) so the interface
5
+ and the machine-learning logic stay separate and extensible.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import dataclasses
11
+ import datetime as _dt
12
+ import os
13
+ from typing import List
14
+
15
+ import numpy as np
16
+ import pandas as pd
17
+
18
+ from . import __version__, CITATION
19
+ from . import ui
20
+ from . import dataset as ds
21
+ from . import preprocessing as pp
22
+ from . import recommend as rec
23
+ from . import reporting as rep
24
+ from . import target as tg
25
+ from . import figures as figs
26
+ from .diagnostics import (
27
+ interpret_comparison,
28
+ interpret_learning_curve,
29
+ learning_curve_data,
30
+ )
31
+ from .importance import honest_importance
32
+ from .latex_report import ReportContext, write_and_compile
33
+ from .help_texts import explain
34
+ from .logbook import LogBook
35
+ from .models import build_registry
36
+ from .distances import (
37
+ DISTANCES,
38
+ DISTANCE_HELP,
39
+ HASSANAT_CITATION,
40
+ HASSANAT_CITATIONS,
41
+ hassanat_form,
42
+ make_knn,
43
+ )
44
+ from .evaluation import (
45
+ METRICS,
46
+ SELECTION,
47
+ SELECTION_LABEL,
48
+ VALIDATION,
49
+ evaluate,
50
+ evaluate_on_test,
51
+ fit_final_model,
52
+ nested_cv,
53
+ recommend_selection,
54
+ split_final_test,
55
+ )
56
+
57
+
58
+ LARGE_DATA_ROWS = 20_000
59
+
60
+
61
+ def figs_saved(fig_files: dict, out_dir: str) -> List[str]:
62
+ """Absolute paths of all figure files, for the list shown at the end."""
63
+ return [os.path.join(out_dir, p) for v in fig_files.values()
64
+ for p in v.values()]
65
+
66
+
67
+ class Wizard:
68
+ def __init__(self) -> None:
69
+ self.log = LogBook()
70
+ self.registry = build_registry()
71
+ self.mode = "beginner"
72
+ # State collected through the wizard.
73
+ self.df = None
74
+ self.source_name = ""
75
+ self.target = ""
76
+ self.insp = None
77
+ self.knn_distance = None
78
+ self.knn_k = None
79
+ self.hassanat_form = None
80
+ self.selection = None
81
+ self.target_display = ""
82
+ self.target_grouping = ""
83
+ self.fig_theme, self.fig_formats = figs.DEFAULT_THEME, "png"
84
+ self.df_loaded = None
85
+ self.comparison_text = ""
86
+ self.learning_text = ""
87
+ # Facts for the report's Methods section.
88
+ self.rows_loaded = self.cols_loaded = 0
89
+ self.data_notes: List[str] = []
90
+ self.left_out = {}
91
+ self.impute_used = None
92
+ self.encoding_used = None
93
+ self._dev = self._test = None
94
+
95
+ # ------------------------------------------------------------------ #
96
+ # Entry point
97
+ # ------------------------------------------------------------------ #
98
+
99
+ def run(self) -> None:
100
+ self.welcome()
101
+ self.choose_mode()
102
+ if not self.load_dataset():
103
+ return
104
+ self.inspect_dataset()
105
+ self.select_target()
106
+ auto = self.choose_auto_or_manual()
107
+ self.drop_useless_columns(auto)
108
+ self.insp = ds.inspect(self.df) # rows/columns may have changed
109
+ self.show_recommendations()
110
+ X, y, class_names = self.preprocess(auto)
111
+ classifiers = self.choose_classifiers()
112
+ if "knn" in classifiers:
113
+ self.configure_knn()
114
+ validation = self.choose_validation(auto, y)
115
+ self.selection = (self.choose_selection(auto, y)
116
+ if len(classifiers) > 1 else None)
117
+ metrics = self.choose_metrics(auto)
118
+ figures = self.choose_figures(auto, y)
119
+ if not self.confirm(classifiers, validation, metrics, figures):
120
+ ui.info("No problem - restart any time by running 'easyclassifier'.")
121
+ return
122
+ self.execute(X, y, class_names, classifiers, validation, metrics,
123
+ figures)
124
+
125
+ # ------------------------------------------------------------------ #
126
+ # Honest final score when several classifiers are compared
127
+ # ------------------------------------------------------------------ #
128
+
129
+ def choose_selection(self, auto: bool, y) -> str:
130
+ recommended = recommend_selection(y)
131
+ if not auto:
132
+ keys = list(SELECTION.keys())
133
+ ui.blank()
134
+ ui.info("You are comparing several classifiers. The score of the "
135
+ "winner is slightly too optimistic, because it partly won "
136
+ "by luck. How should its final score be measured?")
137
+ help_keys = ["selection", "final_test", "nested_cv"]
138
+ idx = ui.menu("Choose one:", list(SELECTION.values()), default=0)
139
+ idx = self._resolve_help(idx, help_keys)
140
+ choice = keys[idx]
141
+ if choice != "auto":
142
+ self.log.add(f"Final score method: {choice} (user)")
143
+ return choice
144
+ ui.note("Final score: " + (
145
+ "20% of rows kept aside as an untouched test set"
146
+ if recommended == "final_test" else
147
+ "nested cross-validation (up to 2,000 rows)")
148
+ + " (auto-selected).")
149
+ self.log.add(f"Final score method: {recommended} (auto)")
150
+ return recommended
151
+
152
+ # ------------------------------------------------------------------ #
153
+ # Phase 3: Welcome
154
+ # ------------------------------------------------------------------ #
155
+
156
+ def welcome(self) -> None:
157
+ ui.banner("Welcome to EasyClassifier",
158
+ "Machine Learning without Programming")
159
+ ui.info("This software will guide you through building and evaluating "
160
+ "classification models - one simple question at a time.")
161
+ ui.blank()
162
+ ui.note(f"Version {__version__}")
163
+ ui.note("Estimated time: 3-10 minutes")
164
+ ui.note("Citation: " + CITATION)
165
+ ui.blank()
166
+ ui.info("Tip: at any menu you can type ?N to get a plain-language "
167
+ "explanation of option N.")
168
+
169
+ # ------------------------------------------------------------------ #
170
+ # Phase 20: Mode
171
+ # ------------------------------------------------------------------ #
172
+
173
+ def choose_mode(self) -> None:
174
+ idx = ui.menu(
175
+ "Select a mode:",
176
+ ["Beginner (essential choices, sensible defaults)",
177
+ "Advanced (more control over each step)",
178
+ "Research (full control, everything shown)"],
179
+ default=0, allow_help=False,
180
+ )
181
+ self.mode = ["beginner", "advanced", "research"][idx]
182
+ self.log.add(f"Mode: {self.mode}")
183
+
184
+ @property
185
+ def is_beginner(self) -> bool:
186
+ return self.mode == "beginner"
187
+
188
+ # ------------------------------------------------------------------ #
189
+ # Phase 4: Load dataset
190
+ # ------------------------------------------------------------------ #
191
+
192
+ def load_dataset(self) -> bool:
193
+ ui.header("Step 1 of 6: Load your data")
194
+ ui.info("Type the name of your data file (.csv or .xlsx), or drag "
195
+ "the file into this window and press ENTER. Type 'demo' to "
196
+ "try a built-in sample dataset, or 'quit' to exit.")
197
+ ui.note(f"Current folder: {os.getcwd()}")
198
+ while True:
199
+ name = ui.ask_text("What is the name of your data file?")
200
+ low = name.strip().lower()
201
+ if low in ("quit", "exit"):
202
+ return False
203
+ if low == "demo":
204
+ from .demo_data import load_demo
205
+ self.df = load_demo()
206
+ self.source_name = "demo (Iris flowers)"
207
+ self.source_stem = "demo_iris"
208
+ ui.success("Loaded built-in demo dataset (Iris flowers).")
209
+ self.log.add("Dataset: demo (Iris)")
210
+ return True
211
+ path = ds.clean_path(name)
212
+ try:
213
+ self.df = ds.load_csv(path)
214
+ self.source_name = os.path.basename(path)
215
+ self.source_stem = os.path.splitext(self.source_name)[0]
216
+ ui.success(f"Dataset found: {self.source_name} "
217
+ f"({self.df.shape[0]} rows, "
218
+ f"{self.df.shape[1]} columns)")
219
+ self.log.add(f"Dataset: {os.path.abspath(path)}")
220
+ return True
221
+ except FileNotFoundError:
222
+ ui.error(f"The file cannot be found: {path}")
223
+ ui.info("Tips: drag the file into this window, or type the "
224
+ "full path, or copy the file into the current folder "
225
+ "shown above. Type 'demo' or 'quit' otherwise.")
226
+ except Exception as exc: # noqa: BLE001
227
+ ui.error(f"Could not read that file: {exc}")
228
+
229
+ # ------------------------------------------------------------------ #
230
+ # Phase 5: Inspect
231
+ # ------------------------------------------------------------------ #
232
+
233
+ def inspect_dataset(self) -> None:
234
+ self.insp = ds.inspect(self.df)
235
+ insp = self.insp
236
+ self.rows_loaded, self.cols_loaded = insp.n_rows, insp.n_cols
237
+ self.df_loaded = self.df.copy() # as loaded, for the missing map
238
+ ui.header("Step 2 of 6: Data inspection")
239
+ ui.note(f"Rows: {insp.n_rows}")
240
+ ui.note(f"Columns: {insp.n_cols}")
241
+ ui.note(f"Numeric columns: {len(insp.numeric_cols)}")
242
+ ui.note(f"Text/categorical columns: {len(insp.categorical_cols)}")
243
+ ui.note(f"Missing values: {insp.missing_total}")
244
+ ui.note(f"Duplicate rows: {insp.duplicate_rows}")
245
+ self.log.section("Inspection")
246
+ self.log.add(f"{insp.n_rows} rows x {insp.n_cols} cols, "
247
+ f"{insp.missing_total} missing, "
248
+ f"{insp.duplicate_rows} duplicates")
249
+
250
+ # ------------------------------------------------------------------ #
251
+ # Phase 6: Target
252
+ # ------------------------------------------------------------------ #
253
+
254
+ def select_target(self) -> None:
255
+ ui.header("Step 3 of 6: Choose what to predict")
256
+ ui.info("Pick the column with the groups (classes) you want to "
257
+ "predict, for example Yes/No, Healthy/Sick, or Species.")
258
+ original = self.df.copy()
259
+ cols = original.columns.tolist()
260
+ infos = [tg.describe_column(original[c]) for c in cols]
261
+ labels = [i.label for i in infos]
262
+ guess = tg.suggest_target(original)
263
+ default = cols.index(guess) if guess is not None else None
264
+ marks = {default: " <- suggested"} if default is not None else {}
265
+
266
+ while True:
267
+ self.df = original
268
+ idx = ui.pick_from_long_list(
269
+ "Which column do you want to predict?", labels, cols,
270
+ default=default, marks=marks)
271
+ if self._accept_target(cols[idx], infos[idx]):
272
+ break
273
+ ui.success(f"Target: {self.target_display}")
274
+ self.log.add(f"Target column: {self.target_display}")
275
+
276
+ def drop_useless_columns(self, auto: bool) -> None:
277
+ """Leave out columns that cannot help predict: IDs/names that differ
278
+ in every row, and columns with a single value."""
279
+ reasons = {}
280
+ for c in self.df.columns:
281
+ if c == self.target:
282
+ continue
283
+ kind = tg.describe_column(self.df[c]).kind
284
+ if kind == tg.ID_LIKE:
285
+ reasons[c] = "different in every row (ID or name)"
286
+ elif kind in (tg.CONSTANT, tg.EMPTY):
287
+ reasons[c] = "only one value"
288
+ if not reasons:
289
+ return
290
+ ui.blank()
291
+ ui.info("These columns cannot help predict and would only add noise "
292
+ "(an ID can even fake good results):")
293
+ for c, why in reasons.items():
294
+ ui.note(f"{c} - {why}")
295
+ leave_out = True
296
+ if not auto:
297
+ leave_out = ui.menu("What would you like to do?",
298
+ ["Leave them out (recommended)",
299
+ "Keep them"],
300
+ default=0, allow_help=False) == 0
301
+ if leave_out:
302
+ self.df = self.df.drop(columns=list(reasons))
303
+ self.left_out = dict(reasons)
304
+ ui.success(f"Left out {len(reasons)} column(s).")
305
+ self.log.add("Left out columns: " + ", ".join(
306
+ f"{c} ({why})" for c, why in reasons.items()))
307
+ else:
308
+ self.log.add("Kept ID/single-value columns: "
309
+ + ", ".join(reasons))
310
+
311
+ def _accept_target(self, col: str, info) -> bool:
312
+ """Check the chosen column; fix or reject it with the user's help.
313
+
314
+ Returns True when the column (possibly grouped) can be used.
315
+ """
316
+ self.target, self.target_display, self.target_grouping = col, col, ""
317
+
318
+ if info.kind in (tg.EMPTY, tg.CONSTANT):
319
+ ui.error(f"'{col}' has only one value, so there is nothing to "
320
+ "predict. Please choose another column.")
321
+ return False
322
+
323
+ if info.kind == tg.ID_LIKE:
324
+ ui.warn(f"'{col}' is different in every row, so it looks like "
325
+ "an ID or a name, not a group. A model cannot learn "
326
+ "groups from it.")
327
+ return ui.menu("What would you like to do?",
328
+ ["Choose another column", "Use it anyway"],
329
+ default=0, allow_help=False) == 1
330
+
331
+ if info.kind == tg.MANY_TEXT:
332
+ ui.warn(f"'{col}' has {info.n_unique} different values. Each "
333
+ "would become its own class, with very few rows each, "
334
+ "so results would be unreliable.")
335
+ return ui.menu("What would you like to do?",
336
+ ["Choose another column", "Use it anyway"],
337
+ default=0, allow_help=False) == 1
338
+
339
+ if info.kind == tg.MEASUREMENT and not self._group_measurement(
340
+ col, info):
341
+ return False
342
+
343
+ return self._check_class_sizes()
344
+
345
+ def _group_measurement(self, col: str, info) -> bool:
346
+ s = self.df[col]
347
+ ui.warn(f"'{col}' looks like a measurement ({info.n_unique} "
348
+ f"different numbers from {tg._fmt(s.min())} to "
349
+ f"{tg._fmt(s.max())}), not a set of groups.")
350
+ ui.info("EasyClassifier predicts groups (classification). Predicting "
351
+ "an exact number is called regression and is not supported "
352
+ "yet. You can turn the numbers into groups instead:")
353
+ choice = ui.menu(
354
+ "What would you like to do?",
355
+ ["Split into 2 equal-sized groups (Low / High)",
356
+ "Split into 3 equal-sized groups (Low / Medium / High)",
357
+ "Split into 4 equal-sized groups",
358
+ "Split at a value I choose (below / at or above it)",
359
+ "Choose another column"],
360
+ allow_help=False)
361
+ if choice == 4:
362
+ return False
363
+ try:
364
+ if choice == 3:
365
+ while True:
366
+ raw = ui.ask_text(
367
+ f"Type the value to split at (between "
368
+ f"{tg._fmt(s.min())} and {tg._fmt(s.max())})")
369
+ try:
370
+ t = float(raw)
371
+ except ValueError:
372
+ ui.error("Please type a number.")
373
+ continue
374
+ if s.min() < t <= s.max():
375
+ break
376
+ ui.error("That value would put every row in one group.")
377
+ grouped, cuts, desc = tg.group_at(s, t)
378
+ else:
379
+ grouped, cuts, desc = tg.group_equal(s, choice + 2)
380
+ except ValueError as exc:
381
+ ui.error(f"Could not form groups: {exc}.")
382
+ return False
383
+
384
+ df = self.df.copy()
385
+ df[col] = grouped
386
+ self.df = df
387
+ counts = grouped.value_counts()
388
+ ui.success("Groups created:")
389
+ for name in counts.index:
390
+ ui.note(f"{name}: {counts[name]} rows")
391
+ if choice < 3 and grouped.nunique() < choice + 2:
392
+ ui.note("Fewer groups than asked, because many rows share the "
393
+ "same value.")
394
+ self.target_grouping = desc
395
+ self.target_display = f"{col} (grouped: {desc})"
396
+ self.log.add(f"'{col}' was a measurement; grouped into: {desc}. "
397
+ "Cut-points were set once on all rows, as part of "
398
+ "defining the question.")
399
+ return True
400
+
401
+ def _check_class_sizes(self) -> bool:
402
+ s = self.df[self.target]
403
+ tiny, small = tg.tiny_and_small_classes(s)
404
+ if tiny:
405
+ n_rows = int(s.astype(str).isin(tiny).sum())
406
+ ui.warn(f"{len(tiny)} class(es) have only one row, so they can "
407
+ f"never be tested: {', '.join(tiny[:5])}"
408
+ + (", ..." if len(tiny) > 5 else ""))
409
+ choice = ui.menu(
410
+ "What would you like to do?",
411
+ [f"Remove those {n_rows} row(s) and continue",
412
+ "Choose another column"],
413
+ default=0, allow_help=False)
414
+ if choice == 1:
415
+ return False
416
+ keep = ~s.astype(str).isin(tiny) | s.isna()
417
+ self.df = self.df[keep].reset_index(drop=True)
418
+ self.data_notes.append(
419
+ f"{n_rows} row(s) belonging to classes with a single row "
420
+ f"({', '.join(tiny)}) were removed, because such classes "
421
+ "cannot be tested.")
422
+ self.log.add(f"Removed {n_rows} rows of single-row classes: "
423
+ f"{', '.join(tiny)}")
424
+ if self.df[self.target].nunique() < 2:
425
+ ui.error("Fewer than two classes are left. Please choose "
426
+ "another column.")
427
+ return False
428
+ if small:
429
+ ui.note(f"Some classes have fewer than {tg.FEW_ROWS_PER_CLASS} "
430
+ f"rows ({', '.join(small[:5])}"
431
+ + (", ..." if len(small) > 5 else "")
432
+ + "). Results for them will be unreliable.")
433
+ self.log.add(f"Small classes (<{tg.FEW_ROWS_PER_CLASS} rows): "
434
+ f"{', '.join(small)}")
435
+ return True
436
+
437
+ # ------------------------------------------------------------------ #
438
+ # Phase 5b: auto vs manual
439
+ # ------------------------------------------------------------------ #
440
+
441
+ def choose_auto_or_manual(self) -> bool:
442
+ if self.is_beginner:
443
+ idx = ui.menu(
444
+ "How would you like to prepare the data?",
445
+ ["Automatically prepare everything (recommended)",
446
+ "Review each step manually"],
447
+ default=0, allow_help=False,
448
+ )
449
+ return idx == 0
450
+ return False
451
+
452
+ def show_recommendations(self) -> None:
453
+ y_raw = self.df[self.target]
454
+ notes = rec.build_notes(self.insp, self.df, self.target, y_raw)
455
+ if notes:
456
+ ui.header("Smart recommendations")
457
+ for n in notes:
458
+ ui.note(n)
459
+ self.log.section("Recommendations")
460
+ for n in notes:
461
+ self.log.add(n)
462
+
463
+ # ------------------------------------------------------------------ #
464
+ # Phase 7-9: Preprocessing
465
+ # ------------------------------------------------------------------ #
466
+
467
+ def preprocess(self, auto: bool):
468
+ ui.header("Step 4 of 6: Prepare the data")
469
+ ui.info("Note: fill values, scaling and encodings are learned from "
470
+ "the training part only, separately for every test round, so "
471
+ "no information from the test data leaks into training.")
472
+ df = self.df
473
+ insp = self.insp
474
+ cfg = pp.PrepConfig()
475
+
476
+ # Rows without a class label can't be used at all.
477
+ df, n = pp.drop_missing_target(df, self.target)
478
+ if n:
479
+ ui.success(f"Removed {n} rows with no value in '{self.target}'.")
480
+ self.log.add(f"Removed {n} rows missing the target")
481
+ self.data_notes.append(f"{n} row(s) with no value for "
482
+ f"{self.target} were removed.")
483
+
484
+ # ---- Missing values ----
485
+ has_missing = bool(df.drop(columns=[self.target]).isna().any().any())
486
+ if has_missing:
487
+ if auto:
488
+ strat = rec.missing_strategy(insp)
489
+ strat = "median" if strat in ("none", "median") else strat
490
+ else:
491
+ choice = ui.menu(
492
+ f"Missing values detected ({insp.missing_total}). "
493
+ "How should they be handled?",
494
+ ["Remove rows with missing values",
495
+ "Replace with mean",
496
+ "Replace with median",
497
+ "Replace with most common value",
498
+ "Automatic recommendation"],
499
+ default=4,
500
+ help_keys=["missing_values"],
501
+ )
502
+ choice = self._resolve_help(choice, ["missing_values"] * 5)
503
+ strat = ["drop", "mean", "median", "mode",
504
+ rec.missing_strategy(insp)][choice]
505
+ strat = "median" if strat == "none" else strat
506
+ if strat == "drop":
507
+ df, n = pp.drop_missing_rows(df)
508
+ desc = f"Removed {n} rows with missing values."
509
+ else:
510
+ cfg.impute = strat
511
+ label = "most common value" if strat == "mode" else strat
512
+ desc = (f"Missing numeric values will be filled with the "
513
+ f"{label} (text columns: most common value), "
514
+ "learned from the training data only.")
515
+ ui.success(desc)
516
+ self.log.add(f"Missing values: {desc}")
517
+ if strat == "drop":
518
+ self.data_notes.append(f"{n} row(s) with missing values "
519
+ "were removed.")
520
+ else:
521
+ self.impute_used = strat
522
+
523
+ # ---- Duplicates ----
524
+ dup = int(df.duplicated().sum())
525
+ if dup:
526
+ by_chance = pp.duplicates_expected_by_chance(df)
527
+ remove = not by_chance
528
+ if not auto:
529
+ remove = ui.menu(
530
+ f"{dup} identical rows detected. "
531
+ + ("With so few possible value combinations, different "
532
+ "cases can easily be identical, so keeping them is "
533
+ "recommended." if by_chance else
534
+ "They are probably accidental copies, so removing "
535
+ "them is recommended."),
536
+ ["Remove them", "Keep them"],
537
+ default=0 if remove else 1, allow_help=False) == 0
538
+ if remove:
539
+ df, n = pp_remove(df)
540
+ ui.success(f"Removed {n} duplicate rows.")
541
+ self.log.add(f"Removed {n} duplicate rows")
542
+ self.data_notes.append(f"{n} duplicate row(s) were "
543
+ "removed.")
544
+ else:
545
+ ui.success(f"Kept {dup} identical rows"
546
+ + (" (expected between different cases with so "
547
+ "few possible value combinations)."
548
+ if by_chance else "."))
549
+ self.log.add(f"Kept {dup} identical rows "
550
+ f"(expected by chance: {by_chance})")
551
+ if by_chance:
552
+ self.data_notes.append(
553
+ f"{dup} identical rows were kept, because with so "
554
+ "few possible value combinations different cases "
555
+ "are expected to coincide.")
556
+
557
+ # Split X / y
558
+ y_raw = df[self.target]
559
+ X = pp.cast_categoricals(df.drop(columns=[self.target]))
560
+
561
+ # ---- Encoding ----
562
+ has_cat = any(not pd.api.types.is_numeric_dtype(X[c])
563
+ for c in X.columns)
564
+ if has_cat:
565
+ if auto:
566
+ method = rec.encoding_method(df, self.target)
567
+ else:
568
+ choice = ui.menu(
569
+ "Text/categorical columns found. Choose how to encode "
570
+ "them into numbers:",
571
+ ["Automatic", "Label Encoding", "One-Hot Encoding"],
572
+ default=0,
573
+ help_keys=["label_encoding", "one_hot_encoding"],
574
+ )
575
+ choice = self._resolve_help(
576
+ choice, ["one_hot_encoding", "label_encoding",
577
+ "one_hot_encoding"])
578
+ method = ["auto", "label", "onehot"][choice]
579
+ cfg.encoding = method
580
+ self.encoding_used = method
581
+ ui.success(f"Text columns will be encoded ({method}).")
582
+ self.log.add(f"Encoding: {method}")
583
+
584
+ y, class_names, _ = pp.encode_target(y_raw)
585
+
586
+ # ---- Scaling ----
587
+ do_scale = None # None = decide per classifier
588
+ method = "standard"
589
+ if not auto:
590
+ choice = ui.menu(
591
+ "Would you like to scale (normalise) the numeric data?",
592
+ ["Yes (standardise)", "Yes (0-1 range)", "No",
593
+ "Automatic (decide from classifier)"],
594
+ default=3,
595
+ help_keys=["scaling"],
596
+ )
597
+ choice = self._resolve_help(choice, ["scaling"] * 4)
598
+ if choice == 0:
599
+ method, do_scale = "standard", True
600
+ elif choice == 1:
601
+ method, do_scale = "minmax", True
602
+ elif choice == 2:
603
+ method, do_scale = "none", False
604
+ else:
605
+ method, do_scale = "standard", None # decide later
606
+ cfg.scale_method = method
607
+ cfg.scale = do_scale
608
+ self.prep_cfg = cfg
609
+ self.log.add(f"Scaling choice: {method if do_scale else do_scale}")
610
+
611
+ n_feat = pp.describe_transformed(X, cfg, scale=False).shape[1]
612
+ ui.success(f"Data ready: {X.shape[0]} rows, {X.shape[1]} columns "
613
+ f"({n_feat} features after encoding).")
614
+ self.log.add(f"Data: {X.shape[0]} rows x {X.shape[1]} columns, "
615
+ f"{n_feat} features after encoding")
616
+ return X, y, class_names
617
+
618
+ # ------------------------------------------------------------------ #
619
+ # Phase 10: Classifiers
620
+ # ------------------------------------------------------------------ #
621
+
622
+ def choose_classifiers(self) -> List[str]:
623
+ ui.header("Step 5 of 6: Choose classifier(s)")
624
+ specs = list(self.registry.values())
625
+ labels = []
626
+ for s in specs:
627
+ labels.append(s.name if s.available
628
+ else f"{s.name} ({s.reason})")
629
+ options = ["All available classifiers"] + labels
630
+ while True:
631
+ picks = ui.multi_select(
632
+ "Which classifier(s) would you like to try? "
633
+ "(You can pick several.)",
634
+ options, default=[0], # ENTER = all available
635
+ )
636
+ if not picks:
637
+ ui.error("Please choose at least one classifier.")
638
+ continue
639
+ keys: List[str] = []
640
+ if 0 in picks:
641
+ keys = [s.key for s in specs if s.available]
642
+ else:
643
+ for p in picks:
644
+ s = specs[p - 1]
645
+ if s.available:
646
+ keys.append(s.key)
647
+ else:
648
+ ui.warn(f"{s.name} is not installed - skipping.")
649
+ if keys:
650
+ self.log.add(f"Classifiers: {', '.join(keys)}")
651
+ slow = [self.registry[k].name for k in keys
652
+ if k in ("svm", "knn", "neural_network")]
653
+ if len(self.df) > LARGE_DATA_ROWS and slow:
654
+ ui.warn(f"Your data has {len(self.df)} rows. "
655
+ + ", ".join(slow) + " can take a long time "
656
+ "(many minutes or more) on data this size. "
657
+ "Random Forest, Logistic Regression, Naive Bayes "
658
+ "and Decision Tree are much faster.")
659
+ if not ui.ask_yes_no("Continue with this choice?",
660
+ default=True):
661
+ continue
662
+ return keys
663
+ ui.error("None of the chosen classifiers are available.")
664
+
665
+ # ------------------------------------------------------------------ #
666
+ # KNN configuration: distance metric and k
667
+ # ------------------------------------------------------------------ #
668
+
669
+ def configure_knn(self) -> None:
670
+ keys = list(DISTANCES.keys())
671
+ ui.header("KNN settings")
672
+ idx = ui.menu(
673
+ "Which distance should KNN use to measure similarity?",
674
+ [DISTANCES[k][0] for k in keys],
675
+ default=0,
676
+ )
677
+ while idx < 0:
678
+ ui.blank()
679
+ ui.info(DISTANCE_HELP[keys[-idx - 1]])
680
+ ui.pause()
681
+ idx = ui.menu(
682
+ "Which distance should KNN use to measure similarity?",
683
+ [DISTANCES[k][0] for k in keys], default=0,
684
+ )
685
+ dist = keys[idx]
686
+
687
+ k = 5
688
+ if not self.is_beginner:
689
+ while True:
690
+ raw = ui.ask_text("How many neighbours (k)?", default="5")
691
+ if raw.isdigit() and int(raw) >= 1:
692
+ k = int(raw)
693
+ break
694
+ ui.error("Please enter a whole number of 1 or more.")
695
+
696
+ self.knn_distance, self.knn_k = dist, k
697
+ spec = self.registry["knn"]
698
+ spec.factory = lambda d=dist, kk=k: make_knn(d, kk)
699
+ spec.name = f"KNN ({DISTANCES[dist][0].split(' (')[0]}, k={k})"
700
+ ui.success(f"KNN will use: {spec.name}")
701
+ self.log.add(f"KNN distance: {dist}, k={k}")
702
+
703
+ if dist == "hassanat":
704
+ ui.blank()
705
+ ui.info("Citation for the Hassanat distance - please cite if you "
706
+ "publish results:")
707
+ for c in HASSANAT_CITATIONS:
708
+ ui.note(c)
709
+ self.log.add("Citation: " + HASSANAT_CITATION)
710
+
711
+ # ------------------------------------------------------------------ #
712
+ # Phase 12: Validation
713
+ # ------------------------------------------------------------------ #
714
+
715
+ def choose_validation(self, auto: bool, y) -> str:
716
+ if auto:
717
+ method = rec.validation_method(y)
718
+ ui.note(f"Validation: {VALIDATION[method]} (auto-selected).")
719
+ self.log.add(f"Validation: {method}")
720
+ return method
721
+ keys = list(VALIDATION.keys())
722
+ idx = ui.menu(
723
+ "How should the models be evaluated?",
724
+ list(VALIDATION.values()),
725
+ default=1,
726
+ help_keys=["hold_out", "cross_validation", "cross_validation",
727
+ "stratified", "leave_one_out"],
728
+ )
729
+ idx = self._resolve_help(idx, ["hold_out", "cross_validation",
730
+ "cross_validation", "stratified",
731
+ "leave_one_out"])
732
+ self.log.add(f"Validation: {keys[idx]}")
733
+ return keys[idx]
734
+
735
+ # ------------------------------------------------------------------ #
736
+ # Phase 13: Metrics
737
+ # ------------------------------------------------------------------ #
738
+
739
+ def choose_metrics(self, auto: bool) -> List[str]:
740
+ keys = list(METRICS.keys())
741
+ if auto:
742
+ chosen = ["accuracy", "precision", "recall", "f1", "roc_auc"]
743
+ ui.note("Metrics: Accuracy, Precision, Recall, F1, ROC AUC "
744
+ "(auto-selected).")
745
+ self.log.add("Metrics: auto set")
746
+ return chosen
747
+ picks = ui.multi_select(
748
+ "Which performance metrics would you like? (Select several.)",
749
+ list(METRICS.values()), default_all=True,
750
+ )
751
+ chosen = [keys[p] for p in picks]
752
+ self.log.add(f"Metrics: {', '.join(chosen)}")
753
+ return chosen
754
+
755
+ # ------------------------------------------------------------------ #
756
+ # Phase 14: Figures
757
+ # ------------------------------------------------------------------ #
758
+
759
+ def choose_figures(self, auto: bool, y) -> List[str]:
760
+ keys = list(figs.FIGURES.keys())
761
+ imbalanced = pp.imbalance_ratio(y) > 1.5
762
+ has_missing = bool(self.df_loaded.isna().values.any())
763
+ defaults = figs.default_figures(imbalanced, has_missing)
764
+ self.fig_theme, self.fig_formats = figs.DEFAULT_THEME, "png"
765
+
766
+ if auto:
767
+ chosen = defaults
768
+ extra = [figs.WHEN_NEEDED[k] for k in figs.WHEN_NEEDED
769
+ if k in chosen]
770
+ ui.note("Figures: the standard set"
771
+ + (" (plus extra because " + " and ".join(extra) + ")"
772
+ if extra else "")
773
+ + f"; {figs.THEMES[self.fig_theme].name.lower()} "
774
+ "colours; PNG at 300 dpi (auto-selected).")
775
+ else:
776
+ picks = ui.multi_select(
777
+ "Which figures would you like? ENTER keeps the suggested "
778
+ "ones (the most used in papers).",
779
+ list(figs.FIGURES.values()),
780
+ default=[keys.index(k) for k in defaults])
781
+ chosen = [keys[p] for p in picks]
782
+ if ("learning_curve" in chosen
783
+ and len(self.df) > LARGE_DATA_ROWS):
784
+ ui.warn("The learning curve trains the model 25 more times; "
785
+ f"with {len(self.df)} rows this can take long.")
786
+ if chosen:
787
+ tkeys = list(figs.THEMES)
788
+ idx = ui.menu(
789
+ "Colour theme for the figures:",
790
+ [f"{figs.THEMES[k].name} - {figs.THEMES[k].description}"
791
+ for k in tkeys], default=0, allow_help=False)
792
+ self.fig_theme = tkeys[idx]
793
+ fkeys = list(figs.FORMATS)
794
+ idx = ui.menu("File format for the figures:",
795
+ list(figs.FORMATS.values()), default=0,
796
+ allow_help=False)
797
+ self.fig_formats = fkeys[idx]
798
+ self.log.add(f"Figures: {', '.join(chosen) or 'none'}; theme "
799
+ f"{self.fig_theme}; format {self.fig_formats}")
800
+ return chosen
801
+
802
+ # ------------------------------------------------------------------ #
803
+ # Phase 16: Confirm
804
+ # ------------------------------------------------------------------ #
805
+
806
+ def confirm(self, classifiers, validation, metrics, figures) -> bool:
807
+ ui.header("Step 6 of 6: Review your choices")
808
+ ui.note(f"Dataset: {self.source_name}")
809
+ ui.note(f"Target: {self.target_display or self.target}")
810
+ ui.note("Classifiers: " +
811
+ ", ".join(self.registry[k].name for k in classifiers))
812
+ if self.knn_distance:
813
+ ui.note(f"KNN: {DISTANCES[self.knn_distance][0]}, "
814
+ f"k={self.knn_k}")
815
+ ui.note(f"Validation: {VALIDATION[validation]}")
816
+ if self.selection:
817
+ ui.note("Final score: " + SELECTION[self.selection])
818
+ ui.note("Metrics: " + ", ".join(METRICS[m] for m in metrics))
819
+ ui.note("Figures: " +
820
+ (", ".join(figs.FIGURES[f] for f in figures) or "none")
821
+ + (f" ({figs.THEMES[self.fig_theme].name.lower()}, "
822
+ f"{self.fig_formats.upper().replace('+', ' + ')})"
823
+ if figures else ""))
824
+ return ui.ask_yes_no("Start the analysis now?", default=True)
825
+
826
+ # ------------------------------------------------------------------ #
827
+ # Phase 17-19: Execute
828
+ # ------------------------------------------------------------------ #
829
+
830
+ def execute(self, X, y, class_names, classifiers, validation, metrics,
831
+ figures) -> None:
832
+ cfg = self.prep_cfg
833
+
834
+ # Decide scaling separately for each classifier.
835
+ # None, "standard" (mean 0, SD 1) or "minmax" (0-1), per classifier.
836
+ self._scale_for = {k: rec.scaling_for(k, self.knn_distance,
837
+ cfg.scale, cfg.scale_method)
838
+ for k in classifiers}
839
+ label = {"standard": "standardised", "minmax": "scaled to 0-1"}
840
+ scaled = [f"{self.registry[k].name} ({label[m]})"
841
+ for k, m in self._scale_for.items() if m]
842
+ if scaled:
843
+ ui.note("Scaling (learned on training data only) used for: "
844
+ + ", ".join(scaled))
845
+ self.log.add("Scaling per classifier: " + ", ".join(
846
+ f"{k}={v or 'none'}" for k, v in self._scale_for.items()))
847
+ if self._scale_for.get("knn") == "minmax" and cfg.scale is None:
848
+ ui.note("KNN with the Hassanat distance uses columns scaled to "
849
+ "0-1, which worked best in EasyClassifier's "
850
+ "benchmarks.")
851
+
852
+ # Report which form of the Hassanat formula applies to this data.
853
+ if "knn" in classifiers and self.knn_distance == "hassanat":
854
+ Xt = pp.describe_transformed(X, cfg, self._scale_for["knn"])
855
+ self.hassanat_form = hassanat_form(Xt)
856
+ ui.note("Hassanat distance: " + self.hassanat_form["text"])
857
+ self.log.add("Hassanat distance: " + self.hassanat_form["text"])
858
+
859
+ ui.header("Running analysis")
860
+ y = np.asarray(y)
861
+ final = None
862
+ self._validation, self._metrics = validation, metrics
863
+ self._classifiers = classifiers
864
+ self._class_names = class_names
865
+
866
+ # Rows used to compare classifiers (all rows, unless a final test
867
+ # set is kept aside first).
868
+ if self.selection == "final_test":
869
+ dev, test = split_final_test(X, y)
870
+ self._dev, self._test = dev, test
871
+ X_cmp, y_cmp = X.iloc[dev], y[dev]
872
+ ui.note(f"{len(test)} rows (20%) set aside as a final test set; "
873
+ f"classifiers are compared on the other {len(dev)} rows.")
874
+ self.log.add(f"Final test set: {len(test)} rows; "
875
+ f"development: {len(dev)} rows")
876
+ else:
877
+ X_cmp, y_cmp = X, y
878
+
879
+ results = []
880
+ for key in classifiers:
881
+ spec = self.registry[key]
882
+ ui.info(f"Training {spec.name} ...")
883
+ try:
884
+ result = evaluate(self._pipeline_spec(key, X), X_cmp, y_cmp,
885
+ class_names, validation, metrics)
886
+ results.append(result)
887
+ self.log.add(f"{spec.name}: " + ", ".join(
888
+ f"{METRICS[m]}={result.metrics[m]:.4f}"
889
+ for m in metrics if m in result.metrics)
890
+ + f", {SELECTION_LABEL}={result.selection_score:.4f}")
891
+ except Exception as exc: # noqa: BLE001
892
+ ui.error(f"{spec.name} failed: {exc}")
893
+ self.log.add(f"{spec.name} FAILED: {exc}")
894
+
895
+ if not results:
896
+ ui.error("No models could be trained. Please check your data.")
897
+ return
898
+
899
+ results.sort(key=lambda r: r.primary_score, reverse=True)
900
+ best = results[0]
901
+ self.comparison_text = interpret_comparison(results)
902
+ if self.comparison_text:
903
+ self.log.add("Comparison: " + self.comparison_text)
904
+ self.log.add(f"Selected: {best.classifier_name} "
905
+ f"(highest {SELECTION_LABEL.lower()})")
906
+
907
+ if self.selection and len(results) > 1:
908
+ try:
909
+ if self.selection == "final_test":
910
+ ui.info(f"Scoring {best.classifier_name} once on the "
911
+ "untouched test set ...")
912
+ final = evaluate_on_test(
913
+ self._pipeline_spec(best.classifier_key, X),
914
+ X.iloc[dev], y[dev], X.iloc[test], y[test],
915
+ class_names, metrics)
916
+ final.note = (f"{best.classifier_name}, trained on "
917
+ f"{len(dev)} rows, tested once on "
918
+ f"{len(test)} unseen rows")
919
+ else:
920
+ ui.info("Running nested cross-validation (repeats the "
921
+ "whole comparison inside each of 5 folds) ...")
922
+ ok = [r.classifier_key for r in results]
923
+ final = nested_cv(
924
+ [self._pipeline_spec(k, X) for k in ok], X, y,
925
+ class_names, validation, metrics,
926
+ progress=lambda f, n: ui.note(
927
+ f"fold {f}: {n} selected"))
928
+ self.log.add("Final estimate: " + ", ".join(
929
+ f"{METRICS[m]}={final.metrics[m]:.4f}"
930
+ for m in metrics if m in final.metrics)
931
+ + f" ({final.note})")
932
+ except Exception as exc: # noqa: BLE001
933
+ ui.error(f"Final scoring failed: {exc}")
934
+ self.log.add(f"Final scoring FAILED: {exc}")
935
+ final = None
936
+
937
+ self.show_results(best, results, metrics, final)
938
+ self.save_outputs(best, results, y, figures, X, class_names, final)
939
+
940
+ @staticmethod
941
+ def _show_metrics(result, metrics) -> None:
942
+ for m in metrics:
943
+ if m in result.metrics:
944
+ val = result.metrics[m]
945
+ shown = (f"{val*100:.2f}%"
946
+ if m in ("accuracy", "precision", "recall", "f1",
947
+ "balanced_accuracy", "specificity")
948
+ else f"{val:.4f}")
949
+ ui.note(f"{METRICS[m]:<18} {shown}")
950
+
951
+ def show_results(self, best, results, metrics, final=None) -> None:
952
+ ui.header("Results")
953
+ ui.info(f"Predicting: {self.target_display or self.target}")
954
+ ui.blank()
955
+ if final is None:
956
+ ui.info(f"Best classifier: {best.classifier_name}")
957
+ ui.blank()
958
+ self._show_metrics(best, metrics)
959
+ else:
960
+ ui.info(f"Selected classifier: {best.classifier_name}")
961
+ ui.blank()
962
+ if self.selection == "final_test":
963
+ ui.info("Final score on the untouched 20% test set - "
964
+ "these are the numbers to report:")
965
+ else:
966
+ ui.info("Nested cross-validation estimate - these are the "
967
+ "numbers to report:")
968
+ self._show_metrics(final, metrics)
969
+ ui.note(final.note)
970
+ if len(results) > 1:
971
+ ui.blank()
972
+ ui.info(f"Comparison of all classifiers ({SELECTION_LABEL.lower()}"
973
+ ", ranked)" + (" - used only to choose the winner; "
974
+ "slightly optimistic, do not report as "
975
+ "the final result:" if final else ":"))
976
+ for r in results:
977
+ score = r.primary_score
978
+ ui.note(f"{r.classifier_name:<34} {score*100:6.2f}%")
979
+ if self.comparison_text:
980
+ ui.blank()
981
+ ui.info(self.comparison_text)
982
+ if self.hassanat_form:
983
+ ui.blank()
984
+ ui.info("Hassanat distance - " + self.hassanat_form["text"])
985
+ ui.info("Please cite:")
986
+ for c in HASSANAT_CITATIONS:
987
+ ui.note(c)
988
+
989
+ def save_outputs(self, best, results, y, figures, X, class_names,
990
+ final=None) -> None:
991
+ # A new folder for every run, so earlier results are never
992
+ # overwritten: Results/<data file>_<date>_<time>/
993
+ stamp = _dt.datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
994
+ safe = "".join(ch if ch.isalnum() or ch in "-_" else "_"
995
+ for ch in getattr(self, "source_stem", "data"))[:40]
996
+ out_dir = rep.ensure_dir(os.path.join(os.getcwd(), "Results",
997
+ f"{safe}_{stamp}"))
998
+ fig_dir = rep.ensure_dir(os.path.join(out_dir, "figures"))
999
+ saved: List[str] = []
1000
+
1001
+ table = list(results)
1002
+ if final is not None:
1003
+ label = ("FINAL - untouched 20% test set"
1004
+ if self.selection == "final_test"
1005
+ else "FINAL - nested cross-validation")
1006
+ table = [dataclasses.replace(
1007
+ final, classifier_name=f"{label} ({best.classifier_name})")
1008
+ ] + table
1009
+ saved.append(rep.save_summary_csv(
1010
+ table, os.path.join(out_dir, "summary.csv")))
1011
+ xl = rep.save_excel(table, os.path.join(out_dir, "results.xlsx"))
1012
+ if xl:
1013
+ saved.append(xl)
1014
+
1015
+ # Figures and predictions come from the honest estimate when present.
1016
+ shown = final if final is not None else best
1017
+
1018
+ # Which columns mattered (measured on held-out rows).
1019
+ importance = None
1020
+ if {"feature_importance", "columns_by_class"} & set(figures):
1021
+ ui.info("Measuring which columns matter ...")
1022
+ try:
1023
+ importance = honest_importance(
1024
+ self._pipeline_spec(best.classifier_key, X), X, y,
1025
+ self._validation, self._dev, self._test)
1026
+ saved.append(rep.save_feature_importance(
1027
+ importance, os.path.join(out_dir,
1028
+ "feature_importance.csv")))
1029
+ top = ", ".join(str(c) for c in importance["column"].head(3))
1030
+ ui.note(f"Most important columns: {top}")
1031
+ self.log.add("Permutation importance (top 5): " + ", ".join(
1032
+ f"{r.column}={r.importance:.4f}"
1033
+ for r in importance.head(5).itertuples()))
1034
+ except Exception as exc: # noqa: BLE001
1035
+ ui.warn(f"Could not measure column importance: {exc}")
1036
+ self.log.add(f"Importance FAILED: {exc}")
1037
+ importance = None
1038
+
1039
+ # Would more data help? (development rows only, if a final test set
1040
+ # was kept aside, so the test rows stay untouched.)
1041
+ curve = None
1042
+ if "learning_curve" in figures:
1043
+ ui.info("Computing the learning curve ...")
1044
+ rows = self._dev if self._dev is not None else np.arange(len(y))
1045
+ try:
1046
+ curve = learning_curve_data(
1047
+ self._pipeline_spec(best.classifier_key, X),
1048
+ X.iloc[rows], y[rows])
1049
+ if curve is not None:
1050
+ self.learning_text = interpret_learning_curve(curve)
1051
+ ui.note(self.learning_text)
1052
+ self.log.add("Learning curve: " + self.learning_text)
1053
+ else:
1054
+ ui.note("Learning curve skipped: too few rows (it needs "
1055
+ "at least 20, with 2 or more in every class).")
1056
+ self.log.add("Learning curve skipped: too few rows")
1057
+ except Exception as exc: # noqa: BLE001
1058
+ ui.warn(f"Could not compute the learning curve: {exc}")
1059
+ self.log.add(f"Learning curve FAILED: {exc}")
1060
+
1061
+ fig_files = self._draw_figures(figures, fig_dir, out_dir, best,
1062
+ results, final, shown, X, y,
1063
+ importance, curve)
1064
+ for paths in figs_saved(fig_files, out_dir):
1065
+ saved.append(paths)
1066
+
1067
+ saved.append(rep.save_predictions(
1068
+ shown, os.path.join(out_dir, "predictions.csv")))
1069
+
1070
+ # Saved model = full pipeline (cleaning + encoding + scaling + model),
1071
+ # so it can be applied directly to new raw data with the same columns.
1072
+ model = fit_final_model(self._pipeline_spec(best.classifier_key, X),
1073
+ X, y)
1074
+ self.log.add(f"Saved model: {best.classifier_name}, refitted on all "
1075
+ f"{len(y)} rows")
1076
+ saved.append(rep.save_model(
1077
+ model, os.path.join(out_dir, "trained_model.pkl")))
1078
+
1079
+ saved.append(self.save_citations(os.path.join(out_dir,
1080
+ "citations.txt")))
1081
+
1082
+ # Report: report.tex, compiled to report.pdf if LaTeX is installed.
1083
+ ui.info("Writing the report ...")
1084
+ try:
1085
+ ctx = self._report_context(best, results, final, X, y,
1086
+ fig_files, importance)
1087
+ tex_path, pdf_path, msg = write_and_compile(ctx, out_dir)
1088
+ saved.append(tex_path)
1089
+ if pdf_path:
1090
+ saved.append(pdf_path)
1091
+ ui.success("Report: " + msg)
1092
+ else:
1093
+ ui.warn(msg)
1094
+ self.log.add("Report: " + msg)
1095
+ except Exception as exc: # noqa: BLE001
1096
+ ui.warn(f"Could not write the report: {exc}")
1097
+ self.log.add(f"Report FAILED: {exc}")
1098
+
1099
+ saved.append(self.log.save(os.path.join(out_dir, "log.txt")))
1100
+
1101
+ ui.header("Saved output")
1102
+ ui.success("All results were saved in this folder:")
1103
+ print(f" {out_dir}")
1104
+ ui.blank()
1105
+ n_fig = sum(1 for s in saved
1106
+ if os.path.dirname(os.path.relpath(s, out_dir)))
1107
+ for s in saved:
1108
+ rel = os.path.relpath(s, out_dir)
1109
+ if not os.path.dirname(rel):
1110
+ ui.note(rel)
1111
+ if n_fig:
1112
+ ui.note(f"figures{os.sep} ({n_fig} files)")
1113
+ ui.blank()
1114
+ main_file = ("report.pdf" if any(s.endswith("report.pdf")
1115
+ for s in saved) else "report.tex")
1116
+ ui.info(f"Finished. Start with {main_file}: it explains the results "
1117
+ "in plain words and includes a Methods paragraph you can "
1118
+ "adapt for a paper. Run EasyClassifier again at any time; "
1119
+ "each run gets its own new folder.")
1120
+
1121
+ def _draw_figures(self, figures, fig_dir, out_dir, best, results, final,
1122
+ shown, X, y, importance, curve) -> dict:
1123
+ """Draw the chosen figures. One failing figure never stops the run.
1124
+
1125
+ Returns {figure key: {format: path relative to out_dir}}.
1126
+ """
1127
+ fm = figs.FigureMaker(fig_dir, theme=self.fig_theme,
1128
+ formats=self.fig_formats,
1129
+ class_names=self._class_names)
1130
+ top_cols = (list(importance["column"]) if importance is not None
1131
+ else list(X.columns))
1132
+ final_label = ("Final score (untouched test set)"
1133
+ if self.selection == "final_test"
1134
+ else "Final score (nested CV)")
1135
+ jobs = {
1136
+ "class_distribution": lambda: fm.class_distribution(y),
1137
+ "missing_values": lambda: fm.missing_values(self.df_loaded),
1138
+ "correlation": lambda: fm.correlation(X, prefer=top_cols),
1139
+ "comparison": lambda: (fm.comparison(
1140
+ results, best.classifier_name,
1141
+ final.selection_score if final is not None else None,
1142
+ final_label) if len(results) > 1 else None),
1143
+ "confusion_matrix": lambda: fm.confusion(shown),
1144
+ "roc_curve": lambda: fm.roc(shown),
1145
+ "pr_curve": lambda: fm.pr(shown),
1146
+ "feature_importance": lambda: (fm.importance(importance)
1147
+ if importance is not None
1148
+ else None),
1149
+ "columns_by_class": lambda: (fm.columns_by_class(
1150
+ X, y, top_cols) if importance is not None else None),
1151
+ "learning_curve": lambda: (fm.learning_curve(curve)
1152
+ if curve is not None else None),
1153
+ }
1154
+ if figures:
1155
+ ui.info("Drawing figures ...")
1156
+ for key in figures:
1157
+ try:
1158
+ jobs[key]()
1159
+ except Exception as exc: # noqa: BLE001
1160
+ ui.warn(f"Figure '{figs.FIGURES[key]}' could not be drawn: "
1161
+ f"{exc}")
1162
+ self.log.add(f"Figure {key} FAILED: {exc}")
1163
+ return {k: {ext: os.path.relpath(p, out_dir) for ext, p in v.items()}
1164
+ for k, v in fm.files.items()}
1165
+
1166
+ def _report_context(self, best, results, final, X, y, fig_files,
1167
+ importance) -> ReportContext:
1168
+ names = {k: self.registry[k].name for k in self._classifiers}
1169
+ y_arr = np.asarray(y)
1170
+ counts = {name: int((y_arr == i).sum())
1171
+ for i, name in enumerate(self._class_names)}
1172
+ knn_text = ""
1173
+ if self.knn_distance:
1174
+ knn_text = (f"K-nearest neighbours used k = {self.knn_k} and the "
1175
+ f"{DISTANCES[self.knn_distance][0].split(' (')[0]}")
1176
+ if self.hassanat_form:
1177
+ form = ("the signed form of the formula, for negative values"
1178
+ if self.hassanat_form["form"] == "signed"
1179
+ else "the standard form of the formula (all values "
1180
+ "non-negative)")
1181
+ knn_text += f", applied in {form}"
1182
+ return ReportContext(
1183
+ software_citation=CITATION,
1184
+ version=__version__,
1185
+ dataset=self.source_name,
1186
+ rows_loaded=self.rows_loaded,
1187
+ cols_loaded=self.cols_loaded,
1188
+ target=self.target,
1189
+ target_grouping=self.target_grouping,
1190
+ class_counts={str(k): int(v) for k, v in counts.items()},
1191
+ rows_used=len(y),
1192
+ predictor_columns=[str(c) for c in X.columns],
1193
+ left_out=self.left_out,
1194
+ data_notes=self.data_notes,
1195
+ impute=self.impute_used,
1196
+ encoding=self.encoding_used,
1197
+ scaled={m: [names[k] for k, v in self._scale_for.items()
1198
+ if v == m] for m in ("standard", "minmax")
1199
+ if m in self._scale_for.values()},
1200
+ classifiers=[names[k] for k in self._classifiers],
1201
+ knn_text=knn_text,
1202
+ hassanat_used=bool(self.hassanat_form),
1203
+ validation=self._validation,
1204
+ selection=self.selection if final is not None else None,
1205
+ n_dev=len(self._dev) if self._dev is not None else 0,
1206
+ n_test=len(self._test) if self._test is not None else 0,
1207
+ best_name=best.classifier_name,
1208
+ final=final,
1209
+ results=results,
1210
+ metrics=self._metrics,
1211
+ figures=fig_files,
1212
+ comparison_text=self.comparison_text,
1213
+ learning_text=self.learning_text,
1214
+ importance=importance,
1215
+ )
1216
+
1217
+ def _pipeline_spec(self, key: str, X):
1218
+ """Classifier spec whose factory builds preprocessing + model."""
1219
+ spec = self.registry[key]
1220
+ cfg, scale = self.prep_cfg, self._scale_for[key]
1221
+ return dataclasses.replace(
1222
+ spec,
1223
+ factory=lambda: pp.build_pipeline(X, cfg, spec.factory(), scale),
1224
+ )
1225
+
1226
+ def save_citations(self, path: str) -> str:
1227
+ lines = ["Please cite the following if you publish these results:",
1228
+ "", "Software:", " " + CITATION]
1229
+ if self.hassanat_form:
1230
+ lines += ["", "Hassanat distance (used by KNN):"]
1231
+ lines += [" " + c for c in HASSANAT_CITATIONS]
1232
+ lines += ["", "Formula form applied: " + self.hassanat_form["text"]]
1233
+ with open(path, "w", encoding="utf-8") as fh:
1234
+ fh.write("\n".join(lines) + "\n")
1235
+ return path
1236
+
1237
+ # ------------------------------------------------------------------ #
1238
+ # Help resolution
1239
+ # ------------------------------------------------------------------ #
1240
+
1241
+ def _resolve_help(self, choice: int, help_keys: List[str]) -> int:
1242
+ """If a menu returned a help sentinel (negative), show help and re-ask.
1243
+
1244
+ Because ``ui.menu`` returns a negative sentinel for ?N, we loop here
1245
+ until the user makes a real choice. The caller passes the same prompt
1246
+ again implicitly by re-invoking; to keep it simple we just show the
1247
+ explanation and ask the user to choose again via a follow-up prompt.
1248
+ """
1249
+ while choice < 0:
1250
+ key_index = (-choice) - 1
1251
+ key = help_keys[key_index] if key_index < len(help_keys) else ""
1252
+ ui.blank()
1253
+ ui.info(explain(key))
1254
+ ui.pause()
1255
+ # Ask again with a minimal numeric prompt.
1256
+ raw = ui.ask_text("Enter your choice number")
1257
+ if raw.isdigit():
1258
+ choice = int(raw) - 1
1259
+ else:
1260
+ choice = -1
1261
+ return choice
1262
+
1263
+
1264
+ def pp_remove(df):
1265
+ """Thin wrapper kept for readability in the wizard."""
1266
+ return pp.remove_duplicates(df)