OSCR

NeuroStat: An Open-Source EEG Connectivity Platform for Randomised Controlled Trials.

Code ↔ Paper

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

The 32 matches
  1. [1] § 2. Materials and Methods › 2.7. Validation Design › 2.7.2. Source-Space Simulation Benchmark ↔ scripts/validation/source_space_simulation.py, lines 1–73 · score 0.95 · Desikan Killiany parcels, Alpha activity, source activity, source signal, phase lag, MNE inverse
  2. [2] § 2. Materials and Methods › 2.7. Validation Design › 2.7.3. Source-Space Physiological Benchmark: PhysioNet EEGBCI Dataset and Procedure ↔ scripts/validation/physionet_validation.py, lines 1–50 · score 0.89 · EEG Motor Movement, Imagery Database, S001 S020, Eyes open, eyes closed, PhysioNet
  3. [3] § 2. Materials and Methods › 2.3. Source Localisation ↔ scripts/validation/source_space_simulation.py, lines 1–73 · score 0.87 · Desikan Killiany parcel, Source activity, inverse solution, MNE inverse, scalp EEG, Source space
  4. [4] § 2. Materials and Methods › 2.2. Automated Preprocessing Pipeline ↔ pli_pipeline.py, lines 94–116 · score 0.85 · Bad channel detection, ICA fitting, score thresholding, FastICA, peak, rejected
  5. [5] § 2. Materials and Methods › 2.7. Validation Design › 2.7.3. Source-Space Physiological Benchmark: PhysioNet EEGBCI Dataset and Procedure ↔ scripts/validation/physionet_source_space_iclabel.py, lines 1–70 · score 0.83 · Desikan Killiany, posterior alpha, Eyes open, eyes closed, PhysioNet, cuneus
  6. [6] § 2. Materials and Methods › 2.7. Validation Design › 2.7.1. Internal Validation: Simulated EEG ↔ scripts/validation/generate_simulated_eeg.py, lines 383–464 · score 0.82 · 20–100 Hz, temporal channels, bandpass filtered, EMG, bursts, min
  7. [7] § 2. Materials and Methods › 2.7. Validation Design › 2.7.2. Source-Space Simulation Benchmark ↔ scripts/validation/source_space_simulation.py, lines 335–423 · score 0.81 · distance matched, control edges, control PLI, random seeds, source space, alpha PLI
  8. [8] § 3. Results › 3.4. Source-Space Physiological Benchmark: PhysioNet EEGBCI ↔ pli_pipeline.py, lines 840–889 · score 0.79 · heart beat, channel noise, eye blink, ICLabel, classified, muscle
  9. [9] § 2. Materials and Methods › 2.4. Functional Connectivity Estimation ↔ app_gui.py, lines 1019–1100 · score 0.76 · 13–30 Hz, 8–13 Hz, 1–4 Hz, 4–8 Hz, frequency bands, delta
  10. [10] § Appendix A. NeuroStat Graphical User Interface Screenshots ↔ app_gui.py, lines 1019–1100 · score 0.75 · signal quality score, band power changes, variance reduction, SNR improvement, tab, Metrics
  11. [11] § Appendix A. NeuroStat Graphical User Interface Screenshots ↔ pli_pipeline.py, lines 967–1076 · score 0.74 · signal quality score, band power changes, variance reduction, SNR improvement, Metrics, component
  12. [12] § 2. Materials and Methods › 2.4. Functional Connectivity Estimation ↔ pli_pipeline.py, lines 128–141 · score 0.72 · 13–30 Hz, 8–13 Hz, 4–8 Hz, frequency bands, gamma, delta
  13. [13] § 2. Materials and Methods › 2.7. Validation Design › 2.7.3. Source-Space Physiological Benchmark: PhysioNet EEGBCI Dataset and Procedure ↔ scripts/validation/physionet_source_space_iclabel.py, lines 1–70 · score 0.72 · PhysioNet source space, Desikan Killiany, posterior alpha, cuneus, inferior, lingual
  14. [14] § 2. Materials and Methods › 2.6. Visual Outputs ↔ app_gui.py, lines 929–976 · score 0.70 · power spectral density, EEG traces, scalp topography, bar, metric, preprocessing
  15. [15] § 3. Results › 3.4. Source-Space Physiological Benchmark: PhysioNet EEGBCI ↔ scripts/validation/physionet_source_space_iclabel.py, lines 197–316 · score 0.68 · frontal transient, EC ICA, eyes open, eyes closed, source space, alpha PLI
  16. [16] § 3. Results › 3.2. Recovery of Known Connectivity ↔ scripts/validation/generate_simulated_eeg.py, lines 68–102 · score 0.67 · weak theta, moderate beta, strong alpha, expected PLI, connectivity pattern, coupling
  17. [17] § 2. Materials and Methods › 2.3. Source Localisation ↔ pli_pipeline.py, lines 1444–1499 · score 0.67 · inverse operator, depth, covariance, fsaverage, model, BEM
  18. [18] § 2. Materials and Methods › 2.7. Validation Design › 2.7.3. Source-Space Physiological Benchmark: PhysioNet EEGBCI Dataset and Procedure › Pipeline Modifications for Short Recordings ↔ scripts/validation/physionet_validation.py, lines 90–115 · score 0.66 · minute recordings, ICLabel brain, relaxed, disabled, ASR, threshold
  19. [19] § 3. Results › 3.3. Source-Space Simulation Recovery ↔ scripts/validation/source_space_simulation.py, lines 335–423 · score 0.65 · source space simulation, control edges, Desikan Killiany parcels, alpha PLI, inverse, anterior
  20. [20] § 3. Results › 3.1. Quality Verification of Simulated Data ↔ scripts/validation/quick_diagnostic.py, lines 155–260 · score 0.65 · sub matrix, frontal parietal, PLI matrix, Diagnostic, alpha band, expected PLI
  21. [21] § 3. Results › 3.3. Source-Space Simulation Recovery ↔ scripts/validation/source_space_simulation.py, lines 307–332 · score 0.61 · anterior control edges, known parcel pair, Source space simulation, Bars, reconstructed, Alpha
  22. [22] § 3. Results › 3.1. Quality Verification of Simulated Data ↔ scripts/validation/generate_simulated_eeg.py, lines 68–102 · score 0.59 · frontal parietal, strong alpha coupling, expected PLI, simulated, band, connectivity
  23. [23] § 3. Results › 3.2. Recovery of Known Connectivity ↔ scripts/validation/analyze_validation_results.py, lines 36–72 · score 0.59 · weak theta, moderate beta, strong alpha, heavy, Traditional, GEDAI
  24. [24] § 2. Materials and Methods › 2.6. Visual Outputs ↔ pli_pipeline.py, lines 1342–1377 · score 0.59 · noise reduction, RMS amplitude, Scalp topography, map, pipeline, EEG
  25. [25] § 2. Materials and Methods › 2.2. Automated Preprocessing Pipeline ↔ pli_pipeline.py, lines 94–116 · score 0.59 · bad channel detection, component rejection, FastICA, notch, brain, ASR
  26. [26] § 3. Results › 3.5. GEDAI Versus Traditional Preprocessing on Real EEG ↔ scripts/validation/physionet_method_comparison.py, lines 226–355 · score 0.58 · EO PLI, real EEG, eyes open, eyes closed, PhysioNet, alpha PLI
  27. [27] § 3. Results › 3.4. Source-Space Physiological Benchmark: PhysioNet EEGBCI ↔ scripts/validation/physionet_source_space_iclabel.py, lines 197–316 · score 0.57 · reversed subjects, eyes open, eyes closed, source space, Shapiro, alpha PLI
  28. [28] § 2. Materials and Methods › 2.7. Validation Design › 2.7.3. Source-Space Physiological Benchmark: PhysioNet EEGBCI Dataset and Procedure ↔ scripts/validation/analyze_validation_results.py, lines 75–118 · score 0.56 · Desikan Killiany, inferior, lingual, pericalcarine, precuneus, superior
  29. [29] § 2. Materials and Methods › 2.2. Automated Preprocessing Pipeline ↔ app_gui.py, lines 624–673 · score 0.55 · notch filtering, generalised eigenvalue decomposition, selection, brain, signal, GEDAI
  30. [30] § 2. Materials and Methods › 2.2. Automated Preprocessing Pipeline ↔ scripts/validation/physionet_validation.py, lines 90–115 · score 0.55 · bad channel, ICLabel, smaller, FastICA, bandpass, notch
  31. [31] § 2. Materials and Methods › 2.7. Validation Design › 2.7.3. Source-Space Physiological Benchmark: PhysioNet EEGBCI Dataset and Procedure › Pipeline Modifications for Short Recordings ↔ pli_pipeline.py, lines 840–889 · score 0.53 · ICLabel brain probability, rejection, threshold, pipeline, components, ICA
  32. [32] § 3. Results › 3.2. Recovery of Known Connectivity ↔ scripts/validation/analyze_validation_results.py, lines 36–72 · score 0.52 · weak theta, strong alpha, heavy, moderate, Traditional

Paper

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

The paper is loaded when this pane is shown.

The authors' code

Python · 2,908 lines · 118 KB · no license · 8 matches

  1. """
  2. EEG PLI pipeline with GUI entry point.
  3. This module loads EEGLAB .set files, performs preprocessing with MNE-Python,
  4. computes Desikan-Killiany source-localized PLI connectivity, and stores
  5. diagnostic visualizations plus connectivity matrices for each recording.
  6. RUN: python .\run_study.py "E:\Dropbox\2024\Analyses\GIT_HUB\EEG_Preprocessing-and-Quality-Check\TEST" "GroupA,GroupB" "pre,post" "E:\Dropbox\2024\Analyses\GIT_HUB\EEG_Preprocessing-and-Quality-Check\processed3"
  7. """
  8. from __future__ import annotations
  9. import json
  10. import logging
  11. import sys
  12. import textwrap
  13. from dataclasses import dataclass, field
  14. import warnings
  15. from pathlib import Path
  16. from typing import Callable, Dict, Iterable, List, Optional, Sequence, Tuple
  17. from itertools import combinations # used later for simple utilities
  18. # New: simple study-design container (user-provided group/session names)
  19. @dataclass
  20. class StudyDesign:
  21. groups: List[str] = field(default_factory=list)
  22. sessions: List[str] = field(default_factory=list)
  23. import matplotlib
  24. # Use non-interactive backend for script/GUI runs
  25. matplotlib.use("Agg")
  26. import matplotlib.pyplot as plt # noqa: E402
  27. import numpy as np # noqa: E402
  28. import pandas as pd # noqa: E402
  29. import scipy.io # noqa: E402
  30. from scipy import sparse # noqa: E402
  31. import mne # noqa: E402
  32. from mne import EpochsArray # noqa: E402
  33. from mne.channels import make_standard_montage # noqa: E402
  34. # connectivity API renamed; fall back to standalone package if needed
  35. try: # pragma: no cover - import paths differ across MNE versions
  36. from mne.connectivity import spectral_connectivity_epochs # type: ignore
  37. except ImportError: # pragma: no cover
  38. from mne_connectivity import spectral_connectivity_epochs # type: ignore
  39. try: # pragma: no cover
  40. from mne.viz import plot_connectivity_circle as _plot_circle # type: ignore
  41. except Exception: # pragma: no cover
  42. try:
  43. from mne_connectivity.viz import plot_connectivity_circle as _plot_circle # type: ignore
  44. except Exception: # pragma: no cover
  45. _plot_circle = None
  46. from mne.datasets import fetch_fsaverage # noqa: E402
  47. from mne.minimum_norm import ( # noqa: E402
  48. apply_inverse_raw,
  49. make_inverse_operator,
  50. )
  51. from mne.preprocessing import ICA # noqa: E402
  52. # Official GEDAI package from https://github.com/neurotuning/gedai
  53. try:
  54. from gedai import Gedai
  55. HAS_OFFICIAL_GEDAI = True
  56. except ImportError:
  57. HAS_OFFICIAL_GEDAI = False
  58. Gedai = None
  59. try:
  60. import tkinter as tk
  61. from tkinter import filedialog, messagebox
  62. except Exception: # pragma: no cover
  63. tk = None
  64. filedialog = None
  65. messagebox = None
  66. LOGGER = logging.getLogger("pli_pipeline")
  67. LOG_FORMAT = "%(asctime)s - %(levelname)s - %(message)s"
  68. try: # pragma: no cover - optional CuPy acceleration
  69. import cupy as cp
  70. from cupy import cuda
  71. GPU_AVAILABLE = cuda.runtime.getDeviceCount() > 0
  72. except Exception: # pragma: no cover - GPU unavailable
  73. cp = None
  74. GPU_AVAILABLE = False
  75. @dataclass
  76. class PreprocessingConfig:
  77. l_freq: float = 1.0
  78. h_freq: float = 40.0
  79. notch_freqs: Sequence[float] = (50.0, 100.0)
  80. ica_method: str = "fastica"
  81. ica_n_components: Optional[int] = None
  82. random_state: int = 42
  83. reject_criteria: Optional[Dict[str, float]] = field(
  84. default_factory=lambda: {"eeg": 150e-6}
  85. )
  86. # Preprocessing method options
  87. use_asr: bool = True
  88. use_iclabel: bool = True
  89. iclabel_brain_threshold: float = 0.85
  90. # GED-based artifact detection (inspired by GEDAI)
  91. use_ged: bool = False
  92. ged_threshold: float = 3.0 # Z-score threshold for artifact component rejection
  93. # Additional preprocessing options
  94. bad_channel_threshold: float = 4.0 # Z-score for bad channel detection
  95. interpolate_bad_channels: bool = True
  96. # Epoch rejection for ICA fitting
  97. ica_reject_threshold: float = 200e-6 # Peak-to-peak threshold in V
  98. @dataclass
  99. class SourceConfig:
  100. trans_path: Path = Path("fsaverage-trans.fif")
  101. subjects_dir: Optional[Path] = None
  102. spacing: str = "ico5"
  103. bem_solution_name: str = "fsaverage-5120-5120-5120-bem-sol.fif"
  104. lambda2: float = 1.0 / 9.0
  105. @dataclass
  106. class ConnectivityConfig:
  107. epoch_length: float = 4.0
  108. frequency_bands: Dict[str, Tuple[float, float]] = field(
  109. default_factory=lambda: {
  110. "delta": (0.5, 4.0),
  111. "theta": (4.0, 8.0),
  112. "alpha": (8.0, 13.0),
  113. "beta": (13.0, 30.0),
  114. "gamma": (30.0, 45.0),
  115. }
  116. )
  117. method: str = "pli"
  118. use_gpu: bool = False # GPU disabled by default for stability; enable manually if needed
  119. @dataclass
  120. class PipelineConfig:
  121. preprocessing: PreprocessingConfig = field(default_factory=PreprocessingConfig)
  122. source: SourceConfig = field(default_factory=SourceConfig)
  123. connectivity: ConnectivityConfig = field(default_factory=ConnectivityConfig)
  124. output_root: Path = Path("processed")
  125. @dataclass
  126. class StatsConfig:
  127. """Configuration for statistical tests/cluster permutation."""
  128. n_permutations: int = 1000
  129. cluster_threshold: Optional[float] = None # t-threshold; None = MNE default
  130. # Extend PipelineConfig to include stats (backward compatible if not referenced)
  131. PipelineConfig.stats = StatsConfig() # type: ignore[attr-defined]
  132. def ensure_subjects_dir(path: Optional[Path]) -> Path:
  133. """Return a valid subjects_dir, fetching fsaverage when necessary."""
  134. if path and Path(path).expanduser().exists():
  135. return Path(path).expanduser()
  136. try:
  137. fetched = Path(fetch_fsaverage(verbose=True))
  138. return fetched.parent
  139. except Exception as exc: # pragma: no cover - fetch may fail offline
  140. raise FileNotFoundError(
  141. "fsaverage subject not found. Set SUBJECTS_DIR env or "
  142. "download fsaverage manually."
  143. ) from exc
  144. def find_set_files(folder: Path) -> List[Path]:
  145. """Return sorted list of EEGLAB .set files in folder (non-recursive).
  146. Excludes files ending with '_cleaned.set' to avoid reprocessing already cleaned data.
  147. """
  148. all_files = sorted(Path(folder).glob("*.set"))
  149. # Exclude already cleaned files to prevent reprocessing
  150. files = [f for f in all_files if not f.stem.endswith("_cleaned")]
  151. if not files:
  152. raise FileNotFoundError(f"No .set files found in {folder}")
  153. return files
  154. def _apply_montage(raw: mne.io.BaseRaw) -> bool:
  155. """
  156. Assign a standard montage if the recording lacks channel positions.
  157. Tries multiple montages in order of preference to maximize channel coverage.
  158. Returns True if a montage was successfully applied with good coverage.
  159. """
  160. # Check if montage already exists with positions
  161. current_montage = raw.get_montage()
  162. if current_montage is not None:
  163. n_with_pos = sum(1 for ch in raw.info['chs'] if any(ch['loc'][:3]))
  164. if n_with_pos >= len(raw.ch_names) * 0.8:
  165. LOGGER.info(f"Montage already set: {n_with_pos}/{len(raw.ch_names)} channels have positions")
  166. return True
  167. # List of montages to try, in order of preference
  168. # BioSemi and extended 10-20 systems first, then basic
  169. montages_to_try = [
  170. 'biosemi64', # BioSemi 64-channel system
  171. 'biosemi32', # BioSemi 32-channel system
  172. 'biosemi128', # BioSemi 128-channel system
  173. 'biosemi256', # BioSemi 256-channel system
  174. 'standard_1005', # Extended 10-20 system (343 channels)
  175. 'standard_1020', # Standard 10-20 system
  176. 'easycap-M1', # EasyCap montage
  177. 'GSN-HydroCel-64_1.0', # EGI 64-channel
  178. 'GSN-HydroCel-128', # EGI 128-channel
  179. ]
  180. best_montage = None
  181. best_coverage = 0
  182. for montage_name in montages_to_try:
  183. try:
  184. montage = make_standard_montage(montage_name)
  185. # Try to apply and count how many channels get positions
  186. raw_test = raw.copy()
  187. raw_test.set_montage(montage, on_missing='ignore', verbose=False)
  188. n_with_pos = sum(1 for ch in raw_test.info['chs'] if any(ch['loc'][:3]))
  189. coverage = n_with_pos / len(raw.ch_names)
  190. if coverage > best_coverage:
  191. best_coverage = coverage
  192. best_montage = montage_name
  193. # If we get 80%+ coverage, use this montage
  194. if coverage >= 0.8:
  195. raw.set_montage(montage, on_missing='ignore', verbose=False)
  196. LOGGER.info(f"Applied montage '{montage_name}': {n_with_pos}/{len(raw.ch_names)} channels ({coverage*100:.0f}%)")
  197. return True
  198. except Exception:
  199. continue
  200. # Use best montage found even if coverage is lower
  201. if best_montage and best_coverage > 0:
  202. try:
  203. montage = make_standard_montage(best_montage)
  204. raw.set_montage(montage, on_missing='ignore', verbose=False)
  205. n_with_pos = sum(1 for ch in raw.info['chs'] if any(ch['loc'][:3]))
  206. LOGGER.warning(f"Applied montage '{best_montage}' with limited coverage: {n_with_pos}/{len(raw.ch_names)} ({best_coverage*100:.0f}%)")
  207. return best_coverage >= 0.5
  208. except Exception as exc:
  209. LOGGER.warning(f"Could not apply best montage '{best_montage}': {exc}")
  210. LOGGER.warning("Could not find suitable montage for channel names")
  211. return False
  212. def _detect_bad_channels(raw: mne.io.BaseRaw, threshold: float = 4.0) -> List[str]:
  213. """Detect bad channels based on amplitude variance z-score."""
  214. data = raw.get_data(picks="eeg")
  215. ch_names = [raw.ch_names[i] for i in mne.pick_types(raw.info, eeg=True)]
  216. # Compute variance per channel
  217. variances = np.var(data, axis=1)
  218. # Z-score of variances
  219. z_scores = (variances - np.mean(variances)) / (np.std(variances) + 1e-10)
  220. # Mark channels with extreme variance as bad
  221. bad_mask = np.abs(z_scores) > threshold
  222. bad_channels = [ch_names[i] for i in np.where(bad_mask)[0]]
  223. return bad_channels
  224. def _gedai_denoise(
  225. raw: mne.io.BaseRaw,
  226. threshold: float = 3.0,
  227. trans_path: Optional[Path] = None,
  228. ) -> Tuple[mne.io.BaseRaw, dict]:
  229. """
  230. GEDAI (Generalized Eigenvalue De-Artifacting Instrument) - Official Implementation.
  231. Uses the official GEDAI package from https://github.com/neurotuning/gedai
  232. Based on: Ros et al. (2025) "Return of the GEDAI: Unsupervised EEG Denoising
  233. based on Leadfield Filtering"
  234. The algorithm uses leadfield-based reference covariance and SENSAI
  235. (Signal & Noise Subspace Alignment Index) for optimal threshold selection.
  236. Parameters:
  237. raw: MNE Raw object (will be modified in-place)
  238. threshold: Artifact rejection threshold (2=aggressive, 3=balanced, 4=conservative)
  239. Maps to noise_multiplier parameter in official GEDAI
  240. trans_path: Path to head-MRI transformation file (optional, not used by official GEDAI)
  241. Returns:
  242. raw: Cleaned raw object
  243. info: Dictionary with diagnostic information
  244. """
  245. from scipy.signal import butter, filtfilt
  246. picks = mne.pick_types(raw.info, eeg=True, meg=False)
  247. data_original = raw.get_data(picks=picks).copy()
  248. n_channels, n_times = data_original.shape
  249. sfreq = raw.info["sfreq"]
  250. # Get GEDAI version
  251. gedai_version = "unknown"
  252. if HAS_OFFICIAL_GEDAI:
  253. try:
  254. import gedai
  255. gedai_version = getattr(gedai, "__version__", "0.1.0")
  256. except Exception:
  257. gedai_version = "0.1.0"
  258. info_dict = {
  259. "method": "Official_GEDAI",
  260. "gedai_version": gedai_version,
  261. "gedai_package": "neurotuning/gedai",
  262. "gedai_repository": "https://github.com/neurotuning/gedai",
  263. "threshold": threshold,
  264. "n_channels": n_channels,
  265. "n_samples": n_times,
  266. "duration_sec": n_times / sfreq,
  267. "sensai_score": 0.0,
  268. "n_components_rejected": 0,
  269. "variance_removed_pct": 0.0,
  270. "snr_improvement_db": 0.0,
  271. "reference_type": "leadfield",
  272. "official_gedai_used": False,
  273. }
  274. # Print prominent banner for verification
  275. print("\n" + "=" * 70)
  276. print("OFFICIAL GEDAI DENOISING")
  277. print("=" * 70)
  278. print(f" Package: gedai v{gedai_version}")
  279. print(f" Repository: https://github.com/neurotuning/gedai")
  280. print(f" Citation: Ros et al. (2025) bioRxiv 10.1101/2025.10.04.680449")
  281. print(f" Channels: {n_channels}, Duration: {n_times/sfreq:.1f}s")
  282. print(f" Threshold (noise_multiplier): {threshold}")
  283. print(f" Reference covariance: leadfield (physics-based)")
  284. print("=" * 70)
  285. LOGGER.info(f"GEDAI: Starting OFFICIAL denoising (v{gedai_version})")
  286. LOGGER.info(f"GEDAI: {n_channels} channels, {n_times/sfreq:.1f}s, threshold={threshold}")
  287. LOGGER.info(f"GEDAI: Using leadfield-based reference covariance")
  288. if not HAS_OFFICIAL_GEDAI:
  289. error_msg = "Official GEDAI package not installed! Install with: pip install gedai"
  290. print(f"\n*** ERROR: {error_msg}")
  291. print("*** Repository: https://github.com/neurotuning/gedai\n")
  292. LOGGER.error(error_msg)
  293. LOGGER.error("Repository: https://github.com/neurotuning/gedai")
  294. info_dict["error"] = "Official GEDAI package not installed"
  295. return raw, info_dict
  296. try:
  297. # Map threshold to official GEDAI parameters
  298. # threshold 2 = aggressive (noise_multiplier=2.0)
  299. # threshold 3 = balanced (noise_multiplier=3.0)
  300. # threshold 4 = conservative (noise_multiplier=4.0)
  301. noise_multiplier = float(threshold)
  302. # Ensure proper montage for leadfield computation
  303. print(" [1/5] Checking/setting electrode montage...")
  304. LOGGER.info("GEDAI: Checking electrode positions for leadfield computation")
  305. raw_for_gedai = raw.copy()
  306. # Standard channel names that GEDAI's leadfield matrix supports
  307. # (10-20, 10-10, 10-05 systems and common variants)
  308. GEDAI_SUPPORTED_CHANNELS = {
  309. # 10-20 system
  310. 'Fp1', 'Fp2', 'F7', 'F3', 'Fz', 'F4', 'F8', 'T3', 'C3', 'Cz', 'C4', 'T4',
  311. 'T5', 'P3', 'Pz', 'P4', 'T6', 'O1', 'O2', 'A1', 'A2',
  312. # Extended 10-20 / 10-10 names
  313. 'AF7', 'AF3', 'AFz', 'AF4', 'AF8', 'F5', 'F1', 'F2', 'F6',
  314. 'FT7', 'FC5', 'FC3', 'FC1', 'FCz', 'FC2', 'FC4', 'FC6', 'FT8',
  315. 'T7', 'C5', 'C1', 'C2', 'C6', 'T8',
  316. 'TP7', 'CP5', 'CP3', 'CP1', 'CPz', 'CP2', 'CP4', 'CP6', 'TP8',
  317. 'P7', 'P5', 'P1', 'P2', 'P6', 'P8', 'P9', 'P10',
  318. 'PO7', 'PO3', 'POz', 'PO4', 'PO8',
  319. 'Oz', 'Iz', 'Fpz',
  320. # BioSemi specific (alternative naming)
  321. 'EXG1', 'EXG2', 'EXG3', 'EXG4', 'EXG5', 'EXG6', 'EXG7', 'EXG8',
  322. }
  323. # Check how many channels match GEDAI's supported names
  324. eeg_picks = mne.pick_types(raw_for_gedai.info, eeg=True, meg=False)
  325. eeg_ch_names = [raw_for_gedai.ch_names[i] for i in eeg_picks]
  326. n_eeg = len(eeg_ch_names)
  327. # Count channels with standard names (case-insensitive match)
  328. n_standard = sum(1 for ch in eeg_ch_names if ch in GEDAI_SUPPORTED_CHANNELS or ch.upper() in GEDAI_SUPPORTED_CHANNELS)
  329. standard_coverage = n_standard / n_eeg if n_eeg > 0 else 0
  330. # Also check position coverage
  331. n_with_pos = sum(1 for ch in raw_for_gedai.info['chs'] if any(ch['loc'][:3]))
  332. pos_coverage = n_with_pos / n_eeg if n_eeg > 0 else 0
  333. print(f" Standard channel names: {n_standard}/{n_eeg} ({standard_coverage*100:.0f}%)")
  334. print(f" Channels with positions: {n_with_pos}/{n_eeg} ({pos_coverage*100:.0f}%)")
  335. # Try to set montage if needed
  336. if pos_coverage < 0.8:
  337. print(f" Attempting to set standard montage...")
  338. LOGGER.info(f"GEDAI: Low position coverage ({pos_coverage*100:.0f}%), attempting to set standard montage")
  339. _apply_montage(raw_for_gedai)
  340. n_with_pos = sum(1 for ch in raw_for_gedai.info['chs'] if any(ch['loc'][:3]))
  341. pos_coverage = n_with_pos / n_eeg if n_eeg > 0 else 0
  342. # GEDAI requires standard channel names for leadfield lookup
  343. if standard_coverage < 0.5:
  344. error_msg = (
  345. f"Insufficient standard channel names for GEDAI leadfield computation. "
  346. f"Only {n_standard}/{n_eeg} channels have standard names ({standard_coverage*100:.0f}%). "
  347. f"GEDAI requires standard 10-20/10-10 channel names (e.g., Fp1, F3, C3, P3, O1). "
  348. f"Your channels: {eeg_ch_names[:5]}{'...' if len(eeg_ch_names) > 5 else ''}"
  349. )
  350. print(f" WARNING: {error_msg}")
  351. LOGGER.warning(f"GEDAI: {error_msg}")
  352. info_dict["warning"] = error_msg
  353. info_dict["standard_coverage"] = standard_coverage
  354. # Don't return - let GEDAI try and fail with a clear error
  355. LOGGER.info(f"GEDAI: Standard names: {n_standard}/{n_eeg} ({standard_coverage*100:.0f}%), Positions: {n_with_pos}/{n_eeg} ({pos_coverage*100:.0f}%)")
  356. info_dict["montage_coverage"] = pos_coverage
  357. info_dict["standard_name_coverage"] = standard_coverage
  358. # Handle dimension mismatch: exclude non-standard channels before GEDAI processing
  359. # GEDAI's leadfield only covers standard 10-20/10-10 channels
  360. non_standard_channels = [ch for ch in eeg_ch_names
  361. if ch not in GEDAI_SUPPORTED_CHANNELS and ch.upper() not in GEDAI_SUPPORTED_CHANNELS]
  362. standard_channels = [ch for ch in eeg_ch_names
  363. if ch in GEDAI_SUPPORTED_CHANNELS or ch.upper() in GEDAI_SUPPORTED_CHANNELS]
  364. excluded_channels = []
  365. if non_standard_channels and n_standard >= 19:
  366. # Have enough standard channels to proceed - exclude non-standard ones
  367. print(f" Excluding {len(non_standard_channels)} non-standard channels for GEDAI: {non_standard_channels}")
  368. LOGGER.info(f"GEDAI: Excluding non-standard channels: {non_standard_channels}")
  369. excluded_channels = non_standard_channels
  370. raw_for_gedai = raw_for_gedai.copy().pick(standard_channels)
  371. info_dict["excluded_channels"] = excluded_channels
  372. info_dict["n_channels_for_gedai"] = len(standard_channels)
  373. print(f" Processing {len(standard_channels)} standard channels")
  374. elif non_standard_channels:
  375. print(f" WARNING: Found non-standard channels but only {n_standard} standard channels (need >=19)")
  376. LOGGER.warning(f"GEDAI: Non-standard channels present but insufficient standard channels ({n_standard}<19)")
  377. # Ensure average reference (required by GEDAI)
  378. print(" [2/5] Applying average reference...")
  379. LOGGER.info("GEDAI: Applying average reference (required by official GEDAI)")
  380. raw_for_gedai.set_eeg_reference("average", projection=False, verbose=False)
  381. # Initialize official GEDAI
  382. print(" [3/5] Initializing official Gedai() class...")
  383. LOGGER.info("GEDAI: Initializing official Gedai() from neurotuning/gedai")
  384. gedai = Gedai()
  385. # Fit the model - try leadfield first, fall back to identity if needed
  386. reference_cov_used = "leadfield"
  387. fit_success = False
  388. # First attempt: Try with leadfield-based reference covariance
  389. print(f" [4/5] Fitting GEDAI model (attempting reference_cov='leadfield')...")
  390. LOGGER.info(f"GEDAI: Attempting fit with reference_cov='leadfield', noise_multiplier={noise_multiplier}")
  391. try:
  392. gedai.fit_raw(
  393. raw_for_gedai,
  394. duration=2.0, # Epoch duration (seconds)
  395. overlap=0.5, # 50% overlap
  396. reject_by_annotation=False, # Don't reject annotated segments
  397. reference_cov="leadfield", # Use leadfield-based reference (core GEDAI feature)
  398. sensai_method="gridsearch", # Grid search for optimal threshold
  399. noise_multiplier=noise_multiplier,
  400. verbose=True,
  401. )
  402. fit_success = True
  403. reference_cov_used = "leadfield"
  404. print(" Leadfield-based model fitted successfully!")
  405. LOGGER.info("GEDAI: Leadfield-based model fitting complete")
  406. except (ValueError, RuntimeError) as e:
  407. # Leadfield failed - try with identity reference
  408. print(f" Leadfield failed ({e}), trying identity reference...")
  409. LOGGER.warning(f"GEDAI: Leadfield reference failed: {e}")
  410. LOGGER.info("GEDAI: Falling back to identity reference covariance")
  411. try:
  412. gedai = Gedai() # Reset
  413. gedai.fit_raw(
  414. raw_for_gedai,
  415. duration=2.0,
  416. overlap=0.5,
  417. reject_by_annotation=False,
  418. reference_cov="identity", # Fallback to identity
  419. sensai_method="gridsearch",
  420. noise_multiplier=noise_multiplier,
  421. verbose=True,
  422. )
  423. fit_success = True
  424. reference_cov_used = "identity"
  425. print(" Identity-based model fitted successfully!")
  426. LOGGER.info("GEDAI: Identity-based model fitting complete (fallback)")
  427. except Exception as e2:
  428. # Identity also failed - try with data-driven approach
  429. print(f" Identity failed ({e2}), trying data-driven reference...")
  430. LOGGER.warning(f"GEDAI: Identity reference failed: {e2}")
  431. try:
  432. gedai = Gedai() # Reset
  433. gedai.fit_raw(
  434. raw_for_gedai,
  435. duration=2.0,
  436. overlap=0.5,
  437. reject_by_annotation=False,
  438. reference_cov="data", # Data-driven reference
  439. sensai_method="gridsearch",
  440. noise_multiplier=noise_multiplier,
  441. verbose=True,
  442. )
  443. fit_success = True
  444. reference_cov_used = "data"
  445. print(" Data-driven model fitted successfully!")
  446. LOGGER.info("GEDAI: Data-driven model fitting complete (fallback)")
  447. except Exception as e3:
  448. raise RuntimeError(f"All GEDAI reference types failed: leadfield={e}, identity={e2}, data={e3}")
  449. info_dict["reference_type"] = reference_cov_used
  450. print(f" Reference covariance used: {reference_cov_used}")
  451. LOGGER.info(f"GEDAI: Reference covariance used: {reference_cov_used}")
  452. # Transform/denoise the data
  453. print(" [5/5] Transforming/denoising data...")
  454. LOGGER.info("GEDAI: Transforming data (applying denoising)")
  455. raw_cleaned = gedai.transform_raw(
  456. raw_for_gedai,
  457. duration=2.0,
  458. overlap=0.5,
  459. verbose=True
  460. )
  461. print(" Transform complete!")
  462. LOGGER.info("GEDAI: Transform complete")
  463. # Mark that official GEDAI was successfully used
  464. info_dict["official_gedai_used"] = True
  465. # Get cleaned data - handle case where channels were excluded
  466. if excluded_channels:
  467. # Get cleaned data for processed channels only
  468. processed_picks = mne.pick_types(raw_cleaned.info, eeg=True, meg=False)
  469. data_clean_partial = raw_cleaned.get_data(picks=processed_picks)
  470. # Create full data array with original data for excluded channels
  471. data_clean = data_original.copy()
  472. # Map processed channel data back to original indices
  473. for i, ch_name in enumerate(standard_channels):
  474. # Find the index in the original channel list
  475. original_idx = eeg_ch_names.index(ch_name)
  476. data_clean[original_idx, :] = data_clean_partial[i, :]
  477. print(f" Merged {len(standard_channels)} processed channels with {len(excluded_channels)} excluded channels")
  478. LOGGER.info(f"GEDAI: Merged processed ({len(standard_channels)}) and excluded ({len(excluded_channels)}) channels")
  479. else:
  480. # All channels were processed
  481. data_clean = raw_cleaned.get_data(picks=mne.pick_types(raw_cleaned.info, eeg=True, meg=False))
  482. # Compute quality metrics
  483. # Variance removed
  484. var_original = np.var(data_original)
  485. var_clean = np.var(data_clean)
  486. var_removed = np.var(data_original - data_clean)
  487. info_dict["variance_removed_pct"] = float(100 * var_removed / (var_original + 1e-10))
  488. info_dict["variance_original"] = float(var_original)
  489. info_dict["variance_clean"] = float(var_clean)
  490. # SNR improvement estimate (using high-freq as noise proxy)
  491. try:
  492. nyq = sfreq / 2
  493. if nyq > 30:
  494. b_hf, a_hf = butter(4, 30 / nyq, btype='high')
  495. noise_before = np.std(filtfilt(b_hf, a_hf, data_original, axis=1))
  496. noise_after = np.std(filtfilt(b_hf, a_hf, data_clean, axis=1))
  497. signal_after = np.std(data_clean)
  498. if noise_after > 0 and noise_before > 0:
  499. snr_before = np.std(data_original) / noise_before
  500. snr_after = signal_after / noise_after
  501. info_dict["snr_improvement_db"] = float(20 * np.log10(snr_after / (snr_before + 1e-10)))
  502. except Exception:
  503. pass
  504. # Extract SENSAI score from fitted model if available
  505. if hasattr(gedai, 'sensai_score_') and gedai.sensai_score_ is not None:
  506. # Official GEDAI SENSAI score (typically 0-1, convert to 0-100)
  507. info_dict["sensai_score"] = float(gedai.sensai_score_ * 100)
  508. info_dict["sensai_score_raw"] = float(gedai.sensai_score_)
  509. print(f" Official SENSAI score: {gedai.sensai_score_:.4f}")
  510. LOGGER.info(f"GEDAI: Official SENSAI score from model: {gedai.sensai_score_:.4f}")
  511. else:
  512. # Compute approximate SENSAI score based on metrics
  513. var_score = np.clip(info_dict["variance_removed_pct"] / 50, 0, 1) * 30
  514. snr_score = np.clip((info_dict["snr_improvement_db"] + 5) / 15, 0, 1) * 40
  515. quality_score = 30 # Base score for successful processing
  516. info_dict["sensai_score"] = float(np.clip(var_score + snr_score + quality_score, 0, 100))
  517. info_dict["sensai_score_computed"] = True
  518. LOGGER.info("GEDAI: SENSAI score computed from metrics (not available from model)")
  519. # Get threshold from fitted model if available
  520. if hasattr(gedai, 'threshold_') and gedai.threshold_ is not None:
  521. info_dict["optimal_threshold"] = float(gedai.threshold_)
  522. print(f" Optimal threshold: {gedai.threshold_:.4f}")
  523. LOGGER.info(f"GEDAI: Optimal threshold from model: {gedai.threshold_:.4f}")
  524. # Check for other useful attributes from the fitted model
  525. for attr in ['n_components_', 'n_rejected_', 'eigenvalues_', 'sensai_scores_']:
  526. if hasattr(gedai, attr):
  527. val = getattr(gedai, attr)
  528. if val is not None:
  529. if isinstance(val, (int, float)):
  530. info_dict[attr] = val
  531. LOGGER.info(f"GEDAI: {attr} = {val}")
  532. elif isinstance(val, np.ndarray) and val.size < 10:
  533. info_dict[attr] = val.tolist()
  534. # Update raw data in place
  535. raw._data[picks] = data_clean
  536. # Print summary
  537. print("\n" + "-" * 70)
  538. print("GEDAI DENOISING COMPLETE")
  539. print("-" * 70)
  540. print(f" SENSAI Score: {info_dict['sensai_score']:.1f}%")
  541. print(f" Variance Removed: {info_dict['variance_removed_pct']:.1f}%")
  542. print(f" SNR Improvement: {info_dict['snr_improvement_db']:.1f} dB")
  543. print(f" Official GEDAI Used: {info_dict['official_gedai_used']}")
  544. print("-" * 70 + "\n")
  545. LOGGER.info(
  546. f"GEDAI complete: SENSAI={info_dict['sensai_score']:.1f}%, "
  547. f"variance_removed={info_dict['variance_removed_pct']:.1f}%, "
  548. f"SNR_improvement={info_dict['snr_improvement_db']:.1f}dB, "
  549. f"official_gedai_used=True"
  550. )
  551. except Exception as exc:
  552. error_msg = f"Official GEDAI denoising failed: {exc}"
  553. print(f"\n*** ERROR: {error_msg}\n")
  554. LOGGER.error(error_msg)
  555. import traceback
  556. traceback.print_exc()
  557. info_dict["error"] = str(exc)
  558. info_dict["official_gedai_used"] = False
  559. info_dict["error"] = str(exc)
  560. return raw, info_dict
  561. def preprocess_raw(
  562. raw: mne.io.BaseRaw,
  563. config: PreprocessingConfig,
  564. trans_path: Optional[Path] = None,
  565. ) -> Tuple[mne.io.BaseRaw, mne.io.BaseRaw, ICA]:
  566. """
  567. Run filtering, referencing, and artifact removal.
  568. Supports multiple artifact removal methods:
  569. - ASR (Artifact Subspace Reconstruction)
  570. - ICLabel (ICA component classification)
  571. - GEDAI (Generalized Eigenvalue De-Artifacting)
  572. Returns:
  573. raw_before: Copy of raw before preprocessing
  574. raw_after: Preprocessed raw object
  575. ica: Fitted ICA object (may be empty if using GEDAI only)
  576. """
  577. preproc_info = {
  578. "asr_applied": False,
  579. "iclabel_applied": False,
  580. "gedai_applied": False,
  581. "ica_excluded": 0,
  582. "bad_channels": [],
  583. }
  584. # Log basic data shape and duration
  585. try:
  586. n_ch = mne.pick_types(raw.info, eeg=True, meg=False).size
  587. sfreq = float(raw.info.get("sfreq", 0.0) or 0.0)
  588. n_times = getattr(raw, "n_times", 0)
  589. duration = (float(n_times - 1) / sfreq) if (sfreq and n_times) else 0.0
  590. LOGGER.info(
  591. "Preprocessing: channels=%d, sfreq=%.2f Hz, duration=%.1f s",
  592. n_ch,
  593. sfreq or 0.0,
  594. duration,
  595. )
  596. except Exception:
  597. pass
  598. raw.load_data()
  599. _apply_montage(raw)
  600. raw_before = raw.copy()
  601. # Step 1: Detect and optionally interpolate bad channels
  602. if getattr(config, "interpolate_bad_channels", True):
  603. try:
  604. bad_ch_threshold = getattr(config, "bad_channel_threshold", 4.0)
  605. bad_channels = _detect_bad_channels(raw, threshold=bad_ch_threshold)
  606. if bad_channels:
  607. LOGGER.info(f"Detected {len(bad_channels)} bad channels: {bad_channels}")
  608. raw.info["bads"] = list(set(raw.info["bads"] + bad_channels))
  609. raw.interpolate_bads(reset_bads=True)
  610. preproc_info["bad_channels"] = bad_channels
  611. LOGGER.info("Interpolated bad channels")
  612. except Exception as exc:
  613. LOGGER.warning(f"Bad channel detection/interpolation failed: {exc}")
  614. # Step 2: GEDAI denoising (if enabled) - applied BEFORE filtering
  615. if getattr(config, "use_ged", False):
  616. try:
  617. LOGGER.info("Applying GEDAI denoising...")
  618. ged_threshold = getattr(config, "ged_threshold", 3.0)
  619. raw, gedai_info = _gedai_denoise(
  620. raw,
  621. threshold=ged_threshold,
  622. trans_path=trans_path,
  623. )
  624. preproc_info["gedai_applied"] = True
  625. preproc_info["gedai_info"] = gedai_info
  626. LOGGER.info(
  627. f"GEDAI applied: SENSAI={gedai_info.get('sensai_score', 0):.1f}%, "
  628. f"variance_removed={gedai_info.get('variance_removed_pct', 0):.1f}%, "
  629. f"components_rejected={gedai_info.get('n_components_rejected', 0)}"
  630. )
  631. except Exception as exc:
  632. LOGGER.warning(f"GEDAI denoising failed: {exc}")
  633. import traceback
  634. traceback.print_exc()
  635. # Step 3: ASR denoising (if enabled and available)
  636. if getattr(config, "use_asr", True):
  637. applied = False
  638. # Try asrpy first (works with MNE Raw objects directly)
  639. try:
  640. import asrpy
  641. LOGGER.info("Applying ASR (asrpy)...")
  642. # ASR parameters: cutoff=20 is a good default (5=conservative, 20=aggressive)
  643. asr = asrpy.ASR(sfreq=raw.info["sfreq"], cutoff=20)
  644. # Fit on the raw data (uses first portion as calibration)
  645. asr.fit(raw)
  646. # Transform returns a new Raw object
  647. raw = asr.transform(raw)
  648. applied = True
  649. preproc_info["asr_applied"] = True
  650. LOGGER.info("Applied ASR (asrpy) successfully")
  651. except ImportError:
  652. LOGGER.info("asrpy not installed, trying meegkit...")
  653. except Exception as exc:
  654. LOGGER.warning(f"asrpy failed: {exc}, trying meegkit...")
  655. # Try meegkit as fallback
  656. if not applied:
  657. try:
  658. from meegkit.asr import ASR as MeegkitASR
  659. LOGGER.info("Applying ASR (meegkit)...")
  660. picks = mne.pick_types(raw.info, eeg=True, meg=False)
  661. data = raw.get_data(picks=picks)
  662. # meegkit ASR expects (n_samples, n_channels) format
  663. asr = MeegkitASR(sfreq=raw.info["sfreq"], cutoff=20)
  664. # fit_transform expects (n_samples, n_channels)
  665. data_clean, _ = asr.fit_transform(data.T)
  666. raw._data[picks] = data_clean.T
  667. applied = True
  668. preproc_info["asr_applied"] = True
  669. LOGGER.info("Applied ASR (meegkit) successfully")
  670. except ImportError:
  671. LOGGER.info("meegkit not installed")
  672. except Exception as exc:
  673. LOGGER.warning(f"meegkit ASR failed: {exc}")
  674. if not applied:
  675. LOGGER.info("ASR not available (install: pip install asrpy or pip install meegkit)")
  676. # Step 4: Filtering
  677. LOGGER.info(f"Applying bandpass filter: {config.l_freq}-{config.h_freq} Hz")
  678. raw.filter(config.l_freq, config.h_freq, fir_design="firwin")
  679. if config.notch_freqs:
  680. valid_notches = [
  681. freq for freq in config.notch_freqs if freq < raw.info["sfreq"] / 2.0
  682. ]
  683. if valid_notches:
  684. LOGGER.info(f"Applying notch filter at: {valid_notches} Hz")
  685. raw.notch_filter(valid_notches)
  686. # Step 5: Re-reference to average
  687. raw.set_eeg_reference("average", projection=True)
  688. # Step 6: ICA fitting
  689. n_components = config.ica_n_components
  690. if n_components is None:
  691. # Auto-determine: use rank or max 25 components
  692. n_eeg = len(mne.pick_types(raw.info, eeg=True, meg=False))
  693. n_components = min(n_eeg - 1, 25)
  694. ica = ICA(
  695. n_components=n_components,
  696. method=config.ica_method,
  697. random_state=config.random_state,
  698. max_iter="auto",
  699. )
  700. LOGGER.info(f"Fitting ICA: method={config.ica_method}, n_components={n_components}")
  701. # Fit ICA with optional rejection threshold
  702. reject_dict = None
  703. ica_reject = getattr(config, "ica_reject_threshold", None)
  704. if ica_reject:
  705. reject_dict = {"eeg": ica_reject}
  706. try:
  707. ica.fit(raw, reject=reject_dict)
  708. except Exception as exc:
  709. LOGGER.warning(f"ICA fit with rejection failed, retrying without: {exc}")
  710. ica.fit(raw)
  711. # Step 7: Component classification and rejection
  712. excluded = []
  713. # Try ICLabel first (if enabled)
  714. if getattr(config, "use_iclabel", True):
  715. try:
  716. from mne_icalabel import label_components
  717. LOGGER.info("Running ICLabel classification...")
  718. # label_components returns a dict with:
  719. # - 'y_pred_proba': array (n_components, 7) - probabilities for each class
  720. # - 'labels': list of predicted labels
  721. # Classes: brain, muscle artifact, eye blink, heart beat, line noise, channel noise, other
  722. labels_dict = label_components(raw, ica, method="iclabel")
  723. # Get probabilities and labels
  724. if isinstance(labels_dict, dict):
  725. ic_probs = labels_dict.get("y_pred_proba", None)
  726. ic_labels = labels_dict.get("labels", None)
  727. else:
  728. # Older versions might return differently
  729. ic_probs = getattr(labels_dict, 'y_pred_proba_', None)
  730. ic_labels = getattr(labels_dict, 'labels_', None)
  731. thr = float(getattr(config, "iclabel_brain_threshold", 0.85))
  732. if ic_probs is not None:
  733. ic_probs = np.array(ic_probs)
  734. LOGGER.info(f"ICLabel: probabilities shape = {ic_probs.shape}")
  735. if ic_labels is not None:
  736. ic_labels = [str(label) for label in ic_labels]
  737. class_counts = {
  738. label: int(ic_labels.count(label))
  739. for label in sorted(set(ic_labels))
  740. }
  741. preproc_info["iclabel_labels"] = ic_labels
  742. preproc_info["iclabel_class_counts"] = class_counts
  743. preproc_info["iclabel_probabilities"] = ic_probs.tolist()
  744. preproc_info["iclabel_brain_threshold"] = thr
  745. for idx in range(ica.n_components_):
  746. try:
  747. # Brain probability is column 0
  748. if ic_probs.ndim == 2 and idx < ic_probs.shape[0]:
  749. prob_brain = float(ic_probs[idx, 0])
  750. elif ic_probs.ndim == 1:
  751. prob_brain = float(ic_probs[idx])
  752. else:
  753. prob_brain = 0.0
  754. # Check label if available
  755. is_brain_label = False
  756. if ic_labels is not None and idx < len(ic_labels):
  757. label = ic_labels[idx]
  758. is_brain_label = (isinstance(label, str) and label.lower() == "brain")
  759. # Exclude if brain probability is below threshold
  760. # OR if label is not "brain" (when available)
  761. if prob_brain < thr:
  762. excluded.append(idx)
  763. LOGGER.debug(f"ICLabel: excluding component {idx} (brain_prob={prob_brain:.3f})")
  764. except Exception as e:
  765. LOGGER.warning(f"ICLabel: error processing component {idx}: {e}")
  766. # Keep the component if we can't classify it
  767. pass
  768. preproc_info["iclabel_applied"] = True
  769. LOGGER.info(f"ICLabel: excluding {len(excluded)}/{ica.n_components_} components (threshold={thr})")
  770. else:
  771. LOGGER.warning("ICLabel: No probabilities returned")
  772. except ImportError:
  773. LOGGER.info("mne-icalabel not installed (install: pip install mne-icalabel)")
  774. except Exception as exc:
  775. LOGGER.warning(f"ICLabel failed: {exc}")
  776. # Fallback: EOG-based detection
  777. if not excluded and not preproc_info.get("iclabel_applied"):
  778. try:
  779. eog_inds, _ = ica.find_bads_eog(raw)
  780. excluded.extend(eog_inds)
  781. LOGGER.info(f"EOG detection: found {len(eog_inds)} EOG-related components")
  782. except Exception as exc:
  783. LOGGER.warning(f"EOG detection failed: {exc}")
  784. # Also try muscle artifact detection
  785. try:
  786. muscle_inds, _ = ica.find_bads_muscle(raw)
  787. excluded.extend(muscle_inds)
  788. LOGGER.info(f"Muscle detection: found {len(muscle_inds)} muscle-related components")
  789. except Exception:
  790. pass
  791. # Apply ICA
  792. ica.exclude = list(sorted(set(excluded)))
  793. preproc_info["ica_excluded"] = len(ica.exclude)
  794. preproc_info["ica_excluded_indices"] = list(ica.exclude)
  795. if ica.exclude:
  796. ica.apply(raw)
  797. LOGGER.info(f"ICA applied: excluded {len(ica.exclude)} components")
  798. else:
  799. LOGGER.info("ICA: no components excluded")
  800. # Log preprocessing summary
  801. LOGGER.info(
  802. f"Preprocessing complete: ASR={preproc_info['asr_applied']}, "
  803. f"GEDAI={preproc_info['gedai_applied']}, "
  804. f"ICLabel={preproc_info['iclabel_applied']}, "
  805. f"ICA excluded={preproc_info['ica_excluded']}"
  806. )
  807. return raw_before, raw, ica, preproc_info
  808. def _infer_paths(file_path: Path, output_root: Path) -> Tuple[str, str, Path]:
  809. """Infer group/session and construct output directory for a file."""
  810. parts = file_path.parts
  811. group_name = parts[-3] if len(parts) >= 3 else "unknown_group"
  812. session_name = parts[-2] if len(parts) >= 2 else "unknown_session"
  813. out_dir = (output_root / group_name / session_name / file_path.stem).expanduser()
  814. out_dir.mkdir(parents=True, exist_ok=True)
  815. return group_name, session_name, out_dir
  816. def _compute_preprocessing_metrics(
  817. raw_before: mne.io.BaseRaw,
  818. raw_after: mne.io.BaseRaw,
  819. preproc_info: dict,
  820. ) -> dict:
  821. """Compute preprocessing quality metrics."""
  822. metrics = {}
  823. try:
  824. # Get data
  825. data_before = raw_before.get_data(picks="eeg")
  826. data_after = raw_after.get_data(picks="eeg")
  827. # Variance reduction (percentage)
  828. var_before = np.var(data_before)
  829. var_after = np.var(data_after)
  830. if var_before > 0:
  831. metrics["variance_reduction"] = float(100 * (1 - var_after / var_before))
  832. else:
  833. metrics["variance_reduction"] = 0.0
  834. # Signal quality score (simplified SNR estimate)
  835. # Higher score = better quality after cleaning
  836. std_before = np.std(data_before, axis=1).mean()
  837. std_after = np.std(data_after, axis=1).mean()
  838. # Estimate noise as high-frequency content (simple approximation)
  839. # Filter to get high-freq component
  840. try:
  841. raw_hf_before = raw_before.copy().filter(30, None, verbose=False)
  842. raw_hf_after = raw_after.copy().filter(30, None, verbose=False)
  843. noise_before = np.std(raw_hf_before.get_data(picks="eeg"))
  844. noise_after = np.std(raw_hf_after.get_data(picks="eeg"))
  845. # SNR improvement in dB
  846. if noise_after > 0 and noise_before > 0:
  847. snr_before = std_before / noise_before
  848. snr_after = std_after / noise_after
  849. metrics["snr_improvement_db"] = float(20 * np.log10(snr_after / snr_before))
  850. else:
  851. metrics["snr_improvement_db"] = 0.0
  852. # Signal quality: percentage based on variance reduction and SNR
  853. snr_factor = min(100, max(0, 50 + metrics["snr_improvement_db"] * 5))
  854. var_factor = min(100, max(0, metrics["variance_reduction"]))
  855. metrics["signal_quality"] = float((snr_factor + var_factor) / 2)
  856. except Exception:
  857. metrics["snr_improvement_db"] = 0.0
  858. metrics["signal_quality"] = max(0, min(100, metrics["variance_reduction"]))
  859. # Band power changes
  860. try:
  861. bands = {
  862. "delta": (1, 4),
  863. "theta": (4, 8),
  864. "alpha": (8, 13),
  865. "beta": (13, 30),
  866. }
  867. band_changes = {}
  868. psd_before = raw_before.compute_psd(fmin=1, fmax=40, picks="eeg", verbose=False)
  869. psd_after = raw_after.compute_psd(fmin=1, fmax=40, picks="eeg", verbose=False)
  870. freqs = psd_before.freqs
  871. psd_data_before = psd_before.get_data().mean(axis=0)
  872. psd_data_after = psd_after.get_data().mean(axis=0)
  873. for band_name, (fmin, fmax) in bands.items():
  874. mask = (freqs >= fmin) & (freqs <= fmax)
  875. power_before = psd_data_before[mask].mean()
  876. power_after = psd_data_after[mask].mean()
  877. if power_before > 0:
  878. change_pct = 100 * (power_after - power_before) / power_before
  879. band_changes[band_name] = float(change_pct)
  880. else:
  881. band_changes[band_name] = 0.0
  882. metrics["band_power_changes"] = band_changes
  883. except Exception as e:
  884. LOGGER.warning(f"Band power computation failed: {e}")
  885. metrics["band_power_changes"] = {}
  886. # Add preprocessing info
  887. metrics["ica_components_rejected"] = preproc_info.get("ica_excluded", 0)
  888. metrics["bad_channels"] = preproc_info.get("bad_channels", [])
  889. metrics["asr_applied"] = preproc_info.get("asr_applied", False)
  890. metrics["iclabel_applied"] = preproc_info.get("iclabel_applied", False)
  891. metrics["gedai_applied"] = preproc_info.get("gedai_applied", False)
  892. # GEDAI-specific metrics
  893. if preproc_info.get("gedai_applied") and "gedai_info" in preproc_info:
  894. gedai_info = preproc_info["gedai_info"]
  895. metrics["sensai_score"] = gedai_info.get("sensai_score", 0)
  896. metrics["gedai_components_rejected"] = gedai_info.get("n_components_rejected", 0)
  897. metrics["gedai_variance_removed"] = gedai_info.get("variance_removed_pct", 0)
  898. metrics["gedai_snr_improvement"] = gedai_info.get("snr_improvement_db", 0)
  899. metrics["gedai_threshold"] = gedai_info.get("threshold", 3.0)
  900. # Override signal quality with SENSAI score for GEDAI
  901. metrics["signal_quality"] = gedai_info.get("sensai_score", metrics.get("signal_quality", 0))
  902. # Also use GEDAI's variance and SNR metrics if available
  903. if gedai_info.get("variance_removed_pct", 0) > 0:
  904. metrics["variance_reduction"] = gedai_info.get("variance_removed_pct", metrics.get("variance_reduction", 0))
  905. if gedai_info.get("snr_improvement_db") is not None:
  906. metrics["snr_improvement_db"] = gedai_info.get("snr_improvement_db", metrics.get("snr_improvement_db", 0))
  907. except Exception as e:
  908. LOGGER.warning(f"Metrics computation failed: {e}")
  909. return metrics
  910. def preprocess_file(
  911. file_path: Path,
  912. config: "PipelineConfig",
  913. progress_cb: Optional[Callable[[str], None]] = None,
  914. ) -> dict:
  915. """Preprocess a single .set file and persist outputs.
  916. Saves diagnostics and a preprocessed FIF file in the subject's output folder.
  917. Returns a summary with key paths.
  918. """
  919. group_name, session_name, out_dir = _infer_paths(file_path, config.output_root)
  920. if progress_cb:
  921. progress_cb(f"Preprocessing {file_path.name} ({group_name}/{session_name}) …")
  922. summary = {
  923. "file": str(file_path),
  924. "group": group_name,
  925. "session": session_name,
  926. "out_dir": str(out_dir),
  927. "preprocessed_fif": None,
  928. "diagnostics": {},
  929. "metrics": {},
  930. }
  931. try:
  932. raw = mne.io.read_raw_eeglab(str(file_path), preload=False, verbose=False)
  933. raw_before, raw_after, _ica, preproc_info = preprocess_raw(raw, config.preprocessing)
  934. # Compute preprocessing metrics
  935. try:
  936. metrics = _compute_preprocessing_metrics(raw_before, raw_after, preproc_info)
  937. metrics["ica_components_total"] = _ica.n_components_ if hasattr(_ica, 'n_components_') else None
  938. summary["metrics"] = metrics
  939. # Save metrics to JSON
  940. metrics_path = out_dir / "preprocessing_metrics.json"
  941. with open(metrics_path, 'w') as f:
  942. json.dump(metrics, f, indent=2)
  943. LOGGER.info(f"Saved preprocessing metrics to {metrics_path}")
  944. except Exception as exc:
  945. LOGGER.warning("Metrics computation failed for %s: %s", file_path.name, exc)
  946. # Diagnostics - visualization plots
  947. try:
  948. traces_path = out_dir / "traces_before_after.png"
  949. topo_path = out_dir / "topomap_before_after.png"
  950. gfp_path = out_dir / "gfp_before_after.png"
  951. psd_path = out_dir / "psd_before_after.png"
  952. _plot_traces_before_after(raw_before, raw_after, traces_path)
  953. _plot_topomap(raw_before, raw_after, topo_path)
  954. _plot_gfp(raw_before, raw_after, gfp_path)
  955. _plot_psd(raw_before, raw_after, "Power Spectrum: Before vs After Preprocessing", psd_path)
  956. summary["diagnostics"] = {
  957. "traces": str(traces_path),
  958. "topomap": str(topo_path),
  959. "gfp": str(gfp_path),
  960. "psd": str(psd_path),
  961. }
  962. except Exception as exc:
  963. LOGGER.warning("Diagnostics failed for %s: %s", file_path.name, exc)
  964. # Save preprocessed raw as FIF format (needed by downstream steps)
  965. fif_path = out_dir / "preprocessed_raw.fif"
  966. try:
  967. raw_after.save(fif_path, overwrite=True, verbose=False)
  968. summary["preprocessed_fif"] = str(fif_path)
  969. if progress_cb:
  970. progress_cb(f"Saved preprocessed FIF: {fif_path}")
  971. except Exception as exc:
  972. if progress_cb:
  973. progress_cb(f"Failed to save preprocessed FIF for {file_path.name}: {exc}")
  974. # Save preprocessed raw as EEGLAB .set format in the processed output folder
  975. cleaned_set_name = f"{file_path.stem}_cleaned.set"
  976. set_path = out_dir / cleaned_set_name
  977. try:
  978. raw_after.export(set_path, overwrite=True, verbose=False)
  979. summary["preprocessed_set"] = str(set_path)
  980. if progress_cb:
  981. progress_cb(f"Saved cleaned SET: {set_path}")
  982. except Exception as exc:
  983. if progress_cb:
  984. progress_cb(f"Failed to save preprocessed SET for {file_path.name}: {exc}")
  985. if progress_cb:
  986. progress_cb(f"Finished preprocessing {file_path.name}")
  987. except Exception as exc:
  988. LOGGER.exception("Preprocessing error for %s", file_path)
  989. if progress_cb:
  990. progress_cb(f"Preprocessing error for {file_path.name}: {exc}")
  991. return summary
  992. def _plot_psd(raw_before: mne.io.BaseRaw, raw_after: mne.io.BaseRaw, title: str, out_file: Path) -> None:
  993. """Plot mean PSD before/after preprocessing with frequency band annotations."""
  994. psd_kwargs = dict(fmin=1.0, fmax=45.0, picks="eeg", tmin=None, tmax=None, n_fft=2048)
  995. psd_before = raw_before.compute_psd(**psd_kwargs)
  996. psd_after = raw_after.compute_psd(**psd_kwargs)
  997. freqs = psd_before.freqs
  998. psd_data_before = 10 * np.log10(psd_before.get_data().mean(axis=0) + 1e-10)
  999. psd_data_after = 10 * np.log10(psd_after.get_data().mean(axis=0) + 1e-10)
  1000. # Create figure with dark theme
  1001. fig, ax = plt.subplots(figsize=(12, 6))
  1002. fig.patch.set_facecolor('#1e1e2e')
  1003. ax.set_facecolor('#252536')
  1004. # Define frequency bands for shading
  1005. bands = [
  1006. ("Delta", 1, 4, "#89b4fa", 0.15),
  1007. ("Theta", 4, 8, "#a6e3a1", 0.15),
  1008. ("Alpha", 8, 13, "#f9e2af", 0.15),
  1009. ("Beta", 13, 30, "#f38ba8", 0.15),
  1010. ("Gamma", 30, 45, "#cba6f7", 0.15),
  1011. ]
  1012. # Add band shading
  1013. y_min, y_max = min(psd_data_before.min(), psd_data_after.min()) - 5, max(psd_data_before.max(), psd_data_after.max()) + 5
  1014. for band_name, fmin, fmax, color, alpha in bands:
  1015. ax.axvspan(fmin, fmax, alpha=alpha, color=color, label=None)
  1016. # Add band label at top
  1017. ax.text((fmin + fmax) / 2, y_max - 2, band_name, fontsize=8, color=color,
  1018. ha='center', va='top', fontweight='bold', alpha=0.8)
  1019. # Plot PSD curves
  1020. ax.plot(freqs, psd_data_before, label="Before Preprocessing", color="#f38ba8",
  1021. lw=2, alpha=0.9)
  1022. ax.plot(freqs, psd_data_after, label="After Preprocessing", color="#a6e3a1",
  1023. lw=2, alpha=0.9)
  1024. # Fill between to show reduction
  1025. ax.fill_between(freqs, psd_data_before, psd_data_after,
  1026. where=psd_data_before > psd_data_after,
  1027. alpha=0.2, color='#a6e3a1', label='Power Reduction')
  1028. # Styling
  1029. ax.set_xlabel("Frequency (Hz)", color='#cdd6f4', fontsize=11)
  1030. ax.set_ylabel("Power Spectral Density (dB)", color='#cdd6f4', fontsize=11)
  1031. ax.set_title(title, color='#89b4fa', fontsize=14, fontweight='bold', pad=15)
  1032. ax.tick_params(colors='#cdd6f4')
  1033. for spine in ax.spines.values():
  1034. spine.set_color('#45475a')
  1035. ax.set_xlim([1, 45])
  1036. ax.set_ylim([y_min, y_max])
  1037. ax.grid(True, alpha=0.2, color='#6c7086')
  1038. # Legend
  1039. legend = ax.legend(loc='upper right', facecolor='#313244', edgecolor='#45475a',
  1040. fontsize=10, framealpha=0.9)
  1041. for text in legend.get_texts():
  1042. text.set_color('#cdd6f4')
  1043. fig.tight_layout()
  1044. fig.savefig(out_file, dpi=200, facecolor='#1e1e2e', edgecolor='none')
  1045. plt.close(fig)
  1046. def _plot_traces_before_after(
  1047. raw_before: mne.io.BaseRaw,
  1048. raw_after: mne.io.BaseRaw,
  1049. out_file: Path,
  1050. n_channels: int = 6,
  1051. duration: float = 10.0,
  1052. ) -> None:
  1053. """Plot a few EEG channels before/after preprocessing with improved visualization."""
  1054. picks = mne.pick_types(raw_after.info, eeg=True, meg=False)
  1055. if len(picks) == 0:
  1056. return
  1057. picks = picks[: min(n_channels, len(picks))]
  1058. ch_names = [raw_after.ch_names[i] for i in picks]
  1059. sfreq = raw_after.info["sfreq"]
  1060. tmax = min(duration, raw_after.times[-1] if len(raw_after.times) else duration)
  1061. n_samp = int(max(1, tmax * sfreq))
  1062. data_b = raw_before.get_data(picks=picks)[:, :n_samp] * 1e6
  1063. data_a = raw_after.get_data(picks=picks)[:, :n_samp] * 1e6
  1064. times = np.arange(n_samp) / sfreq
  1065. # Calculate amplitude stats for display
  1066. std_before = np.std(data_b)
  1067. std_after = np.std(data_a)
  1068. reduction_pct = 100 * (1 - std_after / std_before) if std_before > 0 else 0
  1069. # Create figure with better layout
  1070. fig = plt.figure(figsize=(12, 7))
  1071. fig.patch.set_facecolor('#1e1e2e')
  1072. # Create grid for layout
  1073. gs = fig.add_gridspec(2, 3, width_ratios=[1, 1, 0.15], height_ratios=[1, 1],
  1074. hspace=0.15, wspace=0.05, left=0.08, right=0.92, top=0.88, bottom=0.12)
  1075. ax_before = fig.add_subplot(gs[0, 0:2])
  1076. ax_after = fig.add_subplot(gs[1, 0:2], sharex=ax_before)
  1077. # Set dark background
  1078. for ax in [ax_before, ax_after]:
  1079. ax.set_facecolor('#252536')
  1080. ax.tick_params(colors='#cdd6f4')
  1081. for spine in ax.spines.values():
  1082. spine.set_color('#45475a')
  1083. # Calculate offsets for stacking channels
  1084. max_amp = max(np.max(np.abs(data_b)), np.max(np.abs(data_a)))
  1085. channel_spacing = max_amp * 0.6 if max_amp > 0 else 50.0
  1086. offsets = np.arange(len(picks)) * channel_spacing
  1087. # Plot before
  1088. for ch in range(len(picks)):
  1089. ax_before.plot(times, data_b[ch] + offsets[ch], color="#f38ba8", lw=0.7, alpha=0.9)
  1090. ax_before.set_title("BEFORE Preprocessing", fontsize=12, fontweight='bold',
  1091. color='#f38ba8', pad=10)
  1092. ax_before.set_ylabel("Channels", color='#cdd6f4')
  1093. # Add channel labels on the left
  1094. for ch, name in enumerate(ch_names):
  1095. ax_before.text(-0.5, offsets[ch], name, fontsize=8, color='#6c7086',
  1096. ha='right', va='center', transform=ax_before.get_yaxis_transform())
  1097. # Plot after
  1098. for ch in range(len(picks)):
  1099. ax_after.plot(times, data_a[ch] + offsets[ch], color="#a6e3a1", lw=0.7, alpha=0.9)
  1100. ax_after.set_title("AFTER Preprocessing", fontsize=12, fontweight='bold',
  1101. color='#a6e3a1', pad=10)
  1102. ax_after.set_xlabel("Time (s)", color='#cdd6f4')
  1103. ax_after.set_ylabel("Channels", color='#cdd6f4')
  1104. # Add channel labels
  1105. for ch, name in enumerate(ch_names):
  1106. ax_after.text(-0.5, offsets[ch], name, fontsize=8, color='#6c7086',
  1107. ha='right', va='center', transform=ax_after.get_yaxis_transform())
  1108. # Configure axes
  1109. for ax in [ax_before, ax_after]:
  1110. ax.set_yticks([])
  1111. ax.grid(True, alpha=0.15, color='#6c7086')
  1112. ax.set_xlim([0, tmax])
  1113. plt.setp(ax_before.get_xticklabels(), visible=False)
  1114. # Add summary stats as text box
  1115. stats_text = (
  1116. f"Amplitude Reduction\n"
  1117. f"━━━━━━━━━━━━━━━\n"
  1118. f"Before: {std_before:.1f} µV\n"
  1119. f"After: {std_after:.1f} µV\n"
  1120. f"━━━━━━━━━━━━━━━\n"
  1121. f"Reduction: {reduction_pct:.0f}%"
  1122. )
  1123. fig.text(0.95, 0.5, stats_text, fontsize=9, color='#cdd6f4',
  1124. ha='center', va='center', fontfamily='monospace',
  1125. bbox=dict(boxstyle='round,pad=0.5', facecolor='#313244',
  1126. edgecolor='#45475a', alpha=0.9))
  1127. # Main title
  1128. fig.suptitle("EEG Signal Comparison: Before vs After Preprocessing",
  1129. fontsize=14, fontweight='bold', color='#89b4fa', y=0.96)
  1130. fig.savefig(out_file, dpi=200, facecolor='#1e1e2e', edgecolor='none')
  1131. plt.close(fig)
  1132. def _plot_topomap(raw_before: mne.io.BaseRaw, raw_after: mne.io.BaseRaw, out_file: Path) -> None:
  1133. """Compare RMS scalp maps before/after preprocessing with dark theme."""
  1134. data_before = raw_before.get_data(picks="eeg")
  1135. data_after = raw_after.get_data(picks="eeg")
  1136. rms_before = np.sqrt((data_before**2).mean(axis=1))
  1137. rms_after = np.sqrt((data_after**2).mean(axis=1))
  1138. # Calculate reduction
  1139. reduction = 100 * (1 - rms_after.mean() / rms_before.mean()) if rms_before.mean() > 0 else 0
  1140. fig, axes = plt.subplots(1, 3, figsize=(14, 5))
  1141. fig.patch.set_facecolor('#1e1e2e')
  1142. # Before topomap
  1143. im1, _ = mne.viz.plot_topomap(rms_before, raw_before.info, axes=axes[0], show=False,
  1144. contours=0, cmap='RdYlBu_r')
  1145. axes[0].set_title("BEFORE Preprocessing", color='#f38ba8', fontsize=12, fontweight='bold')
  1146. # After topomap
  1147. im2, _ = mne.viz.plot_topomap(rms_after, raw_after.info, axes=axes[1], show=False,
  1148. contours=0, cmap='RdYlBu_r')
  1149. axes[1].set_title("AFTER Preprocessing", color='#a6e3a1', fontsize=12, fontweight='bold')
  1150. # Difference map (reduction)
  1151. rms_diff = rms_before - rms_after
  1152. im3, _ = mne.viz.plot_topomap(rms_diff, raw_after.info, axes=axes[2], show=False,
  1153. contours=0, cmap='Greens')
  1154. axes[2].set_title(f"Noise Reduction ({reduction:.0f}%)", color='#89b4fa', fontsize=12, fontweight='bold')
  1155. # Main title
  1156. fig.suptitle("Scalp Topography: RMS Amplitude Comparison",
  1157. fontsize=14, fontweight='bold', color='#89b4fa', y=0.98)
  1158. fig.tight_layout(rect=[0, 0, 1, 0.95])
  1159. fig.savefig(out_file, dpi=200, facecolor='#1e1e2e', edgecolor='none')
  1160. plt.close(fig)
  1161. def _plot_gfp(raw_before: mne.io.BaseRaw, raw_after: mne.io.BaseRaw, out_file: Path) -> None:
  1162. """Plot global field power before/after cleaning with dark theme."""
  1163. def _gfp(raw: mne.io.BaseRaw) -> Tuple[np.ndarray, np.ndarray]:
  1164. data = raw.get_data(picks="eeg")
  1165. gfp = np.sqrt((data**2).mean(axis=0)) * 1e6
  1166. return gfp, raw.times
  1167. gfp_before, times = _gfp(raw_before)
  1168. gfp_after, _ = _gfp(raw_after)
  1169. # Calculate stats
  1170. mean_before = np.mean(gfp_before)
  1171. mean_after = np.mean(gfp_after)
  1172. peak_before = np.max(gfp_before)
  1173. peak_after = np.max(gfp_after)
  1174. reduction = 100 * (1 - mean_after / mean_before) if mean_before > 0 else 0
  1175. step = max(1, len(times) // 5000)
  1176. fig, ax = plt.subplots(figsize=(12, 5))
  1177. fig.patch.set_facecolor('#1e1e2e')
  1178. ax.set_facecolor('#252536')
  1179. # Plot GFP traces
  1180. ax.fill_between(times[::step], 0, gfp_before[::step], alpha=0.3, color='#f38ba8', label='Before')
  1181. ax.fill_between(times[::step], 0, gfp_after[::step], alpha=0.3, color='#a6e3a1', label='After')
  1182. ax.plot(times[::step], gfp_before[::step], color='#f38ba8', lw=1, alpha=0.8)
  1183. ax.plot(times[::step], gfp_after[::step], color='#a6e3a1', lw=1, alpha=0.8)
  1184. # Add horizontal lines for means
  1185. ax.axhline(mean_before, color='#f38ba8', linestyle='--', lw=1.5, alpha=0.7)
  1186. ax.axhline(mean_after, color='#a6e3a1', linestyle='--', lw=1.5, alpha=0.7)
  1187. # Styling
  1188. ax.set_xlabel("Time (s)", color='#cdd6f4', fontsize=11)
  1189. ax.set_ylabel("Global Field Power (µV)", color='#cdd6f4', fontsize=11)
  1190. ax.set_title("Global Field Power: Before vs After Preprocessing",
  1191. color='#89b4fa', fontsize=14, fontweight='bold', pad=15)
  1192. ax.tick_params(colors='#cdd6f4')
  1193. for spine in ax.spines.values():
  1194. spine.set_color('#45475a')
  1195. ax.grid(True, alpha=0.2, color='#6c7086')
  1196. # Legend with stats
  1197. legend = ax.legend(loc='upper right', facecolor='#313244', edgecolor='#45475a',
  1198. fontsize=10, framealpha=0.9)
  1199. for text in legend.get_texts():
  1200. text.set_color('#cdd6f4')
  1201. # Add stats text box
  1202. stats_text = (
  1203. f"Mean GFP Reduction: {reduction:.0f}%\n"
  1204. f"Before: {mean_before:.1f} µV (peak: {peak_before:.1f})\n"
  1205. f"After: {mean_after:.1f} µV (peak: {peak_after:.1f})"
  1206. )
  1207. ax.text(0.02, 0.98, stats_text, transform=ax.transAxes, fontsize=9, color='#cdd6f4',
  1208. va='top', ha='left', fontfamily='monospace',
  1209. bbox=dict(boxstyle='round,pad=0.4', facecolor='#313244', edgecolor='#45475a', alpha=0.9))
  1210. fig.tight_layout()
  1211. fig.savefig(out_file, dpi=200, facecolor='#1e1e2e', edgecolor='none')
  1212. plt.close(fig)
  1213. def _prepare_inverse_operator(raw: mne.io.BaseRaw, source_cfg: SourceConfig) -> Tuple[dict, List[mne.Label], mne.SourceSpaces, Path]:
  1214. """Build inverse operator plus label definitions."""
  1215. subjects_dir = ensure_subjects_dir(source_cfg.subjects_dir)
  1216. subject = "fsaverage"
  1217. src_dir = Path(subjects_dir) / subject / "bem"
  1218. src_path = src_dir / f"{subject}-{source_cfg.spacing}-src.fif"
  1219. bem_path = src_dir / source_cfg.bem_solution_name
  1220. if not src_path.exists():
  1221. src = mne.setup_source_space(
  1222. subject=subject,
  1223. spacing=source_cfg.spacing,
  1224. subjects_dir=subjects_dir,
  1225. add_dist=False,
  1226. )
  1227. mne.write_source_spaces(src_path, src, overwrite=True)
  1228. else:
  1229. src = mne.read_source_spaces(src_path)
  1230. if not bem_path.exists():
  1231. mne.make_bem_model(
  1232. subject=subject,
  1233. ico=5,
  1234. conductivity=(0.3, 0.006, 0.3),
  1235. subjects_dir=subjects_dir,
  1236. output=Path(subjects_dir) / subject / "bem" / f"{subject}-5120-5120-5120-bem.fif",
  1237. overwrite=True,
  1238. )
  1239. bem = mne.read_bem_solution(bem_path)
  1240. else:
  1241. bem = mne.read_bem_solution(bem_path)
  1242. trans_path = source_cfg.trans_path
  1243. if not trans_path.exists():
  1244. raise FileNotFoundError(
  1245. f"Transformation file '{trans_path}' not found. Provide a valid head<->MRI transform."
  1246. )
  1247. fwd = mne.make_forward_solution(
  1248. info=raw.info,
  1249. trans=str(trans_path),
  1250. src=src,
  1251. bem=bem,
  1252. meg=False,
  1253. eeg=True,
  1254. mindist=5.0,
  1255. n_jobs=1,
  1256. )
  1257. noise_cov = mne.compute_raw_covariance(raw, method="shrunk", rank="info")
  1258. inverse_operator = make_inverse_operator(raw.info, fwd, noise_cov, loose="auto", depth=0.8)
  1259. labels = [
  1260. label
  1261. for label in mne.read_labels_from_annot(subject, parc="aparc", subjects_dir=subjects_dir)
  1262. if label.name.split("-")[0].lower() != "unknown"
  1263. ]
  1264. return inverse_operator, labels, src, Path(subjects_dir)
  1265. def _run_source_localization(
  1266. raw: mne.io.BaseRaw,
  1267. inverse_operator,
  1268. labels: Sequence[mne.Label],
  1269. src: mne.SourceSpaces,
  1270. source_cfg: SourceConfig,
  1271. ) -> Tuple[np.ndarray, Sequence[str], mne.SourceEstimate]:
  1272. """Apply inverse operator and extract label time courses."""
  1273. stc = apply_inverse_raw(
  1274. raw,
  1275. inverse_operator,
  1276. lambda2=source_cfg.lambda2,
  1277. method="MNE",
  1278. pick_ori="normal",
  1279. )
  1280. label_tc = mne.extract_label_time_course([stc], labels, src=src, mode="pca_flip")[0]
  1281. label_names = [label.name for label in labels]
  1282. return label_tc, label_names, stc
  1283. def _label_tc_to_epochs(
  1284. label_tc: np.ndarray,
  1285. label_names: Sequence[str],
  1286. sfreq: float,
  1287. epoch_length: float,
  1288. ) -> EpochsArray:
  1289. """Convert label time courses into fixed-length epochs."""
  1290. samples_per_epoch = int(epoch_length * sfreq)
  1291. total_samples = label_tc.shape[1]
  1292. usable = (total_samples // samples_per_epoch) * samples_per_epoch
  1293. if usable < samples_per_epoch:
  1294. raise ValueError("Not enough data to create even a single epoch for connectivity.")
  1295. trimmed = label_tc[:, :usable]
  1296. epochs = trimmed.reshape(len(label_names), -1, samples_per_epoch)
  1297. data = np.transpose(epochs, (1, 0, 2)) # (n_epochs, n_labels, n_times)
  1298. info = mne.create_info(ch_names=list(label_names), sfreq=sfreq, ch_types="misc")
  1299. return EpochsArray(data, info, verbose=False)
  1300. def _bandpass_gpu(data: "cp.ndarray", sfreq: float, band: Tuple[float, float]) -> "cp.ndarray":
  1301. fmin, fmax = band
  1302. n_times = data.shape[1]
  1303. freqs = cp.fft.rfftfreq(n_times, d=1.0 / sfreq)
  1304. spectrum = cp.fft.rfft(data, axis=1)
  1305. mask = (freqs >= fmin) & (freqs <= fmax)
  1306. spectrum *= mask
  1307. return cp.fft.irfft(spectrum, n=n_times, axis=1)
  1308. def _hilbert_gpu(data: "cp.ndarray") -> "cp.ndarray":
  1309. n_times = data.shape[1]
  1310. spectrum = cp.fft.fft(data, axis=1)
  1311. h = cp.zeros(n_times, dtype=data.dtype)
  1312. if n_times % 2 == 0:
  1313. h[0] = h[n_times // 2] = 1.0
  1314. h[1:n_times // 2] = 2.0
  1315. else:
  1316. h[0] = 1.0
  1317. h[1:(n_times + 1) // 2] = 2.0
  1318. spectrum *= h
  1319. return cp.fft.ifft(spectrum, axis=1)
  1320. def _compute_pli_gpu(
  1321. epochs: EpochsArray,
  1322. band: Tuple[float, float],
  1323. chunk_size: int = 1024, # Reduced chunk size for memory safety
  1324. ) -> np.ndarray:
  1325. """Compute PLI using GPU acceleration with proper memory management."""
  1326. if not GPU_AVAILABLE or cp is None:
  1327. raise RuntimeError("GPU computation requested but CuPy is unavailable.")
  1328. result = None
  1329. try:
  1330. # Synchronize GPU before starting
  1331. cp.cuda.Stream.null.synchronize()
  1332. data = epochs.get_data(copy=True) # (n_epochs, n_labels, n_times)
  1333. n_epochs, n_labels, n_times = data.shape
  1334. merged = data.transpose(1, 0, 2).reshape(n_labels, -1)
  1335. # Check available GPU memory and adjust if needed
  1336. mempool = cp.get_default_memory_pool()
  1337. try:
  1338. free_mem = cp.cuda.runtime.memGetInfo()[0]
  1339. data_size = merged.nbytes * 4 # float32
  1340. if data_size > free_mem * 0.7: # Use at most 70% of free memory
  1341. LOGGER.warning("GPU memory limited, using smaller chunks")
  1342. chunk_size = max(256, chunk_size // 4)
  1343. except Exception:
  1344. pass
  1345. cp_data = cp.asarray(merged, dtype=cp.float32)
  1346. filtered = _bandpass_gpu(cp_data, epochs.info["sfreq"], band)
  1347. # Free intermediate data
  1348. del cp_data
  1349. mempool.free_all_blocks()
  1350. analytic = _hilbert_gpu(filtered)
  1351. # Free filtered data
  1352. del filtered
  1353. mempool.free_all_blocks()
  1354. samples = analytic.shape[1]
  1355. analytic = analytic.T # (samples, n_labels)
  1356. accum = cp.zeros((n_labels, n_labels), dtype=cp.float32)
  1357. for start in range(0, samples, chunk_size):
  1358. stop = min(samples, start + chunk_size)
  1359. segment = analytic[start:stop]
  1360. if segment.size == 0:
  1361. continue
  1362. phase = cp.angle(segment)
  1363. diff = phase[:, :, None] - phase[:, None, :]
  1364. accum += cp.sign(cp.sin(diff)).sum(axis=0)
  1365. # Free intermediate results each iteration
  1366. del phase, diff
  1367. cp.cuda.Stream.null.synchronize()
  1368. pli_gpu = cp.abs(accum / samples)
  1369. # Copy result to CPU before cleanup
  1370. result = cp.asnumpy(pli_gpu)
  1371. # Clean up GPU memory
  1372. del analytic, accum, pli_gpu
  1373. mempool.free_all_blocks()
  1374. cp.cuda.Stream.null.synchronize()
  1375. except Exception as e:
  1376. LOGGER.error(f"GPU PLI computation failed: {e}")
  1377. # Clean up GPU memory on error
  1378. try:
  1379. mempool = cp.get_default_memory_pool()
  1380. mempool.free_all_blocks()
  1381. cp.cuda.Stream.null.synchronize()
  1382. except Exception:
  1383. pass
  1384. raise
  1385. return result
  1386. def _should_use_gpu(user_pref: bool) -> bool:
  1387. """Determine if GPU should be used for PLI computation."""
  1388. if user_pref:
  1389. if not GPU_AVAILABLE:
  1390. LOGGER.warning("GPU requested but CuPy/CUDA not available, using CPU")
  1391. return False
  1392. return True
  1393. return False
  1394. def compute_pli(
  1395. epochs: EpochsArray,
  1396. band: Tuple[float, float],
  1397. method: str = "pli",
  1398. use_gpu: bool = False,
  1399. ) -> np.ndarray:
  1400. """Compute PLI matrix for a frequency band."""
  1401. if use_gpu and method.lower() == "pli":
  1402. try:
  1403. LOGGER.info("Computing %s PLI on GPU", method.upper())
  1404. result = _compute_pli_gpu(epochs, band)
  1405. if result is not None:
  1406. return result
  1407. LOGGER.warning("GPU PLI returned None, falling back to CPU")
  1408. except Exception as exc: # pragma: no cover - GPU fallback path
  1409. LOGGER.warning("GPU PLI failed (%s); falling back to CPU.", exc)
  1410. # Ensure GPU memory is cleaned up
  1411. try:
  1412. if cp is not None:
  1413. cp.get_default_memory_pool().free_all_blocks()
  1414. cp.cuda.Stream.null.synchronize()
  1415. except Exception:
  1416. pass
  1417. fmin, fmax = band
  1418. con = spectral_connectivity_epochs(
  1419. epochs,
  1420. method=method,
  1421. mode="multitaper",
  1422. sfreq=epochs.info["sfreq"],
  1423. fmin=fmin,
  1424. fmax=fmax,
  1425. faverage=True,
  1426. mt_adaptive=True,
  1427. verbose=False,
  1428. )
  1429. if hasattr(con, "get_data"):
  1430. dense = con.get_data(output="dense")
  1431. if dense.ndim == 4:
  1432. dense = dense[0]
  1433. matrix = np.squeeze(dense, axis=-1)
  1434. else: # pragma: no cover - old tuple output from mne-connectivity<0.6
  1435. matrix = con[0]
  1436. if matrix.ndim == 3:
  1437. matrix = np.squeeze(matrix, axis=-1)
  1438. matrix = np.asarray(matrix, dtype=float)
  1439. if matrix.ndim != 2 or matrix.shape[0] != matrix.shape[1]:
  1440. raise ValueError(
  1441. f"Unexpected PLI matrix shape {matrix.shape}; expected square connectivity matrix."
  1442. )
  1443. matrix = np.maximum(matrix, matrix.T)
  1444. np.fill_diagonal(matrix, 0.0)
  1445. return matrix
  1446. def run_source_loc_and_connectivity(
  1447. preprocessed_fif: Path,
  1448. original_file: Optional[Path],
  1449. config: "PipelineConfig",
  1450. progress_cb: Optional[Callable[[str], None]] = None,
  1451. ) -> dict:
  1452. """Run source localization and PLI on a preprocessed FIF.
  1453. Saves per-band CSV/PNG and returns a summary with pli_outputs mapping.
  1454. """
  1455. # Determine out_dir and group/session from FIF location
  1456. out_dir = preprocessed_fif.parent
  1457. # Try to recover names from parent folders
  1458. parts = out_dir.parts
  1459. group_name = parts[-3] if len(parts) >= 3 else "unknown_group"
  1460. session_name = parts[-2] if len(parts) >= 2 else "unknown_session"
  1461. if progress_cb:
  1462. progress_cb(f"Source/Connectivity for {out_dir.name} ({group_name}/{session_name}) …")
  1463. summary = {
  1464. "file": str(original_file) if original_file else str(preprocessed_fif),
  1465. "group": group_name,
  1466. "session": session_name,
  1467. "pli_outputs": {},
  1468. }
  1469. try:
  1470. raw = mne.io.read_raw_fif(str(preprocessed_fif), preload=True, verbose=False)
  1471. inv_op, labels, src, _subjects_dir = _prepare_inverse_operator(raw, config.source)
  1472. label_tc, label_names, _stc = _run_source_localization(raw, inv_op, labels, src, config.source)
  1473. try:
  1474. _plot_source_power(label_tc, label_names, out_dir / "source_power.png")
  1475. except Exception:
  1476. pass
  1477. epochs = _label_tc_to_epochs(label_tc, label_names, raw.info["sfreq"], config.connectivity.epoch_length)
  1478. try:
  1479. LOGGER.info(
  1480. "Connectivity epochs: n_epochs=%d, epoch_length=%.1f s, labels=%d, sfreq=%.2f",
  1481. len(epochs),
  1482. float(config.connectivity.epoch_length),
  1483. len(label_names),
  1484. float(raw.info.get("sfreq", 0.0) or 0.0),
  1485. )
  1486. except Exception:
  1487. pass
  1488. use_gpu = _should_use_gpu(config.connectivity.use_gpu)
  1489. try:
  1490. LOGGER.info("PLI GPU enabled: %s (GPU_AVAILABLE=%s)", str(use_gpu), str(GPU_AVAILABLE))
  1491. except Exception:
  1492. pass
  1493. for band_name, band in config.connectivity.frequency_bands.items():
  1494. if progress_cb:
  1495. progress_cb(f"Computing {band_name} PLI for {out_dir.name} …")
  1496. try:
  1497. pli_mat = compute_pli(epochs, band, method=config.connectivity.method, use_gpu=use_gpu)
  1498. # Validate result
  1499. if pli_mat is None:
  1500. LOGGER.warning("PLI computation returned None for %s %s", out_dir.name, band_name)
  1501. if progress_cb:
  1502. progress_cb(f"PLI returned None for {band_name}, skipping")
  1503. continue
  1504. if not isinstance(pli_mat, np.ndarray) or pli_mat.ndim != 2:
  1505. LOGGER.warning("PLI returned invalid shape for %s %s", out_dir.name, band_name)
  1506. continue
  1507. except Exception as exc:
  1508. LOGGER.exception("PLI failed for %s %s", out_dir.name, band_name)
  1509. if progress_cb:
  1510. progress_cb(f"PLI failed for {band_name}: {exc}")
  1511. # Clean up GPU memory on error
  1512. try:
  1513. if cp is not None:
  1514. cp.get_default_memory_pool().free_all_blocks()
  1515. except Exception:
  1516. pass
  1517. continue
  1518. try:
  1519. _save_connectivity_outputs(pli_mat, label_names, band_name, out_dir, group_name)
  1520. try:
  1521. _plot_connectivity_circle(pli_mat, label_names, band_name, out_dir / f"{band_name}_circle.png")
  1522. except Exception:
  1523. pass
  1524. summary["pli_outputs"][band_name] = str(out_dir / f"{band_name}_PLI.csv")
  1525. except Exception as exc:
  1526. LOGGER.warning("Saving connectivity outputs failed for %s (%s): %s", out_dir.name, band_name, exc)
  1527. if progress_cb:
  1528. progress_cb(f"Finished source/connectivity for {out_dir.name}")
  1529. except Exception as exc:
  1530. LOGGER.exception("Source/Connectivity error for %s", preprocessed_fif)
  1531. if progress_cb:
  1532. progress_cb(f"Source/Connectivity error: {exc}")
  1533. return summary
  1534. def run_preprocessing_for_design(
  1535. base_folder: Path,
  1536. design: "StudyDesign",
  1537. config: Optional["PipelineConfig"] = None,
  1538. progress_cb: Optional[Callable[[str], None]] = None,
  1539. ) -> List[dict]:
  1540. """Batch preprocessing across study design; returns list of summaries."""
  1541. config = config or PipelineConfig()
  1542. summaries: List[dict] = []
  1543. for group in design.groups:
  1544. for session in design.sessions:
  1545. session_dir = Path(base_folder) / group / session
  1546. if not session_dir.exists():
  1547. if progress_cb:
  1548. progress_cb(f"Missing folder: {session_dir}")
  1549. continue
  1550. try:
  1551. files = find_set_files(session_dir)
  1552. except FileNotFoundError:
  1553. if progress_cb:
  1554. progress_cb(f"No .set files in {session_dir}")
  1555. continue
  1556. for f in files:
  1557. summaries.append(preprocess_file(f, config, progress_cb))
  1558. return summaries
  1559. def run_source_and_connectivity_for_design(
  1560. base_folder: Path,
  1561. design: "StudyDesign",
  1562. config: Optional["PipelineConfig"] = None,
  1563. progress_cb: Optional[Callable[[str], None]] = None,
  1564. ) -> List[dict]:
  1565. """Run source localization + connectivity for all preprocessed FIFs."""
  1566. config = config or PipelineConfig()
  1567. summaries: List[dict] = []
  1568. for group in design.groups:
  1569. for session in design.sessions:
  1570. session_dir = (config.output_root / group / session)
  1571. if not session_dir.exists():
  1572. if progress_cb:
  1573. progress_cb(f"No preprocessed outputs under {session_dir}")
  1574. continue
  1575. for subj_dir in sorted(session_dir.glob("*")):
  1576. fif_path = subj_dir / "preprocessed_raw.fif"
  1577. if fif_path.exists():
  1578. summaries.append(run_source_loc_and_connectivity(fif_path, None, config, progress_cb))
  1579. else:
  1580. if progress_cb:
  1581. progress_cb(f"Missing preprocessed FIF: {fif_path}")
  1582. return summaries
  1583. def aggregate_csv_table(
  1584. design: "StudyDesign",
  1585. config: Optional["PipelineConfig"] = None,
  1586. progress_cb: Optional[Callable[[str], None]] = None,
  1587. ) -> pd.DataFrame:
  1588. """Aggregate existing PLI CSV files across output_root into an Excel table.
  1589. Includes ALL subjects found in the output directories, even those without PLI data.
  1590. """
  1591. config = config or PipelineConfig()
  1592. pseudo_summaries: List[dict] = []
  1593. for group in design.groups:
  1594. for session in design.sessions:
  1595. session_dir = (config.output_root / group / session)
  1596. if not session_dir.exists():
  1597. if progress_cb:
  1598. progress_cb(f"Warning: No output folder for {group}/{session}")
  1599. continue
  1600. # Get all subject directories (not just those with PLI files)
  1601. for subj_dir in sorted(session_dir.glob("*")):
  1602. if not subj_dir.is_dir():
  1603. continue
  1604. pli_outputs = {}
  1605. for csv_file in subj_dir.glob("*_PLI.csv"):
  1606. # band name is prefix before _PLI
  1607. band_name = csv_file.stem.replace("_PLI", "")
  1608. pli_outputs[band_name] = str(csv_file)
  1609. # Include ALL subjects, even those without PLI outputs
  1610. pseudo_summaries.append({
  1611. "file": str(subj_dir / (subj_dir.name + ".set")),
  1612. "group": group,
  1613. "session": session,
  1614. "pli_outputs": pli_outputs, # May be empty dict
  1615. })
  1616. if progress_cb:
  1617. progress_cb(f"Found {len(pseudo_summaries)} subjects across all groups")
  1618. output_excel = config.output_root / "PLI_Table.xlsx"
  1619. return compute_group_stats(pseudo_summaries, {
  1620. "SN": [1, 2, 19, 20],
  1621. "DMN": [15, 16, 21, 22, 29, 30, 31, 32, 35, 36, 47, 48, 51, 52, 53, 54],
  1622. "CEN": [5, 6, 55, 56, 57, 58, 59, 60],
  1623. }, output_excel, progress_cb)
  1624. def run_stats_for_design(
  1625. design: "StudyDesign",
  1626. config: Optional["PipelineConfig"] = None,
  1627. progress_cb: Optional[Callable[[str], None]] = None,
  1628. ) -> None:
  1629. """Run statistical analysis using existing CSVs (reconstructed summaries).
  1630. This runs the cluster-based permutation analysis and also saves simple t-test stats.
  1631. """
  1632. config = config or PipelineConfig()
  1633. # reconstruct summaries like in aggregate_csv_table
  1634. pseudo_summaries: List[dict] = []
  1635. for group in design.groups:
  1636. for session in design.sessions:
  1637. session_dir = (config.output_root / group / session)
  1638. if not session_dir.exists():
  1639. continue
  1640. for subj_dir in sorted(session_dir.glob("*")):
  1641. pli_outputs = {}
  1642. for csv_file in subj_dir.glob("*_PLI.csv"):
  1643. band_name = csv_file.stem.replace("_PLI", "")
  1644. pli_outputs[band_name] = str(csv_file)
  1645. if pli_outputs:
  1646. pseudo_summaries.append({
  1647. "file": str(subj_dir / (subj_dir.name + ".set")),
  1648. "group": group,
  1649. "session": session,
  1650. "pli_outputs": pli_outputs,
  1651. })
  1652. # Classic edge-wise t-tests (saved under stats/..)
  1653. try:
  1654. _perform_statistical_analysis(pseudo_summaries, design, config, progress_cb)
  1655. except Exception as exc:
  1656. LOGGER.warning("Simple statistical analysis failed: %s", exc)
  1657. # Cluster-based permutation (saved under stats/cluster_perm/..)
  1658. try:
  1659. _perform_cluster_based_permutation(pseudo_summaries, design, config, progress_cb)
  1660. except Exception as exc:
  1661. LOGGER.warning("Cluster-based permutation failed: %s", exc)
  1662. def _save_connectivity_outputs(
  1663. pli_matrix: np.ndarray,
  1664. label_names: Sequence[str],
  1665. band_name: str,
  1666. out_dir: Path,
  1667. group_name: str, # Add group_name parameter
  1668. ) -> None:
  1669. """Persist connectivity matrix and heatmap, and save as .mat file."""
  1670. # Create a directory for .mat files if it doesn't exist
  1671. mat_dir = out_dir / "mat files"
  1672. mat_dir.mkdir(parents=True, exist_ok=True)
  1673. # Save the PLI matrix as a .mat file with group name
  1674. mat_file_name = f"{out_dir.stem}_{group_name}_{band_name}_PLI.mat" # e.g., P1_A_Alpha_PLI.mat
  1675. mat_file_path = mat_dir / mat_file_name
  1676. scipy.io.savemat(mat_file_path, {f"{band_name}_PLI": pli_matrix})
  1677. # Save the CSV file
  1678. df = pd.DataFrame(pli_matrix, index=label_names, columns=label_names)
  1679. csv_path = out_dir / f"{band_name}_PLI.csv"
  1680. df.to_csv(csv_path, float_format="%.6f")
  1681. # Plot and save the heatmap
  1682. fig, ax = plt.subplots(figsize=(8, 6))
  1683. im = ax.imshow(pli_matrix, vmin=0.0, vmax=1.0, cmap="viridis")
  1684. ax.set_title(f"{band_name.upper()} band PLI")
  1685. ax.set_xticks(range(len(label_names)))
  1686. ax.set_yticks(range(len(label_names)))
  1687. ax.set_xticklabels(label_names, rotation=90, fontsize=6)
  1688. ax.set_yticklabels(label_names, fontsize=6)
  1689. fig.colorbar(im, ax=ax, shrink=0.6, label="PLI")
  1690. fig.tight_layout()
  1691. fig.savefig(out_dir / f"{band_name}_PLI.png", dpi=200)
  1692. plt.close(fig)
  1693. def _plot_connectivity_circle(
  1694. pli_matrix: np.ndarray,
  1695. label_names: Sequence[str],
  1696. band_name: str,
  1697. out_file: Path,
  1698. max_lines: int = 40,
  1699. ) -> None:
  1700. """Plot top connections using an MNE connectivity circle."""
  1701. if _plot_circle is None or not np.any(pli_matrix):
  1702. return
  1703. lines = min(max_lines, max(5, int(np.count_nonzero(pli_matrix) / pli_matrix.shape[0])))
  1704. fig, _ = _plot_circle(
  1705. pli_matrix,
  1706. label_names,
  1707. n_lines=lines,
  1708. title=f"{band_name.upper()} band PLI",
  1709. colorbar=True,
  1710. linewidth=1.5,
  1711. show=False,
  1712. )
  1713. fig.savefig(out_file, dpi=200, facecolor="black")
  1714. plt.close(fig)
  1715. def _plot_source_power(
  1716. label_tc: np.ndarray,
  1717. label_names: Sequence[str],
  1718. out_file: Path,
  1719. ) -> None:
  1720. """Plot normalized RMS power per Desikan-Killiany region."""
  1721. label_rms = np.sqrt((label_tc**2).mean(axis=1))
  1722. order = np.argsort(label_rms)[::-1]
  1723. sorted_names = [label_names[i] for i in order]
  1724. sorted_vals = label_rms[order]
  1725. fig, ax = plt.subplots(figsize=(10, 6))
  1726. ax.bar(range(len(sorted_vals)), sorted_vals, color="#4c72b0")
  1727. ax.set_xticks(range(len(sorted_vals)))
  1728. ax.set_xticklabels(sorted_names, rotation=90, fontsize=6)
  1729. ax.set_ylabel("RMS (a.u.)")
  1730. ax.set_title("Source-level activity by Desikan-Killiany region")
  1731. fig.tight_layout()
  1732. fig.savefig(out_file, dpi=200)
  1733. plt.close(fig)
  1734. def _compute_group_pli_from_matrix(pli_matrix: np.ndarray, indices: Sequence[int]) -> float:
  1735. """Compute mean PLI for the upper-triangle (excluding diagonal) restricted to given indices.
  1736. Assumes indices are 0-based. Returns NaN if no valid elements."""
  1737. sub = pli_matrix[np.ix_(indices, indices)]
  1738. upper = np.triu(sub, k=1)
  1739. vals = upper[upper > 0]
  1740. if vals.size == 0:
  1741. return float("nan")
  1742. return float(np.mean(vals))
  1743. def process_file(file_path: Path, config: PipelineConfig, progress_cb: Optional[Callable[[str], None]] = None) -> dict:
  1744. """Process a single EEGLAB .set file: preprocess, source-localize, compute PLI per band."""
  1745. if progress_cb:
  1746. progress_cb(f"Processing {file_path.name} ...")
  1747. # Extract group and session from file path
  1748. parts = file_path.parts
  1749. group_name = parts[-3] if len(parts) >= 3 else "unknown_group"
  1750. session_name = parts[-2] if len(parts) >= 2 else "unknown_session"
  1751. # Create output directory maintaining group/session structure
  1752. out_dir = (config.output_root / group_name / session_name / file_path.stem).expanduser()
  1753. out_dir.mkdir(parents=True, exist_ok=True)
  1754. # Rest of the function remains the same
  1755. summary = {"file": str(file_path), "pli_outputs": {}}
  1756. try:
  1757. # Load raw EEGLAB file (do not preload: preprocess_raw will load)
  1758. raw = mne.io.read_raw_eeglab(str(file_path), preload=False, verbose=False)
  1759. # Preprocess (filter, notch, ref, ICA)
  1760. raw_before, raw_after, ica, _ = preprocess_raw(raw, config.preprocessing)
  1761. # Save diagnostic figures
  1762. try:
  1763. _plot_psd(raw_before, raw_after, f"{file_path.stem} - PSD", out_dir / "psd_before_after.png")
  1764. _plot_topomap(raw_before, raw_after, out_dir / "topomap_before_after.png")
  1765. _plot_gfp(raw_before, raw_after, out_dir / "gfp_before_after.png")
  1766. except Exception as exc:
  1767. LOGGER.warning("Failed to write diagnostics for %s: %s", file_path.name, exc)
  1768. # Prepare inverse operator (may download fsaverage if needed)
  1769. inv_op, labels, src, subjects_dir = _prepare_inverse_operator(raw_after, config.source)
  1770. # Run source localization and extract label time-courses
  1771. label_tc, label_names, stc = _run_source_localization(raw_after, inv_op, labels, src, config.source)
  1772. # Save a simple source-power plot
  1773. try:
  1774. _plot_source_power(label_tc, label_names, out_dir / "source_power.png")
  1775. except Exception as exc:
  1776. LOGGER.warning("Failed to plot source power for %s: %s", file_path.name, exc)
  1777. # Convert label time-courses to epochs for connectivity
  1778. epochs = _label_tc_to_epochs(label_tc, label_names, raw_after.info["sfreq"], config.connectivity.epoch_length)
  1779. # Decide GPU usage
  1780. use_gpu = _should_use_gpu(config.connectivity.use_gpu)
  1781. # Compute PLI per band and save outputs
  1782. for band_name, band in config.connectivity.frequency_bands.items():
  1783. if progress_cb:
  1784. progress_cb(f"Computing {band_name} PLI for {file_path.name} ...")
  1785. try:
  1786. pli_mat = compute_pli(epochs, band, method=config.connectivity.method, use_gpu=use_gpu)
  1787. except Exception as exc:
  1788. LOGGER.exception("PLI computation failed for %s %s: %s", file_path.name, band_name, exc)
  1789. continue
  1790. # Persist CSV + heatmap
  1791. try:
  1792. _save_connectivity_outputs(pli_mat, label_names, band_name, out_dir, group_name)
  1793. # plot connectivity circle if available
  1794. try:
  1795. _plot_connectivity_circle(pli_mat, label_names, band_name, out_dir / f"{band_name}_circle.png")
  1796. except Exception:
  1797. pass
  1798. summary["pli_outputs"][band_name] = str(out_dir / f"{band_name}_PLI.csv")
  1799. except Exception as exc:
  1800. LOGGER.warning("Failed to save connectivity outputs for %s (%s): %s", file_path.name, band_name, exc)
  1801. if progress_cb:
  1802. progress_cb(f"Finished processing {file_path.name}")
  1803. return summary
  1804. except Exception as exc:
  1805. LOGGER.exception("Error processing file %s", file_path)
  1806. if progress_cb:
  1807. progress_cb(f"Error processing {file_path.name}: {exc}")
  1808. # return partial summary so compute_group_stats can run (will skip missing CSVs)
  1809. return summary
  1810. def compute_group_stats(
  1811. summaries: List[dict],
  1812. networks: Dict[str, List[int]],
  1813. output_excel: Path,
  1814. progress_cb: Optional[Callable[[str], None]] = None,
  1815. ) -> pd.DataFrame:
  1816. """Aggregate PLI outputs from summaries into a single Excel file.
  1817. summaries: list of summary dicts returned by process_file
  1818. networks: mapping network name -> list of 1-based indices (will convert to 0-based)
  1819. """
  1820. if progress_cb:
  1821. progress_cb("Aggregating connectivity outputs into study-level table...")
  1822. rows = []
  1823. for s in summaries:
  1824. try:
  1825. file_path = Path(s["file"])
  1826. # Attempt to infer participant, group, session from path: expect base/.../Group/Session/filename.set
  1827. parts = file_path.parts
  1828. # Fallbacks
  1829. participant = file_path.stem
  1830. group_name = s.get("group", "")
  1831. session_name = s.get("session", "")
  1832. # If group/session not in summary, try to infer from path
  1833. if not group_name and len(parts) >= 3:
  1834. group_name = parts[-3]
  1835. if not session_name and len(parts) >= 2:
  1836. session_name = parts[-2]
  1837. # Extract subject ID (numeric part from filename)
  1838. import re
  1839. subj_match = re.search(r'(\d+)', participant)
  1840. subject_id = int(subj_match.group(1)) if subj_match else 0
  1841. # each band has CSV path in pli_outputs
  1842. pli_outputs = s.get("pli_outputs", {})
  1843. if not pli_outputs:
  1844. # Subject has no PLI outputs - still include with NaN values
  1845. LOGGER.warning(f"Subject {participant} has no PLI outputs, including with NaN values")
  1846. for band_name in ["theta", "alpha", "beta", "gamma"]:
  1847. for net_name in networks.keys():
  1848. rows.append({
  1849. "SubjectID": subject_id,
  1850. "Participant": participant,
  1851. "Group": group_name,
  1852. "Session": session_name,
  1853. "FrequencyBand": band_name,
  1854. "Network": net_name,
  1855. "MeanPLI": float("nan"),
  1856. "Status": "No PLI data",
  1857. })
  1858. continue
  1859. for band_name, csv_path in pli_outputs.items():
  1860. try:
  1861. df = pd.read_csv(csv_path, index_col=0)
  1862. matrix = df.values.astype(float)
  1863. status = "OK"
  1864. except Exception as e:
  1865. # If CSV fails, include with NaN
  1866. LOGGER.warning(f"Failed to read {csv_path}: {e}")
  1867. for net_name in networks.keys():
  1868. rows.append({
  1869. "SubjectID": subject_id,
  1870. "Participant": participant,
  1871. "Group": group_name,
  1872. "Session": session_name,
  1873. "FrequencyBand": band_name,
  1874. "Network": net_name,
  1875. "MeanPLI": float("nan"),
  1876. "Status": f"Read error: {e}",
  1877. })
  1878. continue
  1879. for net_name, idxs in networks.items():
  1880. # convert 1-based to 0-based if necessary: if any idx == 0 assume already 0-based
  1881. if len(idxs) == 0:
  1882. mean_pli = float("nan")
  1883. status = "Empty network"
  1884. else:
  1885. convert = [i - 1 if min(idxs) > 0 else i for i in idxs]
  1886. # filter out-of-bounds indices
  1887. valid = [i for i in convert if 0 <= i < matrix.shape[0]]
  1888. if not valid:
  1889. mean_pli = float("nan")
  1890. status = "Invalid indices"
  1891. else:
  1892. mean_pli = _compute_group_pli_from_matrix(matrix, valid)
  1893. status = "OK"
  1894. rows.append({
  1895. "SubjectID": subject_id,
  1896. "Participant": participant,
  1897. "Group": group_name,
  1898. "Session": session_name,
  1899. "FrequencyBand": band_name,
  1900. "Network": net_name,
  1901. "MeanPLI": mean_pli,
  1902. "Status": status,
  1903. })
  1904. except Exception as exc:
  1905. LOGGER.warning("Failed to aggregate summary %s: %s", s.get("file", "<unknown>"), exc)
  1906. # Create DataFrame with all columns
  1907. columns = ["SubjectID", "Participant", "Group", "Session", "FrequencyBand", "Network", "MeanPLI", "Status"]
  1908. table = pd.DataFrame(rows, columns=columns)
  1909. # Sort by SubjectID for better readability
  1910. if not table.empty:
  1911. table = table.sort_values(["Group", "SubjectID", "FrequencyBand", "Network"]).reset_index(drop=True)
  1912. try:
  1913. output_excel.parent.mkdir(parents=True, exist_ok=True)
  1914. table.to_excel(output_excel, index=False)
  1915. if progress_cb:
  1916. n_subjects = table["SubjectID"].nunique()
  1917. n_groups = table["Group"].nunique()
  1918. progress_cb(f"Saved PLI table: {n_subjects} subjects, {n_groups} groups -> {output_excel}")
  1919. except Exception as exc:
  1920. LOGGER.error("Failed to save Excel summary: %s", exc)
  1921. if progress_cb:
  1922. progress_cb(f"Failed to save study-level PLI table: {exc}")
  1923. return table
  1924. def run_study_pipeline(
  1925. base_folder: Path,
  1926. design: StudyDesign,
  1927. config: Optional[PipelineConfig] = None,
  1928. progress_cb: Optional[Callable[[str], None]] = None,
  1929. ) -> Tuple[List[dict], pd.DataFrame]:
  1930. """Run pipeline across groups/sessions as defined in design.
  1931. Expects folder structure:
  1932. base_folder / <GroupName> / <SessionName> / *.set
  1933. Returns list of file-level summaries and the aggregated DataFrame.
  1934. """
  1935. if progress_cb:
  1936. progress_cb(f"Running study pipeline in {base_folder}")
  1937. config = config or PipelineConfig()
  1938. summaries = []
  1939. for group in design.groups:
  1940. for session in design.sessions:
  1941. session_dir = Path(base_folder) / group / session
  1942. if not session_dir.exists():
  1943. LOGGER.warning("Session folder not found: %s", session_dir)
  1944. if progress_cb:
  1945. progress_cb(f"Warning: session folder not found: {session_dir}")
  1946. continue
  1947. try:
  1948. set_files = find_set_files(session_dir)
  1949. except FileNotFoundError:
  1950. if progress_cb:
  1951. progress_cb(f"No .set files found in {session_dir}; skipping.")
  1952. continue
  1953. for f in set_files:
  1954. try:
  1955. # process_file will create per-file outputs
  1956. summary = process_file(f, config, progress_cb)
  1957. summaries.append(summary)
  1958. except Exception as exc:
  1959. LOGGER.exception("Failed to process %s: %s", f, exc)
  1960. if progress_cb:
  1961. progress_cb(f"Failed to process {f.name}: {exc}")
  1962. # default networks (same as original notebook). Indices are 1-based here.
  1963. default_networks = {
  1964. "SN": [1, 2, 19, 20],
  1965. "DMN": [15, 16, 21, 22, 29, 30, 31, 32, 35, 36, 47, 48, 51, 52, 53, 54],
  1966. "CEN": [5, 6, 55, 56, 57, 58, 59, 60],
  1967. }
  1968. output_excel = config.output_root / "PLI_Table.xlsx"
  1969. table = compute_group_stats(summaries, default_networks, output_excel, progress_cb)
  1970. # Add statistical analysis after computing PLI tables
  1971. try:
  1972. _perform_statistical_analysis(summaries, design, config, progress_cb)
  1973. except Exception as exc:
  1974. LOGGER.exception("Statistical analysis failed")
  1975. if progress_cb:
  1976. progress_cb(f"Statistical analysis failed: {exc}")
  1977. return summaries, table
  1978. def _build_gui() -> Optional[tk.Tk]:
  1979. """Create a minimal Tkinter GUI with two panels: Study design (left) and Folder/run (right)."""
  1980. if tk is None:
  1981. LOGGER.error("tkinter is not available; GUI cannot be created.")
  1982. return None
  1983. root = tk.Tk()
  1984. root.title("EEG PLI Pipeline - Study Mode")
  1985. root.geometry("900x420")
  1986. root.resizable(True, True)
  1987. # Top instruction
  1988. tk.Label(root, text="Define study design on the left, then select base folder (groups) on the right.").pack(pady=(6, 4))
  1989. container = tk.Frame(root)
  1990. container.pack(fill="both", expand=True, padx=8, pady=6)
  1991. # Left panel: Study design
  1992. left = tk.LabelFrame(container, text="Study design", width=420)
  1993. left.pack(side="left", fill="both", expand=True, padx=(0, 6), pady=4)
  1994. tk.Label(left, text="Groups (comma-separated)").pack(anchor="w", padx=6, pady=(8, 0))
  1995. groups_var = tk.StringVar(value="GroupA,GroupB")
  1996. groups_entry = tk.Entry(left, textvariable=groups_var)
  1997. groups_entry.pack(fill="x", padx=6, pady=4)
  1998. tk.Label(left, text="Sessions (comma-separated)").pack(anchor="w", padx=6, pady=(8, 0))
  1999. sessions_var = tk.StringVar(value="pre,post4W")
  2000. sessions_entry = tk.Entry(left, textvariable=sessions_var)
  2001. sessions_entry.pack(fill="x", padx=6, pady=4)
  2002. tk.Label(left, text="Output root folder (optional)").pack(anchor="w", padx=6, pady=(8, 0))
  2003. out_var = tk.StringVar(value=str(Path.cwd() / "processed"))
  2004. out_entry = tk.Entry(left, textvariable=out_var)
  2005. out_entry.pack(fill="x", padx=6, pady=4)
  2006. # Handy note
  2007. tk.Label(left, text="Folder layout expected:\nbase / <Group> / <Session> / *.set", anchor="w", justify="left", fg="gray").pack(fill="x", padx=6, pady=(8, 4))
  2008. # Right panel: folder selection and run controls
  2009. right = tk.LabelFrame(container, text="Run and folder selection", width=420)
  2010. right.pack(side="left", fill="both", expand=True, padx=(6, 0), pady=4)
  2011. base_var = tk.StringVar()
  2012. tk.Label(right, text="Base folder (contains group folders)").pack(anchor="w", padx=6, pady=(8, 0))
  2013. base_entry = tk.Entry(right, textvariable=base_var)
  2014. base_entry.pack(fill="x", padx=6, pady=4)
  2015. def browse_base() -> None:
  2016. d = filedialog.askdirectory(title="Select base folder (groups)")
  2017. if d:
  2018. base_var.set(d)
  2019. tk.Button(right, text="Browse…", command=browse_base).pack(anchor="e", padx=6, pady=(0, 8))
  2020. log_text = tk.Text(right, height=12, wrap="word", state="disabled")
  2021. log_text.pack(fill="both", padx=6, pady=(4, 6), expand=True)
  2022. def log(msg: str) -> None:
  2023. log_text.configure(state="normal")
  2024. log_text.insert(tk.END, f"{msg}\n")
  2025. log_text.see(tk.END)
  2026. log_text.configure(state="disabled")
  2027. root.update_idletasks()
  2028. def run_clicked() -> None:
  2029. base_path = base_var.get().strip()
  2030. if not base_path:
  2031. messagebox.showwarning("Missing path", "Please select a base folder first.")
  2032. return
  2033. groups = [g.strip() for g in groups_var.get().split(",") if g.strip()]
  2034. sessions = [s.strip() for s in sessions_var.get().split(",") if s.strip()]
  2035. if not groups or not sessions:
  2036. messagebox.showwarning("Design error", "Provide at least one group and one session name (comma-separated).")
  2037. return
  2038. # update config output root
  2039. cfg = PipelineConfig()
  2040. try:
  2041. cfg.output_root = Path(out_var.get()).expanduser()
  2042. cfg.output_root.mkdir(parents=True, exist_ok=True)
  2043. except Exception as exc:
  2044. messagebox.showerror("Output folder error", str(exc))
  2045. return
  2046. design = StudyDesign(groups=groups, sessions=sessions)
  2047. # run study pipeline (blocking)
  2048. try:
  2049. log("Starting study pipeline...")
  2050. summaries, table = run_study_pipeline(Path(base_path), design, cfg, progress_cb=log)
  2051. log(f"Study processing complete. Processed {len(summaries)} files.")
  2052. messagebox.showinfo("Pipeline", f"Processing complete. Saved summary to {cfg.output_root / 'PLI_Table.xlsx'}")
  2053. except Exception as exc:
  2054. LOGGER.exception("Study pipeline failed")
  2055. messagebox.showerror("Pipeline error", str(exc))
  2056. log(f"Pipeline error: {exc}")
  2057. tk.Button(right, text="Run study pipeline", command=run_clicked).pack(pady=(0, 6))
  2058. return root
  2059. def launch_gui() -> None:
  2060. """Expose GUI entry point to end users."""
  2061. logging.basicConfig(level=logging.INFO, format=LOG_FORMAT)
  2062. root = _build_gui()
  2063. if root is None:
  2064. print("tkinter is not available. Run run_pipeline() directly.", file=sys.stderr)
  2065. return
  2066. root.mainloop()
  2067. def main(argv: Optional[Sequence[str]] = None) -> None:
  2068. """CLI entry point - currently launches the GUI."""
  2069. _ = argv # Placeholder for future CLI args
  2070. launch_gui()
  2071. # Add these imports at the top
  2072. from scipy import stats
  2073. import mne.stats
  2074. from itertools import combinations
  2075. import networkx as nx
  2076. def _clean_matrices(matrices: List[np.ndarray]) -> np.ndarray:
  2077. """Clean and validate matrices for statistical testing."""
  2078. if not matrices:
  2079. return np.array([])
  2080. # Stack matrices and ensure float type
  2081. stacked = np.stack(matrices).astype(float)
  2082. # Replace inf values with nan
  2083. stacked[~np.isfinite(stacked)] = np.nan
  2084. # Remove subjects with all NaN values
  2085. valid_subjects = ~np.all(np.isnan(stacked), axis=(1,2))
  2086. cleaned = stacked[valid_subjects]
  2087. # If no valid subjects remain, return empty array
  2088. if len(cleaned) == 0:
  2089. return np.array([])
  2090. # Fill remaining NaNs with mean of non-NaN values
  2091. for i in range(len(cleaned)):
  2092. nan_mask = np.isnan(cleaned[i])
  2093. if np.any(nan_mask):
  2094. valid_vals = cleaned[i][~nan_mask]
  2095. if len(valid_vals) > 0:
  2096. cleaned[i][nan_mask] = np.mean(valid_vals)
  2097. else:
  2098. cleaned[i] = 0 # If no valid values, set to 0
  2099. return cleaned
  2100. def _perform_statistical_comparison(X: np.ndarray, Y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
  2101. """Perform statistical comparison between two groups of matrices."""
  2102. if len(X) < 2 or len(Y) < 2:
  2103. return np.zeros_like(X[0]), np.ones_like(X[0])
  2104. # Perform t-test for each connection
  2105. t_stats = np.zeros_like(X[0])
  2106. p_vals = np.ones_like(X[0])
  2107. for i in range(X.shape[1]):
  2108. for j in range(X.shape[2]):
  2109. if i != j:
  2110. try:
  2111. t_stat, p_val = stats.ttest_ind(
  2112. X[:, i, j],
  2113. Y[:, i, j],
  2114. equal_var=False # Welch's t-test
  2115. )
  2116. t_stats[i, j] = t_stat
  2117. p_vals[i, j] = p_val
  2118. except:
  2119. continue
  2120. return t_stats, p_vals
  2121. def _perform_statistical_analysis(
  2122. summaries: List[dict],
  2123. design: StudyDesign,
  2124. config: PipelineConfig,
  2125. progress_cb: Optional[Callable[[str], None]] = None,
  2126. ) -> None:
  2127. """Perform statistical analyses on PLI matrices and save results."""
  2128. if progress_cb:
  2129. progress_cb("Starting statistical analyses...")
  2130. # Create stats output directories
  2131. stats_dir = config.output_root / "stats"
  2132. between_dir = stats_dir / "between_groups"
  2133. within_dir = stats_dir / "within_groups"
  2134. for d in [stats_dir, between_dir, within_dir]:
  2135. d.mkdir(parents=True, exist_ok=True)
  2136. # Group files by condition and band
  2137. grouped_files: Dict[Tuple[str, str, str], List[str]] = {}
  2138. for s in summaries:
  2139. # Prefer explicit group/session if present
  2140. group = s.get("group")
  2141. session = s.get("session")
  2142. if not group or not session:
  2143. # Fallback: infer from file path
  2144. try:
  2145. file_path = Path(s.get("file", ""))
  2146. parts = file_path.parts
  2147. if len(parts) >= 5:
  2148. group = parts[-5]
  2149. session = parts[-4]
  2150. elif len(parts) >= 3:
  2151. group = parts[-3]
  2152. session = parts[-2]
  2153. except Exception:
  2154. pass
  2155. if not group or not session:
  2156. continue
  2157. for band, csv_path in s.get("pli_outputs", {}).items():
  2158. key = (group, session, band)
  2159. grouped_files.setdefault(key, []).append(csv_path)
  2160. # Between-group analysis
  2161. for session in design.sessions:
  2162. for band in config.connectivity.frequency_bands.keys():
  2163. if progress_cb:
  2164. progress_cb(f"Computing between-group differences for {session} - {band}")
  2165. # Collect and clean matrices for each group
  2166. group_matrices = {}
  2167. for group in design.groups:
  2168. key = (group, session, band)
  2169. if key in grouped_files:
  2170. matrices = []
  2171. for f in grouped_files[key]:
  2172. try:
  2173. df = pd.read_csv(f, index_col=0)
  2174. matrices.append(df.values)
  2175. except Exception as exc:
  2176. LOGGER.warning(f"Could not load {f}: {exc}")
  2177. cleaned_matrices = _clean_matrices(matrices)
  2178. if len(cleaned_matrices) > 0:
  2179. group_matrices[group] = cleaned_matrices
  2180. # Compare groups
  2181. if len(group_matrices) >= 2:
  2182. for g1, g2 in combinations(group_matrices.keys(), 2):
  2183. X = group_matrices[g1]
  2184. Y = group_matrices[g2]
  2185. # Perform statistical comparison
  2186. t_stats, p_vals = _perform_statistical_comparison(X, Y)
  2187. # Create significance mask with FDR correction
  2188. mask = np.zeros_like(p_vals)
  2189. mask[p_vals < 0.05] = 1 # You can adjust threshold
  2190. # Save results
  2191. out_dir = between_dir / f"{session}_{band}"
  2192. out_dir.mkdir(exist_ok=True)
  2193. np.savez(
  2194. out_dir / f"{g1}_vs_{g2}_stats.npz",
  2195. t_stat=t_stats,
  2196. p_values=p_vals,
  2197. significant_mask=mask
  2198. )
  2199. # Plot significant connections
  2200. sig_connections = t_stats * mask
  2201. fig, ax = plt.subplots(figsize=(10, 8))
  2202. vmax = np.max(np.abs(sig_connections))
  2203. im = ax.imshow(sig_connections, cmap='RdBu_r', clim=(-vmax, vmax))
  2204. ax.set_title(f'{band} - {g1} vs {g2} ({session})\nSignificant connections')
  2205. plt.colorbar(im)
  2206. fig.savefig(out_dir / f"{g1}_vs_{g2}_connections.png")
  2207. plt.close(fig)
  2208. # Within-group analysis
  2209. for group in design.groups:
  2210. for band in config.connectivity.frequency_bands.keys():
  2211. if progress_cb:
  2212. progress_cb(f"Computing within-group changes for {group} - {band}")
  2213. session_matrices = {}
  2214. for session in design.sessions:
  2215. key = (group, session, band)
  2216. if key in grouped_files:
  2217. matrices = []
  2218. for f in grouped_files[key]:
  2219. try:
  2220. df = pd.read_csv(f, index_col=0)
  2221. matrices.append(df.values)
  2222. except Exception:
  2223. continue
  2224. cleaned_matrices = _clean_matrices(matrices)
  2225. if len(cleaned_matrices) > 0:
  2226. session_matrices[session] = cleaned_matrices
  2227. if len(session_matrices) >= 2:
  2228. for s1, s2 in combinations(session_matrices.keys(), 2):
  2229. X = session_matrices[s1]
  2230. Y = session_matrices[s2]
  2231. # Perform statistical comparison
  2232. t_stats, p_vals = _perform_statistical_comparison(X, Y)
  2233. # Create significance mask
  2234. mask = np.zeros_like(p_vals)
  2235. mask[p_vals < 0.05] = 1
  2236. # Save results
  2237. out_dir = within_dir / group / band
  2238. out_dir.mkdir(parents=True, exist_ok=True)
  2239. np.savez(
  2240. out_dir / f"{s1}_vs_{s2}_stats.npz",
  2241. t_stat=t_stats,
  2242. p_values=p_vals,
  2243. significant_mask=mask
  2244. )
  2245. # Plot significant connections
  2246. sig_connections = t_stats * mask
  2247. fig, ax = plt.subplots(figsize=(10, 8))
  2248. vmax = np.max(np.abs(sig_connections))
  2249. im = ax.imshow(sig_connections, cmap='RdBu_r', clim=(-vmax, vmax))
  2250. ax.set_title(f'{group} - {band}\n{s1} vs {s2} significant changes')
  2251. plt.colorbar(im)
  2252. fig.savefig(out_dir / f"{s1}_vs_{s2}_connections.png")
  2253. plt.close(fig)
  2254. if progress_cb:
  2255. if not grouped_files:
  2256. progress_cb("No CSVs found to analyze; stats folders may be empty.")
  2257. progress_cb("Statistical analyses complete")
  2258. def _upper_tri_indices(n: int) -> Tuple[np.ndarray, np.ndarray]:
  2259. return np.triu_indices(n, k=1)
  2260. def _edge_adjacency(n_labels: int) -> "sparse.csr_matrix":
  2261. """Adjacency on edges: two edges are neighbors if they share a node.
  2262. Returns a sparse (n_edges x n_edges) matrix.
  2263. """
  2264. iu = _upper_tri_indices(n_labels)
  2265. edges = list(zip(iu[0], iu[1]))
  2266. n_edges = len(edges)
  2267. incident = {i: [] for i in range(n_labels)}
  2268. for idx, (u, v) in enumerate(edges):
  2269. incident[u].append(idx)
  2270. incident[v].append(idx)
  2271. rows = []
  2272. cols = []
  2273. for node, e_list in incident.items():
  2274. # connect all edges incident at this node
  2275. for i in range(len(e_list)):
  2276. for j in range(i + 1, len(e_list)):
  2277. a = e_list[i]
  2278. b = e_list[j]
  2279. rows.extend([a, b])
  2280. cols.extend([b, a])
  2281. data = np.ones(len(rows), dtype=float)
  2282. return sparse.csr_matrix((data, (rows, cols)), shape=(n_edges, n_edges))
  2283. def _cluster_permutation_between(
  2284. X: np.ndarray,
  2285. Y: np.ndarray,
  2286. n_labels: int,
  2287. n_permutations: int = 1000,
  2288. threshold: Optional[float] = None,
  2289. ):
  2290. """Run cluster-based permutation on edge-wise differences using edge adjacency.
  2291. Returns t_obs (vector over upper-tri edges), clusters (list of boolean masks), p_vals, and iu indices.
  2292. """
  2293. iu = _upper_tri_indices(n_labels)
  2294. n_edges = len(iu[0])
  2295. Xv = X[:, iu[0], iu[1]] # (n_subj, n_edges)
  2296. Yv = Y[:, iu[0], iu[1]]
  2297. adjacency = _edge_adjacency(n_labels)
  2298. with warnings.catch_warnings():
  2299. # Suppress benign warning when no clusters are found
  2300. warnings.filterwarnings("ignore", category=RuntimeWarning, message="No clusters found*")
  2301. t_obs, clusters, p_vals, _ = mne.stats.permutation_cluster_test(
  2302. [Xv, Yv],
  2303. n_permutations=n_permutations,
  2304. tail=0,
  2305. adjacency=adjacency,
  2306. out_type="mask",
  2307. threshold=threshold,
  2308. verbose=False,
  2309. )
  2310. return t_obs, clusters, p_vals, iu
  2311. def _perform_cluster_based_permutation(
  2312. summaries: List[dict],
  2313. design: StudyDesign,
  2314. config: PipelineConfig,
  2315. progress_cb: Optional[Callable[[str], None]] = None,
  2316. ) -> None:
  2317. """Compute cluster-based permutation between groups and within groups, save results."""
  2318. if progress_cb:
  2319. progress_cb("Starting cluster-based permutation analyses…")
  2320. stats_dir = config.output_root / "stats" / "cluster_perm"
  2321. between_dir = stats_dir / "between_groups"
  2322. within_dir = stats_dir / "within_groups"
  2323. for d in [stats_dir, between_dir, within_dir]:
  2324. d.mkdir(parents=True, exist_ok=True)
  2325. # Collect files per (group, session, band)
  2326. grouped_files: Dict[Tuple[str, str, str], List[str]] = {}
  2327. for s in summaries:
  2328. group = s.get("group")
  2329. session = s.get("session")
  2330. if not group or not session:
  2331. # Fallback heuristic if not provided
  2332. file_path = Path(s.get("file", "")) if s.get("file") else None
  2333. parts = file_path.parts if file_path else []
  2334. if len(parts) >= 5:
  2335. group = parts[-5]
  2336. session = parts[-4]
  2337. elif len(parts) >= 3:
  2338. group = parts[-3]
  2339. session = parts[-2]
  2340. if not group or not session:
  2341. continue
  2342. for band, csv_path in s.get("pli_outputs", {}).items():
  2343. grouped_files.setdefault((group, session, band), []).append(csv_path)
  2344. # Helper to load/clean matrices -> array
  2345. def _load_clean(paths: List[str]) -> np.ndarray:
  2346. mats = []
  2347. for p in paths:
  2348. try:
  2349. df = pd.read_csv(p, index_col=0)
  2350. mats.append(df.values.astype(float))
  2351. except Exception:
  2352. continue
  2353. return _clean_matrices(mats)
  2354. # Between-group comparisons per session/band
  2355. any_between = False
  2356. for session in design.sessions:
  2357. for band in config.connectivity.frequency_bands.keys():
  2358. if progress_cb:
  2359. progress_cb(f"Cluster perm: between-group {session} - {band}")
  2360. group_mats: Dict[str, np.ndarray] = {}
  2361. for group in design.groups:
  2362. key = (group, session, band)
  2363. if key in grouped_files:
  2364. arr = _load_clean(grouped_files[key])
  2365. if arr.size:
  2366. group_mats[group] = arr
  2367. if len(group_mats) >= 2:
  2368. any_between = True
  2369. # take first two groups pairwise
  2370. from itertools import combinations as _cmb
  2371. for g1, g2 in _cmb(group_mats.keys(), 2):
  2372. X = group_mats[g1]
  2373. Y = group_mats[g2]
  2374. if X.shape[1:] != Y.shape[1:]:
  2375. continue
  2376. n_labels = X.shape[1]
  2377. try:
  2378. n_perm = getattr(getattr(config, "stats", object()), "n_permutations", 1000)
  2379. thr = getattr(getattr(config, "stats", object()), "cluster_threshold", None)
  2380. t_obs, clusters, p_vals, iu = _cluster_permutation_between(
  2381. X, Y, n_labels, n_permutations=int(n_perm), threshold=thr
  2382. )
  2383. except Exception as exc:
  2384. LOGGER.warning("Cluster perm failed %s vs %s: %s", g1, g2, exc)
  2385. continue
  2386. if not len(clusters):
  2387. LOGGER.info(
  2388. "Cluster perm: no clusters found (between) %s vs %s, session=%s, band=%s",
  2389. g1, g2, session, band
  2390. )
  2391. sig_mask_vec = np.zeros_like(t_obs, dtype=bool)
  2392. for cl_mask, pv in zip(clusters, p_vals):
  2393. if pv < 0.05:
  2394. sig_mask_vec = np.logical_or(sig_mask_vec, cl_mask)
  2395. # Reconstruct matrices
  2396. t_mat = np.zeros((n_labels, n_labels))
  2397. sig_mat = np.zeros((n_labels, n_labels), dtype=bool)
  2398. t_mat[iu] = t_obs
  2399. t_mat = t_mat + t_mat.T
  2400. sig_mat[iu] = sig_mask_vec
  2401. sig_mat = sig_mat | sig_mat.T
  2402. out_dir = between_dir / f"{session}_{band}"
  2403. out_dir.mkdir(parents=True, exist_ok=True)
  2404. np.savez(
  2405. out_dir / f"{g1}_vs_{g2}_cluster_perm.npz",
  2406. t_matrix=t_mat,
  2407. significant_mask=sig_mat,
  2408. p_values=p_vals,
  2409. )
  2410. # Plot
  2411. vmax = np.nanmax(np.abs(t_mat)) or 1.0
  2412. fig, ax = plt.subplots(figsize=(9, 7))
  2413. im = ax.imshow(np.where(sig_mat, t_mat, 0.0), cmap="RdBu_r", vmin=-vmax, vmax=vmax)
  2414. ax.set_title(f"{band} - {g1} vs {g2} ({session})\nCluster-based permutation (p<0.05)")
  2415. plt.colorbar(im, ax=ax)
  2416. fig.tight_layout()
  2417. fig.savefig(out_dir / f"{g1}_vs_{g2}_cluster_perm.png", dpi=200)
  2418. plt.close(fig)
  2419. # Within-group comparisons across sessions per band
  2420. any_within = False
  2421. for group in design.groups:
  2422. for band in config.connectivity.frequency_bands.keys():
  2423. if progress_cb:
  2424. progress_cb(f"Cluster perm: within-group {group} - {band}")
  2425. session_mats: Dict[str, np.ndarray] = {}
  2426. for session in design.sessions:
  2427. key = (group, session, band)
  2428. if key in grouped_files:
  2429. arr = _load_clean(grouped_files[key])
  2430. if arr.size:
  2431. session_mats[session] = arr
  2432. if len(session_mats) >= 2:
  2433. any_within = True
  2434. from itertools import combinations as _cmb
  2435. for s1, s2 in _cmb(session_mats.keys(), 2):
  2436. X = session_mats[s1]
  2437. Y = session_mats[s2]
  2438. if X.shape[1:] != Y.shape[1:]:
  2439. continue
  2440. n_labels = X.shape[1]
  2441. try:
  2442. n_perm = getattr(getattr(config, "stats", object()), "n_permutations", 1000)
  2443. thr = getattr(getattr(config, "stats", object()), "cluster_threshold", None)
  2444. t_obs, clusters, p_vals, iu = _cluster_permutation_between(
  2445. X, Y, n_labels, n_permutations=int(n_perm), threshold=thr
  2446. )
  2447. except Exception as exc:
  2448. LOGGER.warning("Cluster perm failed %s %s: %s", group, f"{s1} vs {s2}", exc)
  2449. continue
  2450. if not len(clusters):
  2451. LOGGER.info(
  2452. "Cluster perm: no clusters found (within) %s: %s vs %s, band=%s",
  2453. group, s1, s2, band
  2454. )
  2455. sig_mask_vec = np.zeros_like(t_obs, dtype=bool)
  2456. for cl_mask, pv in zip(clusters, p_vals):
  2457. if pv < 0.05:
  2458. sig_mask_vec = np.logical_or(sig_mask_vec, cl_mask)
  2459. t_mat = np.zeros((n_labels, n_labels))
  2460. sig_mat = np.zeros((n_labels, n_labels), dtype=bool)
  2461. t_mat[iu] = t_obs
  2462. t_mat = t_mat + t_mat.T
  2463. sig_mat[iu] = sig_mask_vec
  2464. sig_mat = sig_mat | sig_mat.T
  2465. out_dir = within_dir / group / band
  2466. out_dir.mkdir(parents=True, exist_ok=True)
  2467. np.savez(
  2468. out_dir / f"{s1}_vs_{s2}_cluster_perm.npz",
  2469. t_matrix=t_mat,
  2470. significant_mask=sig_mat,
  2471. p_values=p_vals,
  2472. )
  2473. vmax = np.nanmax(np.abs(t_mat)) or 1.0
  2474. fig, ax = plt.subplots(figsize=(9, 7))
  2475. im = ax.imshow(np.where(sig_mat, t_mat, 0.0), cmap="RdBu_r", vmin=-vmax, vmax=vmax)
  2476. ax.set_title(f"{group} - {band}\n{s1} vs {s2} Cluster-based permutation (p<0.05)")
  2477. plt.colorbar(im, ax=ax)
  2478. fig.tight_layout()
  2479. fig.savefig(out_dir / f"{s1}_vs_{s2}_cluster_perm.png", dpi=200)
  2480. plt.close(fig)
  2481. if progress_cb:
  2482. if not (any_between or any_within):
  2483. progress_cb("No valid condition pairs for cluster permutation; nothing saved.")
  2484. progress_cb("Cluster-based permutation analyses complete")

pli_pipeline.py at commit 23a1928, no license · at the source

Overview

  1. Centre for Chiropractic Research, New Zealand College of Chiropractic, Auckland 1060, New Zealand; (U.G.); (I.K.N.)
  2. Department of Information Technology, Faculty of Computing and Information Technology, King Abdulaziz University, Jeddah 21589, Saudi Arabia
  3. School of Applied IT, Whitecliffe, Auckland 1010, New Zealand; (S.P.); (S.E.H.)
  4. Health and Rehabilitation Research Institute, Auckland University of Technology, Auckland 1010, New Zealand
  5. Centre for Sensory-Motor Interaction, Department of Health Science and Technology, Aalborg University, 9220 Aalborg, Denmark
Journal: Sensors (Basel, Switzerland), volume 26, issue 13, article 4019
Dates: received 27 May 2026; accepted 22 June 2026; published online 24 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.3390/s26134019 · PMID 42451263 · PMCID PMC13364487 · OpenAlex W7165803648
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: EEG (modality), human (organism)
Methods: Spectral & time-frequency, Preprocessing, Connectivity, Statistics, Smoothing, state filtering, decompositions, Evoked potentials, fMRI & imaging, Physiology & signal measures
Keywords: electroencephalography, functional connectivity, phase lag index, randomised controlled trials, artefact removal, source localisation, open-source software, biomedical signal processing, clinical neurophysiology, Python
MeSH: Electroencephalography*, Randomized Controlled Trials as Topic*, Artifacts, Brain, Humans, Signal Processing, Computer-Assisted, Software (* major topic)
Topic: Functional Brain Connectivity Studies (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: King Abdulaziz University (23425654756854)
Citations: not cited yet (Europe PMC); 42 references in the paper

Abstract

Background: Electroencephalographic (EEG) functional connectivity analysis requires multiple signal-processing, source-modelling, and statistical steps that can limit its adoption in clinician-led randomised controlled trials (RCTs). NeuroStat was developed as a prototype research tool to integrate this workflow; formal usability validation with clinician end-users has not yet been conducted. Methods: NeuroStat is an open-source Python/PyQt6 desktop application that integrates automated artefact removal (a Generalised Eigenvalue Decomposition for Artefact Identification [GEDAI] pathway and a traditional Artefact Subspace Reconstruction (ASR)/Independent Component Analysis (ICA)/ICLabel pathway), boundary element model (BEM) source localisation using the Desikan–Killiany atlas (68 cortical regions), Phase Lag Index (PLI) connectivity estimation across five canonical frequency bands, and RCT-oriented statistical analysis. Evaluation separated sensor-space and source-space claims: a sensor-level simulation (repeated across five independent random seeds) tested preprocessing robustness, a repeated source-space simulation tested recovery of a known cortical parcel-pair contrast after forward projection and inverse reconstruction, a PhysioNet benchmark tested posterior Desikan–Killiany alpha PLI in 20 healthy adults, and an illustrative application to 20 sessions from a published chiropractic RCT demonstrated real-world workflow applicability. Results: In the sensor-level simulation benchmark, the Traditional pathway achieved a mean absolute error of 0.168 ± 0.017 PLI units and root mean squared error of 0.219 ± 0.045 (mean ± SD across five independent random seeds) across all artefact conditions. In the source-space simulation, reconstructed alpha PLI for the known bilateral lateral-occipital parcel pair exceeded anterior control edges across 60 repeated condition runs (mean known-control difference = 0.105 PLI units, 95% CI 0.096–0.114; t(59) = 22.61, p < 0.001). In the PhysioNet source-space benchmark, posterior Desikan–Killiany alpha PLI was higher during eyes-closed than eyes-open rest (Cohen’s d = 0.85, p = 0.001; 16/20 subjects showing the expected direction) after ICLabel-enabled preprocessing. In the pilot RCT application, all 20 sessions completed processing without manual intervention, with default-mode network alpha PLI showing a pre-to-post change of +0.071 in the intervention group versus +0.015 in the active control group. Conclusions: NeuroStat integrates preprocessing, source-space construction, connectivity estimation, and statistical reporting within a parameter-logged desktop workflow for EEG functional connectivity studies. Current evidence supports initial technical feasibility, sensor-level preprocessing robustness for one pathway in controlled simulations, source-space recovery of a known parcel-level contrast, source-space sensitivity to an expected posterior alpha resting-state contrast, and error-free processing across 20 real RCT sessions in a pilot workflow demonstration. Formal usability testing, test–retest reliability analysis, participant-specific source-model validation, and clinical-population validation remain necessary before clinician-facing or trial-deployment claims can be made.

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

Repository

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

ghani097/NeuroStat-for-RCTs

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 23a19289f3e835ac7a00cb136e23108996007de8, 3 August 2026
Languages: Python (20), Shell (1)
Size: 87 files, 21 scripts
Software Heritage: not archived
Found in: “Data Availability Statement”
Holds: README, environment (requirements.txt), documentation
Not found: license file, CITATION.cff, tests, continuous integration
Tools: NumPy (12 files), MNE-Python (10 files), pandas (10 files), Matplotlib (9 files), SciPy (8 files), MNE-Connectivity (3 files), CuPy (2 files), ICLabel (1 file), MEEGkit (1 file), NetworkX (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
22 files

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

Tracing map

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

What the map holds:

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

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

Data

Datasets cited

Data Availability Statement

The NeuroStat application source code, validation scripts, and sample outputs are available at https://github.com/ghani097/NeuroStat-for-RCTs (accessed on 21 June 2026). The PhysioNet EEGBCI dataset used for external validation is publicly available at https://physionet.org/content/eegmmidb/ (accessed on 21 June 2026). Further inquiries can be directed to the corresponding author.

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

Versions

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

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 10 keywords, 7 MeSH terms, 1 funder, 38 references.

Cite

This paper

Ghani, U., Ahmad, I., Pervez, S., Hosseini, S. E., & Niazi, I. K. (2026). NeuroStat: An Open-Source EEG Connectivity Platform for Randomised Controlled Trials. Sensors (Basel, Switzerland), 26(13), 4019. https://doi.org/10.3390/s26134019

BibTeX

@article{ghani2026neurostat,
author = {Ghani, Usman and Ahmad, Iftikhar and Pervez, Shahbaz and Hosseini, Seyed Ebrahim and Niazi, Imran Khan},
title = {{NeuroStat: An Open-Source EEG Connectivity Platform for Randomised Controlled Trials}},
journal = {Sensors (Basel, Switzerland)},
year = {2026},
month = jun,
volume = {26},
number = {13},
pages = {4019},
publisher = {Multidisciplinary Digital Publishing Institute (MDPI)},
issn = {1424-8220},
doi = {10.3390/s26134019},
url = {https://doi.org/10.3390/s26134019},
pmid = {42451263},
pmcid = {PMC13364487}
}

RIS

TY - JOUR
AU - Ghani, Usman
AU - Ahmad, Iftikhar
AU - Pervez, Shahbaz
AU - Hosseini, Seyed Ebrahim
AU - Niazi, Imran Khan
TI - NeuroStat: An Open-Source EEG Connectivity Platform for Randomised Controlled Trials
T2 - Sensors (Basel, Switzerland)
J2 - Sensors (Basel)
PY - 2026
DA - 2026/06/24
VL - 26
IS - 13
SP - 4019
SN - 1424-8220
PB - Multidisciplinary Digital Publishing Institute (MDPI)
DO - 10.3390/s26134019
UR - https://doi.org/10.3390/s26134019
LA - en
ER -

CSL-JSON

{
"id": "10.3390/s26134019",
"type": "article-journal",
"title": "NeuroStat: An Open-Source EEG Connectivity Platform for Randomised Controlled Trials",
"container-title": "Sensors (Basel, Switzerland)",
"author": [
{
"family": "Ghani",
"given": "Usman"
},
{
"family": "Ahmad",
"given": "Iftikhar"
},
{
"family": "Pervez",
"given": "Shahbaz"
},
{
"family": "Hosseini",
"given": "Seyed Ebrahim"
},
{
"family": "Niazi",
"given": "Imran Khan"
}
],
"container-title-short": "Sensors (Basel)",
"volume": "26",
"issue": "13",
"page": "4019",
"DOI": "10.3390/s26134019",
"PMID": "42451263",
"PMCID": "PMC13364487",
"ISSN": "1424-8220",
"publisher": "Multidisciplinary Digital Publishing Institute (MDPI)",
"URL": "https://doi.org/10.3390/s26134019",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
24
]
]
}
}

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

Similar papers

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

[1] doi:10.1371/journal.pcbi.1014043 [code]
EEG-Pype: An accessible MNE-Python pipeline with graphical user interface for preprocessing and analysis of resting-state electroencephalography data.
Journal: PLoS computational biology
In common: ICLabel, MNE-Python, NetworkX, 4 other tools, EEG, 8 references
[2] doi:10.1162/imag.a.1245 [code]
Towards precision EEG connectomics: Evaluating the benefits of dense sampling.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: MNE-Connectivity, ICLabel, MNE-Python, 4 other tools, EEG, 4 references
[3] doi:10.1097/j.pain.0000000000004044 [code]
No effect of rhythmic visual stimulation on experimental pain perception.
Journal: Pain
In common: MNE-Connectivity, ICLabel, MNE-Python, 4 other tools, EEG, 3 references
[4] doi:10.3389/fnhum.2026.1781338 [code]
Golden ratio organization in human EEG is associated with theta-alpha frequency convergence: a multi-dataset validation study.
Journal: Frontiers in human neuroscience
In common: MNE-Python, pandas, SciPy, 2 other tools, physionet.org/content/eegmmidb, EEG, 3 references
[5] doi:10.1038/s41597-025-05174-7 [code]
A large-scale MEG and EEG dataset for object recognition in naturalistic scenes
Journal: n/a
In common: MEEGkit, ICLabel, MNE-Python, 4 other tools, EEG, 2 references
[6] doi:10.1038/s42003-026-10394-7 [code]
Cognitive load weakens neural speech tracking without altering response timing.
Journal: Communications biology
In common: MNE-Connectivity, MNE-Python, NetworkX, 4 other tools, EEG, 3 references
[7] doi:10.3389/fncom.2026.1786996 [code]
Schumann-anchored golden ratio organization of human neural oscillations.
Journal: Frontiers in computational neuroscience
In common: MNE-Connectivity, MNE-Python, NetworkX, 4 other tools, EEG, 3 references
[8] doi:10.1038/s41598-026-56070-y [code]
SSDLabeler: realistic semi-synthetic data generation for multi-label artifact classification in EEG.
Journal: Scientific reports
In common: ICLabel, pandas, Matplotlib, 1 other tool, physionet.org/content/eegmmidb, EEG, 3 references
[9] doi:10.1371/journal.pone.0353371 [code]
Inter-brain functional connectivity: Are we measuring the right thing?
Journal: PloS one
In common: MNE-Connectivity, MNE-Python, SciPy, 2 other tools, EEG, 4 references
[10] doi:10.1093/nc/niag029 [code]
A data-driven approach to identifying and evaluating connectivity-based neural correlates of conscious visual perception.
Journal: Neuroscience of consciousness
In common: MNE-Connectivity, MNE-Python, NetworkX, 4 other tools, 3 references

Contribute

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

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

Request its removal

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

Discussion, reproductions, activity

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

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

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