OSCR

Benchmarking Multimodal Workload Classification: Effects of Modality, Validation Protocol, and Segmentation Contrast on an Open Graded-Arithmetic Dataset.

Code ↔ Paper

21 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 21 matches
  1. [1] § 3. Results › 3.3. The 6 s Baseline vs. 3 s Overlap Run › Statistical Significance Procedures (Stage 7) ↔ analysis_pipeline/stage7_significance.py, lines 1–67 · score 1.00 · outer fold scores, fold confusion matrices, vectors reconstructed, machine readable, sided Wilcoxon signed, task heart rate
  2. [2] § 2. Materials and Methods › 2.5. Feature Extraction ↔ analysis_pipeline/stage4_extract_features.py, lines 1–36 · score 0.96 · 13–30 Hz, 30–40 Hz, 8–13 Hz, 1–4 Hz, 4–8 Hz, occipital
  3. [3] § 2. Materials and Methods › 2.6. Machine Learning ↔ analysis_pipeline/stage6_train_classic_ml.py, lines 641–699 · score 0.84 · input_size, ReLU, nn.GRU, nn.LSTM, Linear, head
  4. [4] § 2. Materials and Methods › 2.6. Machine Learning ↔ analysis_pipeline/stage6_train_classic_ml.py, lines 902–961 · score 0.83 · min_samples_leaf, max depth, Decision tree, class weights, Gini, entropy
  5. [5] § 2. Materials and Methods › 2.6. Machine Learning ↔ analysis_pipeline/stage6_train_classic_ml.py, lines 804–858 · score 0.82 · AdamW, CrossEntropyLoss, weight decay, batch, clipping, zero
  6. [6] § 2. Materials and Methods › 2.6. Machine Learning ↔ analysis_pipeline/stage6_train_classic_ml.py, lines 902–961 · score 0.79 · stage6 train classic, logistic regression, RBF, SVM, ml, kernel
  7. [7] § Appendix A. Feature Definitions ↔ analysis_pipeline/stage4_extract_features.py, lines 358–491 · score 0.79 · frontal alpha asymmetry, Fz theta, band power, ROI, ratios, absolute
  8. [8] § 2. Materials and Methods › 2.3. Preprocessing ↔ analysis_pipeline/stage2_preprocess.py, lines 670–738 · score 0.78 · Butterworth IIR, Bad channels, bandpass filter, notch, ICA, component
  9. [9] § 3. Results › 3.3. The 6 s Baseline vs. 3 s Overlap Run ↔ analysis_pipeline/stage7_significance.py, lines 73–83 · score 0.73 · baseline omit easiest, baseline omit hardest, baseline low high, class scenarios, bins, fused
  10. [10] § 2. Materials and Methods › 2.3. Preprocessing ↔ analysis_pipeline/stage1_qc_summary.py, lines 276–309 · score 0.72 · find_peaks, peak detection, filtfilt, prominence, distance, bandpass
  11. [11] § 2. Materials and Methods › 2.6. Machine Learning ↔ analysis_pipeline/stage6_train_classic_ml.py, lines 641–699 · score 0.63 · nn.GRU, nn.LSTM, PyTorch, head, hidden, layer
  12. [12] § 2. Materials and Methods › 2.2. Quality Checks ↔ analysis_pipeline/stage7_significance.py, lines 1–67 · score 0.62 · paired Wilcoxon signed, heart rate, rank, Quality, baseline
  13. [13] § 3. Results › 3.2. Current 6 s Baseline Run Findings › 3.2.3. Modality and Protocol Comparisons ↔ analysis_pipeline/stage7_significance.py, lines 73–83 · score 0.61 · omit hardest, baseline low high, pooled stratified, holdout, fused, protocol
  14. [14] § 3. Results › 3.3. The 6 s Baseline vs. 3 s Overlap Run ↔ analysis_pipeline/stage7_significance.py, lines 604–696 · score 0.60 · winners changed, matched cells, macro F1, balanced accuracy, overlap, baseline
  15. [15] § 2. Materials and Methods › 2.6. Machine Learning ↔ analysis_pipeline/stage6_train_classic_ml.py, lines 1158–1200 · score 0.56 · Logistic regression, class weights, solver
  16. [16] § 2. Materials and Methods › 2.2. Quality Checks ↔ analysis_pipeline/stage1_qc_summary.py, lines 1845–1925 · score 0.54 · carried forward, low pupil, rejected, downstream, confidence, QC
  17. [17] § 2. Materials and Methods › 2.3. Preprocessing ↔ analysis_pipeline/stage2_preprocess.py, lines 189–199 · score 0.53 · find_peaks, prominence, distance, median, signals, Preprocessing
  18. [18] § 3. Results › 3.2. Current 6 s Baseline Run Findings › 3.2.1. Data Retention and Filtering ↔ analysis_pipeline/stage6_build_publication_report.py, lines 626–689 · score 0.52 · EEG windows, pupil windows, ECG windows, rejection, multimodal, epochs
  19. [19] § 2. Materials and Methods › 2.2. Quality Checks ↔ analysis_pipeline/stage1_qc_summary.py, lines 370–464 · score 0.52 · detected beats, duration, bpm, outside, median, detection
  20. [20] § 2. Materials and Methods › 2.5. Feature Extraction ↔ analysis_pipeline/stage5_build_fused_table.py, lines 496–559 · score 0.52 · stage5 build fused, metadata, fusion, modality
  21. [21] § 2. Materials and Methods › 2.1. Data Source and Pipeline Overview ↔ analysis_pipeline/stage6_build_manuscript_draft.py, lines 76–108 · score 0.51 · validation protocols, unimodal features, classification, matrices, metrics, multimodal

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 2,970 lines · 123 KB · CC0-1.0 · 6 matches

  1. from __future__ import annotations
  2. import argparse
  3. import csv
  4. from functools import partial
  5. import hashlib
  6. import itertools
  7. import json
  8. import math
  9. import os
  10. import random
  11. import subprocess
  12. import time
  13. import warnings
  14. from collections import Counter, defaultdict
  15. from datetime import datetime, timezone
  16. from pathlib import Path
  17. from statistics import mean, pstdev
  18. from typing import Any
  19. import matplotlib
  20. import numpy as np
  21. import pandas as pd
  22. from eeg_condition_psd import build_condition_psd_report
  23. from sklearn.base import BaseEstimator, TransformerMixin, clone
  24. from sklearn.ensemble import ExtraTreesClassifier, RandomForestClassifier
  25. from sklearn.exceptions import ConvergenceWarning
  26. from sklearn.feature_selection import SelectFromModel, SelectPercentile, VarianceThreshold, f_classif, mutual_info_classif
  27. from sklearn.impute import SimpleImputer
  28. from sklearn.linear_model import LogisticRegression
  29. from sklearn.metrics import accuracy_score, balanced_accuracy_score, confusion_matrix, f1_score
  30. from sklearn.model_selection import GroupKFold, StratifiedKFold, StratifiedShuffleSplit
  31. from sklearn.naive_bayes import GaussianNB
  32. from sklearn.neighbors import KNeighborsClassifier
  33. from sklearn.neural_network import MLPClassifier
  34. from sklearn.pipeline import Pipeline
  35. from sklearn.preprocessing import RobustScaler
  36. from sklearn.svm import SVC
  37. from sklearn.tree import DecisionTreeClassifier
  38. matplotlib.use("Agg")
  39. import matplotlib.pyplot as plt
  40. # Avoid torch dynamo/onnx import-time dependency issues in this pipeline runtime.
  41. os.environ.setdefault("TORCH_DISABLE_DYNAMO", "1")
  42. os.environ.setdefault("TORCHDYNAMO_DISABLE", "1")
  43. os.environ.setdefault("PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION", "python")
  44. try:
  45. import torch
  46. import torch.nn as nn
  47. from torch.utils.data import DataLoader, TensorDataset
  48. TORCH_AVAILABLE = True
  49. except Exception: # noqa: BLE001
  50. torch = None # type: ignore[assignment]
  51. nn = None # type: ignore[assignment]
  52. DataLoader = None # type: ignore[assignment]
  53. TensorDataset = None # type: ignore[assignment]
  54. TORCH_AVAILABLE = False
  55. DEEP_MODEL_NAMES = (
  56. "lstm1d",
  57. "gru1d",
  58. "cnn1d",
  59. "transformer",
  60. "bilstm1d",
  61. "bigru1d",
  62. "cnn1d_deep",
  63. "transformer_xl",
  64. )
  65. NON_FEATURE_COLUMNS = {
  66. "fused_row_id",
  67. "ml_row_id",
  68. "modality",
  69. "participant_id",
  70. "analysis_included",
  71. "trial_id",
  72. "epoch_id",
  73. "block",
  74. "is_tutorial",
  75. "difficulty_range",
  76. "difficulty_bin",
  77. "difficulty_bin_index",
  78. "target_label",
  79. "response_accuracy",
  80. "outcome",
  81. "window",
  82. "segment_index",
  83. "is_subwindow",
  84. "epoch_start_s",
  85. "epoch_end_s",
  86. "epoch_duration_s",
  87. "dropout_policy",
  88. "dropout_threshold_value",
  89. "dropout_keep",
  90. "dropped_samples_trial",
  91. "split_group",
  92. "baseline_start_s",
  93. "baseline_end_s",
  94. "preproc_version",
  95. "ml_keep",
  96. "modalities_selected",
  97. "modalities_present",
  98. "n_modalities_present",
  99. "_target_source_label",
  100. "_target_label_effective",
  101. }
  102. def _analysis_root() -> Path:
  103. return Path(__file__).resolve().parent
  104. def _reports_dir() -> Path:
  105. return _analysis_root() / "reports"
  106. def _features_dir() -> Path:
  107. return _analysis_root() / "features"
  108. def _models_root() -> Path:
  109. return _analysis_root() / "models"
  110. def _default_split_manifest() -> Path:
  111. return _features_dir() / "split_manifest.json"
  112. def _default_results_json() -> Path:
  113. return _reports_dir() / "ml_results.json"
  114. def _default_summary_md() -> Path:
  115. return _reports_dir() / "ml_summary.md"
  116. def _default_live_confusion_png_dir() -> Path:
  117. return _reports_dir() / "confusion_pngs_live"
  118. def _default_eeg_condition_psd_dir() -> Path:
  119. return _reports_dir() / "eeg_condition_psd"
  120. def _slug(text: str) -> str:
  121. out = "".join(ch.lower() if ch.isalnum() else "_" for ch in str(text).strip())
  122. out = "_".join(part for part in out.split("_") if part)
  123. return out or "default"
  124. def _resolve_bids_root(bids_root_arg: str) -> Path:
  125. direct = Path(bids_root_arg).expanduser()
  126. if direct.is_absolute():
  127. return direct.resolve()
  128. from_cwd = (Path.cwd() / direct).resolve()
  129. if from_cwd.exists():
  130. return from_cwd
  131. repo_root = Path(__file__).resolve().parent.parent
  132. return (repo_root / direct).resolve()
  133. def _translate_wsl_path(path_text: str) -> Path | None:
  134. normalized = str(path_text).strip().replace("\\", "/")
  135. if not normalized.startswith("/mnt/") or len(normalized) < 8:
  136. return None
  137. drive = normalized[5]
  138. if not drive.isalpha() or normalized[6] != "/":
  139. return None
  140. remainder = normalized[7:]
  141. return Path(f"{drive.upper()}:/{remainder}")
  142. def _translate_windows_drive_path(path_text: str) -> Path | None:
  143. normalized = str(path_text).strip().replace("\\", "/")
  144. if len(normalized) < 4 or normalized[1] != ":" or normalized[2] != "/":
  145. return None
  146. drive = normalized[0]
  147. if not drive.isalpha():
  148. return None
  149. remainder = normalized[3:]
  150. return Path(f"/mnt/{drive.lower()}/{remainder}")
  151. def _resolve_dataset_path(path_text: str, split_manifest_path: Path) -> Path:
  152. direct = Path(path_text).expanduser()
  153. if direct.is_absolute() and direct.exists():
  154. return direct.resolve()
  155. translated_wsl = _translate_wsl_path(path_text)
  156. if translated_wsl is not None and translated_wsl.exists():
  157. return translated_wsl.resolve()
  158. translated_windows = _translate_windows_drive_path(path_text)
  159. if translated_windows is not None and translated_windows.exists():
  160. return translated_windows.resolve()
  161. from_manifest = (split_manifest_path.parent / direct).resolve()
  162. if from_manifest.exists():
  163. return from_manifest
  164. from_cwd = (Path.cwd() / direct).resolve()
  165. if from_cwd.exists():
  166. return from_cwd
  167. repo_root = Path(__file__).resolve().parent.parent
  168. repo_candidate = (repo_root / direct).resolve()
  169. if repo_candidate.exists():
  170. return repo_candidate
  171. if direct.is_absolute():
  172. return direct.resolve()
  173. if translated_wsl is not None:
  174. return translated_wsl.resolve()
  175. return repo_candidate
  176. def _stable_seed_from_text(text: str, base_seed: int) -> int:
  177. digest = hashlib.sha256(text.encode("utf-8")).hexdigest()
  178. return (base_seed + int(digest[:8], 16)) % (2**31 - 1)
  179. def _warning_key(record: warnings.WarningMessage) -> str:
  180. message = str(record.message).strip().splitlines()[0]
  181. if len(message) > 160:
  182. message = message[:157] + "..."
  183. return f"{record.category.__name__}: {message}"
  184. def _capture_warnings(records: list[warnings.WarningMessage], warning_counter: Counter[str] | None) -> None:
  185. if warning_counter is None:
  186. return
  187. for record in records:
  188. warning_counter[_warning_key(record)] += 1
  189. def _as_float(value: Any) -> float | None:
  190. if value is None:
  191. return None
  192. if isinstance(value, float):
  193. if math.isfinite(value):
  194. return value
  195. return None
  196. text = str(value).strip()
  197. if not text or text.lower() == "n/a":
  198. return None
  199. try:
  200. number = float(text)
  201. except ValueError:
  202. return None
  203. if not math.isfinite(number):
  204. return None
  205. return number
  206. def _percent(value: float) -> str:
  207. return f"{100.0 * value:.2f}%"
  208. def _product_grid(space: dict[str, list[Any]]) -> list[dict[str, Any]]:
  209. if not space:
  210. return [{}]
  211. keys = list(space.keys())
  212. values = [space[k] for k in keys]
  213. out: list[dict[str, Any]] = []
  214. for combo in itertools.product(*values):
  215. out.append({k: v for k, v in zip(keys, combo)})
  216. return out
  217. def _coerce_dataset_list(manifest: dict[str, Any], datasets_arg: list[str] | None) -> list[str]:
  218. available = list(manifest.get("datasets", {}).keys())
  219. if not available:
  220. raise ValueError("No datasets listed in split manifest.")
  221. if not datasets_arg:
  222. return available
  223. picked = [d.strip() for d in datasets_arg if d.strip()]
  224. unknown = [d for d in picked if d not in available]
  225. if unknown:
  226. raise ValueError(f"Unknown datasets requested: {', '.join(sorted(set(unknown)))}")
  227. return list(dict.fromkeys(picked))
  228. def _coerce_model_list(models_arg: list[str] | None) -> list[str]:
  229. available = _available_model_names()
  230. if not models_arg:
  231. return available
  232. picked = [m.strip().lower() for m in models_arg if m.strip()]
  233. unknown = [m for m in picked if m not in available]
  234. if unknown:
  235. deep_names = set(DEEP_MODEL_NAMES)
  236. missing_torch = sorted(set([m for m in unknown if m in deep_names]))
  237. if missing_torch and not TORCH_AVAILABLE:
  238. raise ValueError(
  239. "Unknown model names (PyTorch not available for deep models): "
  240. + ", ".join(missing_torch)
  241. + ". Install PyTorch or remove these models."
  242. )
  243. raise ValueError(f"Unknown model names: {', '.join(sorted(set(unknown)))}")
  244. return list(dict.fromkeys(picked))
  245. def _available_feature_selector_names() -> list[str]:
  246. return ["none", "anova", "mutual_info", "l1", "tree"]
  247. def _coerce_feature_selector_list(selectors_arg: list[str] | None) -> list[str]:
  248. available = _available_feature_selector_names()
  249. if not selectors_arg:
  250. return ["none"]
  251. picked = [s.strip().lower() for s in selectors_arg if s.strip()]
  252. unknown = [s for s in picked if s not in available]
  253. if unknown:
  254. raise ValueError(f"Unknown feature selectors: {', '.join(sorted(set(unknown)))}")
  255. return list(dict.fromkeys(picked))
  256. def _coerce_protocol_list(protocols_arg: list[str] | None) -> list[str]:
  257. available = [
  258. "loso",
  259. "group_holdout",
  260. "within_participant",
  261. "pooled_stratified",
  262. "pooled_stratified_holdout",
  263. "pooled_stratified_one_fold",
  264. ]
  265. if not protocols_arg:
  266. return available
  267. picked = [p.strip().lower() for p in protocols_arg if p.strip()]
  268. unknown = [p for p in picked if p not in available]
  269. if unknown:
  270. raise ValueError(f"Unknown protocols: {', '.join(sorted(set(unknown)))}")
  271. return list(dict.fromkeys(picked))
  272. DROP_LABEL_TOKENS = {"drop", "__drop__", "omit", "exclude", "remove"}
  273. def _label_sort_key(label: str, fallback_order: dict[str, int]) -> tuple[float, int]:
  274. text = str(label).strip()
  275. lower = text.lower()
  276. if lower in {"baseline", "rest", "resting"} or lower.startswith("baseline"):
  277. return 10_000.0, fallback_order.get(text, 10_000)
  278. if "-" in text:
  279. left = text.split("-", 1)[0].strip()
  280. try:
  281. return float(left), fallback_order.get(text, 10_000)
  282. except ValueError:
  283. pass
  284. digit_parts: list[str] = []
  285. current = ""
  286. for ch in text:
  287. if ch.isdigit() or ch == ".":
  288. current += ch
  289. else:
  290. if current:
  291. digit_parts.append(current)
  292. current = ""
  293. if current:
  294. digit_parts.append(current)
  295. for token in digit_parts:
  296. try:
  297. return float(token), fallback_order.get(text, 10_000)
  298. except ValueError:
  299. continue
  300. if "low" in lower:
  301. return 1000.0, fallback_order.get(text, 10_000)
  302. if "mid" in lower:
  303. return 2000.0, fallback_order.get(text, 10_000)
  304. if "high" in lower:
  305. return 3000.0, fallback_order.get(text, 10_000)
  306. return 9000.0, fallback_order.get(text, 10_000)
  307. def _reorder_labels_baseline_first(labels: list[str], baseline_label: str | None) -> list[str]:
  308. if not labels:
  309. return labels
  310. baseline_targets = []
  311. if baseline_label and baseline_label.strip():
  312. baseline_targets.append(baseline_label.strip())
  313. baseline_targets.extend(
  314. [label for label in labels if str(label).strip().lower().startswith("baseline") or str(label).strip().lower() == "baseline"]
  315. )
  316. baseline_targets = list(dict.fromkeys(baseline_targets))
  317. baseline_in_labels = [label for label in labels if label in baseline_targets]
  318. if not baseline_in_labels:
  319. return labels
  320. baseline = baseline_in_labels[0]
  321. non_baseline = [label for label in labels if label != baseline]
  322. return [baseline] + non_baseline
  323. def _coerce_label_list(values: list[str] | None) -> list[str]:
  324. if not values:
  325. return []
  326. out: list[str] = []
  327. for value in values:
  328. text = str(value).strip()
  329. if not text:
  330. continue
  331. out.append(text)
  332. return list(dict.fromkeys(out))
  333. def _ordered_counts(values: list[str], preferred_order: list[str] | None = None) -> dict[str, int]:
  334. counts = Counter(str(v) for v in values if str(v).strip())
  335. out: dict[str, int] = {}
  336. if preferred_order:
  337. for label in preferred_order:
  338. key = str(label).strip()
  339. if key in counts:
  340. out[key] = int(counts.pop(key))
  341. # Preserve first-seen order for any remaining labels.
  342. for value in values:
  343. key = str(value).strip()
  344. if key in counts and key not in out:
  345. out[key] = int(counts.pop(key))
  346. # Deterministic fallback if anything remains.
  347. for key in sorted(counts.keys()):
  348. out[key] = int(counts[key])
  349. return out
  350. def _parse_merge_entry(text: str) -> tuple[str, str]:
  351. raw = str(text).strip()
  352. for sep in ("->", "=>", ":", "="):
  353. if sep in raw:
  354. left, right = raw.split(sep, 1)
  355. src = left.strip()
  356. dst = right.strip()
  357. if not src or not dst:
  358. break
  359. return src, dst
  360. raise ValueError(f"Invalid --class-merge entry '{text}'. Use formats like old->new or old=new.")
  361. def _load_merge_map_json(path: Path) -> dict[str, str]:
  362. if not path.exists():
  363. raise FileNotFoundError(path)
  364. payload = json.loads(path.read_text(encoding="utf-8"))
  365. if not isinstance(payload, dict):
  366. raise ValueError(f"--class-merge-json must contain a JSON object of old_label -> new_label: {path}")
  367. out: dict[str, str] = {}
  368. for k, v in payload.items():
  369. src = str(k).strip()
  370. dst = str(v).strip()
  371. if not src or not dst:
  372. raise ValueError(f"Invalid merge map entry in {path}: '{k}' -> '{v}'")
  373. out[src] = dst
  374. return out
  375. def _parse_class_merge_map(args: argparse.Namespace) -> dict[str, str]:
  376. merge_map: dict[str, str] = {}
  377. if args.class_merge_json:
  378. merge_map.update(_load_merge_map_json(Path(args.class_merge_json).resolve()))
  379. if args.class_merge:
  380. for entry in args.class_merge:
  381. src, dst = _parse_merge_entry(entry)
  382. merge_map[src] = dst
  383. return merge_map
  384. def _build_class_scenario(split_manifest: dict[str, Any], args: argparse.Namespace) -> dict[str, Any]:
  385. manifest_labels = [
  386. str(item.get("label", "")).strip()
  387. for item in split_manifest.get("difficulty_bins", [])
  388. if str(item.get("label", "")).strip()
  389. ]
  390. baseline_from_tutorial_label = str(args.baseline_from_tutorial_label).strip() if args.baseline_from_tutorial_label else ""
  391. if baseline_from_tutorial_label:
  392. manifest_labels = [label for label in manifest_labels if label != baseline_from_tutorial_label]
  393. manifest_labels = [baseline_from_tutorial_label] + manifest_labels
  394. manifest_labels = list(dict.fromkeys(manifest_labels))
  395. if not manifest_labels:
  396. raise ValueError("No difficulty bin labels available in split manifest for class scenario setup.")
  397. include_labels = _coerce_label_list(args.class_include_labels)
  398. drop_labels = _coerce_label_list(args.class_drop_labels)
  399. merge_map = _parse_class_merge_map(args)
  400. manifest_set = set(manifest_labels)
  401. unknown_include = sorted(set(include_labels) - manifest_set)
  402. unknown_drop = sorted(set(drop_labels) - manifest_set)
  403. unknown_merge = sorted(set(merge_map.keys()) - manifest_set)
  404. if unknown_include:
  405. raise ValueError(f"Unknown labels in --class-include-labels: {', '.join(unknown_include)}")
  406. if unknown_drop:
  407. raise ValueError(f"Unknown labels in --class-drop-labels: {', '.join(unknown_drop)}")
  408. if unknown_merge:
  409. raise ValueError(f"Unknown source labels in class merge map: {', '.join(unknown_merge)}")
  410. include_set = set(include_labels)
  411. drop_set = set(drop_labels)
  412. allowed_original_labels = [
  413. label for label in manifest_labels if (not include_set or label in include_set) and label not in drop_set
  414. ]
  415. if not allowed_original_labels:
  416. raise ValueError("Class scenario removed all labels before training. Check include/drop settings.")
  417. final_labels: list[str] = []
  418. dropped_by_merge: list[str] = []
  419. for label in allowed_original_labels:
  420. merged = str(merge_map.get(label, label)).strip()
  421. if not merged or merged.lower() in DROP_LABEL_TOKENS:
  422. dropped_by_merge.append(label)
  423. continue
  424. if merged not in final_labels:
  425. final_labels.append(merged)
  426. final_labels = _reorder_labels_baseline_first(
  427. labels=final_labels,
  428. baseline_label=baseline_from_tutorial_label or None,
  429. )
  430. if len(final_labels) < 2:
  431. raise ValueError("Class scenario must keep at least 2 target classes after merge/drop.")
  432. return {
  433. "name": (args.class_scenario_name or "default").strip() or "default",
  434. "baseline_from_tutorial_label": baseline_from_tutorial_label or None,
  435. "manifest_labels": manifest_labels,
  436. "include_labels": include_labels,
  437. "drop_labels": drop_labels,
  438. "merge_map": merge_map,
  439. "allowed_original_labels": allowed_original_labels,
  440. "dropped_by_merge": dropped_by_merge,
  441. "final_labels": final_labels,
  442. }
  443. class QuantileClipper(BaseEstimator, TransformerMixin):
  444. def __init__(self, lower_q: float = 0.01, upper_q: float = 0.99):
  445. self.lower_q = lower_q
  446. self.upper_q = upper_q
  447. self.lower_bounds_: np.ndarray | None = None
  448. self.upper_bounds_: np.ndarray | None = None
  449. def fit(self, X: Any, y: Any = None) -> "QuantileClipper":
  450. arr = np.asarray(X, dtype=np.float64)
  451. if arr.ndim != 2:
  452. raise ValueError("QuantileClipper expects 2D input.")
  453. lowers: list[float] = []
  454. uppers: list[float] = []
  455. for col_idx in range(arr.shape[1]):
  456. col = arr[:, col_idx]
  457. finite = col[np.isfinite(col)]
  458. if finite.size == 0:
  459. lowers.append(np.nan)
  460. uppers.append(np.nan)
  461. continue
  462. lowers.append(float(np.quantile(finite, self.lower_q)))
  463. uppers.append(float(np.quantile(finite, self.upper_q)))
  464. self.lower_bounds_ = np.asarray(lowers, dtype=np.float64)
  465. self.upper_bounds_ = np.asarray(uppers, dtype=np.float64)
  466. return self
  467. def transform(self, X: Any) -> np.ndarray:
  468. if self.lower_bounds_ is None or self.upper_bounds_ is None:
  469. raise RuntimeError("QuantileClipper must be fit before transform.")
  470. arr = np.asarray(X, dtype=np.float64).copy()
  471. for col_idx in range(arr.shape[1]):
  472. lo = self.lower_bounds_[col_idx]
  473. hi = self.upper_bounds_[col_idx]
  474. if not math.isfinite(lo) or not math.isfinite(hi):
  475. continue
  476. arr[:, col_idx] = np.clip(arr[:, col_idx], lo, hi)
  477. return arr
  478. if TORCH_AVAILABLE:
  479. class _TabularCNN1DNet(nn.Module):
  480. def __init__(self, input_dim: int, num_classes: int, hidden_dim: int, dropout: float):
  481. super().__init__()
  482. self.conv = nn.Sequential(
  483. nn.Conv1d(1, 16, kernel_size=5, padding=2),
  484. nn.ReLU(),
  485. nn.BatchNorm1d(16),
  486. nn.Conv1d(16, 32, kernel_size=5, padding=2),
  487. nn.ReLU(),
  488. nn.BatchNorm1d(32),
  489. nn.Conv1d(32, 32, kernel_size=3, padding=1),
  490. nn.ReLU(),
  491. nn.AdaptiveAvgPool1d(16),
  492. )
  493. self.head = nn.Sequential(
  494. nn.Flatten(),
  495. nn.Linear(32 * 16, hidden_dim),
  496. nn.ReLU(),
  497. nn.Dropout(dropout),
  498. nn.Linear(hidden_dim, num_classes),
  499. )
  500. def forward(self, x: Any) -> Any:
  501. x = x.unsqueeze(1)
  502. x = self.conv(x)
  503. return self.head(x)
  504. class _TabularTransformerNet(nn.Module):
  505. def __init__(
  506. self,
  507. input_dim: int,
  508. num_classes: int,
  509. d_model: int,
  510. n_heads: int,
  511. n_layers: int,
  512. dropout: float,
  513. ):
  514. super().__init__()
  515. self.input_proj = nn.Linear(1, d_model)
  516. self.positional = nn.Parameter(torch.zeros(1, input_dim, d_model))
  517. encoder_layer = nn.TransformerEncoderLayer(
  518. d_model=d_model,
  519. nhead=n_heads,
  520. dim_feedforward=max(64, d_model * 2),
  521. dropout=dropout,
  522. batch_first=True,
  523. activation="gelu",
  524. )
  525. self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=n_layers)
  526. self.norm = nn.LayerNorm(d_model)
  527. self.head = nn.Sequential(
  528. nn.Linear(d_model, d_model),
  529. nn.GELU(),
  530. nn.Dropout(dropout),
  531. nn.Linear(d_model, num_classes),
  532. )
  533. def forward(self, x: Any) -> Any:
  534. x = x.unsqueeze(-1)
  535. x = self.input_proj(x)
  536. x = x + self.positional[:, : x.shape[1], :]
  537. x = self.encoder(x)
  538. x = self.norm(torch.mean(x, dim=1))
  539. return self.head(x)
  540. class _TabularRNNNet(nn.Module):
  541. def __init__(
  542. self,
  543. *,
  544. input_dim: int,
  545. num_classes: int,
  546. hidden_dim: int,
  547. n_layers: int,
  548. dropout: float,
  549. rnn_type: str,
  550. bidirectional: bool = False,
  551. ):
  552. super().__init__()
  553. rnn_dropout = float(dropout) if int(n_layers) > 1 else 0.0
  554. hidden_dim_int = int(hidden_dim)
  555. bidirectional_flag = bool(bidirectional)
  556. if rnn_type == "lstm":
  557. self.rnn = nn.LSTM(
  558. input_size=1,
  559. hidden_size=hidden_dim_int,
  560. num_layers=max(1, int(n_layers)),
  561. batch_first=True,
  562. dropout=rnn_dropout,
  563. bidirectional=bidirectional_flag,
  564. )
  565. elif rnn_type == "gru":
  566. self.rnn = nn.GRU(
  567. input_size=1,
  568. hidden_size=hidden_dim_int,
  569. num_layers=max(1, int(n_layers)),
  570. batch_first=True,
  571. dropout=rnn_dropout,
  572. bidirectional=bidirectional_flag,
  573. )
  574. else:
  575. raise ValueError(f"Unknown rnn_type: {rnn_type}")
  576. rnn_output_dim = hidden_dim_int * (2 if bidirectional_flag else 1)
  577. self.head = nn.Sequential(
  578. nn.Linear(rnn_output_dim, hidden_dim_int),
  579. nn.ReLU(),
  580. nn.Dropout(float(dropout)),
  581. nn.Linear(hidden_dim_int, int(num_classes)),
  582. )
  583. def forward(self, x: Any) -> Any:
  584. x = x.unsqueeze(-1)
  585. out, _ = self.rnn(x)
  586. last = out[:, -1, :]
  587. return self.head(last)
  588. class TorchTabularClassifier(BaseEstimator):
  589. def __init__(
  590. self,
  591. architecture: str = "cnn1d",
  592. epochs: int = 30,
  593. batch_size: int = 64,
  594. lr: float = 1e-3,
  595. weight_decay: float = 1e-4,
  596. hidden_dim: int = 64,
  597. dropout: float = 0.2,
  598. n_heads: int = 4,
  599. n_layers: int = 1,
  600. bidirectional: bool = False,
  601. grad_clip: float = 1.0,
  602. random_state: int = 42,
  603. device: str = "cpu",
  604. ):
  605. self.architecture = architecture
  606. self.epochs = epochs
  607. self.batch_size = batch_size
  608. self.lr = lr
  609. self.weight_decay = weight_decay
  610. self.hidden_dim = hidden_dim
  611. self.dropout = dropout
  612. self.n_heads = n_heads
  613. self.n_layers = n_layers
  614. self.bidirectional = bidirectional
  615. self.grad_clip = grad_clip
  616. self.random_state = random_state
  617. self.device = device
  618. self.model_: Any | None = None
  619. self.classes_: np.ndarray | None = None
  620. self.device_: Any | None = None
  621. def _resolve_device(self) -> Any:
  622. if self.device == "auto":
  623. if torch.cuda.is_available():
  624. return torch.device("cuda")
  625. return torch.device("cpu")
  626. return torch.device(self.device)
  627. def _build_model(self, input_dim: int, num_classes: int) -> Any:
  628. if self.architecture == "cnn1d":
  629. return _TabularCNN1DNet(
  630. input_dim=input_dim,
  631. num_classes=num_classes,
  632. hidden_dim=int(self.hidden_dim),
  633. dropout=float(self.dropout),
  634. )
  635. if self.architecture == "transformer":
  636. d_model = int(self.hidden_dim)
  637. n_heads = max(1, int(self.n_heads))
  638. if d_model % n_heads != 0:
  639. raise ValueError(f"hidden_dim ({d_model}) must be divisible by n_heads ({n_heads}).")
  640. return _TabularTransformerNet(
  641. input_dim=input_dim,
  642. num_classes=num_classes,
  643. d_model=d_model,
  644. n_heads=n_heads,
  645. n_layers=max(1, int(self.n_layers)),
  646. dropout=float(self.dropout),
  647. )
  648. if self.architecture in {"lstm1d", "bilstm1d"}:
  649. return _TabularRNNNet(
  650. input_dim=input_dim,
  651. num_classes=num_classes,
  652. hidden_dim=int(self.hidden_dim),
  653. n_layers=max(1, int(self.n_layers)),
  654. dropout=float(self.dropout),
  655. rnn_type="lstm",
  656. bidirectional=(self.architecture == "bilstm1d") or bool(self.bidirectional),
  657. )
  658. if self.architecture in {"gru1d", "bigru1d"}:
  659. return _TabularRNNNet(
  660. input_dim=input_dim,
  661. num_classes=num_classes,
  662. hidden_dim=int(self.hidden_dim),
  663. n_layers=max(1, int(self.n_layers)),
  664. dropout=float(self.dropout),
  665. rnn_type="gru",
  666. bidirectional=(self.architecture == "bigru1d") or bool(self.bidirectional),
  667. )
  668. raise ValueError(f"Unknown architecture: {self.architecture}")
  669. def fit(self, X: Any, y: Any) -> "TorchTabularClassifier":
  670. x_arr = np.asarray(X, dtype=np.float32)
  671. if x_arr.ndim != 2:
  672. raise ValueError("TorchTabularClassifier expects a 2D feature matrix.")
  673. y_arr = np.asarray(y).astype(str)
  674. classes = np.unique(y_arr)
  675. if classes.size < 2:
  676. raise ValueError("TorchTabularClassifier requires at least 2 classes.")
  677. self.classes_ = classes
  678. class_to_idx = {label: idx for idx, label in enumerate(self.classes_)}
  679. y_idx = np.asarray([class_to_idx[str(label)] for label in y_arr], dtype=np.int64)
  680. self.device_ = self._resolve_device()
  681. torch.manual_seed(int(self.random_state))
  682. if torch.cuda.is_available():
  683. torch.cuda.manual_seed_all(int(self.random_state))
  684. self.model_ = self._build_model(input_dim=x_arr.shape[1], num_classes=int(classes.size)).to(self.device_)
  685. class_counts = np.bincount(y_idx, minlength=int(classes.size)).astype(np.float32)
  686. class_weights = class_counts.sum() / np.maximum(class_counts, 1.0)
  687. class_weights = class_weights / np.mean(class_weights)
  688. weight_tensor = torch.tensor(class_weights, dtype=torch.float32, device=self.device_)
  689. criterion = nn.CrossEntropyLoss(weight=weight_tensor)
  690. optimizer = torch.optim.AdamW(
  691. self.model_.parameters(),
  692. lr=float(self.lr),
  693. weight_decay=float(self.weight_decay),
  694. )
  695. dataset = TensorDataset(
  696. torch.from_numpy(x_arr.astype(np.float32)),
  697. torch.from_numpy(y_idx.astype(np.int64)),
  698. )
  699. batch_size = int(max(4, min(int(self.batch_size), len(dataset))))
  700. generator = torch.Generator()
  701. generator.manual_seed(int(self.random_state))
  702. loader = DataLoader(
  703. dataset,
  704. batch_size=batch_size,
  705. shuffle=True,
  706. generator=generator,
  707. drop_last=False,
  708. )
  709. self.model_.train()
  710. for _ in range(int(self.epochs)):
  711. for xb, yb in loader:
  712. xb = xb.to(self.device_)
  713. yb = yb.to(self.device_)
  714. optimizer.zero_grad(set_to_none=True)
  715. logits = self.model_(xb)
  716. loss = criterion(logits, yb)
  717. loss.backward()
  718. if self.grad_clip and float(self.grad_clip) > 0:
  719. nn.utils.clip_grad_norm_(self.model_.parameters(), float(self.grad_clip))
  720. optimizer.step()
  721. return self
  722. def predict(self, X: Any) -> np.ndarray:
  723. if self.model_ is None or self.classes_ is None:
  724. raise RuntimeError("TorchTabularClassifier must be fit before predict.")
  725. x_arr = np.asarray(X, dtype=np.float32)
  726. if x_arr.ndim != 2:
  727. raise ValueError("TorchTabularClassifier expects a 2D feature matrix.")
  728. self.model_.eval()
  729. with torch.no_grad():
  730. tensor = torch.from_numpy(x_arr.astype(np.float32)).to(self.device_)
  731. logits = self.model_(tensor)
  732. pred_idx = torch.argmax(logits, dim=1).detach().cpu().numpy().astype(np.int64)
  733. return np.asarray([self.classes_[int(idx)] for idx in pred_idx], dtype=object)
  734. def _available_model_names() -> list[str]:
  735. names = ["logreg", "knn", "svm", "gaussian_nb", "decision_tree", "mlp", "rf"]
  736. if TORCH_AVAILABLE:
  737. names.extend(list(DEEP_MODEL_NAMES))
  738. return names
  739. def _resolve_torch_runtime_device(requested_device: str) -> str:
  740. wanted = str(requested_device or "auto").strip().lower()
  741. if wanted == "auto":
  742. if TORCH_AVAILABLE and torch is not None and torch.cuda.is_available():
  743. return "cuda"
  744. return "cpu"
  745. return wanted
  746. def _detect_nvidia_gpus(timeout_seconds: float = 3.0) -> tuple[bool, list[str], str | None]:
  747. try:
  748. completed = subprocess.run(
  749. ["nvidia-smi", "--query-gpu=name", "--format=csv,noheader"],
  750. capture_output=True,
  751. text=True,
  752. timeout=max(0.5, float(timeout_seconds)),
  753. check=False,
  754. )
  755. except FileNotFoundError:
  756. return False, [], "nvidia-smi not found"
  757. except Exception as exc: # noqa: BLE001
  758. return False, [], f"nvidia-smi probe failed: {exc}"
  759. if completed.returncode != 0:
  760. stderr = str(completed.stderr or "").strip()
  761. message = stderr if stderr else f"nvidia-smi return code {completed.returncode}"
  762. return False, [], message
  763. names = [line.strip() for line in str(completed.stdout or "").splitlines() if line.strip()]
  764. if not names:
  765. return False, [], "nvidia-smi returned no GPU names"
  766. return True, names, None
  767. def _model_registry(random_state: int, torch_device: str) -> dict[str, dict[str, Any]]:
  768. registry: dict[str, dict[str, Any]] = {
  769. "logreg": {
  770. "label": "LogisticRegression",
  771. "estimator": LogisticRegression(
  772. solver="lbfgs",
  773. max_iter=8000,
  774. tol=1e-3,
  775. class_weight="balanced",
  776. random_state=random_state,
  777. ),
  778. "grid": _product_grid({"C": [0.1, 1.0, 5.0, 10.0]}),
  779. },
  780. "svm": {
  781. "label": "SVM-RBF",
  782. "estimator": SVC(kernel="rbf", class_weight="balanced"),
  783. "grid": _product_grid({"C": [0.5, 1.0, 3.0, 10.0], "gamma": ["scale", "auto"]}),
  784. },
  785. "gaussian_nb": {
  786. "label": "GaussianNB",
  787. "estimator": GaussianNB(),
  788. "grid": _product_grid({"var_smoothing": [1e-9, 1e-8, 1e-7]}),
  789. },
  790. "decision_tree": {
  791. "label": "DecisionTree",
  792. "estimator": DecisionTreeClassifier(
  793. class_weight="balanced",
  794. random_state=random_state,
  795. ),
  796. "grid": _product_grid(
  797. {
  798. "criterion": ["gini", "entropy"],
  799. "max_depth": [None, 8, 16, 24],
  800. "min_samples_leaf": [1, 2, 4],
  801. }
  802. ),
  803. },
  804. "knn": {
  805. "label": "KNN",
  806. "estimator": KNeighborsClassifier(),
  807. "grid": _product_grid(
  808. {"n_neighbors": [3, 5, 9, 13], "weights": ["uniform", "distance"], "p": [1, 2]}
  809. ),
  810. },
  811. "mlp": {
  812. "label": "MLP",
  813. "estimator": MLPClassifier(
  814. max_iter=800,
  815. early_stopping=False,
  816. random_state=random_state,
  817. ),
  818. "grid": _product_grid(
  819. {
  820. "hidden_layer_sizes": [(64,), (128,), (64, 32)],
  821. "alpha": [1e-4, 1e-3],
  822. "learning_rate_init": [1e-3, 5e-4],
  823. }
  824. ),
  825. },
  826. "rf": {
  827. "label": "RandomForest",
  828. "estimator": RandomForestClassifier(
  829. n_estimators=400,
  830. class_weight="balanced_subsample",
  831. random_state=random_state,
  832. n_jobs=-1,
  833. ),
  834. "grid": _product_grid(
  835. {"max_depth": [None, 10, 20], "min_samples_leaf": [1, 2, 4], "max_features": ["sqrt", 0.5]}
  836. ),
  837. },
  838. }
  839. if TORCH_AVAILABLE:
  840. registry["lstm1d"] = {
  841. "label": "Torch-LSTM1D",
  842. "estimator": TorchTabularClassifier(
  843. architecture="lstm1d",
  844. epochs=30,
  845. batch_size=64,
  846. lr=1e-3,
  847. weight_decay=1e-4,
  848. hidden_dim=96,
  849. dropout=0.2,
  850. n_layers=1,
  851. random_state=random_state,
  852. device=torch_device,
  853. ),
  854. "grid": _product_grid(
  855. {
  856. "hidden_dim": [64, 96],
  857. "n_layers": [1, 2],
  858. "lr": [1e-3, 5e-4],
  859. "dropout": [0.1, 0.2],
  860. }
  861. ),
  862. }
  863. registry["bilstm1d"] = {
  864. "label": "Torch-BiLSTM1D",
  865. "estimator": TorchTabularClassifier(
  866. architecture="bilstm1d",
  867. epochs=36,
  868. batch_size=64,
  869. lr=8e-4,
  870. weight_decay=1e-4,
  871. hidden_dim=128,
  872. dropout=0.25,
  873. n_layers=2,
  874. random_state=random_state,
  875. device=torch_device,
  876. ),
  877. "grid": _product_grid(
  878. {
  879. "hidden_dim": [96, 128],
  880. "n_layers": [1, 2],
  881. "lr": [1e-3, 5e-4],
  882. "dropout": [0.2, 0.3],
  883. }
  884. ),
  885. }
  886. registry["gru1d"] = {
  887. "label": "Torch-GRU1D",
  888. "estimator": TorchTabularClassifier(
  889. architecture="gru1d",
  890. epochs=30,
  891. batch_size=64,
  892. lr=1e-3,
  893. weight_decay=1e-4,
  894. hidden_dim=96,
  895. dropout=0.2,
  896. n_layers=1,
  897. random_state=random_state,
  898. device=torch_device,
  899. ),
  900. "grid": _product_grid(
  901. {
  902. "hidden_dim": [64, 96],
  903. "n_layers": [1, 2],
  904. "lr": [1e-3, 5e-4],
  905. "dropout": [0.1, 0.2],
  906. }
  907. ),
  908. }
  909. registry["bigru1d"] = {
  910. "label": "Torch-BiGRU1D",
  911. "estimator": TorchTabularClassifier(
  912. architecture="bigru1d",
  913. epochs=36,
  914. batch_size=64,
  915. lr=8e-4,
  916. weight_decay=1e-4,
  917. hidden_dim=128,
  918. dropout=0.25,
  919. n_layers=2,
  920. random_state=random_state,
  921. device=torch_device,
  922. ),
  923. "grid": _product_grid(
  924. {
  925. "hidden_dim": [96, 128],
  926. "n_layers": [1, 2],
  927. "lr": [1e-3, 5e-4],
  928. "dropout": [0.2, 0.3],
  929. }
  930. ),
  931. }
  932. registry["cnn1d"] = {
  933. "label": "Torch-CNN1D",
  934. "estimator": TorchTabularClassifier(
  935. architecture="cnn1d",
  936. epochs=30,
  937. batch_size=64,
  938. lr=1e-3,
  939. weight_decay=1e-4,
  940. hidden_dim=96,
  941. dropout=0.2,
  942. random_state=random_state,
  943. device=torch_device,
  944. ),
  945. "grid": _product_grid({"hidden_dim": [64, 96], "lr": [1e-3, 5e-4], "dropout": [0.1, 0.2]}),
  946. }
  947. registry["cnn1d_deep"] = {
  948. "label": "Torch-CNN1D-Deep",
  949. "estimator": TorchTabularClassifier(
  950. architecture="cnn1d",
  951. epochs=42,
  952. batch_size=64,
  953. lr=8e-4,
  954. weight_decay=2e-4,
  955. hidden_dim=128,
  956. dropout=0.25,
  957. random_state=random_state,
  958. device=torch_device,
  959. ),
  960. "grid": _product_grid(
  961. {
  962. "hidden_dim": [96, 128, 160],
  963. "lr": [1e-3, 5e-4],
  964. "dropout": [0.15, 0.25],
  965. "weight_decay": [1e-4, 5e-4],
  966. }
  967. ),
  968. }
  969. registry["transformer"] = {
  970. "label": "Torch-Transformer",
  971. "estimator": TorchTabularClassifier(
  972. architecture="transformer",
  973. epochs=30,
  974. batch_size=64,
  975. lr=1e-3,
  976. weight_decay=1e-4,
  977. hidden_dim=64,
  978. dropout=0.1,
  979. n_heads=4,
  980. n_layers=1,
  981. random_state=random_state,
  982. device=torch_device,
  983. ),
  984. "grid": _product_grid({"hidden_dim": [32, 64], "n_layers": [1, 2], "lr": [1e-3], "dropout": [0.1]}),
  985. }
  986. registry["transformer_xl"] = {
  987. "label": "Torch-Transformer-XL",
  988. "estimator": TorchTabularClassifier(
  989. architecture="transformer",
  990. epochs=42,
  991. batch_size=64,
  992. lr=7e-4,
  993. weight_decay=2e-4,
  994. hidden_dim=128,
  995. dropout=0.15,
  996. n_heads=8,
  997. n_layers=2,
  998. random_state=random_state,
  999. device=torch_device,
  1000. ),
  1001. "grid": _product_grid(
  1002. {
  1003. "hidden_dim": [96, 128],
  1004. "n_heads": [4, 8],
  1005. "n_layers": [2, 3],
  1006. "lr": [1e-3, 5e-4],
  1007. "dropout": [0.1, 0.2],
  1008. }
  1009. ),
  1010. }
  1011. return registry
  1012. def _mutual_info_score(
  1013. X: np.ndarray,
  1014. y: np.ndarray,
  1015. *,
  1016. random_state: int,
  1017. ) -> np.ndarray:
  1018. return mutual_info_classif(X, y, random_state=random_state)
  1019. def _feature_selector_registry(random_state: int) -> dict[str, dict[str, Any]]:
  1020. return {
  1021. "none": {
  1022. "label": "NoFeatureSelection",
  1023. "selector": "passthrough",
  1024. "grid": [{}],
  1025. },
  1026. "anova": {
  1027. "label": "ANOVA-SelectPercentile",
  1028. "selector": SelectPercentile(score_func=f_classif),
  1029. "grid": _product_grid({"percentile": [25, 50, 75]}),
  1030. },
  1031. "mutual_info": {
  1032. "label": "MutualInfo-SelectPercentile",
  1033. "selector": SelectPercentile(score_func=partial(_mutual_info_score, random_state=random_state)),
  1034. "grid": _product_grid({"percentile": [25, 50, 75]}),
  1035. },
  1036. "l1": {
  1037. "label": "L1-SelectFromModel",
  1038. "selector": SelectFromModel(
  1039. estimator=LogisticRegression(
  1040. penalty="l1",
  1041. solver="saga",
  1042. class_weight="balanced",
  1043. max_iter=6000,
  1044. random_state=random_state,
  1045. )
  1046. ),
  1047. "grid": _product_grid({"threshold": ["mean", "median"]}),
  1048. },
  1049. "tree": {
  1050. "label": "Tree-SelectFromModel",
  1051. "selector": SelectFromModel(
  1052. estimator=ExtraTreesClassifier(
  1053. n_estimators=300,
  1054. class_weight="balanced_subsample",
  1055. random_state=random_state,
  1056. n_jobs=-1,
  1057. )
  1058. ),
  1059. "grid": _product_grid({"threshold": ["mean", "median"]}),
  1060. },
  1061. }
  1062. def _prefix_param_grid(param_grid: list[dict[str, Any]], prefix: str) -> list[dict[str, Any]]:
  1063. if not param_grid:
  1064. return [{}]
  1065. out: list[dict[str, Any]] = []
  1066. for params in param_grid:
  1067. out.append({f"{prefix}__{key}": value for key, value in params.items()})
  1068. return out
  1069. def _cross_param_grids(left: list[dict[str, Any]], right: list[dict[str, Any]]) -> list[dict[str, Any]]:
  1070. if not left:
  1071. left = [{}]
  1072. if not right:
  1073. right = [{}]
  1074. out: list[dict[str, Any]] = []
  1075. for l_params in left:
  1076. for r_params in right:
  1077. merged = dict(l_params)
  1078. merged.update(r_params)
  1079. out.append(merged)
  1080. return out
  1081. def _split_pipeline_params(params: dict[str, Any]) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any]]:
  1082. model_params: dict[str, Any] = {}
  1083. selector_params: dict[str, Any] = {}
  1084. other_params: dict[str, Any] = {}
  1085. for key, value in params.items():
  1086. if key.startswith("model__"):
  1087. model_params[key[len("model__") :]] = value
  1088. elif key.startswith("selector__"):
  1089. selector_params[key[len("selector__") :]] = value
  1090. else:
  1091. other_params[key] = value
  1092. return model_params, selector_params, other_params
  1093. def _sample_param_grid(combos: list[dict[str, Any]], max_count: int, seed: int) -> list[dict[str, Any]]:
  1094. if len(combos) <= max_count:
  1095. return combos
  1096. rng = random.Random(seed)
  1097. picked_idx = sorted(rng.sample(range(len(combos)), max_count))
  1098. return [combos[i] for i in picked_idx]
  1099. def _sample_plan_items(items: list[dict[str, Any]], max_count: int | None, seed: int) -> list[dict[str, Any]]:
  1100. if max_count is None or len(items) <= max_count:
  1101. return items
  1102. rng = random.Random(seed)
  1103. picked_idx = sorted(rng.sample(range(len(items)), max_count))
  1104. return [items[i] for i in picked_idx]
  1105. def _build_pipeline(
  1106. args: argparse.Namespace,
  1107. estimator: BaseEstimator,
  1108. feature_selector: Any,
  1109. ) -> Pipeline:
  1110. selector_step = feature_selector if feature_selector != "passthrough" else "passthrough"
  1111. return Pipeline(
  1112. steps=[
  1113. ("imputer", SimpleImputer(strategy="median")),
  1114. ("clip", QuantileClipper(lower_q=args.clip_lower_quantile, upper_q=args.clip_upper_quantile)),
  1115. ("scale", RobustScaler(with_centering=True)),
  1116. ("var", VarianceThreshold(threshold=0.0)),
  1117. ("selector", selector_step),
  1118. ("model", estimator),
  1119. ]
  1120. )
  1121. def _metric_bundle(y_true: np.ndarray, y_pred: np.ndarray, labels: list[str]) -> dict[str, Any]:
  1122. return {
  1123. "accuracy": float(accuracy_score(y_true, y_pred)),
  1124. "balanced_accuracy": float(balanced_accuracy_score(y_true, y_pred)),
  1125. "macro_f1": float(f1_score(y_true, y_pred, average="macro", zero_division=0)),
  1126. "weighted_f1": float(f1_score(y_true, y_pred, average="weighted", zero_division=0)),
  1127. "confusion_matrix_labels": labels,
  1128. "confusion_matrix": confusion_matrix(y_true, y_pred, labels=labels).tolist(),
  1129. }
  1130. def _write_live_confusion_png(
  1131. *,
  1132. metrics: dict[str, Any],
  1133. out_root: Path,
  1134. class_scenario_name: str,
  1135. dataset: str,
  1136. protocol: str,
  1137. split_id: str,
  1138. pipeline_id: str,
  1139. eval_index: int,
  1140. cmap: str,
  1141. dpi: int,
  1142. ) -> Path:
  1143. labels = [str(x) for x in metrics.get("confusion_matrix_labels", [])]
  1144. counts = np.asarray(metrics.get("confusion_matrix", []), dtype=np.float64)
  1145. if counts.ndim != 2 or counts.shape[0] != counts.shape[1] or counts.shape[0] != len(labels):
  1146. raise ValueError("Confusion matrix shape/labels mismatch for live confusion PNG export.")
  1147. support = counts.sum(axis=1)
  1148. with np.errstate(divide="ignore", invalid="ignore"):
  1149. norm = np.divide(counts, support[:, None], where=support[:, None] > 0.0)
  1150. norm = np.nan_to_num(norm, nan=0.0, posinf=0.0, neginf=0.0)
  1151. n_classes = max(1, len(labels))
  1152. fig_w = max(6.0, min(16.0, 1.1 * n_classes + 2.0))
  1153. fig_h = max(5.0, min(14.0, 1.0 * n_classes + 2.5))
  1154. fig, ax = plt.subplots(figsize=(fig_w, fig_h))
  1155. im = ax.imshow(norm, interpolation="nearest", cmap=cmap, vmin=0.0, vmax=1.0)
  1156. ax.figure.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
  1157. ax.set(
  1158. xticks=np.arange(len(labels)),
  1159. yticks=np.arange(len(labels)),
  1160. xticklabels=labels,
  1161. yticklabels=labels,
  1162. ylabel="True label",
  1163. xlabel="Predicted label",
  1164. title=(
  1165. f"{dataset} | {protocol} | {pipeline_id}\n"
  1166. f"split={split_id} | scenario={class_scenario_name}"
  1167. ),
  1168. )
  1169. plt.setp(ax.get_xticklabels(), rotation=45, ha="right", rotation_mode="anchor")
  1170. for i in range(norm.shape[0]):
  1171. for j in range(norm.shape[1]):
  1172. val = float(norm[i, j])
  1173. cnt = int(round(float(counts[i, j])))
  1174. text_color = "white" if val >= 0.55 else "black"
  1175. ax.text(j, i, f"{val:.2f}\n({cnt})", ha="center", va="center", color=text_color, fontsize=8)
  1176. fig.tight_layout()
  1177. out_subdir = (
  1178. out_root
  1179. / _slug(class_scenario_name)
  1180. / _slug(dataset)
  1181. / _slug(protocol)
  1182. / _slug(split_id)
  1183. )
  1184. out_subdir.mkdir(parents=True, exist_ok=True)
  1185. out_path = out_subdir / f"{int(eval_index):06d}__{_slug(pipeline_id)}.png"
  1186. fig.savefig(out_path, dpi=max(72, int(dpi)))
  1187. plt.close(fig)
  1188. return out_path
  1189. def _fit_predict(
  1190. estimator: BaseEstimator,
  1191. feature_selector: Any,
  1192. params: dict[str, Any],
  1193. X_train: np.ndarray,
  1194. y_train: np.ndarray,
  1195. X_test: np.ndarray,
  1196. args: argparse.Namespace,
  1197. warning_counter: Counter[str] | None = None,
  1198. ) -> np.ndarray | None:
  1199. model = clone(estimator)
  1200. selector = clone(feature_selector) if feature_selector != "passthrough" else "passthrough"
  1201. pipe = _build_pipeline(args, model, selector)
  1202. if params:
  1203. pipe.set_params(**params)
  1204. with warnings.catch_warnings(record=True) as caught:
  1205. warnings.simplefilter("always")
  1206. warnings.filterwarnings("always", category=ConvergenceWarning)
  1207. try:
  1208. pipe.fit(X_train, y_train)
  1209. y_pred = pipe.predict(X_test)
  1210. except Exception as exc: # noqa: BLE001
  1211. _capture_warnings(caught, warning_counter)
  1212. if warning_counter is not None:
  1213. message = str(exc).strip().splitlines()[0]
  1214. if len(message) > 160:
  1215. message = message[:157] + "..."
  1216. warning_counter[f"FitError: {type(exc).__name__}: {message}"] += 1
  1217. return None
  1218. _capture_warnings(caught, warning_counter)
  1219. return y_pred
  1220. def _inner_score_grouped(
  1221. estimator: BaseEstimator,
  1222. feature_selector: Any,
  1223. params: dict[str, Any],
  1224. X_train: np.ndarray,
  1225. y_train: np.ndarray,
  1226. groups_train: np.ndarray,
  1227. labels: list[str],
  1228. args: argparse.Namespace,
  1229. ) -> tuple[float, dict[str, int]]:
  1230. warning_counter: Counter[str] = Counter()
  1231. unique_groups = np.unique(groups_train)
  1232. if unique_groups.size < 2:
  1233. return -1.0, {}
  1234. n_splits = min(args.inner_folds, int(unique_groups.size))
  1235. if n_splits < 2:
  1236. return -1.0, {}
  1237. splitter = GroupKFold(n_splits=n_splits)
  1238. fold_scores: list[float] = []
  1239. for inner_train_idx, inner_valid_idx in splitter.split(X_train, y_train, groups=groups_train):
  1240. y_pred = _fit_predict(
  1241. estimator=estimator,
  1242. feature_selector=feature_selector,
  1243. params=params,
  1244. X_train=X_train[inner_train_idx],
  1245. y_train=y_train[inner_train_idx],
  1246. X_test=X_train[inner_valid_idx],
  1247. args=args,
  1248. warning_counter=warning_counter,
  1249. )
  1250. if y_pred is None:
  1251. continue
  1252. metrics = _metric_bundle(y_train[inner_valid_idx], y_pred, labels=labels)
  1253. fold_scores.append(float(metrics["balanced_accuracy"]))
  1254. return (float(mean(fold_scores)) if fold_scores else -1.0), dict(sorted(warning_counter.items()))
  1255. def _inner_score_stratified(
  1256. estimator: BaseEstimator,
  1257. feature_selector: Any,
  1258. params: dict[str, Any],
  1259. X_train: np.ndarray,
  1260. y_train: np.ndarray,
  1261. labels: list[str],
  1262. args: argparse.Namespace,
  1263. seed: int,
  1264. ) -> tuple[float, dict[str, int]]:
  1265. warning_counter: Counter[str] = Counter()
  1266. counts = Counter(y_train.tolist())
  1267. if not counts:
  1268. return -1.0, {}
  1269. min_class = min(counts.values())
  1270. n_splits = min(args.inner_folds, int(min_class))
  1271. if n_splits < 2:
  1272. return -1.0, {}
  1273. splitter = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=seed)
  1274. fold_scores: list[float] = []
  1275. for inner_train_idx, inner_valid_idx in splitter.split(X_train, y_train):
  1276. y_pred = _fit_predict(
  1277. estimator=estimator,
  1278. feature_selector=feature_selector,
  1279. params=params,
  1280. X_train=X_train[inner_train_idx],
  1281. y_train=y_train[inner_train_idx],
  1282. X_test=X_train[inner_valid_idx],
  1283. args=args,
  1284. warning_counter=warning_counter,
  1285. )
  1286. if y_pred is None:
  1287. continue
  1288. metrics = _metric_bundle(y_train[inner_valid_idx], y_pred, labels=labels)
  1289. fold_scores.append(float(metrics["balanced_accuracy"]))
  1290. return (float(mean(fold_scores)) if fold_scores else -1.0), dict(sorted(warning_counter.items()))
  1291. def _choose_best_params(
  1292. estimator: BaseEstimator,
  1293. feature_selector: Any,
  1294. param_grid: list[dict[str, Any]],
  1295. X_train: np.ndarray,
  1296. y_train: np.ndarray,
  1297. labels: list[str],
  1298. args: argparse.Namespace,
  1299. mode: str,
  1300. groups_train: np.ndarray | None,
  1301. seed: int,
  1302. ) -> tuple[dict[str, Any], dict[str, Any]]:
  1303. best_params = param_grid[0] if param_grid else {}
  1304. best_score = -1.0
  1305. tried = 0
  1306. details: list[dict[str, Any]] = []
  1307. warning_counter: Counter[str] = Counter()
  1308. for params in param_grid:
  1309. tried += 1
  1310. if mode == "grouped":
  1311. assert groups_train is not None
  1312. score, warning_counts = _inner_score_grouped(
  1313. estimator=estimator,
  1314. feature_selector=feature_selector,
  1315. params=params,
  1316. X_train=X_train,
  1317. y_train=y_train,
  1318. groups_train=groups_train,
  1319. labels=labels,
  1320. args=args,
  1321. )
  1322. else:
  1323. score, warning_counts = _inner_score_stratified(
  1324. estimator=estimator,
  1325. feature_selector=feature_selector,
  1326. params=params,
  1327. X_train=X_train,
  1328. y_train=y_train,
  1329. labels=labels,
  1330. args=args,
  1331. seed=seed + tried,
  1332. )
  1333. warning_counter.update(warning_counts)
  1334. details.append({"params": params, "inner_balanced_accuracy": score, "warning_counts": warning_counts})
  1335. if score > best_score:
  1336. best_score = score
  1337. best_params = params
  1338. return best_params, {
  1339. "best_inner_balanced_accuracy": best_score,
  1340. "n_param_sets": tried,
  1341. "scores": details,
  1342. "warning_counts": dict(sorted(warning_counter.items())),
  1343. }
  1344. def _load_dataset_frame(path: Path) -> pd.DataFrame:
  1345. if not path.exists():
  1346. raise FileNotFoundError(path)
  1347. return pd.read_csv(path, sep="\t", dtype=str, keep_default_na=False)
  1348. def _select_feature_columns(df: pd.DataFrame, target_col: str) -> list[str]:
  1349. feature_cols: list[str] = []
  1350. for col in df.columns:
  1351. if col == target_col:
  1352. continue
  1353. if col in NON_FEATURE_COLUMNS:
  1354. continue
  1355. if col.startswith("preproc_version_"):
  1356. continue
  1357. if col.startswith("baseline_start_s_"):
  1358. continue
  1359. if col.startswith("baseline_end_s_"):
  1360. continue
  1361. if col.startswith("has_"):
  1362. continue
  1363. feature_cols.append(col)
  1364. return feature_cols
  1365. def _prepare_dataset(
  1366. dataset_name: str,
  1367. dataset_path: Path,
  1368. target_col: str,
  1369. labels: list[str],
  1370. class_scenario: dict[str, Any],
  1371. ) -> dict[str, Any]:
  1372. df = _load_dataset_frame(dataset_path)
  1373. if target_col not in df.columns:
  1374. raise ValueError(f"{dataset_name}: missing target column '{target_col}' in {dataset_path}")
  1375. if "participant_id" not in df.columns:
  1376. raise ValueError(f"{dataset_name}: missing participant_id column in {dataset_path}")
  1377. df = df.copy()
  1378. df[target_col] = df[target_col].astype(str).str.strip()
  1379. baseline_label = class_scenario.get("baseline_from_tutorial_label")
  1380. source_label = df[target_col].astype(str).str.strip()
  1381. if baseline_label and "is_tutorial" in df.columns:
  1382. tutorial_mask = df["is_tutorial"].astype(str).str.strip().str.lower() == "true"
  1383. source_label = np.where(tutorial_mask.to_numpy(), str(baseline_label), source_label.to_numpy())
  1384. df["_target_source_label"] = pd.Series(source_label, index=df.index, dtype="string").astype(str)
  1385. rows_before_class_filter = int(df.shape[0])
  1386. manifest_labels_order = list(class_scenario.get("manifest_labels", []))
  1387. full_class_counts = _ordered_counts(df["_target_source_label"].tolist(), preferred_order=manifest_labels_order)
  1388. allowed_original_labels_order = list(class_scenario["allowed_original_labels"])
  1389. allowed_original_labels = set(allowed_original_labels_order)
  1390. merge_map = {str(k): str(v) for k, v in class_scenario["merge_map"].items()}
  1391. final_label_set = set(labels)
  1392. df = df[df["_target_source_label"].isin(allowed_original_labels)].copy()
  1393. rows_after_label_filter = int(df.shape[0])
  1394. class_counts_after_label_filter = _ordered_counts(
  1395. df["_target_source_label"].tolist(),
  1396. preferred_order=allowed_original_labels_order,
  1397. )
  1398. if df.empty:
  1399. raise ValueError(f"{dataset_name}: no rows left after class include/drop filters.")
  1400. mapped = df["_target_source_label"].map(lambda x: str(merge_map.get(x, x)).strip())
  1401. merge_drop_mask = mapped.str.lower().isin(DROP_LABEL_TOKENS) | (mapped == "")
  1402. if bool(merge_drop_mask.any()):
  1403. df = df.loc[~merge_drop_mask].copy()
  1404. mapped = mapped.loc[~merge_drop_mask].copy()
  1405. df["_target_label_effective"] = mapped.values
  1406. df = df[df["_target_label_effective"].isin(final_label_set)].reset_index(drop=True)
  1407. if df.empty:
  1408. raise ValueError(f"{dataset_name}: no rows left after class merge/drop mapping.")
  1409. group_col = "split_group" if "split_group" in df.columns else "participant_id"
  1410. groups = df[group_col].astype(str).to_numpy()
  1411. participants = df["participant_id"].astype(str).to_numpy()
  1412. y = df["_target_label_effective"].astype(str).to_numpy()
  1413. feature_cols = _select_feature_columns(df, target_col=target_col)
  1414. if not feature_cols:
  1415. raise ValueError(f"{dataset_name}: no candidate feature columns.")
  1416. feature_df = df[feature_cols].replace({"n/a": np.nan, "": np.nan})
  1417. for col in feature_cols:
  1418. feature_df[col] = pd.to_numeric(feature_df[col], errors="coerce")
  1419. valid_feature_cols = [col for col in feature_cols if feature_df[col].notna().any()]
  1420. if not valid_feature_cols:
  1421. raise ValueError(f"{dataset_name}: all feature columns are empty/non-numeric after coercion.")
  1422. X = feature_df[valid_feature_cols].to_numpy(dtype=np.float64)
  1423. row_ids = np.arange(df.shape[0], dtype=np.int64)
  1424. trial_ids = df["trial_id"].astype(str).to_numpy() if "trial_id" in df.columns else np.array([""] * df.shape[0])
  1425. return {
  1426. "name": dataset_name,
  1427. "path": str(dataset_path),
  1428. "df": df,
  1429. "X": X,
  1430. "y": y,
  1431. "groups": groups,
  1432. "participants": participants,
  1433. "row_ids": row_ids,
  1434. "trial_ids": trial_ids,
  1435. "feature_columns": valid_feature_cols,
  1436. "class_counts": _ordered_counts(y.tolist(), preferred_order=labels),
  1437. "full_class_counts": full_class_counts,
  1438. "class_counts_after_label_filter": class_counts_after_label_filter,
  1439. "rows_before_class_filter": rows_before_class_filter,
  1440. "rows_after_label_filter": rows_after_label_filter,
  1441. "rows_after_class_mapping": int(df.shape[0]),
  1442. }
  1443. def _subset_by_participants(payload: dict[str, Any], participants: list[str]) -> tuple[np.ndarray, np.ndarray]:
  1444. wanted = set(participants)
  1445. participant_col = payload["participants"]
  1446. mask = np.isin(participant_col, list(wanted))
  1447. idx = np.where(mask)[0]
  1448. return idx, mask
  1449. def _aggregate_metric_rows(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
  1450. grouped: dict[tuple[str, str, str, str, str], list[dict[str, Any]]] = defaultdict(list)
  1451. for row in rows:
  1452. feature_selector = str(row.get("feature_selector", "none"))
  1453. pipeline_id = str(row.get("pipeline_id", f"{row['model']}+{feature_selector}"))
  1454. grouped[(row["dataset"], row["protocol"], row["model"], feature_selector, pipeline_id)].append(row)
  1455. out: list[dict[str, Any]] = []
  1456. for (dataset, protocol, model, feature_selector, pipeline_id), items in sorted(grouped.items()):
  1457. acc = [float(x["metrics"]["accuracy"]) for x in items]
  1458. bacc = [float(x["metrics"]["balanced_accuracy"]) for x in items]
  1459. mf1 = [float(x["metrics"]["macro_f1"]) for x in items]
  1460. wf1 = [float(x["metrics"]["weighted_f1"]) for x in items]
  1461. out.append(
  1462. {
  1463. "dataset": dataset,
  1464. "protocol": protocol,
  1465. "model": model,
  1466. "feature_selector": feature_selector,
  1467. "pipeline_id": pipeline_id,
  1468. "n_evaluations": len(items),
  1469. "accuracy_mean": float(mean(acc)),
  1470. "accuracy_std": float(pstdev(acc)) if len(acc) > 1 else 0.0,
  1471. "balanced_accuracy_mean": float(mean(bacc)),
  1472. "balanced_accuracy_std": float(pstdev(bacc)) if len(bacc) > 1 else 0.0,
  1473. "macro_f1_mean": float(mean(mf1)),
  1474. "macro_f1_std": float(pstdev(mf1)) if len(mf1) > 1 else 0.0,
  1475. "weighted_f1_mean": float(mean(wf1)),
  1476. "weighted_f1_std": float(pstdev(wf1)) if len(wf1) > 1 else 0.0,
  1477. }
  1478. )
  1479. return out
  1480. def _best_model_by_dataset_protocol(aggregate_rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
  1481. grouped: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list)
  1482. for row in aggregate_rows:
  1483. grouped[(row["dataset"], row["protocol"])].append(row)
  1484. out: list[dict[str, Any]] = []
  1485. for (dataset, protocol), items in sorted(grouped.items()):
  1486. ranked = sorted(
  1487. items,
  1488. key=lambda x: (x["balanced_accuracy_mean"], x["macro_f1_mean"], -x["balanced_accuracy_std"]),
  1489. reverse=True,
  1490. )
  1491. out.append(
  1492. {
  1493. "dataset": dataset,
  1494. "protocol": protocol,
  1495. "best_model": ranked[0]["model"],
  1496. "best_feature_selector": ranked[0].get("feature_selector", "none"),
  1497. "best_pipeline": ranked[0].get("pipeline_id", ranked[0]["model"]),
  1498. "balanced_accuracy_mean": ranked[0]["balanced_accuracy_mean"],
  1499. "macro_f1_mean": ranked[0]["macro_f1_mean"],
  1500. "n_evaluations": ranked[0]["n_evaluations"],
  1501. }
  1502. )
  1503. return out
  1504. def _aggregate_warning_counts(evaluations: list[dict[str, Any]]) -> dict[str, int]:
  1505. warning_counter: Counter[str] = Counter()
  1506. for row in evaluations:
  1507. warning_payload = row.get("warning_counts", {})
  1508. if not isinstance(warning_payload, dict):
  1509. continue
  1510. for section in ("tuning", "outer_fit_predict"):
  1511. section_payload = warning_payload.get(section, {})
  1512. if not isinstance(section_payload, dict):
  1513. continue
  1514. for key, count in section_payload.items():
  1515. warning_counter[str(key)] += int(count)
  1516. return {k: warning_counter[k] for k in sorted(warning_counter, key=lambda x: (-warning_counter[x], x))}
  1517. def _protocol_debug_label(protocol: str) -> str:
  1518. mapping = {
  1519. "within_participant": "single_participant",
  1520. "group_holdout": "group",
  1521. "loso": "loso",
  1522. "pooled_stratified": "pooled_mixed",
  1523. "pooled_stratified_holdout": "pooled_mixed",
  1524. "pooled_stratified_one_fold": "pooled_mixed",
  1525. }
  1526. return mapping.get(protocol, protocol)
  1527. def _dataset_protocol_coverage(
  1528. aggregate_rows: list[dict[str, Any]],
  1529. datasets: list[str],
  1530. protocols: list[str],
  1531. min_aggregate_rows_per_pair: int,
  1532. ) -> dict[str, Any]:
  1533. pair_counts: Counter[tuple[str, str]] = Counter()
  1534. for row in aggregate_rows:
  1535. dataset = str(row.get("dataset", ""))
  1536. protocol = str(row.get("protocol", ""))
  1537. if not dataset or not protocol:
  1538. continue
  1539. pair_counts[(dataset, protocol)] += 1
  1540. pairs_summary: list[dict[str, Any]] = []
  1541. missing_pairs: list[dict[str, Any]] = []
  1542. for dataset in datasets:
  1543. for protocol in protocols:
  1544. count = int(pair_counts.get((dataset, protocol), 0))
  1545. entry = {"dataset": dataset, "protocol": protocol, "aggregate_rows": count}
  1546. pairs_summary.append(entry)
  1547. if min_aggregate_rows_per_pair > 0 and count < min_aggregate_rows_per_pair:
  1548. missing_pairs.append(entry)
  1549. return {
  1550. "min_aggregate_rows_per_pair": int(min_aggregate_rows_per_pair),
  1551. "pairs": pairs_summary,
  1552. "missing_pairs": missing_pairs,
  1553. }
  1554. def _build_summary_markdown(
  1555. config: dict[str, Any],
  1556. dataset_stats: list[dict[str, Any]],
  1557. aggregate_rows: list[dict[str, Any]],
  1558. best_rows: list[dict[str, Any]],
  1559. warning_counts_total: dict[str, int],
  1560. coverage: dict[str, Any],
  1561. skip_counts: dict[str, int],
  1562. eeg_condition_psd: dict[str, Any] | None,
  1563. ) -> str:
  1564. lines: list[str] = []
  1565. lines.append("# Stage 6 ML Summary")
  1566. lines.append("")
  1567. lines.append("## Config")
  1568. lines.append(f"- Target: `{config['target']}` (`{config['target_column']}`)")
  1569. lines.append(f"- Protocols: `{', '.join(config['protocols'])}`")
  1570. lines.append(f"- Models: `{', '.join(config['models'])}`")
  1571. lines.append(f"- Feature selectors: `{', '.join(config['feature_selectors'])}`")
  1572. lines.append(f"- Datasets: `{', '.join(config['datasets'])}`")
  1573. lines.append(f"- Inner folds: `{config['inner_folds']}`")
  1574. lines.append(f"- Max param combos/model: `{config['max_param_combos']}`")
  1575. lines.append(f"- Pooled stratified folds: `{config.get('pooled_stratified_folds', 'n/a')}`")
  1576. lines.append(
  1577. f"- Pooled no-folds mode: "
  1578. f"`{config.get('pooled_stratified_no_folds', False)}`"
  1579. )
  1580. lines.append(
  1581. f"- Pooled test size: "
  1582. f"`{config.get('pooled_stratified_test_size', 'n/a')}`"
  1583. )
  1584. lines.append(
  1585. f"- Within-participant no-folds mode: "
  1586. f"`{config.get('within_participant_no_folds', False)}`"
  1587. )
  1588. lines.append(
  1589. f"- Within-participant test size: "
  1590. f"`{config.get('within_participant_test_size', 'n/a')}`"
  1591. )
  1592. lines.append(f"- Torch available: `{config.get('torch_available', False)}`")
  1593. lines.append(
  1594. f"- Torch device requested/effective: "
  1595. f"`{config.get('torch_device_requested', 'auto')}` -> `{config.get('torch_device_effective', 'cpu')}`"
  1596. )
  1597. lines.append(f"- NVIDIA GPU detected: `{config.get('nvidia_gpu_detected', False)}`")
  1598. gpu_names = [str(x) for x in config.get("nvidia_gpu_names", []) if str(x).strip()]
  1599. if gpu_names:
  1600. lines.append(f"- NVIDIA GPU names: `{', '.join(gpu_names)}`")
  1601. lines.append(f"- Live confusion PNGs during training: `{config.get('live_confusion_pngs', False)}`")
  1602. lines.append(f"- EEG condition PSD plots: `{config.get('eeg_condition_psd_plots', False)}`")
  1603. class_cfg = config.get("class_scenario", {})
  1604. lines.append(f"- Class scenario: `{class_cfg.get('name', 'default')}`")
  1605. lines.append(f"- Final labels: `{', '.join(class_cfg.get('final_labels', []))}`")
  1606. lines.append("")
  1607. lines.append("## EEG PSD / Topomap QC")
  1608. if eeg_condition_psd:
  1609. lines.append(f"- Status: `{eeg_condition_psd.get('status', 'unknown')}`")
  1610. manifest_json = str(eeg_condition_psd.get("manifest_json", "") or "").strip()
  1611. if manifest_json:
  1612. lines.append(f"- Manifest: `{manifest_json}`")
  1613. all_participants_png = str(eeg_condition_psd.get("all_participants_png", "") or "").strip()
  1614. if all_participants_png:
  1615. lines.append(f"- All participants PSD PNG: `{all_participants_png}`")
  1616. all_participants_topomap_png = str(eeg_condition_psd.get("all_participants_topomap_png", "") or "").strip()
  1617. if all_participants_topomap_png:
  1618. lines.append(f"- All participants topomap PNG: `{all_participants_topomap_png}`")
  1619. roi_reference_topomap_png = str(eeg_condition_psd.get("roi_reference_topomap_png", "") or "").strip()
  1620. if roi_reference_topomap_png:
  1621. lines.append(f"- EEG ROI reference PNG: `{roi_reference_topomap_png}`")
  1622. if eeg_condition_psd.get("status") == "ok":
  1623. lines.append(
  1624. f"- Participants plotted: `{int(eeg_condition_psd.get('n_participants_plotted', 0))}`"
  1625. )
  1626. label_order_plotted = [str(x) for x in eeg_condition_psd.get("label_order_plotted", []) if str(x).strip()]
  1627. if label_order_plotted:
  1628. lines.append(f"- Conditions plotted: `{', '.join(label_order_plotted)}`")
  1629. reason = str(eeg_condition_psd.get("reason", "") or "").strip()
  1630. if reason:
  1631. lines.append(f"- Note: `{reason}`")
  1632. else:
  1633. lines.append("- Status: disabled")
  1634. lines.append("")
  1635. lines.append("## Dataset Snapshot")
  1636. for ds in dataset_stats:
  1637. lines.append(
  1638. f"- `{ds['dataset']}`: rows={ds['rows']}, features={ds['n_features']}, "
  1639. f"participants={ds['n_participants']}, classes={ds['n_classes']}"
  1640. )
  1641. lines.append("")
  1642. lines.append("## Coverage")
  1643. lines.append(
  1644. f"- Minimum aggregate rows required per dataset+protocol: `{coverage.get('min_aggregate_rows_per_pair', 0)}`"
  1645. )
  1646. missing_pairs = list(coverage.get("missing_pairs", []))
  1647. if missing_pairs:
  1648. lines.append("- Missing pairs:")
  1649. for pair in missing_pairs:
  1650. lines.append(
  1651. f" - `{pair['dataset']}` + `{pair['protocol']}` "
  1652. f"(aggregate_rows={int(pair.get('aggregate_rows', 0))})"
  1653. )
  1654. else:
  1655. lines.append("- Missing pairs: none")
  1656. lines.append("")
  1657. lines.append("| dataset | protocol | aggregate_rows |")
  1658. lines.append("|---|---|---:|")
  1659. for pair in coverage.get("pairs", []):
  1660. lines.append(
  1661. f"| {pair['dataset']} | {pair['protocol']} | {int(pair.get('aggregate_rows', 0))} |"
  1662. )
  1663. lines.append("")
  1664. lines.append("## Best By Dataset/Protocol")
  1665. if best_rows:
  1666. for row in best_rows:
  1667. lines.append(
  1668. f"- `{row['dataset']}` + `{row['protocol']}` -> `{row.get('best_pipeline', row['best_model'])}` "
  1669. f"(balanced_acc={row['balanced_accuracy_mean']:.4f}, macro_f1={row['macro_f1_mean']:.4f}, "
  1670. f"n={row['n_evaluations']})"
  1671. )
  1672. else:
  1673. lines.append("- none")
  1674. lines.append("")
  1675. lines.append("## Aggregates")
  1676. lines.append("| dataset | protocol | model | feature_selector | pipeline_id | n | balanced_acc_mean | macro_f1_mean |")
  1677. lines.append("|---|---|---|---|---|---:|---:|---:|")
  1678. if aggregate_rows:
  1679. for row in aggregate_rows:
  1680. lines.append(
  1681. f"| {row['dataset']} | {row['protocol']} | {row['model']} | {row.get('feature_selector', 'none')} | "
  1682. f"{row.get('pipeline_id', row['model'])} | {row['n_evaluations']} | "
  1683. f"{row['balanced_accuracy_mean']:.4f} | {row['macro_f1_mean']:.4f} |"
  1684. )
  1685. else:
  1686. lines.append("| n/a | n/a | n/a | n/a | n/a | 0 | 0.0000 | 0.0000 |")
  1687. lines.append("")
  1688. lines.append("## Skip Summary")
  1689. if skip_counts:
  1690. for key, count in sorted(skip_counts.items(), key=lambda x: (-x[1], x[0])):
  1691. lines.append(f"- {key}: {count}")
  1692. else:
  1693. lines.append("- none")
  1694. lines.append("")
  1695. lines.append("## Warning Summary")
  1696. if warning_counts_total:
  1697. for key, count in warning_counts_total.items():
  1698. lines.append(f"- {key}: {count}")
  1699. else:
  1700. lines.append("- none")
  1701. lines.append("")
  1702. return "\n".join(lines) + "\n"
  1703. def main() -> None:
  1704. parser = argparse.ArgumentParser(
  1705. description=(
  1706. "Stage 6 classic ML battery with leakage-safe preprocessing and "
  1707. "multiple evaluation protocols (within-participant, LOSO, grouped holdout, pooled stratified)."
  1708. )
  1709. )
  1710. parser.add_argument("--bids-root", required=True, help="Path to BIDS root.")
  1711. parser.add_argument("--split-manifest", default=None, help="Path to split_manifest.json.")
  1712. parser.add_argument("--datasets", nargs="*", default=None, help="Datasets to evaluate (from split manifest).")
  1713. parser.add_argument(
  1714. "--protocols",
  1715. nargs="*",
  1716. default=None,
  1717. help=(
  1718. "Evaluation protocols: loso group_holdout within_participant "
  1719. "pooled_stratified pooled_stratified_holdout pooled_stratified_one_fold "
  1720. "(default: all)."
  1721. ),
  1722. )
  1723. parser.add_argument(
  1724. "--models",
  1725. nargs="*",
  1726. default=None,
  1727. help=(
  1728. "Models (default: all available): logreg knn svm gaussian_nb decision_tree mlp rf "
  1729. "[lstm1d gru1d cnn1d transformer bilstm1d bigru1d cnn1d_deep transformer_xl if torch installed]."
  1730. ),
  1731. )
  1732. parser.add_argument(
  1733. "--feature-selectors",
  1734. nargs="*",
  1735. default=None,
  1736. help="Feature selectors (default: none): none anova mutual_info l1 tree.",
  1737. )
  1738. parser.add_argument("--inner-folds", type=int, default=4, help="Inner CV folds for hyperparameter tuning.")
  1739. parser.add_argument(
  1740. "--max-param-combos",
  1741. type=int,
  1742. default=12,
  1743. help="Max hyperparameter combinations evaluated per model+selector pipeline per outer split.",
  1744. )
  1745. parser.add_argument(
  1746. "--max-outer-splits-per-protocol",
  1747. type=int,
  1748. default=None,
  1749. help="Optional cap on outer splits per protocol for quick runs.",
  1750. )
  1751. parser.add_argument(
  1752. "--pooled-stratified-folds",
  1753. type=int,
  1754. default=5,
  1755. help=(
  1756. "Number of folds for pooled_stratified protocol "
  1757. "(all participants mixed; row-level stratified CV)."
  1758. ),
  1759. )
  1760. parser.add_argument(
  1761. "--pooled-stratified-no-folds",
  1762. action="store_true",
  1763. help=(
  1764. "Use one stratified holdout split for pooled_stratified "
  1765. "(instead of K-fold CV)."
  1766. ),
  1767. )
  1768. parser.add_argument(
  1769. "--pooled-stratified-test-size",
  1770. type=float,
  1771. default=0.20,
  1772. help=(
  1773. "Test fraction for --pooled-stratified-no-folds "
  1774. "(default: 0.20)."
  1775. ),
  1776. )
  1777. parser.add_argument(
  1778. "--within-participant-no-folds",
  1779. action="store_true",
  1780. help=(
  1781. "Use one stratified holdout split per participant for within_participant "
  1782. "(instead of K-fold CV)."
  1783. ),
  1784. )
  1785. parser.add_argument(
  1786. "--within-participant-test-size",
  1787. type=float,
  1788. default=0.20,
  1789. help=(
  1790. "Test fraction for --within-participant-no-folds "
  1791. "(default: 0.20)."
  1792. ),
  1793. )
  1794. parser.add_argument(
  1795. "--min-aggregate-rows-per-dataset-protocol",
  1796. type=int,
  1797. default=1,
  1798. help=(
  1799. "Minimum aggregate rows required per dataset+protocol pair. "
  1800. "Set to 0 to disable coverage checks."
  1801. ),
  1802. )
  1803. parser.add_argument(
  1804. "--allow-incomplete-coverage",
  1805. action="store_true",
  1806. help="Do not fail if dataset+protocol coverage checks are not met.",
  1807. )
  1808. parser.add_argument("--clip-lower-quantile", type=float, default=0.01)
  1809. parser.add_argument("--clip-upper-quantile", type=float, default=0.99)
  1810. parser.add_argument("--random-seed", type=int, default=42)
  1811. parser.add_argument(
  1812. "--baseline-from-tutorial-label",
  1813. default=None,
  1814. help=(
  1815. "Optional label to assign tutorial rows before class filtering/merging "
  1816. "(for tutorial-as-baseline proxy experiments)."
  1817. ),
  1818. )
  1819. parser.add_argument(
  1820. "--class-scenario-name",
  1821. default="default",
  1822. help="Name for this class scenario (recorded in run metadata).",
  1823. )
  1824. parser.add_argument(
  1825. "--class-include-labels",
  1826. nargs="*",
  1827. default=None,
  1828. help="Optional subset of original labels to keep before merging.",
  1829. )
  1830. parser.add_argument(
  1831. "--class-drop-labels",
  1832. nargs="*",
  1833. default=None,
  1834. help="Optional original labels to omit before merging.",
  1835. )
  1836. parser.add_argument(
  1837. "--class-merge",
  1838. action="append",
  1839. default=None,
  1840. help=(
  1841. "Repeatable label merge mapping. Formats: old->new, old:new, old=new. "
  1842. "Example: --class-merge 0.6-1.5->low --class-merge 1.5-2.4->low"
  1843. ),
  1844. )
  1845. parser.add_argument(
  1846. "--class-merge-json",
  1847. default=None,
  1848. help="Optional JSON file with merge mapping object: {\"old_label\": \"new_label\", ...}.",
  1849. )
  1850. parser.add_argument("--run-tag", default=None, help="Optional model run tag.")
  1851. parser.add_argument("--results-json", default=None, help="Output JSON path (default: reports/ml_results.json).")
  1852. parser.add_argument("--summary-md", default=None, help="Output markdown path (default: reports/ml_summary.md).")
  1853. parser.add_argument(
  1854. "--eeg-condition-psd-plots",
  1855. dest="eeg_condition_psd_plots",
  1856. action="store_true",
  1857. default=True,
  1858. help=(
  1859. "Write scenario-aware EEG power spectrum QC plots by ROI for all participants "
  1860. "and each single participant (default: enabled)."
  1861. ),
  1862. )
  1863. parser.add_argument(
  1864. "--no-eeg-condition-psd-plots",
  1865. dest="eeg_condition_psd_plots",
  1866. action="store_false",
  1867. help="Disable scenario-aware EEG power spectrum QC plots.",
  1868. )
  1869. parser.add_argument(
  1870. "--eeg-condition-psd-dir",
  1871. default=None,
  1872. help="Output root for scenario-aware EEG PSD plots (default: analysis_pipeline/reports/eeg_condition_psd).",
  1873. )
  1874. parser.add_argument(
  1875. "--eeg-condition-psd-fmin",
  1876. type=float,
  1877. default=2.0,
  1878. help="Minimum frequency for EEG condition PSD plots (default: 2.0 Hz).",
  1879. )
  1880. parser.add_argument(
  1881. "--eeg-condition-psd-fmax",
  1882. type=float,
  1883. default=40.0,
  1884. help="Maximum frequency for EEG condition PSD plots (default: 40.0 Hz).",
  1885. )
  1886. parser.add_argument(
  1887. "--live-confusion-pngs",
  1888. dest="live_confusion_pngs",
  1889. action="store_true",
  1890. default=True,
  1891. help="Write a confusion PNG for each evaluated model/split as training runs (default: enabled).",
  1892. )
  1893. parser.add_argument(
  1894. "--no-live-confusion-pngs",
  1895. dest="live_confusion_pngs",
  1896. action="store_false",
  1897. help="Disable per-evaluation live confusion PNG writing.",
  1898. )
  1899. parser.add_argument(
  1900. "--live-confusion-png-dir",
  1901. default=None,
  1902. help="Directory for per-evaluation live confusion PNGs (default: analysis_pipeline/reports/confusion_pngs_live).",
  1903. )
  1904. parser.add_argument(
  1905. "--live-confusion-png-cmap",
  1906. default="Blues",
  1907. help="Matplotlib colormap for live confusion PNGs.",
  1908. )
  1909. parser.add_argument(
  1910. "--live-confusion-png-dpi",
  1911. type=int,
  1912. default=150,
  1913. help="PNG DPI for live confusion PNGs (default: 150).",
  1914. )
  1915. parser.add_argument("--models-root", default=None, help="Model artifact root (default: analysis_pipeline/models).")
  1916. parser.add_argument(
  1917. "--torch-device",
  1918. default="auto",
  1919. help="Torch device for deep models: auto (default), cpu, cuda, cuda:0, etc.",
  1920. )
  1921. args = parser.parse_args()
  1922. if args.inner_folds < 2:
  1923. raise ValueError("--inner-folds must be >= 2.")
  1924. if args.max_param_combos < 1:
  1925. raise ValueError("--max-param-combos must be >= 1.")
  1926. if args.pooled_stratified_folds < 2:
  1927. raise ValueError("--pooled-stratified-folds must be >= 2.")
  1928. if args.pooled_stratified_test_size <= 0 or args.pooled_stratified_test_size >= 1:
  1929. raise ValueError("--pooled-stratified-test-size must be within (0,1).")
  1930. if args.within_participant_test_size <= 0 or args.within_participant_test_size >= 1:
  1931. raise ValueError("--within-participant-test-size must be within (0,1).")
  1932. if args.min_aggregate_rows_per_dataset_protocol < 0:
  1933. raise ValueError("--min-aggregate-rows-per-dataset-protocol must be >= 0.")
  1934. if args.live_confusion_png_dpi < 72:
  1935. raise ValueError("--live-confusion-png-dpi must be >= 72.")
  1936. if args.clip_lower_quantile < 0 or args.clip_lower_quantile >= 1:
  1937. raise ValueError("--clip-lower-quantile must be within [0,1).")
  1938. if args.clip_upper_quantile <= 0 or args.clip_upper_quantile > 1:
  1939. raise ValueError("--clip-upper-quantile must be within (0,1].")
  1940. if args.clip_lower_quantile >= args.clip_upper_quantile:
  1941. raise ValueError("clip lower quantile must be less than clip upper quantile.")
  1942. if args.eeg_condition_psd_fmin < 0:
  1943. raise ValueError("--eeg-condition-psd-fmin must be >= 0.")
  1944. if args.eeg_condition_psd_fmax <= args.eeg_condition_psd_fmin:
  1945. raise ValueError("--eeg-condition-psd-fmax must be greater than --eeg-condition-psd-fmin.")
  1946. args.torch_device = str(args.torch_device or "auto").strip() or "auto"
  1947. if not TORCH_AVAILABLE and args.torch_device.lower() != "auto":
  1948. raise ValueError("PyTorch is not installed, so --torch-device cannot be set.")
  1949. if TORCH_AVAILABLE and args.torch_device.lower().startswith("cuda"):
  1950. if torch is None or not torch.cuda.is_available():
  1951. raise ValueError(f"Requested torch device '{args.torch_device}' but CUDA is not available.")
  1952. bids_root = _resolve_bids_root(args.bids_root)
  1953. if not bids_root.exists():
  1954. raise FileNotFoundError(bids_root)
  1955. split_manifest_path = Path(args.split_manifest).resolve() if args.split_manifest else _default_split_manifest()
  1956. if not split_manifest_path.exists():
  1957. raise FileNotFoundError(split_manifest_path)
  1958. split_manifest = json.loads(split_manifest_path.read_text(encoding="utf-8"))
  1959. target_col = str(split_manifest.get("target_column", "target_label"))
  1960. class_scenario = _build_class_scenario(split_manifest, args)
  1961. labels = list(class_scenario["final_labels"])
  1962. datasets = _coerce_dataset_list(split_manifest, args.datasets)
  1963. protocols = _coerce_protocol_list(args.protocols)
  1964. models = _coerce_model_list(args.models)
  1965. feature_selectors = _coerce_feature_selector_list(args.feature_selectors)
  1966. torch_device_effective = _resolve_torch_runtime_device(args.torch_device)
  1967. nvidia_gpu_detected, nvidia_gpu_names, nvidia_probe_message = _detect_nvidia_gpus()
  1968. registry = _model_registry(random_state=args.random_seed, torch_device=args.torch_device)
  1969. selector_registry = _feature_selector_registry(random_state=args.random_seed)
  1970. stage_started = time.time()
  1971. print("Stage 6 starting.")
  1972. print(f" Datasets: {', '.join(datasets)}")
  1973. print(f" Protocols: {', '.join(protocols)}")
  1974. print(f" Models: {', '.join(models)}")
  1975. print(f" Feature selectors: {', '.join(feature_selectors)}")
  1976. print(f" Class scenario: {class_scenario['name']} ({len(labels)} classes)")
  1977. print(f" Torch available: {TORCH_AVAILABLE}")
  1978. print(f" NVIDIA GPU detected: {nvidia_gpu_detected}")
  1979. if nvidia_gpu_detected:
  1980. print(f" NVIDIA GPUs: {', '.join(nvidia_gpu_names)}")
  1981. elif nvidia_probe_message:
  1982. print(f" NVIDIA probe note: {nvidia_probe_message}")
  1983. if TORCH_AVAILABLE:
  1984. print(f" Torch device requested: {args.torch_device}")
  1985. print(f" Torch device effective: {torch_device_effective}")
  1986. if torch_device_effective.startswith("cuda") and torch is not None and torch.cuda.is_available():
  1987. print(f" CUDA device name: {torch.cuda.get_device_name(0)}")
  1988. elif nvidia_gpu_detected and (torch is None or not torch.cuda.is_available()):
  1989. print(
  1990. " WARNING: NVIDIA GPU detected but torch.cuda.is_available() is False. "
  1991. "Install a CUDA-enabled PyTorch build to run deep models on GPU."
  1992. )
  1993. elif nvidia_gpu_detected:
  1994. print(
  1995. " WARNING: NVIDIA GPU detected but PyTorch is unavailable. "
  1996. "Deep models require a CUDA-enabled PyTorch installation."
  1997. )
  1998. results_json = Path(args.results_json).resolve() if args.results_json else _default_results_json()
  1999. summary_md = Path(args.summary_md).resolve() if args.summary_md else _default_summary_md()
  2000. models_root = Path(args.models_root).resolve() if args.models_root else _models_root()
  2001. live_confusion_png_dir = (
  2002. Path(args.live_confusion_png_dir).resolve()
  2003. if args.live_confusion_png_dir
  2004. else _default_live_confusion_png_dir()
  2005. )
  2006. eeg_condition_psd_dir = (
  2007. Path(args.eeg_condition_psd_dir).resolve()
  2008. if args.eeg_condition_psd_dir
  2009. else _default_eeg_condition_psd_dir()
  2010. )
  2011. run_stamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")
  2012. run_tag = args.run_tag.strip() if args.run_tag else "default"
  2013. run_dir = models_root / f"stage6_{run_stamp}_{run_tag}"
  2014. run_dir.mkdir(parents=True, exist_ok=True)
  2015. print(f" Run dir: {run_dir}")
  2016. print(f" Live confusion PNGs: {bool(args.live_confusion_pngs)}")
  2017. if args.live_confusion_pngs:
  2018. live_confusion_png_dir.mkdir(parents=True, exist_ok=True)
  2019. print(f" Live confusion PNG dir: {live_confusion_png_dir}")
  2020. print(f" EEG condition PSD plots: {bool(args.eeg_condition_psd_plots)}")
  2021. if args.eeg_condition_psd_plots:
  2022. print(f" EEG condition PSD dir: {eeg_condition_psd_dir}")
  2023. dataset_payloads: dict[str, dict[str, Any]] = {}
  2024. dataset_stats: list[dict[str, Any]] = []
  2025. for dataset_idx, dataset_name in enumerate(datasets, start=1):
  2026. print(f"[Dataset {dataset_idx}/{len(datasets)}] Loading '{dataset_name}'...")
  2027. dataset_path = _resolve_dataset_path(
  2028. path_text=str(split_manifest["datasets"][dataset_name]["path"]),
  2029. split_manifest_path=split_manifest_path,
  2030. )
  2031. payload = _prepare_dataset(
  2032. dataset_name=dataset_name,
  2033. dataset_path=dataset_path,
  2034. target_col=target_col,
  2035. labels=labels,
  2036. class_scenario=class_scenario,
  2037. )
  2038. dataset_payloads[dataset_name] = payload
  2039. dataset_stats.append(
  2040. {
  2041. "dataset": dataset_name,
  2042. "path": payload["path"],
  2043. "rows": int(payload["X"].shape[0]),
  2044. "n_features": int(payload["X"].shape[1]),
  2045. "n_participants": int(len(set(payload["participants"].tolist()))),
  2046. "n_classes": int(len(set(payload["y"].tolist()))),
  2047. "class_counts": payload["class_counts"],
  2048. "full_class_counts": payload["full_class_counts"],
  2049. "class_counts_after_label_filter": payload["class_counts_after_label_filter"],
  2050. "rows_before_class_filter": payload["rows_before_class_filter"],
  2051. "rows_after_label_filter": payload["rows_after_label_filter"],
  2052. "rows_after_class_mapping": payload["rows_after_class_mapping"],
  2053. }
  2054. )
  2055. print(
  2056. f" Loaded '{dataset_name}': rows={payload['X'].shape[0]} "
  2057. f"features={payload['X'].shape[1]} participants={len(set(payload['participants'].tolist()))}"
  2058. )
  2059. eeg_condition_psd: dict[str, Any] | None = None
  2060. if args.eeg_condition_psd_plots:
  2061. epoch_manifest_path: Path | None = None
  2062. epoch_manifest_value = split_manifest.get("epoch_manifest")
  2063. if epoch_manifest_value:
  2064. epoch_manifest_path = _resolve_dataset_path(str(epoch_manifest_value), split_manifest_path)
  2065. eeg_payload = dataset_payloads.get("eeg")
  2066. eeg_condition_psd = build_condition_psd_report(
  2067. eeg_df=(eeg_payload["df"] if eeg_payload is not None else None),
  2068. epoch_manifest_path=epoch_manifest_path,
  2069. out_dir=eeg_condition_psd_dir / _slug(class_scenario["name"]),
  2070. scenario_name=class_scenario["name"],
  2071. label_order=labels,
  2072. fmin=float(args.eeg_condition_psd_fmin),
  2073. fmax=float(args.eeg_condition_psd_fmax),
  2074. )
  2075. print(f" EEG condition PSD status: {eeg_condition_psd.get('status', 'unknown')}")
  2076. print(f" EEG condition PSD manifest: {eeg_condition_psd.get('manifest_json', '')}")
  2077. if eeg_condition_psd.get("all_participants_png"):
  2078. print(f" EEG condition PSD all participants PNG: {eeg_condition_psd['all_participants_png']}")
  2079. if eeg_condition_psd.get("all_participants_topomap_png"):
  2080. print(f" EEG condition topomap all participants PNG: {eeg_condition_psd['all_participants_topomap_png']}")
  2081. evaluations: list[dict[str, Any]] = []
  2082. skip_counts: Counter[str] = Counter()
  2083. live_confusion_png_count = 0
  2084. global_strategies = split_manifest.get("strategies", {}) or {}
  2085. dataset_strategy_map = split_manifest.get("dataset_strategies", {}) or {}
  2086. dataset_protocol_total = len(datasets) * len(protocols)
  2087. dataset_protocol_idx = 0
  2088. for dataset_name in datasets:
  2089. payload = dataset_payloads[dataset_name]
  2090. X = payload["X"]
  2091. y = payload["y"]
  2092. groups = payload["groups"]
  2093. participants = payload["participants"]
  2094. dataset_strategies = dataset_strategy_map.get(dataset_name, {}) or {}
  2095. strategy_map = {
  2096. "loso": dataset_strategies.get(
  2097. "leave_one_participant_out",
  2098. global_strategies.get("leave_one_participant_out", []),
  2099. ),
  2100. "group_holdout": dataset_strategies.get(
  2101. "group_holdout",
  2102. global_strategies.get("group_holdout", []),
  2103. ),
  2104. "within_participant": dataset_strategies.get(
  2105. "within_participant",
  2106. global_strategies.get("within_participant", []),
  2107. ),
  2108. }
  2109. for protocol in protocols:
  2110. dataset_protocol_idx += 1
  2111. protocol_debug = _protocol_debug_label(protocol)
  2112. print(
  2113. f"[Evaluation {dataset_protocol_idx}/{dataset_protocol_total}] "
  2114. f"dataset={dataset_name} protocol={protocol} "
  2115. f"classification_set={protocol_debug} class_scenario={class_scenario['name']}"
  2116. )
  2117. if protocol in ("loso", "group_holdout"):
  2118. split_plan = _sample_plan_items(
  2119. items=list(strategy_map.get(protocol, [])),
  2120. max_count=args.max_outer_splits_per_protocol,
  2121. seed=args.random_seed + len(dataset_name) + len(protocol),
  2122. )
  2123. print(f" Planned outer splits: {len(split_plan)}")
  2124. for split_index, split_item in enumerate(split_plan, start=1):
  2125. train_participants = list(split_item.get("train_participants", []))
  2126. test_participants = list(split_item.get("test_participants", []))
  2127. train_idx, _ = _subset_by_participants(payload, train_participants)
  2128. test_idx, _ = _subset_by_participants(payload, test_participants)
  2129. if train_idx.size < 10 or test_idx.size < 2:
  2130. skip_counts[f"{dataset_name}:{protocol}:split_insufficient_rows"] += 1
  2131. continue
  2132. y_train = y[train_idx]
  2133. y_test = y[test_idx]
  2134. if len(set(y_train.tolist())) < 2:
  2135. skip_counts[f"{dataset_name}:{protocol}:split_single_train_class"] += 1
  2136. continue
  2137. X_train = X[train_idx]
  2138. X_test = X[test_idx]
  2139. groups_train = groups[train_idx]
  2140. split_name = str(split_item.get("split_id", f"{protocol}_{split_index:03d}"))
  2141. print(
  2142. f" Split {split_index}/{len(split_plan)} ({split_name}): "
  2143. f"train_rows={train_idx.size} test_rows={test_idx.size} "
  2144. f"model_selectors={len(models)}x{len(feature_selectors)}"
  2145. )
  2146. combo_total = len(models) * len(feature_selectors)
  2147. combo_idx = 0
  2148. for model_key in models:
  2149. model_spec = registry[model_key]
  2150. estimator = model_spec["estimator"]
  2151. model_label = str(model_spec.get("label", model_key))
  2152. model_grid = _prefix_param_grid(list(model_spec["grid"]), prefix="model")
  2153. for selector_key in feature_selectors:
  2154. combo_idx += 1
  2155. print(
  2156. f" [Train {combo_idx}/{combo_total}] "
  2157. f"dataset={dataset_name} class_scenario={class_scenario['name']} "
  2158. f"classification_set={protocol_debug} protocol={protocol} "
  2159. f"split={split_name} model={model_key} selector={selector_key}"
  2160. )
  2161. selector_spec = selector_registry[selector_key]
  2162. selector = selector_spec["selector"]
  2163. selector_label = str(selector_spec.get("label", selector_key))
  2164. selector_grid = _prefix_param_grid(list(selector_spec["grid"]), prefix="selector")
  2165. all_combos = _cross_param_grids(model_grid, selector_grid)
  2166. param_grid = _sample_param_grid(
  2167. combos=all_combos,
  2168. max_count=args.max_param_combos,
  2169. seed=(
  2170. args.random_seed
  2171. + split_index
  2172. + len(dataset_name)
  2173. + len(model_key)
  2174. + len(selector_key)
  2175. ),
  2176. )
  2177. best_params, tune_summary = _choose_best_params(
  2178. estimator=estimator,
  2179. feature_selector=selector,
  2180. param_grid=param_grid,
  2181. X_train=X_train,
  2182. y_train=y_train,
  2183. labels=labels,
  2184. args=args,
  2185. mode="grouped",
  2186. groups_train=groups_train,
  2187. seed=args.random_seed + split_index,
  2188. )
  2189. outer_warning_counter: Counter[str] = Counter()
  2190. y_pred = _fit_predict(
  2191. estimator=estimator,
  2192. feature_selector=selector,
  2193. params=best_params,
  2194. X_train=X_train,
  2195. y_train=y_train,
  2196. X_test=X_test,
  2197. args=args,
  2198. warning_counter=outer_warning_counter,
  2199. )
  2200. if y_pred is None:
  2201. skip_counts[f"{dataset_name}:{protocol}:fit_predict_failed"] += 1
  2202. continue
  2203. metrics = _metric_bundle(y_test, y_pred, labels=labels)
  2204. best_model_params, best_selector_params, best_other_params = _split_pipeline_params(best_params)
  2205. pipeline_id = f"{model_key}+{selector_key}"
  2206. eval_row: dict[str, Any] = {
  2207. "dataset": dataset_name,
  2208. "protocol": protocol,
  2209. "model": model_key,
  2210. "model_label": model_label,
  2211. "feature_selector": selector_key,
  2212. "feature_selector_label": selector_label,
  2213. "pipeline_id": pipeline_id,
  2214. "pipeline_label": f"{model_label} + {selector_label}",
  2215. "split_id": split_name,
  2216. "n_train_rows": int(train_idx.size),
  2217. "n_test_rows": int(test_idx.size),
  2218. "n_train_participants": int(len(set(participants[train_idx].tolist()))),
  2219. "n_test_participants": int(len(set(participants[test_idx].tolist()))),
  2220. "best_params": best_params,
  2221. "best_model_params": best_model_params,
  2222. "best_feature_selector_params": best_selector_params,
  2223. "best_pipeline_other_params": best_other_params,
  2224. "tuning": tune_summary,
  2225. "warning_counts": {
  2226. "tuning": tune_summary.get("warning_counts", {}),
  2227. "outer_fit_predict": dict(sorted(outer_warning_counter.items())),
  2228. },
  2229. "metrics": metrics,
  2230. }
  2231. if args.live_confusion_pngs:
  2232. live_png = _write_live_confusion_png(
  2233. metrics=metrics,
  2234. out_root=live_confusion_png_dir,
  2235. class_scenario_name=class_scenario["name"],
  2236. dataset=dataset_name,
  2237. protocol=protocol,
  2238. split_id=split_name,
  2239. pipeline_id=pipeline_id,
  2240. eval_index=len(evaluations) + 1,
  2241. cmap=str(args.live_confusion_png_cmap),
  2242. dpi=int(args.live_confusion_png_dpi),
  2243. )
  2244. live_confusion_png_count += 1
  2245. eval_row["confusion_png_live"] = str(live_png)
  2246. print(f" Live confusion PNG: {live_png}")
  2247. evaluations.append(eval_row)
  2248. elif protocol == "within_participant":
  2249. within_plan = _sample_plan_items(
  2250. items=list(strategy_map.get("within_participant", [])),
  2251. max_count=args.max_outer_splits_per_protocol,
  2252. seed=args.random_seed + len(dataset_name) + 17,
  2253. )
  2254. print(f" Candidate within-participant entries: {len(within_plan)}")
  2255. for within_item in within_plan:
  2256. participant_id = str(within_item.get("participant_id", ""))
  2257. if not participant_id:
  2258. skip_counts[f"{dataset_name}:{protocol}:missing_participant_id"] += 1
  2259. continue
  2260. if not bool(within_item.get("eligible_for_within_participant_cv", False)):
  2261. skip_counts[f"{dataset_name}:{protocol}:participant_not_eligible"] += 1
  2262. continue
  2263. participant_mask = participants == participant_id
  2264. idx_all = np.where(participant_mask)[0]
  2265. if idx_all.size < 8:
  2266. skip_counts[f"{dataset_name}:{protocol}:participant_rows_lt8"] += 1
  2267. continue
  2268. y_participant = y[idx_all]
  2269. counts = Counter(y_participant.tolist())
  2270. if len(counts) < 2:
  2271. skip_counts[f"{dataset_name}:{protocol}:participant_single_class"] += 1
  2272. continue
  2273. participant_seed = _stable_seed_from_text(participant_id, args.random_seed)
  2274. if args.within_participant_no_folds:
  2275. min_class_support = min(counts.values())
  2276. if min_class_support < 2:
  2277. skip_counts[f"{dataset_name}:{protocol}:class_support_lt2_no_fold"] += 1
  2278. continue
  2279. splitter = StratifiedShuffleSplit(
  2280. n_splits=1,
  2281. test_size=float(args.within_participant_test_size),
  2282. random_state=participant_seed,
  2283. )
  2284. split_iter = splitter.split(np.zeros(idx_all.size), y_participant)
  2285. split_name_prefix = "holdout"
  2286. print(
  2287. f" Participant {participant_id}: holdout_splits=1 "
  2288. f"test_size={args.within_participant_test_size:.2f} "
  2289. f"model_selectors={len(models)}x{len(feature_selectors)}"
  2290. )
  2291. else:
  2292. rec_splits = int(within_item.get("recommended_n_splits", 0))
  2293. if rec_splits < 2:
  2294. skip_counts[f"{dataset_name}:{protocol}:recommended_splits_lt2"] += 1
  2295. continue
  2296. max_splits_by_class = min(counts.values())
  2297. n_splits = min(rec_splits, int(max_splits_by_class))
  2298. if n_splits < 2:
  2299. skip_counts[f"{dataset_name}:{protocol}:folds_lt2_after_class_cap"] += 1
  2300. continue
  2301. splitter = StratifiedKFold(
  2302. n_splits=n_splits,
  2303. shuffle=True,
  2304. random_state=participant_seed,
  2305. )
  2306. split_iter = splitter.split(np.zeros(idx_all.size), y_participant)
  2307. split_name_prefix = "fold"
  2308. print(
  2309. f" Participant {participant_id}: folds={n_splits} "
  2310. f"model_selectors={len(models)}x{len(feature_selectors)}"
  2311. )
  2312. for split_idx, (inner_train_pos, inner_test_pos) in enumerate(
  2313. split_iter,
  2314. start=1,
  2315. ):
  2316. train_idx = idx_all[inner_train_pos]
  2317. test_idx = idx_all[inner_test_pos]
  2318. X_train = X[train_idx]
  2319. y_train = y[train_idx]
  2320. X_test = X[test_idx]
  2321. y_test = y[test_idx]
  2322. if len(set(y_train.tolist())) < 2:
  2323. skip_counts[f"{dataset_name}:{protocol}:fold_single_train_class"] += 1
  2324. continue
  2325. combo_total = len(models) * len(feature_selectors)
  2326. combo_idx = 0
  2327. for model_key in models:
  2328. model_spec = registry[model_key]
  2329. estimator = model_spec["estimator"]
  2330. model_label = str(model_spec.get("label", model_key))
  2331. model_grid = _prefix_param_grid(list(model_spec["grid"]), prefix="model")
  2332. for selector_key in feature_selectors:
  2333. combo_idx += 1
  2334. split_name = f"within_{participant_id}_{split_name_prefix}{split_idx:02d}"
  2335. print(
  2336. f" [Train {combo_idx}/{combo_total}] "
  2337. f"dataset={dataset_name} class_scenario={class_scenario['name']} "
  2338. f"classification_set={protocol_debug} protocol={protocol} "
  2339. f"split={split_name} model={model_key} selector={selector_key}"
  2340. )
  2341. selector_spec = selector_registry[selector_key]
  2342. selector = selector_spec["selector"]
  2343. selector_label = str(selector_spec.get("label", selector_key))
  2344. selector_grid = _prefix_param_grid(list(selector_spec["grid"]), prefix="selector")
  2345. all_combos = _cross_param_grids(model_grid, selector_grid)
  2346. param_grid = _sample_param_grid(
  2347. combos=all_combos,
  2348. max_count=args.max_param_combos,
  2349. seed=args.random_seed + split_idx + len(model_key) + len(selector_key),
  2350. )
  2351. best_params, tune_summary = _choose_best_params(
  2352. estimator=estimator,
  2353. feature_selector=selector,
  2354. param_grid=param_grid,
  2355. X_train=X_train,
  2356. y_train=y_train,
  2357. labels=labels,
  2358. args=args,
  2359. mode="stratified",
  2360. groups_train=None,
  2361. seed=args.random_seed + split_idx,
  2362. )
  2363. outer_warning_counter = Counter()
  2364. y_pred = _fit_predict(
  2365. estimator=estimator,
  2366. feature_selector=selector,
  2367. params=best_params,
  2368. X_train=X_train,
  2369. y_train=y_train,
  2370. X_test=X_test,
  2371. args=args,
  2372. warning_counter=outer_warning_counter,
  2373. )
  2374. if y_pred is None:
  2375. skip_counts[f"{dataset_name}:{protocol}:fit_predict_failed"] += 1
  2376. continue
  2377. metrics = _metric_bundle(y_test, y_pred, labels=labels)
  2378. best_model_params, best_selector_params, best_other_params = _split_pipeline_params(
  2379. best_params
  2380. )
  2381. pipeline_id = f"{model_key}+{selector_key}"
  2382. eval_row = {
  2383. "dataset": dataset_name,
  2384. "protocol": protocol,
  2385. "model": model_key,
  2386. "model_label": model_label,
  2387. "feature_selector": selector_key,
  2388. "feature_selector_label": selector_label,
  2389. "pipeline_id": pipeline_id,
  2390. "pipeline_label": f"{model_label} + {selector_label}",
  2391. "split_id": split_name,
  2392. "participant_id": participant_id,
  2393. "n_train_rows": int(train_idx.size),
  2394. "n_test_rows": int(test_idx.size),
  2395. "n_train_participants": 1,
  2396. "n_test_participants": 1,
  2397. "best_params": best_params,
  2398. "best_model_params": best_model_params,
  2399. "best_feature_selector_params": best_selector_params,
  2400. "best_pipeline_other_params": best_other_params,
  2401. "tuning": tune_summary,
  2402. "warning_counts": {
  2403. "tuning": tune_summary.get("warning_counts", {}),
  2404. "outer_fit_predict": dict(sorted(outer_warning_counter.items())),
  2405. },
  2406. "metrics": metrics,
  2407. }
  2408. if args.live_confusion_pngs:
  2409. live_png = _write_live_confusion_png(
  2410. metrics=metrics,
  2411. out_root=live_confusion_png_dir,
  2412. class_scenario_name=class_scenario["name"],
  2413. dataset=dataset_name,
  2414. protocol=protocol,
  2415. split_id=split_name,
  2416. pipeline_id=pipeline_id,
  2417. eval_index=len(evaluations) + 1,
  2418. cmap=str(args.live_confusion_png_cmap),
  2419. dpi=int(args.live_confusion_png_dpi),
  2420. )
  2421. live_confusion_png_count += 1
  2422. eval_row["confusion_png_live"] = str(live_png)
  2423. print(f" Live confusion PNG: {live_png}")
  2424. evaluations.append(eval_row)
  2425. elif protocol in ("pooled_stratified", "pooled_stratified_holdout", "pooled_stratified_one_fold"):
  2426. counts = Counter(y.tolist())
  2427. if len(counts) < 2:
  2428. skip_counts[f"{dataset_name}:{protocol}:single_class_dataset"] += 1
  2429. continue
  2430. n_splits = 1
  2431. split_total = 1
  2432. pooled_seed = _stable_seed_from_text(f"{dataset_name}:{class_scenario['name']}:pooled", args.random_seed)
  2433. protocol_forces_holdout = protocol == "pooled_stratified_holdout"
  2434. protocol_forces_one_fold = protocol == "pooled_stratified_one_fold"
  2435. use_pooled_holdout = bool(args.pooled_stratified_no_folds) or protocol_forces_holdout
  2436. use_single_pooled_fold = protocol_forces_one_fold
  2437. if use_pooled_holdout:
  2438. min_class_support = min(counts.values())
  2439. if min_class_support < 2:
  2440. skip_counts[f"{dataset_name}:{protocol}:class_support_lt2_no_fold"] += 1
  2441. continue
  2442. splitter = StratifiedShuffleSplit(
  2443. n_splits=1,
  2444. test_size=float(args.pooled_stratified_test_size),
  2445. random_state=pooled_seed,
  2446. )
  2447. split_name_prefix = "pooled_holdout"
  2448. else:
  2449. min_class_support = min(counts.values())
  2450. n_splits = min(int(args.pooled_stratified_folds), int(min_class_support))
  2451. if n_splits < 2:
  2452. skip_counts[f"{dataset_name}:{protocol}:folds_lt2_after_class_cap"] += 1
  2453. continue
  2454. splitter = StratifiedKFold(
  2455. n_splits=n_splits,
  2456. shuffle=True,
  2457. random_state=pooled_seed,
  2458. )
  2459. split_name_prefix = "pooled_fold"
  2460. split_total = n_splits
  2461. fold_plan = [{"fold_index": i + 1} for i in range(n_splits)]
  2462. fold_plan = _sample_plan_items(
  2463. items=fold_plan,
  2464. max_count=(1 if use_single_pooled_fold else args.max_outer_splits_per_protocol),
  2465. seed=args.random_seed + len(dataset_name) + 31,
  2466. )
  2467. fold_indices = [int(item.get("fold_index", 0)) for item in fold_plan if int(item.get("fold_index", 0)) > 0]
  2468. selected_fold_set = set(fold_indices)
  2469. if not selected_fold_set:
  2470. skip_counts[f"{dataset_name}:{protocol}:no_outer_folds_selected"] += 1
  2471. continue
  2472. if use_pooled_holdout:
  2473. print(
  2474. " Planned pooled holdout splits: "
  2475. f"{len(selected_fold_set)} (of {split_total}, test_size={args.pooled_stratified_test_size:.2f})"
  2476. )
  2477. else:
  2478. print(f" Planned pooled stratified folds: {len(selected_fold_set)} (of {split_total})")
  2479. for fold_idx, (train_idx, test_idx) in enumerate(splitter.split(np.zeros(X.shape[0]), y), start=1):
  2480. if fold_idx not in selected_fold_set:
  2481. continue
  2482. if train_idx.size < 10 or test_idx.size < 2:
  2483. skip_counts[f"{dataset_name}:{protocol}:split_insufficient_rows"] += 1
  2484. continue
  2485. y_train = y[train_idx]
  2486. y_test = y[test_idx]
  2487. if len(set(y_train.tolist())) < 2:
  2488. skip_counts[f"{dataset_name}:{protocol}:split_single_train_class"] += 1
  2489. continue
  2490. X_train = X[train_idx]
  2491. X_test = X[test_idx]
  2492. split_name = f"{split_name_prefix}{fold_idx:02d}"
  2493. print(
  2494. f" Split {fold_idx}/{split_total} ({split_name}): "
  2495. f"train_rows={train_idx.size} test_rows={test_idx.size} "
  2496. f"model_selectors={len(models)}x{len(feature_selectors)}"
  2497. )
  2498. combo_total = len(models) * len(feature_selectors)
  2499. combo_idx = 0
  2500. for model_key in models:
  2501. model_spec = registry[model_key]
  2502. estimator = model_spec["estimator"]
  2503. model_label = str(model_spec.get("label", model_key))
  2504. model_grid = _prefix_param_grid(list(model_spec["grid"]), prefix="model")
  2505. for selector_key in feature_selectors:
  2506. combo_idx += 1
  2507. print(
  2508. f" [Train {combo_idx}/{combo_total}] "
  2509. f"dataset={dataset_name} class_scenario={class_scenario['name']} "
  2510. f"classification_set={protocol_debug} protocol={protocol} "
  2511. f"split={split_name} model={model_key} selector={selector_key}"
  2512. )
  2513. selector_spec = selector_registry[selector_key]
  2514. selector = selector_spec["selector"]
  2515. selector_label = str(selector_spec.get("label", selector_key))
  2516. selector_grid = _prefix_param_grid(list(selector_spec["grid"]), prefix="selector")
  2517. all_combos = _cross_param_grids(model_grid, selector_grid)
  2518. param_grid = _sample_param_grid(
  2519. combos=all_combos,
  2520. max_count=args.max_param_combos,
  2521. seed=args.random_seed + fold_idx + len(model_key) + len(selector_key),
  2522. )
  2523. best_params, tune_summary = _choose_best_params(
  2524. estimator=estimator,
  2525. feature_selector=selector,
  2526. param_grid=param_grid,
  2527. X_train=X_train,
  2528. y_train=y_train,
  2529. labels=labels,
  2530. args=args,
  2531. mode="stratified",
  2532. groups_train=None,
  2533. seed=args.random_seed + fold_idx,
  2534. )
  2535. outer_warning_counter = Counter()
  2536. y_pred = _fit_predict(
  2537. estimator=estimator,
  2538. feature_selector=selector,
  2539. params=best_params,
  2540. X_train=X_train,
  2541. y_train=y_train,
  2542. X_test=X_test,
  2543. args=args,
  2544. warning_counter=outer_warning_counter,
  2545. )
  2546. if y_pred is None:
  2547. skip_counts[f"{dataset_name}:{protocol}:fit_predict_failed"] += 1
  2548. continue
  2549. metrics = _metric_bundle(y_test, y_pred, labels=labels)
  2550. best_model_params, best_selector_params, best_other_params = _split_pipeline_params(
  2551. best_params
  2552. )
  2553. pipeline_id = f"{model_key}+{selector_key}"
  2554. eval_row = {
  2555. "dataset": dataset_name,
  2556. "protocol": protocol,
  2557. "model": model_key,
  2558. "model_label": model_label,
  2559. "feature_selector": selector_key,
  2560. "feature_selector_label": selector_label,
  2561. "pipeline_id": pipeline_id,
  2562. "pipeline_label": f"{model_label} + {selector_label}",
  2563. "split_id": split_name,
  2564. "n_train_rows": int(train_idx.size),
  2565. "n_test_rows": int(test_idx.size),
  2566. "n_train_participants": int(len(set(participants[train_idx].tolist()))),
  2567. "n_test_participants": int(len(set(participants[test_idx].tolist()))),
  2568. "best_params": best_params,
  2569. "best_model_params": best_model_params,
  2570. "best_feature_selector_params": best_selector_params,
  2571. "best_pipeline_other_params": best_other_params,
  2572. "tuning": tune_summary,
  2573. "warning_counts": {
  2574. "tuning": tune_summary.get("warning_counts", {}),
  2575. "outer_fit_predict": dict(sorted(outer_warning_counter.items())),
  2576. },
  2577. "metrics": metrics,
  2578. }
  2579. if args.live_confusion_pngs:
  2580. live_png = _write_live_confusion_png(
  2581. metrics=metrics,
  2582. out_root=live_confusion_png_dir,
  2583. class_scenario_name=class_scenario["name"],
  2584. dataset=dataset_name,
  2585. protocol=protocol,
  2586. split_id=split_name,
  2587. pipeline_id=pipeline_id,
  2588. eval_index=len(evaluations) + 1,
  2589. cmap=str(args.live_confusion_png_cmap),
  2590. dpi=int(args.live_confusion_png_dpi),
  2591. )
  2592. live_confusion_png_count += 1
  2593. eval_row["confusion_png_live"] = str(live_png)
  2594. print(f" Live confusion PNG: {live_png}")
  2595. evaluations.append(eval_row)
  2596. aggregate_rows = _aggregate_metric_rows(evaluations)
  2597. best_rows = _best_model_by_dataset_protocol(aggregate_rows)
  2598. warning_counts_total = _aggregate_warning_counts(evaluations)
  2599. skip_counts_total = {k: skip_counts[k] for k in sorted(skip_counts, key=lambda x: (-skip_counts[x], x))}
  2600. coverage = _dataset_protocol_coverage(
  2601. aggregate_rows=aggregate_rows,
  2602. datasets=datasets,
  2603. protocols=protocols,
  2604. min_aggregate_rows_per_pair=args.min_aggregate_rows_per_dataset_protocol,
  2605. )
  2606. config = {
  2607. "bids_root": str(bids_root),
  2608. "split_manifest": str(split_manifest_path),
  2609. "target": split_manifest.get("target"),
  2610. "target_column": target_col,
  2611. "class_scenario": class_scenario,
  2612. "datasets": datasets,
  2613. "protocols": protocols,
  2614. "models": models,
  2615. "feature_selectors": feature_selectors,
  2616. "available_models": _available_model_names(),
  2617. "available_feature_selectors": _available_feature_selector_names(),
  2618. "torch_available": TORCH_AVAILABLE,
  2619. "torch_device_requested": args.torch_device,
  2620. "torch_device_effective": torch_device_effective,
  2621. "nvidia_gpu_detected": bool(nvidia_gpu_detected),
  2622. "nvidia_gpu_names": nvidia_gpu_names,
  2623. "nvidia_probe_message": nvidia_probe_message,
  2624. "inner_folds": args.inner_folds,
  2625. "max_param_combos": args.max_param_combos,
  2626. "max_outer_splits_per_protocol": args.max_outer_splits_per_protocol,
  2627. "pooled_stratified_folds": int(args.pooled_stratified_folds),
  2628. "pooled_stratified_no_folds": bool(args.pooled_stratified_no_folds),
  2629. "pooled_stratified_test_size": float(args.pooled_stratified_test_size),
  2630. "within_participant_no_folds": bool(args.within_participant_no_folds),
  2631. "within_participant_test_size": float(args.within_participant_test_size),
  2632. "min_aggregate_rows_per_dataset_protocol": args.min_aggregate_rows_per_dataset_protocol,
  2633. "allow_incomplete_coverage": bool(args.allow_incomplete_coverage),
  2634. "clip_lower_quantile": args.clip_lower_quantile,
  2635. "clip_upper_quantile": args.clip_upper_quantile,
  2636. "live_confusion_pngs": bool(args.live_confusion_pngs),
  2637. "live_confusion_png_dir": str(live_confusion_png_dir),
  2638. "live_confusion_png_cmap": args.live_confusion_png_cmap,
  2639. "live_confusion_png_dpi": int(args.live_confusion_png_dpi),
  2640. "eeg_condition_psd_plots": bool(args.eeg_condition_psd_plots),
  2641. "eeg_condition_psd_dir": str(eeg_condition_psd_dir),
  2642. "eeg_condition_psd_fmin": float(args.eeg_condition_psd_fmin),
  2643. "eeg_condition_psd_fmax": float(args.eeg_condition_psd_fmax),
  2644. "random_seed": args.random_seed,
  2645. "run_dir": str(run_dir),
  2646. "run_tag": run_tag,
  2647. "timestamp_utc": run_stamp,
  2648. }
  2649. results = {
  2650. "config": config,
  2651. "dataset_stats": dataset_stats,
  2652. "evaluations": evaluations,
  2653. "aggregates": aggregate_rows,
  2654. "best_models": best_rows,
  2655. "warning_counts_total": warning_counts_total,
  2656. "skip_counts_total": skip_counts_total,
  2657. "coverage": coverage,
  2658. "eeg_condition_psd": eeg_condition_psd,
  2659. "counts": {
  2660. "n_evaluations": len(evaluations),
  2661. "n_aggregate_rows": len(aggregate_rows),
  2662. "n_live_confusion_pngs": int(live_confusion_png_count),
  2663. },
  2664. }
  2665. results_json.parent.mkdir(parents=True, exist_ok=True)
  2666. results_json.write_text(json.dumps(results, indent=2) + "\n", encoding="utf-8")
  2667. summary_text = _build_summary_markdown(
  2668. config=config,
  2669. dataset_stats=dataset_stats,
  2670. aggregate_rows=aggregate_rows,
  2671. best_rows=best_rows,
  2672. warning_counts_total=warning_counts_total,
  2673. coverage=coverage,
  2674. skip_counts=skip_counts_total,
  2675. eeg_condition_psd=eeg_condition_psd,
  2676. )
  2677. summary_md.parent.mkdir(parents=True, exist_ok=True)
  2678. summary_md.write_text(summary_text, encoding="utf-8")
  2679. run_meta = {
  2680. "results_json": str(results_json),
  2681. "summary_md": str(summary_md),
  2682. "n_evaluations": len(evaluations),
  2683. "n_aggregate_rows": len(aggregate_rows),
  2684. "n_live_confusion_pngs": int(live_confusion_png_count),
  2685. "live_confusion_png_dir": str(live_confusion_png_dir),
  2686. "best_models": best_rows,
  2687. "warning_counts_total": warning_counts_total,
  2688. "skip_counts_total": skip_counts_total,
  2689. "coverage": coverage,
  2690. "eeg_condition_psd": eeg_condition_psd,
  2691. }
  2692. (run_dir / "run_meta.json").write_text(json.dumps(run_meta, indent=2) + "\n", encoding="utf-8")
  2693. (run_dir / "run_config.json").write_text(json.dumps(config, indent=2) + "\n", encoding="utf-8")
  2694. missing_pairs = list(coverage.get("missing_pairs", []))
  2695. if missing_pairs:
  2696. print(" Coverage gaps detected:")
  2697. for pair in missing_pairs:
  2698. print(
  2699. f" - dataset={pair['dataset']} protocol={pair['protocol']} "
  2700. f"aggregate_rows={int(pair.get('aggregate_rows', 0))}"
  2701. )
  2702. if not args.allow_incomplete_coverage:
  2703. raise RuntimeError(
  2704. "Stage 6 coverage check failed: missing dataset+protocol aggregate rows. "
  2705. "Use --allow-incomplete-coverage or set --min-aggregate-rows-per-dataset-protocol 0 to bypass."
  2706. )
  2707. print("Stage 6 complete.")
  2708. print(f" Evaluations: {len(evaluations)}")
  2709. print(f" Aggregate rows: {len(aggregate_rows)}")
  2710. print(f" Results JSON: {results_json}")
  2711. print(f" Summary Markdown: {summary_md}")
  2712. print(f" Model run dir: {run_dir}")
  2713. print(f" Class scenario: {class_scenario['name']} ({len(labels)} classes)")
  2714. if eeg_condition_psd is not None:
  2715. print(f" EEG condition PSD manifest: {eeg_condition_psd.get('manifest_json', '')}")
  2716. if args.live_confusion_pngs:
  2717. print(f" Live confusion PNGs written: {live_confusion_png_count}")
  2718. print(f" Live confusion PNG dir: {live_confusion_png_dir}")
  2719. print(
  2720. " Coverage: "
  2721. f"missing_pairs={len(missing_pairs)} "
  2722. f"(min_aggregate_rows_per_pair={coverage['min_aggregate_rows_per_pair']})"
  2723. )
  2724. print(f" Elapsed seconds: {time.time() - stage_started:.1f}")
  2725. if warning_counts_total:
  2726. print(f" Captured warnings: {sum(warning_counts_total.values())}")
  2727. if skip_counts_total:
  2728. print(f" Skip events: {sum(skip_counts_total.values())}")
  2729. if __name__ == "__main__":
  2730. main()

stage6_train_classic_ml.py at commit a26b91d, under CC0-1.0 · at the source

Overview

  1. Faculty of Science and Engineering, University of Hull, Hull HU6 7RX, UK
  2. School of Digital and Physical Sciences, Faculty of Science and Engineering, University of Hull, Hull HU6 7RX, UK
Institutions: University of Hull (United Kingdom)
Journal: Bioengineering (Basel, Switzerland), volume 13, issue 7, article 820
Dates: received 20 May 2026; accepted 3 July 2026; published online 16 July 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.3390/bioengineering13070820 · PMID 42510485 · PMCID PMC13405904 · OpenAlex W7169591137
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: EEG (modality), other (modality), human (organism), methods / tools (subfield)
Methods: Spectral & time-frequency, Preprocessing, Connectivity, Statistics, Smoothing, state filtering, decompositions, Machine learning, fMRI & imaging, Evoked potentials, Physiology & signal measures
Keywords: mental workload, EEG, ECG, pupillometry, multimodal fusion, machine learning benchmark, segmentation, cross-validation, open data, ds007262
Topic: EEG and Brain-Computer Interfaces (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Citations: not cited yet (Europe PMC); 30 references in the paper

Abstract

Physiology-based mental workload classification is hard to compare across studies because task design, preprocessing, segmentation, and validation protocols vary widely. Using OpenNeuro ds007262, an open multimodal arithmetic dataset of synchronised 19-channel 10–20 system electroencephalography (EEG), electrocardiography (ECG), and pupillometry data from 18 released participants (16 retained after participant-level quality control for downstream modelling) spanning seven objective difficulty bands plus baseline fixation, we present a reproducible end-to-end machine learning pipeline for graded workload classification. The pipeline standardises participant-level quality control, trial-aligned windowing, modality-specific preprocessing and feature extraction (153 EEG, 18 ECG, and 25 pupillometry features), and supervised evaluation under three validation protocols (within-participant, pooled-stratified, and group-holdout) over eleven models and five class scenarios. Fused representations generally performed best; EEG was the strongest unimodal modality, and classical models outperformed deep models in most feature-based conditions. The best 6 s pipeline reached a balanced accuracy of 0.635; under denser 3 s overlap segmentation, the best pipeline reached 0.718. Mean balanced accuracy across 60 matched cells rose from 0.324 to 0.380, with gains concentrated in within-participant and pooled-stratified evaluation rather than strict unseen-participant transfer. The pipeline provides a transparent benchmark framework for fine-grained physiological workload modelling.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repositories

Its files are read in the Code ↔ Paper reader above, with 21 matches between paragraphs and lines of code.

LMBooth/Arithmetic_Workload_Estimation

License: CC0-1.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: a26b91d6a212a49d78affc2e58253924097d8768, 8 June 2026
Languages: Python (18)
Size: 30 files, 18 scripts
Software Heritage: not archived
Found in: “Data Availability Statement”
Holds: README, license file, environment (requirements.txt), documentation
Not found: CITATION.cff, tests, continuous integration
Tools: NumPy (10 files), Matplotlib (6 files), pandas (6 files), SciPy (5 files), MNE-Python (4 files), PyTorch (1 file), scikit-learn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
20 files

Zenodo 20162864

License: CC0-1.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: the references
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (10 files), Matplotlib (6 files), pandas (6 files), SciPy (5 files), MNE-Python (4 files), PyTorch (1 file), scikit-learn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)
20 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 36 scripts, each with its path and the digest of its content;
  • 21 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

Data Availability Statement

The original physiological recordings analysed in this study are openly available in OpenNeuro at https://doi.org/10.18112/openneuro.ds007262.v1.1.0 (accession ds007262, version 1.1.0; accessed on 2 July 2026) [7]. The complete analysis pipeline is openly available on GitHub (https://github.com/LMBooth/Arithmetic_Workload_Estimation, accessed on 2 July 2026) and archived on Zenodo at version 0.0.4, DOI https://doi.org/10.5281/zenodo.20162864 (accessed on 8 June 2026) [15], released under the CC0 1.0 Universal licence. The two checked-in YAML configuration profiles reproduce the Stage 0 through Stage 6 results from a single command, and Stage 7 (stage7_significance) reproduces every inferential statistic reported in Section 3.3, including the paired Wilcoxon W=1636, the bootstrap CIs in Table 7, and the label-shuffle permutation p-values.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 3 authors, 10 keywords, 27 references.

Cite

This paper

Booth, L., Mehmood, A., & Zeinali, M. (2026). Benchmarking Multimodal Workload Classification: Effects of Modality, Validation Protocol, and Segmentation Contrast on an Open Graded-Arithmetic Dataset. Bioengineering (Basel, Switzerland), 13(7), 820. https://doi.org/10.3390/bioengineering13070820

BibTeX

@article{booth2026benchmarking,
author = {Booth, Liam and Mehmood, Adeel and Zeinali, Mehdi},
title = {{Benchmarking Multimodal Workload Classification: Effects of Modality, Validation Protocol, and Segmentation Contrast on an Open Graded-Arithmetic Dataset}},
journal = {Bioengineering (Basel, Switzerland)},
year = {2026},
month = jul,
volume = {13},
number = {7},
pages = {820},
publisher = {Multidisciplinary Digital Publishing Institute (MDPI)},
issn = {2306-5354},
doi = {10.3390/bioengineering13070820},
url = {https://doi.org/10.3390/bioengineering13070820},
pmid = {42510485},
pmcid = {PMC13405904}
}

RIS

TY - JOUR
AU - Booth, Liam
AU - Mehmood, Adeel
AU - Zeinali, Mehdi
TI - Benchmarking Multimodal Workload Classification: Effects of Modality, Validation Protocol, and Segmentation Contrast on an Open Graded-Arithmetic Dataset
T2 - Bioengineering (Basel, Switzerland)
J2 - Bioengineering (Basel)
PY - 2026
DA - 2026/07/16
VL - 13
IS - 7
SP - 820
SN - 2306-5354
PB - Multidisciplinary Digital Publishing Institute (MDPI)
DO - 10.3390/bioengineering13070820
UR - https://doi.org/10.3390/bioengineering13070820
LA - en
ER -

CSL-JSON

{
"id": "10.3390/bioengineering13070820",
"type": "article-journal",
"title": "Benchmarking Multimodal Workload Classification: Effects of Modality, Validation Protocol, and Segmentation Contrast on an Open Graded-Arithmetic Dataset",
"container-title": "Bioengineering (Basel, Switzerland)",
"author": [
{
"family": "Booth",
"given": "Liam"
},
{
"family": "Mehmood",
"given": "Adeel"
},
{
"family": "Zeinali",
"given": "Mehdi"
}
],
"container-title-short": "Bioengineering (Basel)",
"volume": "13",
"issue": "7",
"page": "820",
"DOI": "10.3390/bioengineering13070820",
"PMID": "42510485",
"PMCID": "PMC13405904",
"ISSN": "2306-5354",
"publisher": "Multidisciplinary Digital Publishing Institute (MDPI)",
"URL": "https://doi.org/10.3390/bioengineering13070820",
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
16
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1038/s41597-026-07120-7 [code]
PhysioMotion Artifact: A task-driven EEG dataset with point-wise motion artifact annotations.
Journal: Scientific data
In common: MNE-Python, PyTorch, scikit-learn, 3 other tools, methods / tools, EEG, 2 references
[2] doi:10.3389/fpsyg.2026.1774068 [code]
Analysis of cognitive mechanisms in phoneme perception and pronunciation errors among Korean language learners.
Journal: Frontiers in psychology
In common: MNE-Python, PyTorch, scikit-learn, 4 other tools, EEG, 1 reference
[3] doi:10.1038/s41597-025-05174-7 [code]
A large-scale MEG and EEG dataset for object recognition in naturalistic scenes
Journal: n/a
In common: MNE-Python, PyTorch, scikit-learn, 4 other tools, methods / tools, EEG, 1 reference
[4] doi:10.1038/s41598-026-52330-z [code]
SHAP analysis of an improved EEG-based mental workload classification framework: utilizing data augmentation and explainable AI.
Journal: Scientific reports
In common: MNE-Python, scikit-learn, pandas, 3 other tools, methods / tools, EEG, 1 reference
[5] doi:10.1038/s41597-026-07077-7 [code]
Everyday Activity Science and Engineering Table Setting Dataset.
Journal: Scientific data
In common: PyTorch, scikit-learn, pandas, 3 other tools, other, methods / tools, EEG, 1 reference
[6] doi:10.1167/jov.26.8.4 [code]
The neural processes of illusory occlusion in object recognition.
Journal: Journal of vision
In common: MNE-Python, scikit-learn, pandas, 3 other tools, EEG, 2 references
[7] doi:10.1371/journal.pone.0351872 [code]
Decoding visual object recognition from EEG signals.
Journal: PloS one
In common: MNE-Python, PyTorch, scikit-learn, 4 other tools, EEG, 1 reference
[8] doi:10.1371/journal.pcbi.1014302 [code]
Trial-level sequence modeling reveals hidden dynamics of dual-task interference.
Journal: PLoS computational biology
In common: MNE-Python, PyTorch, scikit-learn, 4 other tools, EEG, 1 reference
[9] doi:10.1162/imag.a.1207 [code]
Investigating the temporal dynamics and modeling of mid-level feature representations in humans.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: MNE-Python, PyTorch, scikit-learn, 4 other tools, EEG, 1 reference
[10] doi:10.1038/s41598-026-41532-0 [code]
Prediction, syntax and semantic grounding in the brain and large language models.
Journal: Scientific reports
In common: MNE-Python, PyTorch, scikit-learn, 4 other tools, EEG, 1 reference

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.