diff --git a/machine-learning-pack/__init__.py b/machine-learning-pack/__init__.py new file mode 100644 index 0000000..bd201d4 --- /dev/null +++ b/machine-learning-pack/__init__.py @@ -0,0 +1,2804 @@ +import base64 +import io +import json +import os +import time +from itertools import combinations +from pathlib import Path + +import numpy as np +import pandas as pd +from sklearn.metrics import accuracy_score, balanced_accuracy_score + +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt + +from PySide6.QtGui import QAction + +from PySide6.QtCore import Qt, QThread, Signal +from PySide6.QtWidgets import ( QApplication, QCheckBox, QComboBox, QFileDialog, QFormLayout, QGroupBox, QHBoxLayout, QHeaderView, QLabel, QLineEdit, + QMainWindow, QMenu, QMessageBox, QPushButton, QScrollArea, QTabWidget, QTableWidget, QTableWidgetItem, QTextEdit, QVBoxLayout, QWidget,) +from scipy import stats + +from sklearn.pipeline import Pipeline +from sklearn.preprocessing import StandardScaler +from sklearn.feature_selection import SelectKBest, f_classif +from sklearn.linear_model import LogisticRegression +from sklearn.svm import SVC +from sklearn.discriminant_analysis import LinearDiscriminantAnalysis +from sklearn.model_selection import LeaveOneGroupOut, GridSearchCV, StratifiedKFold +from sklearn.metrics import accuracy_score, balanced_accuracy_score + + + +# Metric Registry mapping short keys to CSV column suffixes +METRIC_REGISTRY = { + 'Peak_Amp': 'Peak_Amp', + 'TTP': 'Time_to_Peak', + 'AUC': 'AUC', + 'Rising_Slope': 'Rising_Slope', + 'Recovery_Slope': 'Recovery_Slope', + 'FWHM': 'FWHM', + 'Onset_Latency': 'Onset_Latency', + 'P2P_Amp': 'Peak_to_Peak_Amp', + 'Signal_Std': 'Signal_Std', + 'RMS': 'RMS' +} + +# Neutral, factual definitions of each metric (used for user-facing explanations) +METRIC_DESCRIPTIONS = { + 'Peak_Amp': "the maximum amplitude reached during the hemodynamic response", + 'TTP': "how long it took the signal to reach its peak amplitude", + 'AUC': "the total area under the response curve (overall response magnitude over time)", + 'Rising_Slope': "how steeply the signal rises on its way up to the peak", + 'Recovery_Slope': "how steeply the signal declines from its peak back toward baseline", + 'FWHM': "the width of the response at half its maximum height (how long the response stays elevated)", + 'Onset_Latency': "how long after stimulus onset the response begins to rise", + 'P2P_Amp': "the difference between the highest and lowest points of the response", + 'Signal_Std': "the overall variability (standard deviation) of the signal", + 'RMS': "the root-mean-square magnitude of the signal, a measure of overall signal energy", +} + +# --- Modeling Settings --- +K_FEATURES = 5 +MODEL_RANDOM_STATE = 42 +MODEL_C = 0.5 +TUNE_HYPERPARAMETERS = True +CANDIDATE_C_VALUES = [5.0, 10.0, 50.0] +INNER_CV_FOLDS = 3 + +CLASSIFIER_FACTORIES = { + 'L1_Logistic': lambda: LogisticRegression( + l1_ratio=1.0, solver='liblinear', max_iter=2000, C=MODEL_C, random_state=MODEL_RANDOM_STATE, + class_weight='balanced', + ), + 'Linear_SVM': lambda: SVC( + kernel='linear', C=MODEL_C, random_state=MODEL_RANDOM_STATE, + class_weight='balanced', + ), + 'LDA': lambda: LinearDiscriminantAnalysis(solver='lsqr', shrinkage='auto'), +} + + +CONFIG_FILE = "splitter_config.json" +DEFAULT_CSV_PATH = "" +SWEEP_RESULTS_CSV = "metric_sweep_leaderboard.csv" +ALL_METRIC_KEYS = tuple(METRIC_REGISTRY.keys()) + + +def sigmoid(x): + return 1.0 / (1.0 + np.exp(-np.clip(x, -50, 50))) + + +def _fmt_confidence(conf): + """ + Standard 1-decimal formatting rounds anything above 99.95% to a flat '100.0%', which + hides the actual precision the sigmoid produced. Switch to higher precision near the + extremes so 'why is this 100%' has a real number to look at instead of a rounded one. + """ + if conf >= 99.95 or conf <= 0.05: + return f"{conf:.4f}%" + return f"{conf:.1f}%" + + +def _fmt_val(v): + """Format small/large feature values readably; tiny magnitudes fall back to scientific notation.""" + if v == 0: + return "0" + if abs(v) < 1e-3: + return f"{v:.3e}" + return f"{v:.4f}" + + +def _split_feature_name(feat_name): + """ + Splits a raw column name like 'FIR_Reach_ROI_Somatosensory_S1_BA40_Part_2_Recovery_Slope' + into (readable_roi_label, metric_key). Falls back gracefully if no metric suffix matches. + """ + metric_map = {} + for k, v in METRIC_REGISTRY.items(): + metric_map[k] = k + metric_map[v] = k + sorted_patterns = sorted(metric_map.keys(), key=len, reverse=True) + + matched_key = None + matched_pattern = None + for pattern in sorted_patterns: + if feat_name.endswith(f"_{pattern}"): + matched_key = metric_map[pattern] + matched_pattern = pattern + break + + if matched_pattern: + stem = feat_name[: -(len(matched_pattern) + 1)] + else: + stem = feat_name + + parts = stem.split('_') + # Drop leading FIR__ prefix if present + if len(parts) >= 2 and parts[0] == 'FIR': + parts = parts[2:] + + roi_label = " ".join(parts) if parts else stem + return roi_label, matched_key + + +def _render_feature_distribution_png(class_a_values, class_b_values, subject_val, + class_a_mean, class_b_mean, actual_label, + metric_label, roi_label, z_score): + """ + Renders a single feature's training-set distribution for both classes, with this test + subject's value overlaid on top, as a base64-encoded PNG. This is deliberately not a bar + chart: individual training subjects are shown as points so you can see spread and overlap, + not just a single average, and the subject's true-class row is tinted green so it's obvious + at a glance how far off they were from where they were "supposed" to land. + """ + rng = np.random.default_rng(0) + fig, ax = plt.subplots(figsize=(7.0, 2.5), dpi=130) + + row_a_y, row_b_y = 1.0, 0.0 + band_h = 0.42 + + correct_y = row_a_y if actual_label == "Class_A" else row_b_y + incorrect_y = row_b_y if actual_label == "Class_A" else row_a_y + ax.axhspan(correct_y - band_h, correct_y + band_h, color="#2e7d32", alpha=0.08, zorder=0) + ax.axhspan(incorrect_y - band_h, incorrect_y + band_h, color="#c62828", alpha=0.05, zorder=0) + + if len(class_a_values) > 0: + ja = rng.uniform(-0.18, 0.18, size=len(class_a_values)) + ax.scatter(class_a_values, row_a_y + ja, s=26, color="#1f77b4", alpha=0.6, zorder=3) + if len(class_b_values) > 0: + jb = rng.uniform(-0.18, 0.18, size=len(class_b_values)) + ax.scatter(class_b_values, row_b_y + jb, s=26, color="#ff7f0e", alpha=0.6, zorder=3) + + ax.axvline(class_a_mean, color="#1f77b4", linestyle="--", linewidth=1.3, zorder=2) + ax.axvline(class_b_mean, color="#ff7f0e", linestyle="--", linewidth=1.3, zorder=2) + + ax.axvline(subject_val, color="black", linewidth=1.6, zorder=4) + ax.scatter([subject_val], [row_a_y], marker="*", s=260, color="gold", edgecolor="black", linewidth=0.8, zorder=5) + ax.scatter([subject_val], [row_b_y], marker="*", s=260, color="gold", edgecolor="black", linewidth=0.8, zorder=5) + + ymin, ymax = row_b_y - band_h - 0.18, row_a_y + band_h + 0.4 + ax.set_ylim(ymin, ymax) + + a_true = actual_label == "Class_A" + b_true = actual_label == "Class_B" + a_label = "Class_A" + (" \u2190 true class" if a_true else "") + b_label = "Class_B" + (" \u2190 true class" if b_true else "") + ax.set_yticks([row_a_y, row_b_y]) + ax.set_yticklabels([a_label, b_label], fontsize=9) + ax.get_yticklabels()[0].set_color("#2e7d32" if a_true else "#333333") + ax.get_yticklabels()[1].set_color("#2e7d32" if b_true else "#333333") + + ax.text(class_a_mean, ymax - 0.06, f"A mean\n{_fmt_val(class_a_mean)}", color="#1f77b4", + fontsize=7.5, ha="center", va="top") + ax.text(class_b_mean, ymax - 0.06, f"B mean\n{_fmt_val(class_b_mean)}", color="#ff7f0e", + fontsize=7.5, ha="center", va="top") + ax.text(subject_val, ymin + 0.03, f"this subject: {_fmt_val(subject_val)} (z={z_score:+.2f} SD)", + color="black", fontsize=7.5, ha="center", va="bottom", fontweight="bold") + + ax.set_title(f"{metric_label} \u2014 {roi_label}", fontsize=9.5, loc="left") + ax.set_xlabel("Raw feature value \u2014 each dot is one training subject", fontsize=8) + for spine in ("top", "right", "left"): + ax.spines[spine].set_visible(False) + ax.tick_params(axis="x", labelsize=8) + fig.tight_layout() + + buf = io.BytesIO() + fig.savefig(buf, format="png") + plt.close(fig) + return base64.b64encode(buf.getvalue()).decode("ascii") + + +class PermutationTestWorker(QThread): + log_signal = Signal(str) + finished_signal = Signal(dict) + error_signal = Signal(str) + + def __init__(self, X_raw, y, groups, feature_names, feature_channels, quality_list, + model_name, k_features, scoring='accuracy', n_permutations=200): + super().__init__() + self.X_raw = X_raw + self.y = y + self.groups = groups + self.feature_names = feature_names + self.feature_channels = feature_channels + self.quality_list = quality_list + self.model_name = model_name + self.k_features = k_features + self.scoring = scoring + self.n_permutations = n_permutations + + def run(self): + try: + self.log_signal.emit( + f"Running permutation test for {self.model_name} " + f"({self.n_permutations} label shuffles, scoring={self.scoring})..." + ) + + def progress_cb(done, total): + if done % 20 == 0 or done == total: + self.log_signal.emit(f" Permutation {done}/{total}") + + result = run_permutation_test( + X_raw=self.X_raw, y=self.y, groups=self.groups, + feature_names=self.feature_names, feature_channels=self.feature_channels, + quality_list=self.quality_list, + classifier_factories={self.model_name: CLASSIFIER_FACTORIES[self.model_name]}, + k_features=self.k_features, tune_hyperparameters=False, scoring=self.scoring, + n_permutations=self.n_permutations, progress_callback=progress_cb, + ) + self.log_signal.emit("Permutation test complete.") + self.finished_signal.emit(result[self.model_name]) + except Exception as e: + self.error_signal.emit(str(e)) + + +class PairedPermutationWorker(QThread): + log_signal = Signal(str) + finished_signal = Signal(dict) + error_signal = Signal(str) + + def __init__(self, X_raw, y, pair_ids, feature_names, feature_channels, quality_list, + model_name, k_features, scoring='accuracy', n_permutations=200): + super().__init__() + self.X_raw = X_raw + self.y = y + self.pair_ids = pair_ids + self.feature_names = feature_names + self.feature_channels = feature_channels + self.quality_list = quality_list + self.model_name = model_name + self.k_features = k_features + self.scoring = scoring + self.n_permutations = n_permutations + + def run(self): + try: + n_pairs = len(np.unique(self.pair_ids)) + self.log_signal.emit( + f"Running paired permutation test for {self.model_name} " + f"({n_pairs} subject pairs, {self.n_permutations} sign-flip shuffles, " + f"scoring={self.scoring})..." + ) + + def progress_cb(done, total): + if done % 20 == 0 or done == total: + self.log_signal.emit(f" Permutation {done}/{total}") + + result = run_paired_permutation_test( + X_raw=self.X_raw, y=self.y, pair_ids=self.pair_ids, + feature_names=self.feature_names, feature_channels=self.feature_channels, + quality_list=self.quality_list, + classifier_factories={self.model_name: CLASSIFIER_FACTORIES[self.model_name]}, + k_features=self.k_features, tune_hyperparameters=False, scoring=self.scoring, + n_permutations=self.n_permutations, progress_callback=progress_cb, + ) + self.log_signal.emit("Paired permutation test complete.") + self.finished_signal.emit(result[self.model_name]) + except Exception as e: + self.error_signal.emit(str(e)) + + +class PairedDiagnosticsWorker(QThread): + log_signal = Signal(str) + finished_signal = Signal(dict) + error_signal = Signal(str) + + def __init__(self, X_raw, y, pair_ids, feature_names, feature_channels, quality_list, + model_name, k_features, scoring='accuracy', n_bootstrap=500): + super().__init__() + self.X_raw = X_raw + self.y = y + self.pair_ids = pair_ids + self.feature_names = feature_names + self.feature_channels = feature_channels + self.quality_list = quality_list + self.model_name = model_name + self.k_features = k_features + self.scoring = scoring + self.n_bootstrap = n_bootstrap + + def run(self): + try: + n_pairs = len(np.unique(self.pair_ids)) + self.log_signal.emit( + f"Running paired diagnostics for {self.model_name} " + f"({n_pairs} pairs, {self.n_bootstrap} bootstrap resamples, " + f"scoring={self.scoring})..." + ) + + def progress_cb(done, total): + if done % 50 == 0 or done == total: + self.log_signal.emit(f" Bootstrap {done}/{total}") + + result = paired_accuracy_diagnostics( + X_raw=self.X_raw, y=self.y, pair_ids=self.pair_ids, + feature_names=self.feature_names, feature_channels=self.feature_channels, + quality_list=self.quality_list, + classifier_factory=CLASSIFIER_FACTORIES[self.model_name], + k_features=self.k_features, n_bootstrap=self.n_bootstrap, scoring=self.scoring, + progress_callback=progress_cb, + ) + self.log_signal.emit("Paired diagnostics complete.") + self.finished_signal.emit(result) + except Exception as e: + self.error_signal.emit(str(e)) + + +class MaxStatWorker(QThread): + log_signal = Signal(str) + finished_signal = Signal(dict) + error_signal = Signal(str) + + def __init__(self, X_raw, y, groups, feature_combos, combo_labels, feature_names, + model_name, k_features, scoring='accuracy', n_permutations=100): + super().__init__() + self.X_raw = X_raw + self.y = y + self.groups = groups + self.feature_combos = feature_combos + self.combo_labels = combo_labels + self.feature_names = feature_names + self.model_name = model_name + self.k_features = k_features + self.scoring = scoring + self.n_permutations = n_permutations + + def run(self): + try: + self.log_signal.emit( + f"Running max-stat permutation test: {len(self.feature_combos)} combos x " + f"{self.n_permutations} shuffles (this can take a while)..." + ) + + def progress_cb(done, total): + if done % 10 == 0 or done == total: + self.log_signal.emit(f" Permutation {done}/{total}") + + result = run_max_stat_permutation_test( + X_raw=self.X_raw, y=self.y, groups=self.groups, + feature_combos=self.feature_combos, feature_names=self.feature_names, + classifier_factories={self.model_name: CLASSIFIER_FACTORIES[self.model_name]}, + k_features=self.k_features, tune_hyperparameters=False, scoring=self.scoring, + n_permutations=self.n_permutations, progress_callback=progress_cb, + ) + result['combo_labels'] = self.combo_labels + self.log_signal.emit("Max-stat permutation test complete.") + self.finished_signal.emit(result) + except Exception as e: + self.error_signal.emit(str(e)) + + +class MetricSweepWorker(QThread): + log_signal = Signal(str) + finished_signal = Signal(pd.DataFrame) + error_signal = Signal(str) + + def __init__(self, df, subject_col, target_condition, train_label_map, test_label_map, + data_root=None, scoring='accuracy'): + super().__init__() + self.df = df + self.subject_col = subject_col + self.target_condition = target_condition + self.train_label_map = train_label_map + self.test_label_map = test_label_map + self.data_root = data_root + self.scoring = scoring + + def run(self): + try: + self.log_signal.emit("Initializing metric-combination sweep from memory...") + + extractor = FNIRSClassificationPipeline( + training_subjects=list(self.train_label_map.keys()), + data_root=self.data_root, + target_condition=self.target_condition, + selected_metrics=ALL_METRIC_KEYS + ) + extractor.load_from_dataframe(self.df, self.subject_col, self.train_label_map) + + unseen_extractor = FNIRSClassificationPipeline( + training_subjects=list(self.test_label_map.keys()), + data_root=self.data_root, + target_condition=self.target_condition, + selected_metrics=ALL_METRIC_KEYS + ) + unseen_extractor.load_from_dataframe(self.df, self.subject_col, self.test_label_map) + + feature_keys = build_metric_keys(extractor.feature_names) + unseen_keys = build_metric_keys(unseen_extractor.feature_names) + + singles = list(combinations(ALL_METRIC_KEYS, 1)) + pairs = list(combinations(ALL_METRIC_KEYS, 2)) + all_combos = singles + pairs + + self.log_signal.emit(f"Running Stage 1 & 2 ({len(all_combos)} runs, scoring={self.scoring})...") + results = [] + + for idx, combo in enumerate(all_combos, 1): + res = self._eval_combo( + combo, extractor, unseen_extractor, feature_keys, unseen_keys + ) + results.append(res) + if idx % 10 == 0 or idx == len(all_combos): + self.log_signal.emit(f" Completed {idx}/{len(all_combos)} combos") + + ranked = self._rank_metrics(results, ALL_METRIC_KEYS) + top_5 = [m for m, _ in ranked[:5]] + self.log_signal.emit(f"Top 5 metrics: {', '.join(top_5)}") + + stage3_combos = [] + for size in (3, 4, 5): + stage3_combos.extend(list(combinations(top_5, size))) + + self.log_signal.emit(f"Running Stage 3 ({len(stage3_combos)} runs)...") + for combo in stage3_combos: + res = self._eval_combo( + combo, extractor, unseen_extractor, feature_keys, unseen_keys + ) + results.append(res) + + res_df = pd.DataFrame(results).sort_values( + ['avg_test_acc', 'avg_train_acc'], ascending=False + ).reset_index(drop=True) + res_df.to_csv(SWEEP_RESULTS_CSV, index=False) + + self.log_signal.emit("Sweep complete!") + self.finished_signal.emit(res_df) + + except Exception as e: + self.error_signal.emit(str(e)) + + def _eval_combo(self, keys, ex, uex, f_keys, u_keys): + col_idx = np.where(np.isin(f_keys, keys))[0] + unseen_col_idx = np.where(np.isin(u_keys, keys))[0] + + X_sub = ex.X_raw[:, col_idx] + X_unseen_sub = uex.X_raw[:, unseen_col_idx] + + model_accs, fitted_pipes, _ = run_loso_cv( + X_raw=X_sub, y=ex.y, groups=ex.groups, + feature_names=[ex.feature_names[i] for i in col_idx], + feature_channels=[ex.feature_channels[i] for i in col_idx], + quality_list=ex.quality_list, + classifier_factories=CLASSIFIER_FACTORIES, + k_features=min(K_FEATURES, X_sub.shape[1]), + verbose=False, + tune_hyperparameters=TUNE_HYPERPARAMETERS, + scoring=self.scoring, + ) + + score_fn = accuracy_score if self.scoring == 'accuracy' else balanced_accuracy_score + + unseen_accs = {} + if len(uex.y) > 0 and X_unseen_sub.shape[1] > 0: + for name, pipe in fitted_pipes.items(): + unseen_accs[name] = float(score_fn(uex.y, pipe.predict(X_unseen_sub))) + else: + unseen_accs = {name: 0.0 for name in fitted_pipes} + + return { + 'combo_label': "+".join(keys), + 'size': len(keys), + 'avg_train_acc': float(np.mean(list(model_accs.values()))), + 'avg_test_acc': float(np.mean(list(unseen_accs.values()))), + 'L1_Logistic_train': model_accs.get('L1_Logistic', 0), + 'Linear_SVM_train': model_accs.get('Linear_SVM', 0), + 'LDA_train': model_accs.get('LDA', 0), + 'L1_Logistic_test': unseen_accs.get('L1_Logistic', 0), + 'Linear_SVM_test': unseen_accs.get('Linear_SVM', 0), + 'LDA_test': unseen_accs.get('LDA', 0), + } + + def _rank_metrics(self, results, all_keys): + scores = {m: [] for m in all_keys} + for r in results: + for m in r['combo_label'].split('+'): + scores[m].append(r['avg_train_acc']) + return sorted([(m, np.mean(v) if v else 0) for m, v in scores.items()], key=lambda x: x[1], reverse=True) + + +class DataSplitterGUI(QWidget): + + def __init__(self, parent: QWidget | None = None) -> None: + super().__init__(parent) + self.setWindowTitle("fNIRS Data Splitter & Model Inspector") + self.resize(1200, 850) + + self.df = None + self.subject_col = None + self.all_subjects = [] + self.sweep_thread = None + self.sweep_results_df = None + self.test_combos = {} + + # Leaderboard sort state + self.leaderboard_sort_col = None + self.leaderboard_sort_asc = True + self.sweep_results_df = None + + self.breakdown_sort_col = None + self.breakdown_sort_asc = True + self.breakdown_df = None + + self.perm_thread = None + self.current_train_pipe = None + self.current_k_actual = None + + self.paired_mode = False + self.paired_thread = None + + self.diagnostics_thread = None + self.maxstat_thread = None + + self.use_balanced_accuracy = False + + # Per-row diagnostic details for the currently displayed Test File Breakdown table, + # aligned by row index. Populated in update_test_breakdown(), consumed by Tab 4. + self.current_breakdown_details = [] + + self.init_ui() + self.load_initial_data() + + def _scoring(self): + return 'balanced_accuracy' if self.use_balanced_accuracy else 'accuracy' + + def init_ui(self): + top_layout = QVBoxLayout(self) + + self.tabs = QTabWidget() + self.tabs.currentChanged.connect(self.on_tab_changed) + top_layout.addWidget(self.tabs) + + # TAB 1: SPLITTER & SETUP + tab_split = QWidget() + main_layout = QVBoxLayout(tab_split) + + file_box = QGroupBox("1. Dataset Path & Target Condition") + file_layout = QVBoxLayout(file_box) + + mode_row = QHBoxLayout() + self.paired_mode_checkbox = QCheckBox("Paired (Within-Subjects) Mode") + self.paired_mode_checkbox.toggled.connect(self.toggle_paired_mode) + mode_row.addWidget(self.paired_mode_checkbox) + + self.balanced_acc_checkbox = QCheckBox("Use Balanced Accuracy (recommended for imbalanced test sets)") + self.balanced_acc_checkbox.toggled.connect(self.toggle_balanced_accuracy) + mode_row.addWidget(self.balanced_acc_checkbox) + mode_row.addStretch() + file_layout.addLayout(mode_row) + + path_row = QHBoxLayout() + self.path_input = QLineEdit() + self.path_input.setText(DEFAULT_CSV_PATH if os.path.exists(DEFAULT_CSV_PATH) else "") + path_row.addWidget(self.path_input) + + browse_btn = QPushButton("Browse...") + browse_btn.clicked.connect(self.browse_file) + path_row.addWidget(browse_btn) + + load_btn = QPushButton("Load CSV") + load_btn.setStyleSheet("font-weight: bold; background-color: #2b5b84; color: white;") + load_btn.clicked.connect(self.load_csv) + path_row.addWidget(load_btn) + file_layout.addLayout(path_row) + + cond_row = QHBoxLayout() + cond_row.addWidget(QLabel("Target Condition:")) + self.condition_combo = QComboBox() + cond_row.addWidget(self.condition_combo) + cond_row.addStretch() + file_layout.addLayout(cond_row) + + main_layout.addWidget(file_box) + + table_box = QGroupBox("2. Assign Training & Testing Groups") + table_layout = QVBoxLayout(table_box) + self.table = QTableWidget() + self.table.setColumnCount(5) + self.table.setHorizontalHeaderLabels( + ["Subject ID / Path", "Class A (Train)", "Class B (Train)", "Testing Pool", "Pair ID"] + ) + self.table.horizontalHeader().setSectionResizeMode(0, QHeaderView.Stretch) + for col in range(1, 5): + self.table.horizontalHeader().setSectionResizeMode(col, QHeaderView.Interactive) + self.table.setColumnWidth(col, 130) + table_layout.addWidget(self.table) + main_layout.addWidget(table_box) + + self.test_box = QGroupBox("3. Assign Testing Ground Truth") + self.test_layout = QFormLayout(self.test_box) + self.test_scroll = QScrollArea() + self.test_scroll.setWidgetResizable(True) + self.test_scroll.setWidget(self.test_box) + self.test_scroll.setMaximumHeight(160) + main_layout.addWidget(self.test_scroll) + + btn_layout = QHBoxLayout() + load_config_btn = QPushButton("Load Saved Config") + load_config_btn.clicked.connect(self.load_config_to_gui) + btn_layout.addWidget(load_config_btn) + btn_layout.addStretch() + + self.run_sweep_btn = QPushButton("Run Metric-Combination Sweep") + self.run_sweep_btn.setStyleSheet("font-weight: bold; background-color: #d32f2f; color: white; padding: 8px 12px;") + self.run_sweep_btn.clicked.connect(self.start_sweep) + btn_layout.addWidget(self.run_sweep_btn) + + self.paired_model_select = QComboBox() + self.paired_model_select.addItems(list(CLASSIFIER_FACTORIES.keys())) + self.paired_model_select.setVisible(False) + btn_layout.addWidget(self.paired_model_select) + + self.run_paired_btn = QPushButton("Run Paired Analysis") + self.run_paired_btn.setStyleSheet("font-weight: bold; background-color: #6a1b9a; color: white; padding: 8px 12px;") + self.run_paired_btn.clicked.connect(self.start_paired_analysis) + self.run_paired_btn.setVisible(False) + btn_layout.addWidget(self.run_paired_btn) + + self.run_diagnostics_btn = QPushButton("Paired Diagnostics (CI + Power + Bayes)") + self.run_diagnostics_btn.setStyleSheet("font-weight: bold; background-color: #00695c; color: white; padding: 8px 12px;") + self.run_diagnostics_btn.clicked.connect(self.start_paired_diagnostics) + self.run_diagnostics_btn.setVisible(False) + btn_layout.addWidget(self.run_diagnostics_btn) + + main_layout.addLayout(btn_layout) + + btn_layout2 = QHBoxLayout() + btn_layout2.addStretch() + + self.run_univariate_btn = QPushButton("Univariate Paired Tests (FDR)") + self.run_univariate_btn.setStyleSheet("font-weight: bold; background-color: #ef6c00; color: white; padding: 8px 12px;") + self.run_univariate_btn.clicked.connect(self.start_univariate_test) + self.run_univariate_btn.setVisible(False) + btn_layout2.addWidget(self.run_univariate_btn) + + self.run_stability_btn = QPushButton("Feature Selection Stability") + self.run_stability_btn.setStyleSheet("font-weight: bold; background-color: #455a64; color: white; padding: 8px 12px;") + self.run_stability_btn.clicked.connect(self.start_feature_stability) + self.run_stability_btn.setVisible(False) + btn_layout2.addWidget(self.run_stability_btn) + + main_layout.addLayout(btn_layout2) + self.tabs.addTab(tab_split, "Dataset Splitter") + + # TAB 2: SWEEP LOGS & LEADERBOARD + tab_sweep = QWidget() + sweep_layout = QVBoxLayout(tab_sweep) + + sweep_layout.addWidget(QLabel("Sweep Output & Real-time Logs:")) + self.log_output = QTextEdit() + self.log_output.setReadOnly(True) + self.log_output.setMaximumHeight(130) + sweep_layout.addWidget(self.log_output) + + leaderboard_header_row = QHBoxLayout() + leaderboard_header_row.addWidget(QLabel( + "Leaderboard (click a row to inspect file breakdown in Tab 3; click a column header to sort by it):" + )) + leaderboard_header_row.addStretch() + self.run_maxstat_btn = QPushButton("Test Search Significance (Max-Stat)") + self.run_maxstat_btn.setStyleSheet("font-weight: bold; background-color: #ad1457; color: white; padding: 4px 10px;") + self.run_maxstat_btn.clicked.connect(self.start_max_stat_test) + leaderboard_header_row.addWidget(self.run_maxstat_btn) + sweep_layout.addLayout(leaderboard_header_row) + + self.leaderboard_table = QTableWidget() + self.leaderboard_table.verticalHeader().setVisible(False) + self.leaderboard_table.horizontalHeader().setSectionResizeMode(QHeaderView.Interactive) + self.leaderboard_table.horizontalHeader().setSectionsClickable(True) + self.leaderboard_table.horizontalHeader().sectionClicked.connect(self.on_leaderboard_header_clicked) + self.leaderboard_table.setSelectionBehavior(QTableWidget.SelectRows) + self.leaderboard_table.cellClicked.connect(self.on_leaderboard_row_clicked) + sweep_layout.addWidget(self.leaderboard_table) + + self.tabs.addTab(tab_sweep, "Metric Sweep Results") + + # TAB 3: TEST FILE BREAKDOWN & EXPLANATIONS + tab_breakdown = QWidget() + bd_layout = QVBoxLayout(tab_breakdown) + + sel_box = QGroupBox("Select Model & Metric Combination to Inspect") + sel_form = QHBoxLayout(sel_box) + + sel_form.addWidget(QLabel("Metric Combination:")) + self.inspect_combo_select = QComboBox() + self.inspect_combo_select.currentIndexChanged.connect(self.update_test_breakdown) + sel_form.addWidget(self.inspect_combo_select, stretch=2) + + sel_form.addWidget(QLabel("Model Architecture:")) + self.inspect_model_select = QComboBox() + self.inspect_model_select.addItems(list(CLASSIFIER_FACTORIES.keys())) + self.inspect_model_select.currentIndexChanged.connect(self.update_test_breakdown) + sel_form.addWidget(self.inspect_model_select, stretch=1) + + self.perm_test_btn = QPushButton("Test Significance") + self.perm_test_btn.setStyleSheet("font-weight: bold; background-color: #6a1b9a; color: white;") + self.perm_test_btn.clicked.connect(self.start_permutation_test) + sel_form.addWidget(self.perm_test_btn) + + bd_layout.addWidget(sel_box) + + self.summary_banner = QLabel("Select a completed sweep or configuration to view test predictions.") + self.summary_banner.setStyleSheet("font-size: 14px; font-weight: bold; color: #2b5b84; padding: 6px;") + bd_layout.addWidget(self.summary_banner) + + bd_layout.addWidget(QLabel("Click any row to see a full per-feature explanation in Tab 4:")) + + self.breakdown_table = QTableWidget() + self.breakdown_table.setColumnCount(8) + self.breakdown_table.verticalHeader().setVisible(False) + self.breakdown_table.horizontalHeader().setSectionResizeMode(QHeaderView.Interactive) + self.breakdown_table.horizontalHeader().setSectionsClickable(True) + self.breakdown_table.horizontalHeader().sectionClicked.connect(self.on_breakdown_header_clicked) + self.breakdown_table.setSelectionBehavior(QTableWidget.SelectRows) + self.breakdown_table.cellClicked.connect(self.on_breakdown_row_clicked) + bd_layout.addWidget(self.breakdown_table) + + self.tabs.addTab(tab_breakdown, "Test File Breakdown") + + # TAB 4: FILE DEEP DIVE + tab_deep_dive = QWidget() + dd_layout = QVBoxLayout(tab_deep_dive) + + self.deep_dive_banner = QLabel("Click a subject row in Tab 3 to see a detailed feature-by-feature explanation here.") + self.deep_dive_banner.setStyleSheet("font-size: 14px; font-weight: bold; color: #2b5b84; padding: 6px;") + self.deep_dive_banner.setWordWrap(True) + dd_layout.addWidget(self.deep_dive_banner) + + self.deep_dive_text = QTextEdit() + self.deep_dive_text.setReadOnly(True) + dd_layout.addWidget(self.deep_dive_text) + + self.tabs.addTab(tab_deep_dive, "File Deep Dive") + + def resizeEvent(self, event): + super().resizeEvent(event) + self.redistribute_active_tab_columns() + + def on_tab_changed(self, index): + self.redistribute_active_tab_columns() + + def redistribute_active_tab_columns(self): + current_index = self.tabs.currentIndex() + if current_index == 1 and hasattr(self, 'leaderboard_table'): + self.distribute_table_columns_evenly(self.leaderboard_table) + elif current_index == 2 and hasattr(self, 'breakdown_table'): + self.distribute_table_columns_evenly(self.breakdown_table) + + def distribute_table_columns_evenly(self, table): + total_width = table.viewport().width() + col_count = table.columnCount() + if col_count > 0 and total_width > 50: # Ensures tab layout has actually rendered + col_width = total_width // col_count + for col in range(col_count): + table.setColumnWidth(col, col_width) + return True + return False + + def detect_conditions(self): + conditions = set() + if self.df is not None: + for col in self.df.columns: + parts = col.split('_') + if len(parts) >= 3 and parts[0] == 'FIR': + conditions.add(parts[1]) + return sorted(list(conditions)) if conditions else ["Reach"] + + def browse_file(self): + filename, _ = QFileDialog.getOpenFileName(self, "Select CSV", "", "CSV Files (*.csv)") + if filename: + self.path_input.setText(filename) + self.load_csv() + + def load_initial_data(self): + filepath = self.path_input.text().strip() + if filepath and os.path.exists(filepath): + self.load_csv() + + def load_csv(self): + filepath = self.path_input.text().strip() + if not os.path.exists(filepath): + QMessageBox.warning(self, "Error", f"File path does not exist:\n{filepath}") + return + + try: + self.df = pd.read_csv(filepath) + self.subject_col = "Participant" if "Participant" in self.df.columns else ("Subject" if "Subject" in self.df.columns else self.df.columns[0]) + self.all_subjects = sorted(self.df[self.subject_col].astype(str).unique().tolist()) + + self.condition_combo.clear() + self.condition_combo.addItems(self.detect_conditions()) + + self.populate_table() + + if os.path.exists(CONFIG_FILE): + self.load_config_to_gui(silent=True) + if os.path.exists(SWEEP_RESULTS_CSV): + try: + self.sweep_results_df = pd.read_csv(SWEEP_RESULTS_CSV) + self.leaderboard_sort_col = None + self.leaderboard_sort_asc = False + self.populate_leaderboard(self.sweep_results_df) + except Exception: + pass + except Exception as e: + QMessageBox.critical(self, "Error", f"Failed to load CSV file:\n{str(e)}") + + def populate_table(self): + self.table.blockSignals(True) + + self.table.setRowCount(0) + self.table.setRowCount(len(self.all_subjects)) + + for row, sub in enumerate(self.all_subjects): + sub_item = QTableWidgetItem(sub) + sub_item.setFlags(sub_item.flags() ^ Qt.ItemIsEditable) + self.table.setItem(row, 0, sub_item) + + chk_a = QTableWidgetItem() + chk_a.setFlags(Qt.ItemIsUserCheckable | Qt.ItemIsEnabled) + chk_a.setCheckState(Qt.Unchecked) + self.table.setItem(row, 1, chk_a) + + chk_b = QTableWidgetItem() + chk_b.setFlags(Qt.ItemIsUserCheckable | Qt.ItemIsEnabled) + chk_b.setCheckState(Qt.Unchecked) + self.table.setItem(row, 2, chk_b) + + test_item = QTableWidgetItem("Testing Pool") + test_item.setFlags(test_item.flags() ^ Qt.ItemIsEditable) + self.table.setItem(row, 3, test_item) + + pair_item = QTableWidgetItem("") + self.table.setItem(row, 4, pair_item) + + self.table.itemChanged.connect(self.on_table_cell_changed) + self.table.blockSignals(False) + self.update_testing_labels() + + def on_table_cell_changed(self, item): + row, col = item.row(), item.column() + self.table.blockSignals(True) + + if col == 1 and item.checkState() == Qt.Checked: + self.table.item(row, 2).setCheckState(Qt.Unchecked) + elif col == 2 and item.checkState() == Qt.Checked: + self.table.item(row, 1).setCheckState(Qt.Unchecked) + + is_a = self.table.item(row, 1).checkState() == Qt.Checked + is_b = self.table.item(row, 2).checkState() == Qt.Checked + + if self.paired_mode: + if is_a: + self.table.item(row, 3).setText("Condition 1") + elif is_b: + self.table.item(row, 3).setText("Condition 2") + else: + self.table.item(row, 3).setText("Unassigned") + else: + if is_a: + self.table.item(row, 3).setText("Class A (Train)") + elif is_b: + self.table.item(row, 3).setText("Class B (Train)") + else: + self.table.item(row, 3).setText("Testing Pool") + + self.table.blockSignals(False) + self.update_testing_labels() + + def update_testing_labels(self): + while self.test_layout.count(): + child = self.test_layout.takeAt(0) + if child.widget(): + child.widget().deleteLater() + + self.test_combos = {} + + for row in range(self.table.rowCount()): + sub_id = self.table.item(row, 0).text() + is_a = self.table.item(row, 1).checkState() == Qt.Checked + is_b = self.table.item(row, 2).checkState() == Qt.Checked + + if not is_a and not is_b: + combo = QComboBox() + combo.addItems(["Class_A", "Class_B"]) + self.test_combos[sub_id] = combo + self.test_layout.addRow(QLabel(f"Test Subject {os.path.basename(sub_id)} Ground Truth:"), combo) + + def get_groups(self): + class_a_subs, class_b_subs, testing_subs = [], [], [] + for row in range(self.table.rowCount()): + sub_id = self.table.item(row, 0).text() + if self.table.item(row, 1).checkState() == Qt.Checked: + class_a_subs.append(sub_id) + elif self.table.item(row, 2).checkState() == Qt.Checked: + class_b_subs.append(sub_id) + else: + testing_subs.append(sub_id) + + test_labels = {} + for sub_id in testing_subs: + if sub_id in self.test_combos: + test_labels[sub_id] = self.test_combos[sub_id].currentText() + + return class_a_subs, class_b_subs, testing_subs, test_labels + + def toggle_paired_mode(self, checked): + self.paired_mode = checked + headers = ( + ["Subject ID / Path", "Condition 1 (No Blindfold)", "Condition 2 (Blindfold)", "Status", "Pair ID"] + if checked else + ["Subject ID / Path", "Class A (Train)", "Class B (Train)", "Testing Pool", "Pair ID"] + ) + self.table.setHorizontalHeaderLabels(headers) + self.run_sweep_btn.setVisible(not checked) + self.paired_model_select.setVisible(checked) + self.run_paired_btn.setVisible(checked) + self.run_diagnostics_btn.setVisible(checked) + self.run_univariate_btn.setVisible(checked) + self.run_stability_btn.setVisible(checked) + self.refresh_status_column() + + def toggle_balanced_accuracy(self, checked): + self.use_balanced_accuracy = checked + + def refresh_status_column(self): + self.table.blockSignals(True) + for row in range(self.table.rowCount()): + is_a = self.table.item(row, 1).checkState() == Qt.Checked + is_b = self.table.item(row, 2).checkState() == Qt.Checked + if self.paired_mode: + text = "Condition 1" if is_a else ("Condition 2" if is_b else "Unassigned") + else: + text = "Class A (Train)" if is_a else ("Class B (Train)" if is_b else "Testing Pool") + self.table.item(row, 3).setText(text) + self.table.blockSignals(False) + + def get_paired_groups(self): + """ + Reads the table in paired mode and returns (cond1_subs, cond2_subs, subject_to_pair, + skipped). cond1_subs[i] and cond2_subs[i] are always the same physical subject + (they are appended together within the same loop iteration below), which matters + for any code that needs row-to-row alignment between the two conditions, not just + correct group labeling. Only pair IDs with exactly one Condition-1 row and one + Condition-2 row are kept -- there is no diff/pairing to compute for a subject who + only has one condition. + """ + pair_rows = {} + for row in range(self.table.rowCount()): + sub_id = self.table.item(row, 0).text() + pair_id = self.table.item(row, 4).text().strip() + if not pair_id: + continue + is_c1 = self.table.item(row, 1).checkState() == Qt.Checked + is_c2 = self.table.item(row, 2).checkState() == Qt.Checked + if not is_c1 and not is_c2: + continue + entry = pair_rows.setdefault(pair_id, {}) + entry['cond1' if is_c1 else 'cond2'] = sub_id + + cond1_subs, cond2_subs, subject_to_pair = [], [], {} + skipped = [] + for pair_id, entry in pair_rows.items(): + if 'cond1' in entry and 'cond2' in entry: + cond1_subs.append(entry['cond1']) + cond2_subs.append(entry['cond2']) + subject_to_pair[entry['cond1']] = pair_id + subject_to_pair[entry['cond2']] = pair_id + else: + skipped.append(pair_id) + + return cond1_subs, cond2_subs, subject_to_pair, skipped + + def _build_paired_pipes(self, cond1_subs, cond2_subs): + target_condition = self.condition_combo.currentText() + cond1_map = {s: "Condition_1" for s in cond1_subs} + cond2_map = {s: "Condition_2" for s in cond2_subs} + + pipe1 = FNIRSClassificationPipeline( + training_subjects=cond1_subs, target_condition=target_condition, + selected_metrics=ALL_METRIC_KEYS + ).load_from_dataframe(self.df, self.subject_col, cond1_map) + + pipe2 = FNIRSClassificationPipeline( + training_subjects=cond2_subs, target_condition=target_condition, + selected_metrics=ALL_METRIC_KEYS + ).load_from_dataframe(self.df, self.subject_col, cond2_map) + + return pipe1, pipe2 + + def start_paired_analysis(self): + cond1_subs, cond2_subs, subject_to_pair, skipped = self.get_paired_groups() + + if len(cond1_subs) < 3: + QMessageBox.warning( + self, "Error", + f"Need at least 3 complete pairs (Pair ID present on both a Condition 1 " + f"and Condition 2 row). Found {len(cond1_subs)} complete pair(s), " + f"{len(skipped)} incomplete." + ) + return + + pipe1, pipe2 = self._build_paired_pipes(cond1_subs, cond2_subs) + + if pipe1.X_raw.shape[0] == 0 or pipe2.X_raw.shape[0] == 0: + QMessageBox.warning(self, "Error", "No feature data found for the paired subjects.") + return + + X_raw = np.vstack([pipe1.X_raw, pipe2.X_raw]) + y = np.concatenate([np.zeros(pipe1.X_raw.shape[0], dtype=int), + np.ones(pipe2.X_raw.shape[0], dtype=int)]) + pair_ids = np.array( + [subject_to_pair[s] for s in pipe1.groups] + [subject_to_pair[s] for s in pipe2.groups] + ) + + model_name = self.paired_model_select.currentText() + k_actual = max(1, min(K_FEATURES, X_raw.shape[1])) + + self.tabs.setCurrentIndex(1) + self.run_paired_btn.setEnabled(False) + self.append_log( + f"--- Paired permutation test: {model_name} | {len(cond1_subs)} pairs " + f"({len(skipped)} incomplete pairs skipped) | scoring={self._scoring()} ---" + ) + + self.paired_thread = PairedPermutationWorker( + X_raw=X_raw, y=y, pair_ids=pair_ids, + feature_names=pipe1.feature_names, feature_channels=pipe1.feature_channels, + quality_list=pipe1.quality_list, model_name=model_name, + k_features=k_actual, scoring=self._scoring(), n_permutations=200, + ) + self.paired_thread.log_signal.connect(self.append_log) + self.paired_thread.finished_signal.connect(self.on_paired_finished) + self.paired_thread.error_signal.connect(self.on_paired_error) + self.paired_thread.start() + + def on_paired_finished(self, result): + self.run_paired_btn.setEnabled(True) + obs = result['observed_acc'] + null = result['null_accs'] + p = result['p_value'] + null_mean = float(np.mean(null)) if null else 0.0 + + self.append_log(f"Paired LOSO acc: {obs:.2%} | null mean: {null_mean:.2%} | p={p:.4f}") + + verdict = "distinguishable from chance (p < 0.05)" if p < 0.05 else "NOT distinguishable from chance" + QMessageBox.information( + self, "Paired Permutation Test Result", + f"Observed leave-one-pair-out accuracy: {obs:.2%}\n" + f"Null distribution mean ({len(null)} sign-flip shuffles): {null_mean:.2%}\n" + f"p-value: {p:.4f}\n\n" + f"Result is {verdict}." + ) + + def on_paired_error(self, err_msg): + self.run_paired_btn.setEnabled(True) + QMessageBox.critical(self, "Paired Permutation Test Error", f"An error occurred:\n{err_msg}") + + def start_paired_diagnostics(self): + cond1_subs, cond2_subs, subject_to_pair, skipped = self.get_paired_groups() + + if len(cond1_subs) < 3: + QMessageBox.warning( + self, "Error", + f"Need at least 3 complete pairs. Found {len(cond1_subs)} complete pair(s), " + f"{len(skipped)} incomplete." + ) + return + + pipe1, pipe2 = self._build_paired_pipes(cond1_subs, cond2_subs) + + if pipe1.X_raw.shape[0] == 0 or pipe2.X_raw.shape[0] == 0: + QMessageBox.warning(self, "Error", "No feature data found for the paired subjects.") + return + + X_raw = np.vstack([pipe1.X_raw, pipe2.X_raw]) + y = np.concatenate([np.zeros(pipe1.X_raw.shape[0], dtype=int), + np.ones(pipe2.X_raw.shape[0], dtype=int)]) + pair_ids = np.array( + [subject_to_pair[s] for s in pipe1.groups] + [subject_to_pair[s] for s in pipe2.groups] + ) + + model_name = self.paired_model_select.currentText() + k_actual = max(1, min(K_FEATURES, X_raw.shape[1])) + + self.tabs.setCurrentIndex(1) + self.run_diagnostics_btn.setEnabled(False) + self.append_log( + f"--- Paired diagnostics: {model_name} | {len(cond1_subs)} pairs " + f"({len(skipped)} incomplete pairs skipped) | scoring={self._scoring()} ---" + ) + + self.diagnostics_thread = PairedDiagnosticsWorker( + X_raw=X_raw, y=y, pair_ids=pair_ids, + feature_names=pipe1.feature_names, feature_channels=pipe1.feature_channels, + quality_list=pipe1.quality_list, model_name=model_name, + k_features=k_actual, scoring=self._scoring(), n_bootstrap=500, + ) + self.diagnostics_thread.log_signal.connect(self.append_log) + self.diagnostics_thread.finished_signal.connect(self.on_diagnostics_finished) + self.diagnostics_thread.error_signal.connect(self.on_diagnostics_error) + self.diagnostics_thread.start() + + def on_diagnostics_finished(self, result): + self.run_diagnostics_btn.setEnabled(True) + obs = result['observed_acc'] + ci_low = result['ci_low'] + ci_high = result['ci_high'] + n_pairs = result['n_pairs'] + mde = result['minimum_detectable_accuracy'] + n_boot_ok = result['n_bootstrap_successful'] + bf01 = result.get('bayes_factor_bf01', float('nan')) + n_correct = result.get('n_correct') + n_total = result.get('n_total_predictions') + + self.append_log( + f"Observed acc: {obs:.2%} | 95% CI: [{ci_low:.2%}, {ci_high:.2%}] " + f"| MDE @80% power: {mde:.2%} | BF01={bf01:.3f} | n_pairs={n_pairs}" + ) + + clears_mde = obs >= mde + power_note = ( + "Observed accuracy meets or exceeds the minimum detectable effect." + if clears_mde else + "Observed accuracy does NOT exceed the minimum detectable effect \u2014 " + "a real effect at or below this size could exist but not be reliably caught by this design." + ) + + if np.isnan(bf01): + bf_note = "Bayes factor unavailable." + elif bf01 > 3: + bf_note = f"BF01 = {bf01:.2f} \u2014 substantial-to-strong evidence FOR the null (no effect)." + elif bf01 > 1: + bf_note = f"BF01 = {bf01:.2f} \u2014 weak evidence favoring the null, not strong enough to claim equivalence." + elif bf01 > 1 / 3: + bf_note = f"BF01 = {bf01:.2f} \u2014 weak evidence favoring a real effect, inconclusive." + else: + bf_note = f"BF01 = {bf01:.2f} \u2014 substantial-to-strong evidence AGAINST the null (favors a real effect)." + + QMessageBox.information( + self, "Paired Diagnostics Result", + f"Subject pairs: {n_pairs}\n" + f"Observed paired LOSO accuracy: {obs:.2%} ({n_correct}/{n_total} predictions correct)\n" + f"95% Bootstrap CI ({n_boot_ok} successful resamples): [{ci_low:.2%}, {ci_high:.2%}]\n" + f"Minimum detectable accuracy (~80% power, alpha=0.05): {mde:.2%}\n\n" + f"{power_note}\n\n" + f"Bayes factor (accuracy vs. chance): {bf_note}" + ) + + def on_diagnostics_error(self, err_msg): + self.run_diagnostics_btn.setEnabled(True) + QMessageBox.critical(self, "Paired Diagnostics Error", f"An error occurred:\n{err_msg}") + + def start_univariate_test(self): + """ + Runs a Wilcoxon signed-rank test per feature (condition 2 vs condition 1), with + Benjamini-Hochberg FDR correction across all features. Unlike the classifier-based + tests, this needs the two conditions' rows explicitly aligned by pair -- row order + from load_from_dataframe follows the source CSV's row order, not the order pairs + were entered in the table, so alignment is rebuilt here by pair_id rather than + assumed from list order. + """ + cond1_subs, cond2_subs, subject_to_pair, skipped = self.get_paired_groups() + + if len(cond1_subs) < 3: + QMessageBox.warning( + self, "Error", + f"Need at least 3 complete pairs. Found {len(cond1_subs)} complete pair(s), " + f"{len(skipped)} incomplete." + ) + return + + pipe1, pipe2 = self._build_paired_pipes(cond1_subs, cond2_subs) + + if pipe1.X_raw.shape[0] == 0 or pipe2.X_raw.shape[0] == 0: + QMessageBox.warning(self, "Error", "No feature data found for the paired subjects.") + return + if list(pipe1.feature_names) != list(pipe2.feature_names): + QMessageBox.warning(self, "Error", "Feature sets differ between conditions; cannot align.") + return + + pair_to_row1 = {subject_to_pair[s]: i for i, s in enumerate(pipe1.groups) if s in subject_to_pair} + pair_to_row2 = {subject_to_pair[s]: i for i, s in enumerate(pipe2.groups) if s in subject_to_pair} + common_pairs = sorted(set(pair_to_row1) & set(pair_to_row2)) + + if len(common_pairs) < 3: + QMessageBox.warning(self, "Error", "Fewer than 3 subjects have data in both conditions.") + return + + row1_idx = [pair_to_row1[p] for p in common_pairs] + row2_idx = [pair_to_row2[p] for p in common_pairs] + X_cond1_aligned = pipe1.X_raw[row1_idx] + X_cond2_aligned = pipe2.X_raw[row2_idx] + + results = run_univariate_paired_tests(X_cond1_aligned, X_cond2_aligned, pipe1.feature_names) + results.sort(key=lambda r: r['p_value_fdr']) + n_sig = sum(1 for r in results if r['significant_fdr']) + + lines = [ + f"{len(common_pairs)} aligned pairs | {n_sig}/{len(results)} features significant " + f"after FDR correction (alpha=0.05)\n" + ] + for r in results[:15]: + roi_label, metric_key = _split_feature_name(r['feature_name']) + marker = " *" if r['significant_fdr'] else "" + lines.append( + f"{metric_key or '?'} - {roi_label}: mean diff={r['mean_diff']:+.4f}, " + f"p={r['p_value']:.4f}, FDR q={r['p_value_fdr']:.4f}{marker}" + ) + if len(results) > 15: + lines.append(f"... and {len(results) - 15} more features not shown.") + + self.append_log(f"Univariate paired tests: {n_sig}/{len(results)} significant after FDR.") + QMessageBox.information(self, "Univariate Paired Test Results (Wilcoxon + BH-FDR)", "\n".join(lines)) + + def start_feature_stability(self): + """ + Runs a single (non-permuted) paired LOSO pass with feature-selection tracking + enabled, to check whether SelectKBest picks roughly the same features fold to + fold. Wildly inconsistent selection across folds is itself evidence against a + stable underlying effect, independent of the accuracy/p-value numbers. + """ + cond1_subs, cond2_subs, subject_to_pair, skipped = self.get_paired_groups() + + if len(cond1_subs) < 3: + QMessageBox.warning( + self, "Error", + f"Need at least 3 complete pairs. Found {len(cond1_subs)} complete pair(s), " + f"{len(skipped)} incomplete." + ) + return + + pipe1, pipe2 = self._build_paired_pipes(cond1_subs, cond2_subs) + + if pipe1.X_raw.shape[0] == 0 or pipe2.X_raw.shape[0] == 0: + QMessageBox.warning(self, "Error", "No feature data found for the paired subjects.") + return + + X_raw = np.vstack([pipe1.X_raw, pipe2.X_raw]) + y = np.concatenate([np.zeros(pipe1.X_raw.shape[0], dtype=int), + np.ones(pipe2.X_raw.shape[0], dtype=int)]) + pair_ids = np.array( + [subject_to_pair[s] for s in pipe1.groups] + [subject_to_pair[s] for s in pipe2.groups] + ) + + model_name = self.paired_model_select.currentText() + k_actual = max(1, min(K_FEATURES, X_raw.shape[1])) + n_folds = len(np.unique(pair_ids)) + + _, _, selection_counts = run_loso_cv( + X_raw=X_raw, y=y, groups=pair_ids, + feature_names=pipe1.feature_names, + classifier_factories={model_name: CLASSIFIER_FACTORIES[model_name]}, + k_features=k_actual, tune_hyperparameters=False, scoring=self._scoring(), + track_feature_selection=True, + ) + + counts = selection_counts[model_name] + sorted_counts = sorted(counts.items(), key=lambda kv: kv[1], reverse=True) + + lines = [f"Feature selection frequency across {n_folds} LOSO folds ({model_name}):\n"] + any_selected = False + for feat_name, cnt in sorted_counts[:15]: + if cnt == 0: + continue + any_selected = True + roi_label, metric_key = _split_feature_name(feat_name) + lines.append(f"{metric_key or '?'} - {roi_label}: selected in {cnt}/{n_folds} folds") + if not any_selected: + lines.append("No feature was selected in any fold.") + + self.append_log(f"Feature stability check complete for {model_name} ({n_folds} folds).") + QMessageBox.information(self, "Feature Selection Stability", "\n".join(lines)) + + def start_max_stat_test(self): + """ + Reruns the entire metric-combination search under label permutation, tracking only + the best accuracy per shuffle, so the leaderboard's winning combo can be judged + against "best of a search this size on pure noise" rather than just "better than + chance." Operates on the current (non-paired) Class A / Class B leaderboard. + """ + if self.sweep_results_df is None or self.sweep_results_df.empty: + QMessageBox.warning(self, "Error", "Run the metric-combination sweep first.") + return + + target_condition = self.condition_combo.currentText() + class_a, class_b, testing_subs, test_labels = self.get_groups() + if not class_a or not class_b: + QMessageBox.warning(self, "Error", "Class A and Class B need at least one subject.") + return + + train_label_map = {s: "Class_A" for s in class_a} + train_label_map.update({s: "Class_B" for s in class_b}) + + extractor = FNIRSClassificationPipeline( + training_subjects=list(train_label_map.keys()), target_condition=target_condition, + selected_metrics=ALL_METRIC_KEYS + ).load_from_dataframe(self.df, self.subject_col, train_label_map) + + if extractor.X_raw.shape[0] == 0: + QMessageBox.warning(self, "Error", "No feature data found for the current training groups.") + return + + feature_keys = build_metric_keys(extractor.feature_names) + combo_labels = self.sweep_results_df['combo_label'].dropna().unique().tolist() + + feature_combos, kept_labels = [], [] + for combo_label in combo_labels: + keys = str(combo_label).split('+') + col_idx = np.where(np.isin(feature_keys, keys))[0] + if len(col_idx) > 0: + feature_combos.append(col_idx) + kept_labels.append(combo_label) + + if not feature_combos: + QMessageBox.warning(self, "Error", "Could not map leaderboard combos back to features.") + return + + model_name = self.inspect_model_select.currentText() or next(iter(CLASSIFIER_FACTORIES)) + + self.tabs.setCurrentIndex(1) + self.run_maxstat_btn.setEnabled(False) + self.append_log( + f"--- Max-stat permutation test across {len(feature_combos)} combos " + f"({model_name}, scoring={self._scoring()}) ---" + ) + + self.maxstat_thread = MaxStatWorker( + X_raw=extractor.X_raw, y=extractor.y, groups=extractor.groups, + feature_combos=feature_combos, combo_labels=kept_labels, + feature_names=extractor.feature_names, model_name=model_name, + k_features=K_FEATURES, scoring=self._scoring(), n_permutations=100, + ) + self.maxstat_thread.log_signal.connect(self.append_log) + self.maxstat_thread.finished_signal.connect(self.on_maxstat_finished) + self.maxstat_thread.error_signal.connect(self.on_maxstat_error) + self.maxstat_thread.start() + + def on_maxstat_finished(self, result): + self.run_maxstat_btn.setEnabled(True) + obs = result['observed_max_acc'] + p = result['p_value'] + null_max_accs = result['null_max_accs'] + null_mean = float(np.mean(null_max_accs)) if null_max_accs else 0.0 + best_idx = result['observed_best_combo_idx'] + combo_labels = result.get('combo_labels', []) + best_label = combo_labels[best_idx] if 0 <= best_idx < len(combo_labels) else "?" + + self.append_log( + f"Max-stat: best combo '{best_label}' acc={obs:.2%} | " + f"null (best-of-search) mean={null_mean:.2%} | p={p:.4f}" + ) + + verdict = ( + "survives multiple-comparisons correction (p < 0.05)" if p < 0.05 + else "does NOT survive multiple-comparisons correction" + ) + QMessageBox.information( + self, "Max-Stat Permutation Test Result", + f"Best combo in real search: {best_label}\n" + f"Observed best accuracy: {obs:.2%}\n" + f"Null 'best-of-{len(combo_labels)}-combos' mean: {null_mean:.2%}\n" + f"p-value: {p:.4f}\n\n" + f"This result {verdict}, accounting for having searched {len(combo_labels)} combinations " + f"before picking a winner." + ) + + def on_maxstat_error(self, err_msg): + self.run_maxstat_btn.setEnabled(True) + QMessageBox.critical(self, "Max-Stat Test Error", f"An error occurred:\n{err_msg}") + + def load_config_to_gui(self, silent=False): + if not os.path.exists(CONFIG_FILE): + if not silent: + QMessageBox.information(self, "Config", "No saved configuration file found.") + return + + with open(CONFIG_FILE, "r") as f: + config = json.load(f) + + class_a = set(config.get("class_a", [])) + class_b = set(config.get("class_b", [])) + test_labels = config.get("test_labels", {}) + + self.table.blockSignals(True) + for row in range(self.table.rowCount()): + sub_id = self.table.item(row, 0).text() + if sub_id in class_a: + self.table.item(row, 1).setCheckState(Qt.Checked) + self.table.item(row, 2).setCheckState(Qt.Unchecked) + elif sub_id in class_b: + self.table.item(row, 1).setCheckState(Qt.Unchecked) + self.table.item(row, 2).setCheckState(Qt.Checked) + else: + self.table.item(row, 1).setCheckState(Qt.Unchecked) + self.table.item(row, 2).setCheckState(Qt.Unchecked) + + self.table.blockSignals(False) + self.refresh_status_column() + self.update_testing_labels() + + for sub_id, label in test_labels.items(): + if sub_id in self.test_combos: + combo = self.test_combos[sub_id] + idx = combo.findText(label) + if idx >= 0: + combo.setCurrentIndex(idx) + + def start_sweep(self): + target_condition = self.condition_combo.currentText() + if not target_condition: + QMessageBox.warning(self, "Error", "No target condition selected.") + return + + class_a, class_b, testing_subs, test_labels = self.get_groups() + + if not class_a or not class_b: + QMessageBox.warning(self, "Error", "Class A and Class B need at least one subject.") + return + + train_label_map = {sub_id: "Class_A" for sub_id in class_a} + train_label_map.update({sub_id: "Class_B" for sub_id in class_b}) + + config_data = {"class_a": class_a, "class_b": class_b, "test_labels": test_labels} + with open(CONFIG_FILE, "w") as f: + json.dump(config_data, f, indent=4) + + self.tabs.setCurrentIndex(1) + self.log_output.clear() + self.run_sweep_btn.setEnabled(False) + + self.sweep_thread = MetricSweepWorker( + df=self.df, + subject_col=self.subject_col, + target_condition=target_condition, + train_label_map=train_label_map, + test_label_map=test_labels, + scoring=self._scoring(), + ) + self.sweep_thread.log_signal.connect(self.append_log) + self.sweep_thread.finished_signal.connect(self.on_sweep_finished) + self.sweep_thread.error_signal.connect(self.on_sweep_error) + self.sweep_thread.start() + + def append_log(self, text): + self.log_output.append(f"[{time.strftime('%H:%M:%S')}] {text}") + + def on_sweep_finished(self, df): + self.run_sweep_btn.setEnabled(True) + self.leaderboard_sort_col = None + self.leaderboard_sort_asc = False + self.sweep_results_df = df + self.populate_leaderboard(df) + + def on_sweep_error(self, err_msg): + self.run_sweep_btn.setEnabled(True) + QMessageBox.critical(self, "Sweep Error", f"An error occurred:\n{err_msg}") + + def start_permutation_test(self): + if self.current_train_pipe is None or self.current_train_pipe.X_raw.shape[0] == 0: + QMessageBox.warning(self, "Error", "Select a valid combination in Tab 3 first.") + return + + model_name = self.inspect_model_select.currentText() + if not model_name: + return + + tp = self.current_train_pipe + self.tabs.setCurrentIndex(1) + self.perm_test_btn.setEnabled(False) + self.append_log( + f"--- Permutation test: {model_name} on {self.inspect_combo_select.currentText()} " + f"| scoring={self._scoring()} ---" + ) + + self.perm_thread = PermutationTestWorker( + X_raw=tp.X_raw, y=tp.y, groups=tp.groups, + feature_names=tp.feature_names, feature_channels=tp.feature_channels, + quality_list=tp.quality_list, model_name=model_name, + k_features=self.current_k_actual, scoring=self._scoring(), n_permutations=200, + ) + self.perm_thread.log_signal.connect(self.append_log) + self.perm_thread.finished_signal.connect(self.on_permutation_finished) + self.perm_thread.error_signal.connect(self.on_permutation_error) + self.perm_thread.start() + + def on_permutation_finished(self, result): + self.perm_test_btn.setEnabled(True) + obs = result['observed_acc'] + null = result['null_accs'] + p = result['p_value'] + null_mean = float(np.mean(null)) if null else 0.0 + + self.append_log(f"Observed LOSO acc: {obs:.2%} | null mean: {null_mean:.2%} | p={p:.4f}") + + verdict = "distinguishable from chance (p < 0.05)" if p < 0.05 else "NOT distinguishable from chance" + QMessageBox.information( + self, "Permutation Test Result", + f"Observed LOSO accuracy: {obs:.2%}\n" + f"Null distribution mean ({len(null)} shuffles): {null_mean:.2%}\n" + f"p-value: {p:.4f}\n\n" + f"Result is {verdict} at n={len(self.current_train_pipe.groups)} subjects." + ) + + def on_permutation_error(self, err_msg): + self.perm_test_btn.setEnabled(True) + QMessageBox.critical(self, "Permutation Test Error", f"An error occurred:\n{err_msg}") + + def populate_leaderboard(self, df): + if "Rank" not in df.columns: + df = df.copy() + df.insert(0, "Rank", range(1, len(df) + 1)) + + self.sweep_results_df = df.reset_index(drop=True) + + self.leaderboard_table.setRowCount(0) + self.leaderboard_table.setColumnCount(len(df.columns)) + self.leaderboard_table.setHorizontalHeaderLabels(list(df.columns)) + + self.leaderboard_table.setRowCount(len(df)) + for r_idx, row in df.iterrows(): + for c_idx, val in enumerate(row): + col_name = df.columns[c_idx] + if isinstance(val, float): + val_str = f"{val:.2%}" if "acc" in col_name.lower() else f"{val:.4f}" + else: + val_str = str(val) + item = QTableWidgetItem(val_str) + item.setFlags(item.flags() ^ Qt.ItemIsEditable) + self.leaderboard_table.setItem(r_idx, c_idx, item) + + for c_idx, col_name in enumerate(df.columns): + label = col_name + if col_name == self.leaderboard_sort_col: + label += " \u25b2" if self.leaderboard_sort_asc else " \u25bc" + self.leaderboard_table.horizontalHeaderItem(c_idx).setText(label) + + self.distribute_table_columns_evenly(self.leaderboard_table) + + if 'combo_label' in df.columns: + combos = df['combo_label'].tolist() + self.inspect_combo_select.blockSignals(True) + self.inspect_combo_select.clear() + self.inspect_combo_select.addItems(combos) + self.inspect_combo_select.blockSignals(False) + + if combos: + self.update_test_breakdown() + + def on_leaderboard_header_clicked(self, col_idx): + if self.sweep_results_df is None or self.sweep_results_df.empty: + return + + col_name = self.sweep_results_df.columns[col_idx] + if self.leaderboard_sort_col == col_name: + self.leaderboard_sort_asc = not self.leaderboard_sort_asc + else: + self.leaderboard_sort_col = col_name + self.leaderboard_sort_asc = True + + sorted_df = self.sweep_results_df.sort_values( + col_name, ascending=self.leaderboard_sort_asc, kind="mergesort" + ).reset_index(drop=True) + + self.populate_leaderboard(sorted_df) + + # ========================================================================= + # TAB 3: BREAKDOWN POPULATION & SORTING + # ========================================================================= + def populate_breakdown_table(self, df): + headers = [ + "Subject / File", "Actual Class", "Predicted Class", "Correct?", + "Confidence (%)", "Top Driving Feature", "Feature Value", "Feature Impact" + ] + + if df is None or df.empty: + self.breakdown_table.setRowCount(0) + self.breakdown_table.setColumnCount(len(headers)) + self.breakdown_table.setHorizontalHeaderLabels(headers) + self.breakdown_df = None + return + + self.breakdown_df = df.reset_index(drop=True) + col_keys = ['sub_id', 'actual', 'pred', 'correct', 'conf', 'top_feat', 'raw_val', 'impact'] + + # Build column header labels with active sort arrow + display_headers = [] + for key, header in zip(col_keys, headers): + lbl = header + if key == self.breakdown_sort_col or header == self.breakdown_sort_col: + lbl += " \u25b2" if self.breakdown_sort_asc else " \u25bc" + display_headers.append(lbl) + + self.breakdown_table.setRowCount(0) + self.breakdown_table.setColumnCount(len(headers)) + self.breakdown_table.setHorizontalHeaderLabels(display_headers) + + self.breakdown_table.setRowCount(len(df)) + for r_idx, row in df.iterrows(): + detail_obj = row.get('_detail_obj') + + for c_idx, key in enumerate(col_keys): + val_str = str(row.get(key, "")) + item = QTableWidgetItem(val_str) + item.setFlags(item.flags() ^ Qt.ItemIsEditable) + + # Attach detail object directly to cell 0 for sort-safe lookup + if c_idx == 0 and detail_obj is not None: + item.setData(Qt.UserRole, detail_obj) + + self.breakdown_table.setItem(r_idx, c_idx, item) + + self.distribute_table_columns_evenly(self.breakdown_table) + + def on_breakdown_header_clicked(self, col_idx): + if self.breakdown_df is None or self.breakdown_df.empty: + return + + # Map display column index to sorting field + col_keys = ['sub_id', 'actual', 'pred', 'correct', '_conf_num', 'top_feat', '_raw_val_num', '_impact_num'] + sort_field = col_keys[col_idx] if col_idx < len(col_keys) else self.breakdown_df.columns[col_idx] + + if self.breakdown_sort_col == sort_field: + self.breakdown_sort_asc = not self.breakdown_sort_asc + else: + self.breakdown_sort_col = sort_field + self.breakdown_sort_asc = True + + sorted_df = self.breakdown_df.sort_values( + sort_field, ascending=self.breakdown_sort_asc, kind="mergesort" + ).reset_index(drop=True) + + self.populate_breakdown_table(sorted_df) + + def on_leaderboard_row_clicked(self, row, col): + if self.sweep_results_df is None or row < 0 or row >= len(self.sweep_results_df): + return + + combo_name = self.sweep_results_df.iloc[row]['combo_label'] + idx = self.inspect_combo_select.findText(combo_name) + if idx >= 0: + self.inspect_combo_select.setCurrentIndex(idx) + + self.tabs.setCurrentIndex(2) + self.leaderboard_table.clearSelection() + + def update_test_breakdown(self): + if self.df is None: + return + + target_condition = self.condition_combo.currentText() + combo_label = self.inspect_combo_select.currentText() + model_name = self.inspect_model_select.currentText() + + if not combo_label or not model_name: + return + + class_a, class_b, testing_subs, test_labels = self.get_groups() + if not class_a or not class_b or not test_labels: + self.summary_banner.setText("Assign Training and Testing subjects in Tab 1 first.") + self.populate_breakdown_table(None) + self.breakdown_details_map = {} + return + + train_label_map = {sub_id: "Class_A" for sub_id in class_a} + train_label_map.update({sub_id: "Class_B" for sub_id in class_b}) + + metrics = combo_label.split('+') + + # Fit training pipeline on full train set + train_pipe = FNIRSClassificationPipeline( + training_subjects=list(train_label_map.keys()), + target_condition=target_condition, + selected_metrics=metrics + ).load_from_dataframe(self.df, self.subject_col, train_label_map) + + # Extract test subjects + test_pipe = FNIRSClassificationPipeline( + training_subjects=list(test_labels.keys()), + target_condition=target_condition, + selected_metrics=metrics + ).load_from_dataframe(self.df, self.subject_col, test_labels) + + if train_pipe.X_raw.shape[0] == 0 or test_pipe.X_raw.shape[0] == 0: + self.summary_banner.setText("No data points available for evaluation.") + self.populate_breakdown_table(None) + self.breakdown_details_map = {} + return + + k_actual = max(1, min(K_FEATURES, train_pipe.X_raw.shape[1])) + self.current_train_pipe = train_pipe + self.current_k_actual = k_actual + factory = CLASSIFIER_FACTORIES[model_name] + pipeline_template, grid = _make_pipeline_and_grid(factory, k_actual, TUNE_HYPERPARAMETERS) + fitted_pipe = _fit_with_optional_tuning(pipeline_template, grid, train_pipe.X_raw, train_pipe.y) + + # Inspect weights and scaler + scaler = fitted_pipe.named_steps['scaler'] + select = fitted_pipe.named_steps['select'] + clf = fitted_pipe.named_steps['classifier'] + + support_mask = select.get_support() + selected_indices = np.where(support_mask)[0] + selected_feature_names = [train_pipe.feature_names[i] for i in selected_indices] + + has_weights = hasattr(clf, 'coef_') and clf.coef_ is not None + weights = clf.coef_[0] if has_weights else np.ones(len(selected_indices)) + + has_intercept = hasattr(clf, 'intercept_') and clf.intercept_ is not None + intercept = float(clf.intercept_[0]) if has_intercept else 0.0 + + unique_labels = sorted(list(set(train_label_map.values()))) + int_to_label = {idx: lbl for idx, lbl in enumerate(unique_labels)} + + class_raw_values = {} + class_means = {} + for cls_int, cls_label in int_to_label.items(): + mask = train_pipe.y == cls_int + if mask.any(): + class_raw_values[cls_label] = train_pipe.X_raw[mask][:, selected_indices] + class_means[cls_label] = class_raw_values[cls_label].mean(axis=0) + else: + class_raw_values[cls_label] = np.zeros((0, len(selected_indices))) + class_means[cls_label] = np.zeros(len(selected_indices)) + + rows = [] + self.breakdown_details_map = {} + correct_count = 0 + + for idx, sub_id in enumerate(test_pipe.groups): + clean_sub_id = os.path.basename(sub_id) + x_raw_sample = test_pipe.X_raw[idx] + actual_int = test_pipe.y[idx] + actual_label = int_to_label.get(actual_int, str(actual_int)) + + pred_int = fitted_pipe.predict([x_raw_sample])[0] + pred_label = int_to_label.get(pred_int, str(pred_int)) + is_correct = (pred_int == actual_int) + if is_correct: + correct_count += 1 + + decision_score = None + if hasattr(clf, "decision_function"): + try: + decision_score = float(fitted_pipe.decision_function([x_raw_sample])[0]) + except Exception: + decision_score = None + + if hasattr(clf, "predict_proba"): + try: + probs = fitted_pipe.predict_proba([x_raw_sample])[0] + conf = probs[pred_int] * 100.0 + except Exception: + conf = 50.0 + elif decision_score is not None: + p1 = sigmoid(decision_score) + p0 = 1.0 - p1 + conf = (p1 if pred_int == 1 else p0) * 100.0 + else: + conf = 100.0 if is_correct else 50.0 + + x_scaled_all = scaler.transform([x_raw_sample])[0] + x_scaled_sel = x_scaled_all[selected_indices] + x_raw_sel = x_raw_sample[selected_indices] + + contribs = weights * x_scaled_sel + abs_contribs = np.abs(contribs) + top_feat_idx = np.argmax(abs_contribs) + top_feat_name = selected_feature_names[top_feat_idx] + top_feat_raw_val = x_raw_sel[top_feat_idx] + top_feat_contrib = contribs[top_feat_idx] + + total_abs = np.sum(abs_contribs) + impact_pct_top = (abs_contribs[top_feat_idx] / total_abs * 100.0) if total_abs > 0 else 0.0 + + feature_details = [] + for j, feat_name in enumerate(selected_feature_names): + impact_pct = (abs_contribs[j] / total_abs * 100.0) if total_abs > 0 else 0.0 + feature_details.append({ + 'feat_name': feat_name, + 'col_idx': j, + 'raw_val': x_raw_sel[j], + 'z': x_scaled_sel[j], + 'weight': weights[j], + 'contribution': contribs[j], + 'impact_pct': impact_pct, + 'class_a_mean': class_means.get('Class_A', np.zeros(len(selected_indices)))[j], + 'class_b_mean': class_means.get('Class_B', np.zeros(len(selected_indices)))[j], + }) + feature_details.sort(key=lambda d: abs(d['contribution']), reverse=True) + + # Store deep dive detail in dictionary keyed by clean subject ID + self.breakdown_details_map[clean_sub_id] = { + 'sub_id_full': sub_id, + 'sub_id': clean_sub_id, + 'actual': actual_label, + 'pred': pred_label, + 'correct': is_correct, + 'conf': conf, + 'decision_score': decision_score, + 'contribution_sum': float(np.sum(contribs)), + 'intercept': intercept, + 'model_name': model_name, + 'combo_label': combo_label, + 'features': feature_details, + 'class_a_values': class_raw_values.get('Class_A', np.zeros((0, len(selected_indices)))), + 'class_b_values': class_raw_values.get('Class_B', np.zeros((0, len(selected_indices)))), + } + + rows.append({ + 'sub_id': clean_sub_id, + 'actual': actual_label, + 'pred': pred_label, + 'correct': "Yes" if is_correct else "No", + 'conf': _fmt_confidence(conf), + '_conf_num': conf, + 'top_feat': top_feat_name, + 'raw_val': f"{top_feat_raw_val:.4f}", + '_raw_val_num': top_feat_raw_val, + 'impact': f"{top_feat_contrib:+.3f} ({impact_pct_top:.1f}% rel)", + '_impact_num': abs(top_feat_contrib), + }) + + # ========================================================================= + # METRIC CONTRIBUTION BREAKDOWN (For Multi-Metric Combinations) + # ========================================================================= + metric_summary_text = "" + if len(metrics) > 1 and len(selected_feature_names) > 0: + metric_totals = {m: 0.0 for m in metrics} + metric_feat_counts = {m: 0 for m in metrics} + total_abs_weight = 0.0 + + # Option A: Based on Model Weight Importance + for j, feat_name in enumerate(selected_feature_names): + _, metric_key = _split_feature_name(feat_name) + + # Match feature back to one of the active combo metrics + matched_metric = next((m for m in metrics if m.lower() in metric_key.lower()), None) + if matched_metric is None: + matched_metric = metrics[0] # Fallback match + + abs_w = abs(weights[j]) + metric_totals[matched_metric] += abs_w + metric_feat_counts[matched_metric] += 1 + total_abs_weight += abs_w + + # Build summary string + if total_abs_weight > 0: + parts = [] + for m in metrics: + pct = (metric_totals[m] / total_abs_weight) * 100.0 + n_feats = metric_feat_counts[m] + parts.append(f"{m}: {pct:.1f}% ({n_feats} feats)") + + metric_summary_text = "Metric Share: " + "  \u2022  ".join(parts) + + # Update Tab 3 Banner with accuracy AND metric contribution breakdown + total_test = len(test_pipe.groups) + acc_pct = (correct_count / total_test * 100.0) if total_test > 0 else 0.0 + + self.summary_banner.setText( + f"Model: {model_name} | Combination: {combo_label} | " + f"Test Accuracy: {acc_pct:.1f}% ({correct_count}/{total_test} Correct)" + "
" + f"{metric_summary_text}" + ) + + df = pd.DataFrame(rows) + self.populate_breakdown_table(df) + + def on_breakdown_row_clicked(self, row, col): + if row < 0: + return + + item = self.breakdown_table.item(row, 0) + if item is not None: + sub_id = item.text() + details_map = getattr(self, 'breakdown_details_map', {}) + detail = details_map.get(sub_id) or details_map.get(os.path.basename(sub_id)) + if detail: + self.render_deep_dive(detail) + + self.tabs.setCurrentIndex(3) + self.breakdown_table.clearSelection() + + def render_deep_dive(self, detail): + """ + Builds Tab 4's HTML explanation: one section per feature used by the model for this + subject, ranked by how much it influenced the decision, explaining what the metric + measures, how this subject's value compares to each class's training average (with an + overlaid distribution chart, not just numbers), and which direction the model's weight + pushed the decision. Also explains the *net* effect of all features together. + """ + conf_str = _fmt_confidence(detail['conf']) + verdict_color = "#2e7d32" if detail['correct'] else "#c62828" + verdict_text = "Correct" if detail['correct'] else "Incorrect" + + self.deep_dive_banner.setText( + f"{detail['sub_id']} | Model: {detail['model_name']} | Metrics: {detail['combo_label']} | " + f"Actual: {detail['actual']} \u2192 Predicted: {detail['pred']} ({conf_str} confidence) | {verdict_text}" + ) + self.deep_dive_banner.setStyleSheet( + f"font-size: 14px; font-weight: bold; color: {verdict_color}; padding: 6px;" + ) + + contribution_sum = detail['contribution_sum'] + intercept = detail.get('intercept', 0.0) + decision_score = detail['decision_score'] + sum_push = "Class_B" if contribution_sum > 1e-9 else ("Class_A" if contribution_sum < -1e-9 else "neither class") + intercept_push = "Class_B" if intercept > 1e-9 else ("Class_A" if intercept < -1e-9 else "neutral (0.0)") + decision_note = "" + if decision_score is not None: + decision_note = ( + f" Including the model's intercept term, the full decision score was " + f"{decision_score:+.4f}." + ) + + html_parts = [ + "", + f"

{detail['sub_id']}

", + f"

Actual class: {detail['actual']}    " + f"Predicted class: {detail['pred']}    " + f"Confidence: {conf_str}    " + f"{verdict_text}

", + "

Features are listed from most to least influential on this particular prediction. " + "\u201cContribution\u201d is the model weight multiplied by this subject's standardized " + "(z-scored) value for that feature: positive contributions push the decision toward " + "Class_B, negative contributions push it toward Class_A.

", + f"

" + f"Why the confidence looks so extreme: the features below don't all agree \u2014 some " + f"push toward Class_A and some push toward Class_B. But confidence is driven by the " + f"sum of every feature's contribution, not a vote count. Here that sum is " + f"{contribution_sum:+.4f} (net pull toward {sum_push}).{decision_note} " + f"A logistic/linear model turns that score into a probability with a sigmoid curve, which " + f"saturates very quickly \u2014 once the dominant features push the total score this far from " + f"zero, a couple of small disagreeing features (a few percent of the total influence each) " + f"aren't enough to meaningfully move the probability, so it rounds to a number very close to " + f"0% or 100%.

", + f"
", + f"Decision Score & Intercept Breakdown", + f"

", + f"\u2022 Total Feature Contribution: {contribution_sum:+.4f} (net pull toward {sum_push})
", + f"\u2022 Model Intercept (Baseline Bias): {intercept:+.4f} (inherent preference toward {intercept_push})
", + f"\u2022 Final Decision Score: {decision_score:+.4f} (= Feature Sum + Intercept)", + f"

", + f"

What is the Intercept? The intercept is the model's baseline starting point before considering any of this subject's specific sensor features. If a subject had perfectly average values across all features (z-scores of 0), the model would assign a decision score of {intercept:+.4f}.

", + f"

Why confidence can look extreme: A linear/logistic model passes the final decision score ({decision_score:+.4f}) through an S-shaped sigmoid function to calculate confidence. Because this curve saturates rapidly far from zero, a strong combined score produces near-0% or near-100% confidence, even if a few minor features disagree with the overall verdict.

", + f"
", + "
", + ] + + class_a_matrix = detail.get('class_a_values') + class_b_matrix = detail.get('class_b_values') + + for rank, feat in enumerate(detail['features'], 1): + roi_label, metric_key = _split_feature_name(feat['feat_name']) + metric_desc = METRIC_DESCRIPTIONS.get(metric_key, "this measurement of the response") + metric_label = metric_key.replace('_', ' ') if metric_key else "Unknown metric" + + raw_val = feat['raw_val'] + class_a_mean = feat['class_a_mean'] + class_b_mean = feat['class_b_mean'] + diff_a = abs(raw_val - class_a_mean) + diff_b = abs(raw_val - class_b_mean) + if abs(diff_a - diff_b) < 1e-12: + closer_text = "roughly equidistant between the two classes' training averages" + else: + closer_class = "Class_A" if diff_a < diff_b else "Class_B" + closer_text = f"closer to the {closer_class} training average" + + contribution = feat['contribution'] + if contribution > 1e-9: + push_text = "Class_B" + elif contribution < -1e-9: + push_text = "Class_A" + else: + push_text = "neither class strongly" + + img_tag = "" + col_idx = feat.get('col_idx') + if class_a_matrix is not None and class_b_matrix is not None and col_idx is not None: + try: + a_vals = class_a_matrix[:, col_idx] + b_vals = class_b_matrix[:, col_idx] + png_b64 = _render_feature_distribution_png( + a_vals, b_vals, raw_val, class_a_mean, class_b_mean, + detail['actual'], metric_label, roi_label, feat['z'] + ) + img_tag = ( + f"

" + ) + except Exception as e: + img_tag = f"

(chart unavailable: {e})

" + + html_parts.append( + f"

#{rank}. {metric_label} \u2014 {roi_label}

" + f"

{metric_label} measures {metric_desc}.

" + f"{img_tag}" + f"" + f"" + f"" + f"" + f"" + f"" + f"" + f"" + f"
This subject's value{_fmt_val(raw_val)}
Class_A training average{_fmt_val(class_a_mean)}
Class_B training average{_fmt_val(class_b_mean)}
Standardized value (z-score){feat['z']:+.2f} SD from the overall training mean
Model weight{feat['weight']:+.4f}
Contribution to decision score{contribution:+.4f}
Share of this prediction's total feature influence{feat['impact_pct']:.1f}%
" + f"

This subject's value is {closer_text}. Combined with the model's learned weight for " + f"this feature, it pushed the prediction toward {push_text}.

" + f"
" + ) + + html_parts.append("") + self.deep_dive_text.setHtml("".join(html_parts)) + + + + + + + + + + + + + + + + + + + + + + + + +def _make_pipeline_and_grid(make_clf, k, tune_hyperparameters): + pipeline = Pipeline([ + ('scaler', StandardScaler()), + ('select', SelectKBest(score_func=f_classif, k=k)), + ('classifier', make_clf()), + ]) + if not tune_hyperparameters or not hasattr(pipeline.named_steps['classifier'], 'C'): + return pipeline, None + return pipeline, {'classifier__C': CANDIDATE_C_VALUES} + + +def _fit_with_optional_tuning(pipeline, param_grid, X, y): + if param_grid is None: + pipeline.fit(X, y) + return pipeline + + min_class_count = np.bincount(y).min() if len(y) > 0 else 0 + inner_splits = min(INNER_CV_FOLDS, min_class_count) + # Below ~3 samples/class/fold the inner CV score is not statistically meaningful; + # tuning on noise tends to select the smallest C (heaviest regularization) on ties + # and silently collapses the model to near-zero coefficients. Skip tuning instead + # of pretending the search told us anything. + if inner_splits < 2 or min_class_count < 2 * INNER_CV_FOLDS: + pipeline.fit(X, y) + return pipeline + + inner_cv = StratifiedKFold(n_splits=inner_splits, shuffle=True, random_state=MODEL_RANDOM_STATE) + search = GridSearchCV(pipeline, param_grid, cv=inner_cv, scoring='roc_auc', n_jobs=1) + try: + search.fit(X, y) + return search.best_estimator_ + except Exception: + pipeline.fit(X, y) + return pipeline + + +def run_loso_cv(X_raw, y, groups, feature_names=None, feature_channels=None, quality_list=None, + classifier_factories=None, k_features=5, verbose=False, tune_hyperparameters=True, + scoring='accuracy', track_feature_selection=False): + """ + Executes Leave-One-Group-Out Cross Validation across subject groups. + + scoring: 'accuracy' (raw) or 'balanced_accuracy' (mean of per-class recall). Balanced + accuracy is the more honest metric whenever the evaluated set is class-imbalanced -- + raw accuracy on an imbalanced set can look inflated even from a model that leans + toward the majority class. + + track_feature_selection: if True, counts how many LOSO folds selected each feature + (via SelectKBest's support mask) per classifier. A real, stable effect should recruit + roughly the same features fold to fold; wildly inconsistent selection is itself + evidence against a stable underlying signal. Adds negligible cost since the support + mask is already computed during fitting. + + Returns: (model_accs, fitted_pipes, feature_selection_counts) + feature_selection_counts is None unless track_feature_selection=True, in which case + it is {classifier_name: {feature_name: n_folds_selected}}. + """ + if classifier_factories is None: + classifier_factories = CLASSIFIER_FACTORIES + + score_fn = accuracy_score if scoring == 'accuracy' else balanced_accuracy_score + + logo = LeaveOneGroupOut() + model_preds = {name: [] for name in classifier_factories} + model_trues = {name: [] for name in classifier_factories} + + feature_selection_counts = None + if track_feature_selection and feature_names is not None: + feature_selection_counts = { + name: {fn: 0 for fn in feature_names} for name in classifier_factories + } + + k_actual = max(1, min(k_features, X_raw.shape[1])) + + # Cross-validation loop + for train_idx, val_idx in logo.split(X_raw, y, groups): + X_tr, y_tr = X_raw[train_idx], y[train_idx] + X_va, y_va = X_raw[val_idx], y[val_idx] + + for name, factory in classifier_factories.items(): + pipe, grid = _make_pipeline_and_grid(factory, k_actual, tune_hyperparameters) + fitted_pipe = _fit_with_optional_tuning(pipe, grid, X_tr, y_tr) + + if feature_selection_counts is not None: + select = fitted_pipe.named_steps.get('select') + if select is not None: + support = select.get_support() + for f_idx, selected in enumerate(support): + if selected and f_idx < len(feature_names): + feature_selection_counts[name][feature_names[f_idx]] += 1 + + preds = fitted_pipe.predict(X_va) + model_preds[name].extend(preds) + model_trues[name].extend(y_va) + + model_accs = { + name: score_fn(model_trues[name], model_preds[name]) if len(model_trues[name]) > 0 else 0.0 + for name in classifier_factories + } + + # Fit final models on all training data + fitted_pipes = {} + for name, factory in classifier_factories.items(): + pipe, grid = _make_pipeline_and_grid(factory, k_actual, tune_hyperparameters) + fitted_pipes[name] = _fit_with_optional_tuning(pipe, grid, X_raw, y) + + return model_accs, fitted_pipes, feature_selection_counts + + +class FNIRSClassificationPipeline: + """ + Extracts features and targets directly from an loaded DataFrame or raw files. + """ + def __init__(self, training_subjects=None, data_root=None, target_condition="Reach", + fir_delays=None, selected_metrics=None, cache_file=None, n_jobs=1): + self.training_subjects = training_subjects or [] + self.data_root = data_root + self.target_condition = target_condition + self.fir_delays = fir_delays + self.selected_metrics = selected_metrics + self.cache_file = cache_file + self.n_jobs = n_jobs + + self.X_raw = np.array([]) + self.y = np.array([]) + self.groups = np.array([]) + self.feature_names = [] + self.feature_channels = [] + self.quality_list = [] + + def load_from_dataframe(self, df, subject_col, label_map): + """ + Populates dataset matrices from pre-computed master DataFrame based on subject label map. + """ + # Filter dataframe for matching subjects + df_sub = df[df[subject_col].astype(str).isin(label_map.keys())].copy() + df_sub[subject_col] = df_sub[subject_col].astype(str) + + # Identify FIR feature columns matching target condition + fir_cols = [c for c in df_sub.columns if c.startswith('FIR_')] + if self.target_condition: + cond_str = f"FIR_{self.target_condition}_" + matched = [c for c in fir_cols if cond_str.lower() in c.lower()] + if matched: + fir_cols = matched + + # Filter by selected metrics if requested + if self.selected_metrics: + valid_cols = [] + for col in fir_cols: + for met in self.selected_metrics: + met_label = METRIC_REGISTRY.get(met, met) + if met == 'Peak_Amp' and col.endswith('_Peak_to_Peak_Amp'): + continue + if col.endswith(f"_{met}") or col.endswith(f"_{met_label}"): + valid_cols.append(col) + break + fir_cols = valid_cols if valid_cols else fir_cols + + self.feature_names = fir_cols + self.feature_channels = [ + col.split('_')[2] if len(col.split('_')) >= 3 else "ch" for col in fir_cols + ] + + # Map unique labels to numeric integer targets (0, 1) + unique_labels = sorted(list(set(label_map.values()))) + label_to_int = {lbl: idx for idx, lbl in enumerate(unique_labels)} + + X_list, y_list, group_list = [], [], [] + for _, row in df_sub.iterrows(): + sub_id = str(row[subject_col]) + target_class = label_map[sub_id] + X_list.append(row[fir_cols].values.astype(float)) + y_list.append(label_to_int[target_class]) + group_list.append(sub_id) + + self.X_raw = np.array(X_list) if X_list else np.empty((0, len(fir_cols))) + self.y = np.array(y_list) + self.groups = np.array(group_list) + self.quality_list = [1.0] * len(self.feature_names) + return self + + +def run_permutation_test(X_raw, y, groups, feature_names=None, feature_channels=None, + quality_list=None, classifier_factories=None, k_features=5, + tune_hyperparameters=False, n_permutations=200, + random_state=MODEL_RANDOM_STATE, scoring='accuracy', + progress_callback=None): + """ + Builds a null distribution for LOSO-CV accuracy by repeatedly shuffling class labels + at the subject level (each subject keeps one consistent shuffled label across its + rows) and rerunning the *entire* LOSO-CV pipeline -- including SelectKBest feature + selection -- on the shuffled labels. Feature selection is re-run inside the loop, + not just the classifier, or the null distribution is invalid: SelectKBest fit on + real labels would leak real signal into a "null" run. + + tune_hyperparameters defaults to False: a full inner-CV search inside every outer + LOSO fit, repeated n_permutations times, is prohibitively slow, and is unstable at + n=15 anyway -- it would add noise to the null estimate rather than remove it. + + scoring: 'accuracy' or 'balanced_accuracy' -- passed through to run_loso_cv so the + null distribution is built with the same scoring rule as the observed statistic. + + progress_callback(done, total), if provided, is invoked after each permutation so + a caller (e.g. a QThread) can report status without this function knowing about UI. + + Returns: {model_name: {'observed_acc', 'null_accs', 'p_value'}} + p_value uses the standard +1/+1 correction so it's never reported as exactly 0. + """ + if classifier_factories is None: + classifier_factories = CLASSIFIER_FACTORIES + + rng = np.random.default_rng(random_state) + + observed_accs, _, _ = run_loso_cv( + X_raw=X_raw, y=y, groups=groups, + feature_names=feature_names, feature_channels=feature_channels, + quality_list=quality_list, classifier_factories=classifier_factories, + k_features=k_features, tune_hyperparameters=tune_hyperparameters, scoring=scoring, + ) + + unique_groups = np.unique(groups) + group_to_label = {g: y[groups == g][0] for g in unique_groups} + label_values = np.array(list(group_to_label.values())) + + null_accs = {name: [] for name in classifier_factories} + + for i in range(n_permutations): + shuffled = rng.permutation(label_values) + shuffled_map = dict(zip(unique_groups, shuffled)) + y_perm = np.array([shuffled_map[g] for g in groups]) + + perm_accs, _, _ = run_loso_cv( + X_raw=X_raw, y=y_perm, groups=groups, + feature_names=feature_names, feature_channels=feature_channels, + quality_list=quality_list, classifier_factories=classifier_factories, + k_features=k_features, tune_hyperparameters=tune_hyperparameters, scoring=scoring, + ) + for name in classifier_factories: + null_accs[name].append(perm_accs[name]) + + if progress_callback is not None: + progress_callback(i + 1, n_permutations) + + results = {} + for name in classifier_factories: + obs = observed_accs[name] + null = np.array(null_accs[name]) + p_value = (1 + np.sum(null >= obs)) / (1 + n_permutations) + results[name] = {'observed_acc': obs, 'null_accs': null.tolist(), 'p_value': float(p_value)} + + return results + + +def run_paired_permutation_test(X_raw, y, pair_ids, feature_names=None, feature_channels=None, + quality_list=None, classifier_factories=None, k_features=5, + tune_hyperparameters=False, n_permutations=200, + random_state=MODEL_RANDOM_STATE, scoring='accuracy', + progress_callback=None): + """ + Paired (within-subjects) variant of run_permutation_test. Each physical subject + contributes exactly two rows in X_raw/y -- one per condition -- sharing the same + pair_id. LOSO-CV groups on pair_id (not row identity), so both of a subject's rows + are always held out together, preventing same-person leakage across train/test. + + The null distribution is built by sign-flip permutation, not free label shuffling: + for each pair, its two condition labels are swapped with 50% probability. This is + the null appropriate for paired designs (condition assignment within a subject is + exchangeable) rather than the between-subjects null (subject identity is + exchangeable), which would be invalid here -- it could assign both of a subject's + rows the same label, a configuration the real data can never produce. + """ + if classifier_factories is None: + classifier_factories = CLASSIFIER_FACTORIES + + rng = np.random.default_rng(random_state) + pair_ids = np.asarray(pair_ids) + unique_pairs = np.unique(pair_ids) + + observed_accs, _, _ = run_loso_cv( + X_raw=X_raw, y=y, groups=pair_ids, + feature_names=feature_names, feature_channels=feature_channels, + quality_list=quality_list, classifier_factories=classifier_factories, + k_features=k_features, tune_hyperparameters=tune_hyperparameters, scoring=scoring, + ) + + pair_row_idx = {p: np.where(pair_ids == p)[0] for p in unique_pairs} + null_accs = {name: [] for name in classifier_factories} + + for i in range(n_permutations): + y_perm = y.copy() + for p, idxs in pair_row_idx.items(): + if len(idxs) == 2 and rng.random() < 0.5: + y_perm[idxs[0]], y_perm[idxs[1]] = y_perm[idxs[1]], y_perm[idxs[0]] + + perm_accs, _, _ = run_loso_cv( + X_raw=X_raw, y=y_perm, groups=pair_ids, + feature_names=feature_names, feature_channels=feature_channels, + quality_list=quality_list, classifier_factories=classifier_factories, + k_features=k_features, tune_hyperparameters=tune_hyperparameters, scoring=scoring, + ) + for name in classifier_factories: + null_accs[name].append(perm_accs[name]) + + if progress_callback is not None: + progress_callback(i + 1, n_permutations) + + results = {} + for name in classifier_factories: + obs = observed_accs[name] + null = np.array(null_accs[name]) + p_value = (1 + np.sum(null >= obs)) / (1 + n_permutations) + results[name] = {'observed_acc': obs, 'null_accs': null.tolist(), 'p_value': float(p_value)} + + return results + + +def binomial_bayes_factor(n_correct, n_total): + """ + Closed-form Bayes factor (BF01) comparing "accuracy is exactly chance (p=0.5)" (the + null) against "accuracy is some unknown value between 0 and 1" (the alternative, with + a flat/uniform Beta(1,1) prior on that unknown accuracy). Because the alternative's + prior is uniform, its marginal likelihood integrates to a constant (1/(n+1)) + regardless of the observed count, which makes this Bayes factor exact and cheap -- + no numerical integration needed. + + BF01 > 1 favors the null (performance indistinguishable from chance); BF01 < 1 favors + the alternative (performance differs from chance). This directly answers "is there + positive evidence for no effect," which a p-value alone cannot -- failing to reject + the null is not the same as evidence for it. + + Returns float('nan') if n_total is 0. + """ + if n_total <= 0: + return float('nan') + log_binom_null = ( + stats.binom.logpmf(n_correct, n_total, 0.5) + ) + log_marginal_alt = -np.log(n_total + 1) # log(1/(n+1)) under a uniform prior + return float(np.exp(log_binom_null - log_marginal_alt)) + + +def paired_bic_bayes_factor(diffs): + """ + Approximate Bayes factor (BF01, evidence for the null of "no mean difference" over + the alternative "nonzero mean difference") for a paired/one-sample design, using the + BIC approximation described in Wagenmakers (2007), "A practical solution to the + pervasive problems of p values." This needs only the paired differences -- no priors + to specify, no numerical integration -- at the cost of being an approximation rather + than an exact Bayes factor. + + BF01 > 1 favors the null (no difference); BF01 < 1 favors a real difference. + Conventional rough guide (Jeffreys / Kass & Raftery): 1-3 "barely worth mentioning", + 3-10 "substantial" evidence, >10 "strong" evidence; reciprocal thresholds apply in + favor of the alternative when BF01 < 1. + + diffs: 1D array of per-subject paired differences (e.g. condition2 - condition1, for + one feature, or any other paired continuous quantity). + Returns: (bf01, t_stat, p_value) + """ + diffs = np.asarray(diffs, dtype=float) + n = len(diffs) + if n < 3 or np.allclose(diffs, diffs[0]): + return float('nan'), float('nan'), float('nan') + + t_stat, p_value = stats.ttest_1samp(diffs, popmean=0.0) + delta_bic = n * np.log(1.0 + (t_stat ** 2) / (n - 1)) - np.log(n) + bf01 = float(np.exp(delta_bic / 2.0)) + return bf01, float(t_stat), float(p_value) + + +def paired_accuracy_diagnostics(X_raw, y, pair_ids, feature_names=None, feature_channels=None, + quality_list=None, classifier_factory=None, k_features=5, + n_bootstrap=500, random_state=MODEL_RANDOM_STATE, + scoring='accuracy', progress_callback=None): + """ + Companion diagnostics for a paired LOSO result: a bootstrap confidence interval on + accuracy, the minimum detectable effect (MDE) this sample size supports, and a + binomial Bayes factor comparing "accuracy is chance" against "accuracy is unknown." + + Bootstrap CI: resamples pairs (not rows) with replacement, refits + reruns LOSO each + time. Resampling at the pair level preserves the paired structure -- each bootstrap + draw still has both condition-rows for every included subject. + + MDE: with n pairs, treats each pair's correct/incorrect LOSO classification as a + Bernoulli trial and asks: assuming the true underlying accuracy were p, what is the + narrowest one-sided margin above 50% (chance) that this design would detect with + ~80% power at alpha=0.05, using a normal approximation to the binomial. This is a + rough guide, not a substitute for a full simulation-based power analysis, but it is + enough to state "this design could reliably detect an accuracy of X% or higher." + + Bayes factor: computed from the observed LOSO predictions (n_correct out of n_total, + where n_total = 2 * n_pairs since each pair contributes two held-out rows), via + binomial_bayes_factor. This is the piece that can positively support "no effect," + rather than only failing to find one. + """ + if classifier_factory is None: + classifier_factory = CLASSIFIER_FACTORIES['LDA'] + + rng = np.random.default_rng(random_state) + pair_ids = np.asarray(pair_ids) + unique_pairs = np.unique(pair_ids) + n_pairs = len(unique_pairs) + + pair_row_idx = {p: np.where(pair_ids == p)[0] for p in unique_pairs} + + # Rerun LOSO once more, manually, to get raw correct/total counts for the Bayes factor + # (run_loso_cv only returns the aggregated score, not the raw prediction counts). + logo_preds, logo_trues = [], [] + k_actual = max(1, min(k_features, X_raw.shape[1])) + for train_idx, val_idx in LeaveOneGroupOut().split(X_raw, y, pair_ids): + pipe, grid = _make_pipeline_and_grid(classifier_factory, k_actual, False) + fitted = _fit_with_optional_tuning(pipe, grid, X_raw[train_idx], y[train_idx]) + logo_preds.extend(fitted.predict(X_raw[val_idx])) + logo_trues.extend(y[val_idx]) + + score_fn = accuracy_score if scoring == 'accuracy' else balanced_accuracy_score + observed_acc = float(score_fn(logo_trues, logo_preds)) if logo_trues else 0.0 + n_total = len(logo_trues) + n_correct = int(np.sum(np.array(logo_preds) == np.array(logo_trues))) + bf01 = binomial_bayes_factor(n_correct, n_total) + + boot_accs = [] + for i in range(n_bootstrap): + sampled_pairs = rng.choice(unique_pairs, size=n_pairs, replace=True) + row_idx = np.concatenate([pair_row_idx[p] for p in sampled_pairs]) + boot_group_ids = np.concatenate([ + np.full(len(pair_row_idx[p]), f"{p}__{j}") for j, p in enumerate(sampled_pairs) + ]) + + try: + boot_accs_dict, _, _ = run_loso_cv( + X_raw=X_raw[row_idx], y=y[row_idx], groups=boot_group_ids, + feature_names=feature_names, feature_channels=feature_channels, + quality_list=quality_list, classifier_factories={'model': classifier_factory}, + k_features=k_features, tune_hyperparameters=False, scoring=scoring, + ) + boot_accs.append(boot_accs_dict['model']) + except Exception: + pass + + if progress_callback is not None: + progress_callback(i + 1, n_bootstrap) + + boot_accs = np.array(boot_accs) + ci_low, ci_high = (np.percentile(boot_accs, 2.5), np.percentile(boot_accs, 97.5)) if len(boot_accs) > 0 else (np.nan, np.nan) + + # Minimum detectable effect: smallest true accuracy p (>0.5) such that a one-sided + # binomial test at n_pairs trials, alpha=0.05, has ~80% power to exceed the alpha + # threshold. Uses a normal approximation; z_alpha=1.645, z_power=0.8416 (80% power). + z_alpha, z_power = 1.645, 0.8416 + p0 = 0.5 + se0 = np.sqrt(p0 * (1 - p0) / n_pairs) if n_pairs > 0 else np.nan + mde_accuracy = p0 + (z_alpha + z_power) * se0 + + return { + 'observed_acc': observed_acc, + 'ci_low': float(ci_low), + 'ci_high': float(ci_high), + 'n_bootstrap_successful': len(boot_accs), + 'n_pairs': n_pairs, + 'minimum_detectable_accuracy': float(min(mde_accuracy, 1.0)), + 'n_correct': n_correct, + 'n_total_predictions': n_total, + 'bayes_factor_bf01': bf01, + } + + +def _benjamini_hochberg(p_values, alpha=0.05): + """ + Standard Benjamini-Hochberg step-up FDR correction. Returns adjusted p-values + (q-values), one per input p-value, in the original input order. + """ + p_values = np.asarray(p_values, dtype=float) + n = len(p_values) + if n == 0: + return p_values + + order = np.argsort(p_values) + ranked = p_values[order] + + adjusted = np.empty(n) + running_min = 1.0 + for i in range(n - 1, -1, -1): + rank = i + 1 + val = ranked[i] * n / rank + running_min = min(running_min, val) + adjusted[i] = running_min + + out = np.empty(n) + out[order] = np.clip(adjusted, 0, 1) + return out + + +def run_univariate_paired_tests(X_cond1, X_cond2, feature_names, alpha=0.05): + """ + Runs a paired Wilcoxon signed-rank test independently for each feature, comparing + condition 1 vs condition 2 values for the same subjects. X_cond1[i] and X_cond2[i] + must already be aligned to the same subject/pair before calling this -- this function + does no subject matching itself. + + This is a different, complementary question from a multivariate classifier: it asks + whether any single feature shifts consistently within-subject, rather than whether + some combination of features can separate the two conditions. It can catch a real, + small, single-feature effect that a classifier's feature-selection step dilutes or + steps over when choosing a small combination of features. + + Falls back to a one-sample t-test for a feature if Wilcoxon cannot be computed (e.g. + all paired differences are exactly zero for that feature, which scipy raises on). + + Multiple-comparisons correction uses Benjamini-Hochberg FDR across all features + tested in this call, since testing ~10-20 features independently and reporting only + nominal p < 0.05 without correction would inflate the false-positive rate. + + Returns a list of dicts, one per feature: feature_name, mean_diff, statistic, + p_value, p_value_fdr, significant_fdr (bool). + """ + X_cond1 = np.asarray(X_cond1, dtype=float) + X_cond2 = np.asarray(X_cond2, dtype=float) + diffs = X_cond2 - X_cond1 + n_features = X_cond1.shape[1] + + raw_p = np.full(n_features, np.nan) + stat_out = np.full(n_features, np.nan) + mean_diffs = diffs.mean(axis=0) + + for j in range(n_features): + col_diffs = diffs[:, j] + if np.allclose(col_diffs, 0.0): + raw_p[j] = 1.0 + stat_out[j] = 0.0 + continue + try: + stat, p = stats.wilcoxon(col_diffs) + except ValueError: + stat, p = stats.ttest_1samp(col_diffs, popmean=0.0) + raw_p[j] = p + stat_out[j] = stat + + p_fdr = _benjamini_hochberg(raw_p, alpha) + + results = [] + for j in range(n_features): + results.append({ + 'feature_name': feature_names[j], + 'mean_diff': float(mean_diffs[j]), + 'statistic': float(stat_out[j]), + 'p_value': float(raw_p[j]), + 'p_value_fdr': float(p_fdr[j]), + 'significant_fdr': bool(p_fdr[j] < alpha), + }) + return results + + +def run_max_stat_permutation_test(X_raw, y, groups, feature_combos, feature_names=None, + classifier_factories=None, k_features=5, + tune_hyperparameters=False, n_permutations=200, + random_state=MODEL_RANDOM_STATE, scoring='accuracy', + progress_callback=None): + """ + Multiple-comparisons-corrected significance test for a leaderboard search: instead of + asking "is my single best combo's accuracy better than chance," this asks "is my best + combo's accuracy better than the best accuracy a search across this many combos finds + on pure noise." For each of n_permutations label shuffles, every combo in + feature_combos is evaluated and only the single best (max) accuracy across all combos + for that shuffle is kept, building a null distribution of "best-of-search" accuracies. + The observed max accuracy (from the real, unshuffled labels) is then compared against + that null. This corrects for having searched many combinations before selecting a + winner, which a per-combo permutation test does not. + + feature_combos: list of arrays of column indices into X_raw, one per combination to + evaluate (e.g. every single-metric and pair-metric combo from a sweep). + + This is expensive: each permutation reruns LOSO-CV for every combo, so cost scales as + n_permutations * len(feature_combos) * n_subjects * n_classifiers. Keep n_permutations + and/or the combo list modest, and keep tune_hyperparameters=False (default) to keep + each individual fit cheap. + + Returns: {'observed_max_acc', 'observed_best_combo_idx', 'null_max_accs', 'p_value'} + """ + if classifier_factories is None: + classifier_factories = {'LDA': CLASSIFIER_FACTORIES['LDA']} + + rng = np.random.default_rng(random_state) + + def _best_acc_across_combos(y_labels): + best_acc = -1.0 + best_idx = -1 + for combo_i, col_idx in enumerate(feature_combos): + X_sub = X_raw[:, col_idx] + sub_feature_names = [feature_names[i] for i in col_idx] if feature_names else None + accs, _, _ = run_loso_cv( + X_raw=X_sub, y=y_labels, groups=groups, + feature_names=sub_feature_names, + classifier_factories=classifier_factories, + k_features=min(k_features, X_sub.shape[1]), + tune_hyperparameters=tune_hyperparameters, scoring=scoring, + ) + combo_acc = float(np.mean(list(accs.values()))) + if combo_acc > best_acc: + best_acc = combo_acc + best_idx = combo_i + return best_acc, best_idx + + observed_max_acc, observed_best_idx = _best_acc_across_combos(y) + + unique_groups = np.unique(groups) + group_to_label = {g: y[groups == g][0] for g in unique_groups} + label_values = np.array(list(group_to_label.values())) + + null_max_accs = [] + for i in range(n_permutations): + shuffled = rng.permutation(label_values) + shuffled_map = dict(zip(unique_groups, shuffled)) + y_perm = np.array([shuffled_map[g] for g in groups]) + + best_acc, _ = _best_acc_across_combos(y_perm) + null_max_accs.append(best_acc) + + if progress_callback is not None: + progress_callback(i + 1, n_permutations) + + null_max_accs = np.array(null_max_accs) + p_value = (1 + np.sum(null_max_accs >= observed_max_acc)) / (1 + n_permutations) + + return { + 'observed_max_acc': float(observed_max_acc), + 'observed_best_combo_idx': int(observed_best_idx), + 'null_max_accs': null_max_accs.tolist(), + 'p_value': float(p_value), + } + + +def build_metric_keys(feature_names): + """ + Maps each raw feature column name to its short metric key (e.g. 'Peak_Amp'), by + matching the longest suffix pattern from METRIC_REGISTRY. Shared by the sweep worker + and the max-stat permutation test so combo labels (e.g. "Peak_Amp+RMS") can be mapped + back to column indices consistently in both places. + """ + metric_map = {} + for k, v in METRIC_REGISTRY.items(): + metric_map[k] = k + metric_map[v] = k + sorted_patterns = sorted(metric_map.keys(), key=len, reverse=True) + + keys = [] + for name in feature_names: + matched = None + for pattern in sorted_patterns: + if name.endswith(f"_{pattern}"): + matched = metric_map[pattern] + break + keys.append(matched) + return np.array(keys) + + +class Plugin: + """Plugin entry point contract loaded by PluginManager.""" + + def __init__(self, main_window: QWidget) -> None: + self.main_window = main_window + self.name = "ML Channel Builder" + self.widget_instance: DataSplitterGUI | None = None + + def register_menu(self, plugin_menu: QMenu) -> None: + """Registers plugin options into the application's top menubar.""" + open_action = QAction("Open ML Tool", self.main_window) + open_action.triggered.connect(self.show_widget) + + about_action = QAction("About ML Builder", self.main_window) + about_action.triggered.connect(self.show_about) + + plugin_menu.addAction(open_action) + plugin_menu.addAction(about_action) + + def show_widget(self) -> None: + """Instantiates or focuses the ML Builder window.""" + if self.widget_instance is None or not self.widget_instance.isVisible(): + self.widget_instance = DataSplitterGUI() + self.widget_instance.show() + else: + self.widget_instance.raise_() + self.widget_instance.activateWindow() + + def show_about(self) -> None: + """Displays plugin information dynamically from plugin.json if available.""" + # Fallback default values + title = "ROI Channel Builder" + version = "Unknown" + author = "Unknown" + description = ( + "This plugin loads SNIRF binary files using h5py, extracts source-detector " + "channel pairs, and exports custom ROI channel group JSON files." + ) + + # Look for plugin.json in the same directory as this file + json_path = Path(__file__).resolve().parent / "plugin.json" + + if json_path.is_file(): + try: + with open(json_path, "r", encoding="utf-8") as f: + data = json.load(f) + if isinstance(data, dict): + title = data.get("name", title) + version = data.get("version", version) + author = data.get("author", author) + description = data.get("description", description) + except Exception: + # Silently catch file read/parse errors to prevent crashing + pass + + # Build formatted display message + message = ( + f"{title}
" + f"Version: {version}  |  Author: {author}

" + f"{description}" + ) + + QMessageBox.about( + self.main_window, + f"About {title}", + message, + ) \ No newline at end of file diff --git a/machine-learning-pack/plugin.json b/machine-learning-pack/plugin.json new file mode 100644 index 0000000..2db1284 --- /dev/null +++ b/machine-learning-pack/plugin.json @@ -0,0 +1,10 @@ +{ + "id": "machine-learning-pack", + "name": "Machine Learning Pack", + "version": "0.1.0", + "author": "Tyler de Zeeuw", + "description": "Machine Learning on Processed fNIRS Data.", + "min_app_version": "1.7.0", + "platforms": ["win_x64"], + "download_url": "https://git.research.dezeeuw.ca/tyler/flares-plugins/raw/branch/main/packages/machine-learning-pack-win-x64.zip" +} \ No newline at end of file diff --git a/packages/machine-learning-pack-win-x64.zip b/packages/machine-learning-pack-win-x64.zip new file mode 100644 index 0000000..abf3194 Binary files /dev/null and b/packages/machine-learning-pack-win-x64.zip differ diff --git a/packages/roi-builder-pack.zip b/packages/roi-builder-pack.zip index c8f1844..d678450 100644 Binary files a/packages/roi-builder-pack.zip and b/packages/roi-builder-pack.zip differ diff --git a/packages/test-plugin.zip b/packages/test-plugin.zip deleted file mode 100644 index 195a99a..0000000 Binary files a/packages/test-plugin.zip and /dev/null differ diff --git a/plugins.json b/plugins.json index a1db863..bab2b67 100644 --- a/plugins.json +++ b/plugins.json @@ -9,30 +9,10 @@ "platforms": ["win_x64"], "download_url": "https://git.research.dezeeuw.ca/tyler/flares-plugins/raw/branch/main/packages/machine-learning-pack-win-x64.zip" }, - { - "id": "theme-pack", - "name": "Dark Theme Pack", - "version": "1.0.0", - "author": "Tyler de Zeeuw", - "description": "Pure Python themes (cross-platform).", - "min_app_version": "1.8.0", - "platforms": ["darwin_arm64", "darwin_x64", "linux_x64"], - "download_url": "https://git.research.dezeeuw.ca/tyler/flares-plugins/raw/branch/main/packages/theme-pack-multi.zip" - }, - { - "id": "test-plugin", - "name": "Test Plugin", - "version": "1.0.0", - "description": "A basic test plugin demonstrating menu registration and custom widgets.", - "author": "Tyler de Zeeuw", - "min_app_version": "1.8.0", - "platforms": ["win_x64", "darwin_arm64", "darwin_x64", "linux_x64"], - "download_url": "https://git.research.dezeeuw.ca/tyler/flares-plugins/raw/branch/main/packages/test-plugin.zip" - }, { "id": "roi-builder-pack", "name": "ROI Channel Builder Pack", - "version": "1.0.0", + "version": "0.1.0", "description": "Extracts channel labels from SNIRF files and generates custom ROI grouping JSON files.", "author": "Tyler de Zeeuw", "min_app_version": "1.7.0", diff --git a/roi-builder-pack/__init__.py b/roi-builder-pack/__init__.py index af358d6..402ca00 100644 --- a/roi-builder-pack/__init__.py +++ b/roi-builder-pack/__init__.py @@ -9,6 +9,7 @@ License: GPL-3.0 # Built-in imports import os import json +from pathlib import Path from typing import Optional, Any # External library imports @@ -502,12 +503,43 @@ class Plugin: self.widget_instance.activateWindow() def show_about(self) -> None: - """Displays plugin information.""" + """Displays plugin information dynamically from plugin.json if available.""" + # Fallback default values + title = "ROI Channel Builder" + version = "Unknown" + author = "Unknown" + description = ( + "This plugin loads SNIRF binary files using h5py, extracts source-detector " + "channel pairs, and exports custom ROI channel group JSON files." + ) + + # Look for plugin.json in the same directory as this file + json_path = Path(__file__).resolve().parent / "plugin.json" + + if json_path.is_file(): + try: + with open(json_path, "r", encoding="utf-8") as f: + data = json.load(f) + if isinstance(data, dict): + title = data.get("name", title) + version = data.get("version", version) + author = data.get("author", author) + description = data.get("description", description) + except Exception: + # Silently catch file read/parse errors to prevent crashing + pass + + # Build formatted display message + message = ( + f"{title}
" + f"Version: {version}  |  Author: {author}

" + f"{description}" + ) + QMessageBox.about( self.main_window, - "About ROI Channel Builder", - "This plugin loads SNIRF binary files using h5py, extracts source-detector channel pairs, " - "and exports custom ROI channel group JSON files.", + f"About {title}", + message, ) diff --git a/roi-builder-pack/plugin.json b/roi-builder-pack/plugin.json index 776b385..2b320f0 100644 --- a/roi-builder-pack/plugin.json +++ b/roi-builder-pack/plugin.json @@ -1,7 +1,7 @@ { "id": "roi-builder-pack", "name": "ROI Channel Builder Pack", - "version": "1.0.0", + "version": "0.1.0", "description": "Extracts channel labels from SNIRF files and generates custom ROI grouping JSON files.", "author": "Tyler de Zeeuw", "min_app_version": "1.7.0", diff --git a/test-plugin/__init__.py b/test-plugin/__init__.py deleted file mode 100644 index c5c24ff..0000000 --- a/test-plugin/__init__.py +++ /dev/null @@ -1,68 +0,0 @@ -""" -Basic Test Plugin -Creates a widget window and registers submenu items under the Plugins menu. -""" - -from PySide6.QtGui import QAction -from PySide6.QtWidgets import QLabel, QMenu, QMessageBox, QPushButton, QVBoxLayout, QWidget - - -class PluginWidget(QWidget): - """Custom widget UI owned by the plugin.""" - - def __init__(self, parent: QWidget | None = None) -> None: - super().__init__(parent) - self.setWindowTitle("Test Plugin Window") - self.resize(350, 200) - - layout = QVBoxLayout(self) - - label = QLabel("Hello from the Test Plugin!", self) - btn_action = QPushButton("Click Me", self) - btn_action.clicked.connect(self._on_button_clicked) - - layout.addWidget(label) - layout.addWidget(btn_action) - - def _on_button_clicked(self) -> None: - QMessageBox.information(self, "Plugin Interactive", "Button inside the plugin QWidget was clicked!") - - -class Plugin: - """Plugin entry point contract loaded by the main application.""" - - def __init__(self, main_window: QWidget) -> None: - self.main_window = main_window - self.name = "Test Plugin" - self.widget_instance: PluginWidget | None = None - - def register_menu(self, plugin_menu: QMenu) -> None: - """ - Populates the plugin's dedicated submenu in the main menubar. - Called by the main app after creating the submenu header. - """ - open_action = QAction("Open Tool Window", self.main_window) - open_action.triggered.connect(self.show_widget) - - about_action = QAction("About Test Plugin", self.main_window) - about_action.triggered.connect(self.show_about) - - plugin_menu.addAction(open_action) - plugin_menu.addAction(about_action) - - def show_widget(self) -> None: - """Instantiates and displays the plugin's QWidget window.""" - if self.widget_instance is None or not self.widget_instance.isVisible(): - self.widget_instance = PluginWidget() - self.widget_instance.show() - else: - self.widget_instance.raise_() - self.widget_instance.activateWindow() - - def show_about(self) -> None: - """Secondary menu action example.""" - QMessageBox.about( - self.main_window, - "About Test Plugin", - "This is a basic test plugin demonstrating submenu registration and custom QWidget window spawning.", - ) \ No newline at end of file diff --git a/test-plugin/plugin.json b/test-plugin/plugin.json deleted file mode 100644 index 89fdd28..0000000 --- a/test-plugin/plugin.json +++ /dev/null @@ -1,10 +0,0 @@ -{ - "id": "test_plugin", - "name": "Test Plugin", - "version": "1.0.0", - "description": "A basic test plugin demonstrating menu registration and custom widgets.", - "author": "Tyler de Zeeuw", - "min_app_version": "1.7.0", - "platforms": ["win_x64", "darwin_arm64", "darwin_x64", "linux_x64"], - "download_url": "https://git.research.dezeeuw.ca/tyler/flares-plugins/raw/branch/main/packages/test-plugin.zip" -} \ No newline at end of file