OSCR

No effect of rhythmic visual stimulation on experimental pain perception.

Code ↔ Paper

8 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 8 matches
  1. [1] § 2. Methods › 2.1. Experiment 1 › 2.1.6. Data processing and feature extraction › 2.1.6.1. Electroencephalography preprocessing ↔ experiment1_eeg/code/coll_lab_eeg_pipeline.py, lines 2343–2419 · score 0.91 · MNE BIDS pipeline, peak amplitude, Independent components, mne icalabel, resampled, ZapLine
  2. [2] § 2. Methods › 2.1. Experiment 1 › 2.1.6. Data processing and feature extraction › 2.1.6.2. Frequency spectrum and signal-to-noise ratio ↔ experiment1_eeg/code/coll_lab_eeg_pipeline.py, lines 1022–1070 · score 0.75 · zero padding, power spectral density, frequency resolution, multitaper, epoch
  3. [3] § 2. Methods › 2.1. Experiment 1 › 2.1.4. Electroencephalography recordings ↔ experiment1_eeg/code/00_eegsvr_rawtobids.py, lines 35–57 · score 0.73 · actiChamp, Brain Vision, EGI, cap, filtered, EEG
  4. [4] § 2. Methods › 2.1. Experiment 1 › 2.1.7. Statistical analyses › 2.1.7.3. Electroencephalography analyses in regions of interest ↔ experiment1_eeg/code/03_eegsvr_tfr_group.py, lines 45–53 · score 0.71 · 30–80 Hz, 13–30 Hz, 8–13 Hz, frequency bands, beta, gamma
  5. [5] § 2. Methods › 2.1. Experiment 1 › 2.1.7. Statistical analyses › 2.1.7.3. Electroencephalography analyses in regions of interest ↔ experiment1_eeg/code/03s_eegsvr_tfr_group_noflip_suppmat.py, lines 47–55 · score 0.71 · 30–80 Hz, 13–30 Hz, 8–13 Hz, frequency bands, beta, gamma
  6. [6] § 2. Methods › 2.1. Experiment 1 › 2.1.6. Data processing and feature extraction › 2.1.6.2. Frequency spectrum and signal-to-noise ratio ↔ experiment1_eeg/code/04_eegsvr_snr_epo.py, lines 80–102 · score 0.68 · power spectral density, neighboring frequency, SNR, noise, padding, epoch
  7. [7] § 2. Methods › 2.1. Experiment 1 › 2.1.7. Statistical analyses › 2.1.7.1. Statistical approach ↔ experiment1_eeg/code/05_postulate_check.py, lines 477–542 · score 0.65 · Greenhouse Geisser correction, Mauchly, sphericity, residuals, outliers, ANOVA
  8. [8] § 2. Methods › 2.1. Experiment 1 › 2.1.6. Data processing and feature extraction › 2.1.6.3. Time-frequency decomposition ↔ experiment1_eeg/code/02-eegsvr_preprocess_config.py, lines 127–155 · score 0.55 · 1–80 Hz, 2–6 seconds, cropped, cycles, TFR

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 · 3,678 lines · 139 KB · no license · 2 matches

  1. # -*- coding: utf-8 -*-
  2. import mne
  3. import pyprep
  4. import importlib
  5. from mne_bids import BIDSPath, get_entities_from_fname
  6. import pandas as pd
  7. from mne_icalabel import label_components
  8. import os
  9. import sys
  10. import numpy as np
  11. import scipy
  12. from specparam import SpectralModel
  13. from mne.beamformer import make_lcmv, apply_lcmv_epochs
  14. from mne_connectivity import (
  15. spectral_connectivity_epochs,
  16. envelope_correlation,
  17. )
  18. import scipy
  19. import warnings
  20. from joblib import Parallel, delayed
  21. from mne_bids_pipeline._logging import gen_log_kwargs, logger
  22. from mne.datasets import fetch_fsaverage
  23. import bct
  24. import matplotlib.pyplot as plt
  25. from mne_bids_pipeline import _logging
  26. from nilearn.plotting import plot_connectome, plot_markers
  27. import seaborn as sns
  28. import sys
  29. import seaborn as sns
  30. from typing import List, Tuple, Optional
  31. from mne.time_frequency import AverageTFR
  32. from functools import partial
  33. def reject_bad_segs(raw, annot_to_reject="", delete_non_bads=True):
  34. """This function rejects all time spans annotated as annot_to_reject and concatenates the rest"""
  35. # this implementation seemed buggy, modified it here
  36. raw_segs = []
  37. if delete_non_bads:
  38. # Find idx of annotations not in annot_to_reject
  39. del_idx = [
  40. i for i, x in enumerate(raw.annotations.description) if x != annot_to_reject
  41. ]
  42. raw.annotations.delete(del_idx) # delete the good annotations
  43. # Add the first segment before the first bad annotation if recording does not start with a bad annotation
  44. if len(raw.annotations.onset) > 0 and raw.annotations.onset[0] != 0:
  45. tmin = 0
  46. tmax = raw.annotations.onset[0] # start at beginning of raw
  47. raw_segs.append(
  48. raw.copy().crop( # this retains raw between tmin and tmax
  49. tmin=tmin,
  50. tmax=tmax,
  51. include_tmax=False, # this is onset of first bad annot
  52. )
  53. )
  54. for jsegment in range(1, len(raw.annotations)):
  55. # print(raw.annotations.description[jsegment], annot_to_reject)
  56. if (
  57. raw.annotations.description[jsegment] == annot_to_reject
  58. ): # Append all other than 'bad_ITI'
  59. tmin = (
  60. raw.annotations.onset[jsegment - 1]
  61. + raw.annotations.duration[jsegment - 1]
  62. ) # start at ending of last bad annot
  63. tmax = raw.annotations.onset[jsegment] # end at onset of current bad annot
  64. raw_segs.append(
  65. raw.copy().crop( # this retains raw between tmin and tmax
  66. tmin=tmin,
  67. tmax=tmax,
  68. include_tmax=False, # this is onset of bad annot
  69. )
  70. )
  71. # Add the last segment after the last bad annotation if recording does not end with a bad annotation
  72. if (
  73. len(raw.annotations.onset) > 0
  74. and raw.annotations.onset[-1] + raw.annotations.duration[-1] < raw.times[-1]
  75. ):
  76. tmin = (
  77. raw.annotations.onset[-1] + raw.annotations.duration[-1]
  78. ) # start at ending of last bad annot
  79. tmax = raw.times[-1]
  80. raw_segs.append(
  81. raw.copy().crop( # this retains raw between tmin and tmax
  82. tmin=tmin,
  83. tmax=tmax,
  84. include_tmax=True, # this is end of last bad annot
  85. )
  86. )
  87. return mne.concatenate_raws(raw_segs), len(raw_segs)
  88. def run_bads_detection(
  89. bids_path,
  90. pipeline_path,
  91. task,
  92. session=None,
  93. ransac=False,
  94. repeats=3,
  95. average_reref=False,
  96. file_extension=".vhdr",
  97. montage="easycap-M1",
  98. delete_breaks=False,
  99. rename_anot_dict=None,
  100. breaks_min_length=20,
  101. t_start_after_previous=2,
  102. t_stop_before_next=2,
  103. overwrite_chans_tsv=True,
  104. consider_previous_bads=False,
  105. n_jobs=1,
  106. l_pass=100,
  107. notch=None,
  108. subjects="all",
  109. custom_bad_dict=None,
  110. ):
  111. """
  112. Run pyprep to detect bad channels and update the corresponding channels.tsv file.
  113. bids_path : str
  114. Path to the BIDS dataset.
  115. task : str
  116. The task to process.
  117. eog_chans : list
  118. List of EOG channels.
  119. misc_chans : list
  120. List of misc channels.
  121. ransac : bool
  122. Whether to use RANSAC to detect bad channels.
  123. repeats : int
  124. Number of times to repeat the bad channel detection.
  125. average_reref : bool
  126. Whether to re-reference to average.
  127. extension : str
  128. The extension of the EEG files.
  129. montage : str
  130. The montage to use.
  131. annot_breaks : bool
  132. Whether to annotate breaks.
  133. rename_anot_dict : dict or None
  134. Dictionary to rename annotations.
  135. overwrite_chans_tsv : bool
  136. Whether to overwrite the channels.tsv file. If False, a new file with the bad channels will be created.
  137. clear_previous_bads : bool
  138. Whether to clear previously marked bad channels.
  139. n_jobs : int
  140. Number of jobs to use. If 1, the process will be sequential.
  141. l_pass : float
  142. Low pass filter frequency to apply before detection. If None, no filtering will be applied.
  143. Note that pyprep applies a high pass filter at 1 Hz by default.
  144. notch: float
  145. Notch filter frequency to apply before detection. If None, no filtering will be applied.
  146. custom_bad_dict:
  147. Dictionary to identify bad channels. Keys are participant IDs and values are lists of bad channels selected. These will be added to the bad channels found by pyprep.
  148. """
  149. # Find the corresponding eeg file
  150. eeg_files = list(
  151. set(
  152. str(f)
  153. for f in BIDSPath(
  154. root=bids_path,
  155. task=task,
  156. session=session,
  157. datatype="eeg",
  158. suffix="eeg",
  159. extension=file_extension,
  160. ).match()
  161. )
  162. )
  163. eeg_files.sort()
  164. # For some datasets, same file is in sourcedata, root or derivatives because dataset is a modified BIDS
  165. # Remove sourcedata and derivatives files
  166. eeg_files = [
  167. f for f in eeg_files if "sourcedata" not in f and "derivatives" not in f
  168. ]
  169. if subjects != "all":
  170. # If subjects is not 'all', filter the eeg_files list
  171. eeg_files = [
  172. f
  173. for f in eeg_files
  174. if get_entities_from_fname(f).get("subject") in subjects
  175. ]
  176. logger.title(f"Custom step - Find bad channels in {len(eeg_files)} files.")
  177. # Initialize a data frame with file name and participants to store the results
  178. def _run_bads_detection(
  179. file,
  180. ransac=ransac,
  181. repeats=repeats,
  182. average_reref=average_reref,
  183. montage=montage,
  184. delete_breaks=delete_breaks,
  185. rename_anot_dict=rename_anot_dict,
  186. overwrite_chans_tsv=overwrite_chans_tsv,
  187. breaks_min_length=breaks_min_length,
  188. t_start_after_previous=t_start_after_previous,
  189. t_stop_before_next=t_stop_before_next,
  190. consider_previous_bads=consider_previous_bads,
  191. l_pass=l_pass,
  192. notch=notch,
  193. custom_bad_dict=custom_bad_dict,
  194. ):
  195. from mne_bids_pipeline._logging import logger
  196. # Initialize the bads_frame
  197. bads_frame = pd.DataFrame(
  198. data=None,
  199. columns=[
  200. "file_name",
  201. "participant_id",
  202. "session",
  203. "n_bads",
  204. "bad_channels",
  205. "n_breaks_found",
  206. "recording_duration",
  207. "success",
  208. "error",
  209. ],
  210. )
  211. # Add filename to bads_frame using only the file name remove rest of path
  212. bads_frame.loc[file, "file_name"] = os.path.basename(file)
  213. # Try to find bad channels using pypred. Catch errors and add to bads_frame
  214. try:
  215. with mne.utils.use_log_level(False):
  216. msg = f"Finding bad channels using pyprep."
  217. # Get subject from file
  218. logger.info(
  219. **gen_log_kwargs(
  220. message=msg,
  221. subject=get_entities_from_fname(file)["subject"],
  222. session=get_entities_from_fname(file)["session"],
  223. )
  224. )
  225. # Load corresponding chan file
  226. chan_file = pd.read_csv(
  227. file.replace("_eeg.vhdr", "_channels.tsv"), sep="\t"
  228. )
  229. # Add participant_id and session to bads_frame
  230. bads_frame.loc[file, "participant_id"] = get_entities_from_fname(file)[
  231. "subject"
  232. ]
  233. if get_entities_from_fname(file).get("session") is not None:
  234. bads_frame.loc[file, "session"] = get_entities_from_fname(file)[
  235. "session"
  236. ]
  237. previous_bads = chan_file[
  238. (chan_file["status"] == "bad")
  239. & (chan_file["type"].isin(["eeg", "EEG"]))
  240. ]["name"].tolist()
  241. # Get eog, ecg, emg and misc channels if type column is EEG, EOG, EMG, ECG
  242. if previous_bads:
  243. # Log how many bad channels were found
  244. if not consider_previous_bads:
  245. msg = f"Found {len(previous_bads)} bad channels already marked. THOSE WILL BE IGNORED AND CLEARED BECAUSE consider_previous_bads=False."
  246. else:
  247. msg = f"Found {len(previous_bads)} bad channels already marked. THOSE WILL BE CONSIDERED BECAUSE consider_previous_bads=True."
  248. logger.info(
  249. **gen_log_kwargs(
  250. message=msg,
  251. subject=get_entities_from_fname(file)["subject"],
  252. session=get_entities_from_fname(file)["session"],
  253. emoji="⚠️",
  254. )
  255. )
  256. eog_chans = chan_file.loc[
  257. chan_file["type"].isin(["EOG", "eog"]), "name"
  258. ].tolist()
  259. ecg_chans = chan_file.loc[
  260. chan_file["type"].isin(["ecg", "ECG"]), "name"
  261. ].tolist()
  262. emg_chans = chan_file.loc[
  263. chan_file["type"].isin(["EMG", "emg"]), "name"
  264. ].tolist()
  265. misc_chans = chan_file.loc[
  266. chan_file["type"].isin(["MISC", "misc"]), "name"
  267. ].tolist()
  268. # For painlaval quick fix, remove later. IF GSR in type, replace with misc
  269. if "GSR" in chan_file["name"].tolist():
  270. chan_file.loc[chan_file["type"] == "GSR", "type"] = "MISC"
  271. misc_chans = misc_chans + ["GSR"]
  272. # read input data
  273. raw = mne.io.read_raw(
  274. file,
  275. preload=True,
  276. verbose=False,
  277. eog=eog_chans,
  278. misc=misc_chans + ecg_chans + emg_chans,
  279. )
  280. assert isinstance(raw, mne.io.BaseRaw)
  281. raw.set_montage(montage)
  282. if l_pass:
  283. raw.filter(None, l_pass)
  284. if notch:
  285. raw.notch_filter(notch, picks="eeg", verbose=False)
  286. # Set previous bads
  287. if consider_previous_bads:
  288. raw.info["bads"] = list(set(raw.info["bads"] + previous_bads))
  289. # Annotate breaks
  290. if delete_breaks:
  291. annot_breaks = mne.preprocessing.annotate_break(
  292. raw=raw,
  293. min_break_duration=breaks_min_length,
  294. t_start_after_previous=t_start_after_previous,
  295. t_stop_before_next=t_stop_before_next,
  296. ignore=(
  297. "bad",
  298. "edge",
  299. "New Segment",
  300. ),
  301. )
  302. raw.set_annotations(raw.annotations + annot_breaks)
  303. original_dur = raw.times[-1]
  304. # We remove the breaks from the data because pyprep
  305. # is not yet able to handle "BAD_break" annotations
  306. # This creates discontinuities in the data but should
  307. # not affect the bad channel detection too much
  308. raw, n_blocks = reject_bad_segs(raw, annot_to_reject="BAD_break")
  309. new_dur = raw.times[-1]
  310. msg = f"Found {len(annot_breaks)} breaks and in the data and {n_blocks} valid segments."
  311. logger.info(
  312. **gen_log_kwargs(
  313. message=msg,
  314. subject=get_entities_from_fname(file)["subject"],
  315. session=get_entities_from_fname(file)["session"],
  316. emoji="⚠️",
  317. )
  318. )
  319. msg = f"Removed {original_dur - new_dur} s of breaks from the data."
  320. logger.info(
  321. **gen_log_kwargs(
  322. message=msg,
  323. subject=get_entities_from_fname(file)["subject"],
  324. session=get_entities_from_fname(file)["session"],
  325. emoji="⚠️",
  326. )
  327. )
  328. bads_frame.loc[file, "n_breaks_found"] = len(annot_breaks)
  329. bads_frame.loc[file, "recording_duration"] = original_dur - new_dur
  330. # Rename bad boundary annotations
  331. if rename_anot_dict:
  332. raw.annotations.rename(rename_anot_dict)
  333. # In some datasets referenced to FCz, channels around the ref are
  334. # considered flat by pyprep, so we migh want to re-reference to average
  335. if average_reref:
  336. raw.set_eeg_reference("average")
  337. all_bads: list[str] = []
  338. for _ in range(repeats):
  339. # Find noisy channels, already detrended
  340. nc = pyprep.NoisyChannels(raw=raw, random_state=42)
  341. nc.find_bad_by_deviation()
  342. nc.find_bad_by_correlation()
  343. if ransac:
  344. nc.find_bad_by_ransac()
  345. bads = nc.get_bads()
  346. all_bads.extend(bads)
  347. all_bads = sorted(all_bads)
  348. raw.info["bads"] = all_bads
  349. # Add custom bad channels if provided
  350. if custom_bad_dict is not None:
  351. task = get_entities_from_fname(file)["task"]
  352. sub = get_entities_from_fname(file)["subject"]
  353. # Check if task is in custom_bad_dict:
  354. if task in custom_bad_dict:
  355. if sub in custom_bad_dict[task]:
  356. all_bads.extend(custom_bad_dict[task][sub])
  357. all_bads = sorted(set(all_bads))
  358. # Log how many custom bad channels were found
  359. msg = f"Found {len(custom_bad_dict[task][sub])} custom bad channels: {custom_bad_dict[task][sub]}."
  360. logger.info(
  361. **gen_log_kwargs(
  362. message=msg,
  363. subject=get_entities_from_fname(file)["subject"],
  364. session=get_entities_from_fname(file)["session"],
  365. emoji="⚠️",
  366. )
  367. )
  368. removed_custom_bads = [
  369. ch
  370. for ch in custom_bad_dict[task][sub]
  371. if ch not in raw.info["bads"]
  372. ]
  373. else:
  374. removed_custom_bads = []
  375. else:
  376. removed_custom_bads = []
  377. else:
  378. removed_custom_bads = []
  379. # Log how many bad channels were found
  380. bad_chans = ", ".join(sorted(all_bads))
  381. # Convert to list
  382. bad_chans = bad_chans.replace(" ", "").split(",")
  383. # Set type of column description to string
  384. if "description" in chan_file.columns:
  385. chan_file["description"] = chan_file["description"].astype(str)
  386. if not consider_previous_bads:
  387. # Set all EEG channels to good
  388. chan_file.loc[chan_file["type"].isin(["EEG", "eeg"]), "status"] = (
  389. "good"
  390. )
  391. chan_file.loc[
  392. chan_file["type"].isin(["EEG", "eeg"]), "description"
  393. ] = ""
  394. # Flag bad channels
  395. for ch in bad_chans:
  396. chan_file.loc[chan_file["name"] == ch, "status"] = "bad"
  397. chan_file.loc[chan_file["name"] == ch, "description"] = (
  398. "Bad channel detected by pyprep"
  399. )
  400. # Different description if ch in custom_bad_dict
  401. if custom_bad_dict is not None and ch in custom_bad_dict.get(
  402. get_entities_from_fname(file)["subject"], []
  403. ):
  404. chan_file.loc[chan_file["name"] == ch, "description"] = (
  405. "Bad channel from custom bad channel list"
  406. )
  407. # Save the file
  408. if overwrite_chans_tsv:
  409. chan_file.to_csv(
  410. file.replace("_eeg.vhdr", "_channels.tsv"),
  411. sep="\t",
  412. index=False,
  413. )
  414. else:
  415. chan_file.to_csv(
  416. file.replace("_eeg.vhdr", "_channels.tsv").replace(
  417. ".tsv", "_bad_channels.tsv"
  418. ),
  419. sep="\t",
  420. index=False,
  421. )
  422. msg = f"Found {len(raw.info['bads'])} bad channels using pyprep: {raw.info['bads']} and {len(removed_custom_bads)} custom bad channels that were not detected by pyprep: {removed_custom_bads} for a total of {len(all_bads)} bad channels."
  423. bads_frame.loc[file, "n_bads"] = len(all_bads)
  424. bads_frame.loc[file, "bad_channels"] = bad_chans
  425. # Get subject from file
  426. logger.info(
  427. **gen_log_kwargs(
  428. message=msg,
  429. subject=get_entities_from_fname(file)["subject"],
  430. session=get_entities_from_fname(file)["session"],
  431. emoji="✅",
  432. )
  433. )
  434. except Exception as e:
  435. bads_frame.loc[file, "success"] = 0
  436. bads_frame.loc[file, "error"] = str(e)
  437. # Log the error with a danger emoji
  438. logger.error(
  439. **gen_log_kwargs(
  440. message=f"Error while finding bad channels in {file}: {e}",
  441. subject=get_entities_from_fname(file)["subject"],
  442. session=get_entities_from_fname(file)["session"],
  443. emoji="❌",
  444. )
  445. )
  446. # Add all the functions input parameters to the bads_frame
  447. bads_frame.loc[file, "ransac"] = ransac
  448. bads_frame.loc[file, "repeats"] = repeats
  449. bads_frame.loc[file, "average_reref"] = average_reref
  450. bads_frame.loc[file, "montage"] = montage
  451. bads_frame.loc[file, "delete_breaks"] = delete_breaks
  452. bads_frame.loc[file, "rename_anot_dict"] = str(rename_anot_dict)
  453. bads_frame.loc[file, "overwrite_chans_tsv"] = overwrite_chans_tsv
  454. bads_frame.loc[file, "breaks_min_length"] = breaks_min_length
  455. bads_frame.loc[file, "t_start_after_previous"] = t_start_after_previous
  456. bads_frame.loc[file, "t_stop_before_next"] = t_stop_before_next
  457. bads_frame.loc[file, "consider_previous_bads"] = consider_previous_bads
  458. bads_frame.loc[file, "l_pass"] = l_pass
  459. bads_frame.loc[file, "notch"] = notch
  460. if custom_bad_dict is not None:
  461. bads_frame.loc[file, "custom_bad_dict"] = str(custom_bad_dict)
  462. else:
  463. bads_frame.loc[file, "custom_bad_dict"] = "None"
  464. # Return the bads_frame
  465. bads_frame.loc[file, "success"] = 1
  466. bads_frame.loc[file, "error_log"] = ""
  467. return bads_frame
  468. # Use joblib to parallelize the process
  469. if n_jobs != 1:
  470. bads_frame_list = Parallel(n_jobs=n_jobs)(
  471. delayed(_run_bads_detection)(file) for file in eeg_files
  472. )
  473. else:
  474. bads_frame_list = []
  475. for file in eeg_files:
  476. bframe = _run_bads_detection(file)
  477. bads_frame_list.append(bframe)
  478. if len(bads_frame_list) > 1:
  479. # Concatenate the bad_ica_frames
  480. bads_frame = pd.concat(bads_frame_list, ignore_index=False)
  481. else:
  482. # If there is only one bad_ica_frame, just use it
  483. bads_frame = bads_frame_list[0]
  484. # Save the bads_frame to a file in the root derivatives folder
  485. if not os.path.exists(pipeline_path):
  486. os.makedirs(pipeline_path)
  487. bads_frame.to_csv(
  488. os.path.join(pipeline_path, f"pyprep_task_{task}_log.csv"), index=False
  489. )
  490. def run_ica_label(
  491. pipeline_path,
  492. task,
  493. prob_threshold=0.8,
  494. labels_to_keep=["brain", "other"],
  495. n_jobs=1,
  496. keep_mnebids_bads=False,
  497. subjects="all",
  498. ):
  499. """
  500. Run ICA label and flag bad components.
  501. pipeline_path : str
  502. Path to the pipeline.
  503. p : str
  504. The participant.
  505. task : str
  506. The task to process.
  507. prob_threshold : float
  508. The probability threshold to flag bad components.
  509. labels_to_keep : list
  510. The labels to keep. If a component is labeled as one of these labels, it will not be flagged as bad.
  511. n_jobs : int
  512. Number of jobs to use. If 1, the process will be sequential.
  513. keep_mnebids_bads : bool
  514. Whether to keep the bad components flagged by mne-bids pipeline. If False, the status will be set to good and the description will be empty if
  515. not flagged as bad by mne_icalabel.
  516. """
  517. ica_files = list(
  518. set(
  519. str(f)
  520. for f in BIDSPath(
  521. root=pipeline_path,
  522. task=task,
  523. session=None,
  524. suffix="ica",
  525. processing="icafit",
  526. extension=".fif",
  527. check=False,
  528. ).match()
  529. )
  530. )
  531. if subjects != "all":
  532. # If subjects is not 'all', filter the ica_files list
  533. ica_files = [
  534. f
  535. for f in ica_files
  536. if get_entities_from_fname(f).get("subject") in subjects
  537. ]
  538. ica_files.sort()
  539. logger.title(
  540. "Custom step - Find bad ICs using mne_icalabel in %d files." % len(ica_files)
  541. )
  542. def _run_ica_label(
  543. p,
  544. prob_threshold=prob_threshold,
  545. labels_to_keep=labels_to_keep,
  546. keep_mnebids_bads=keep_mnebids_bads,
  547. ):
  548. from mne_bids_pipeline._logging import logger
  549. # Createa bad_ica dataframe to store the results
  550. bad_ica_frame = pd.DataFrame(
  551. index=None,
  552. columns=[
  553. "file_name",
  554. "participant_id",
  555. "session",
  556. "n_bad_icas",
  557. "bad_icas",
  558. ],
  559. )
  560. # Check if session is in the path
  561. if "ses-" in p:
  562. ses_num = get_entities_from_fname(p)["session"]
  563. else:
  564. ses_num = None
  565. sub_num = get_entities_from_fname(p)["subject"]
  566. # Add file name, participant_id and session to the bad_ica_frame
  567. bad_ica_frame.loc[p, "file_name"] = os.path.basename(p)
  568. bad_ica_frame.loc[p, "participant_id"] = sub_num
  569. bad_ica_frame.loc[p, "session"] = ses_num
  570. with mne.utils.use_log_level(False):
  571. # Load ica file
  572. msg = f"Finding bad icas using mne-icalabel."
  573. # Get subject from file
  574. logger.info(**gen_log_kwargs(message=msg, subject=sub_num, session=ses_num))
  575. ica = mne.preprocessing.read_ica(p)
  576. ica_epo = mne.read_epochs(
  577. p.replace("_proc-icafit_ica.fif", "_proc-icafit_epo.fif")
  578. )
  579. # Set average reference (required for iclabel)
  580. ica_epo.set_eeg_reference("average")
  581. # Get ICA labels and probabilities
  582. icalabel = label_components(ica_epo, ica, method="iclabel")
  583. icalabel["labels"] = np.asanyarray(icalabel["labels"])
  584. # Flag bad components with probability > 0.7 and labels that are not im labels_to_keep
  585. bad_comps = np.where(
  586. (icalabel["y_pred_proba"] > prob_threshold)
  587. & (~np.isin(icalabel["labels"], labels_to_keep))
  588. )[0]
  589. # Log how many components were flagged
  590. msg = f"Found {len(bad_comps)} bad components."
  591. logger.info(
  592. **gen_log_kwargs(
  593. message=msg, subject=sub_num, session=ses_num, emoji="✅"
  594. )
  595. )
  596. # Add the number of bad components to the bad_ica_frame
  597. bad_ica_frame.loc[p, "n_bad_icas"] = len(bad_comps)
  598. bad_ica_frame.loc[p, "bad_icas"] = ", ".join([str(c) for c in bad_comps])
  599. # Load corresponding ica file
  600. try:
  601. # Check if the components.tsv file exists or create an empty one
  602. if os.path.exists(
  603. p.replace("_proc-icafit_ica.fif", "_proc-ica_components.tsv")
  604. ):
  605. ica_frame = pd.read_csv(
  606. p.replace("_proc-icafit_ica.fif", "_proc-ica_components.tsv"),
  607. sep="\t",
  608. )
  609. else:
  610. n_components = ica.n_components_
  611. ica_frame = pd.DataFrame(
  612. columns=[
  613. "component",
  614. "type",
  615. "status",
  616. "status_description",
  617. "mne_icalabel_labels",
  618. "mne_icalabel_proba",
  619. ]
  620. )
  621. ica_frame["component"] = np.arange(n_components)
  622. # Reset the status (the EOG/ECG detection in the pipeline is not playing well with some datasets)
  623. if not keep_mnebids_bads:
  624. ica_frame["status"] = "good"
  625. ica_frame["status_description"] = ""
  626. # Flag bad components with columns status as bad and description as "Bad component detected by mne_icalabel"
  627. for comp in bad_comps:
  628. ica_frame.loc[ica_frame["component"] == comp, "status"] = "bad"
  629. ica_frame.loc[
  630. ica_frame["component"] == comp, "status_description"
  631. ] = "Bad component detected by mne_icalabel"
  632. # Add the labels and probabilities to the ica_frame
  633. ica_frame["mne_icalabel_labels"] = icalabel["labels"]
  634. ica_frame["mne_icalabel_proba"] = icalabel["y_pred_proba"]
  635. # Save the file
  636. ica_frame.to_csv(
  637. p.replace("_proc-icafit_ica.fif", "_proc-ica_components.tsv"),
  638. sep="\t",
  639. index=False,
  640. )
  641. bad_ica_frame.loc[p, "success"] = 1
  642. bad_ica_frame.loc[p, "error_log"] = ""
  643. # Save the ica file with the new bad components flagged
  644. ica.exclude = bad_comps.tolist()
  645. ica.save(
  646. p.replace("_proc-icafit_ica.fif", "_proc-ica_ica.fif"),
  647. overwrite=True,
  648. )
  649. except Exception as e:
  650. # Log the error with a danger emoji
  651. logger.error(
  652. **gen_log_kwargs(
  653. message=f"Error while finding bad components in {p}: {e}, skipping this file",
  654. subject=sub_num,
  655. session=ses_num,
  656. emoji="❌",
  657. )
  658. )
  659. bad_ica_frame.loc[p, "success"] = 0
  660. bad_ica_frame.loc[p, "error_log"] = str(e)
  661. return bad_ica_frame
  662. # Use joblib to parallelize the process
  663. if n_jobs != 1:
  664. bad_ica_frames = Parallel(n_jobs=n_jobs)(
  665. delayed(_run_ica_label)(p) for p in ica_files
  666. )
  667. else:
  668. bad_ica_frames = []
  669. for p in ica_files:
  670. _run_ica_label(p)
  671. bad_ica_frames.append(_run_ica_label(p))
  672. if len(bad_ica_frames) > 1:
  673. # Concatenate the bad_ica_frames
  674. bad_ica_frame = pd.concat(bad_ica_frames, ignore_index=False)
  675. else:
  676. # If there is only one bad_ica_frame, just use it
  677. bad_ica_frame = bad_ica_frames[0]
  678. # Save the bad_ica_frame to a file in the root derivatives folder
  679. bad_ica_frame.to_csv(
  680. os.path.join(pipeline_path, f"icalabel_task_{task}_log.csv"),
  681. sep="\t",
  682. index=False,
  683. )
  684. def update_config(config, new_values, outfile=None):
  685. # Read .py file
  686. with open(config, "r") as file:
  687. lines = file.readlines()
  688. # Find the lines to replace
  689. for key, value in new_values.items():
  690. for i, line in enumerate(lines):
  691. if line.replace(" ", "").startswith(key + "=") and "no update" not in line:
  692. # Replace the line
  693. if isinstance(value, str):
  694. lines[i] = f"{key} = '{value}'\n"
  695. else:
  696. lines[i] = f"{key} = {value}\n"
  697. found = True
  698. if not found:
  699. # If the key was not found, create a new line
  700. lines.append("\n")
  701. lines.append(f"{key} = {value}\n")
  702. if outfile is not None:
  703. # Write the new file
  704. with open(outfile, "w") as file:
  705. file.writelines(lines)
  706. else:
  707. # Overwrite the original file
  708. with open(config, "w") as file:
  709. file.writelines(lines)
  710. def get_specific_config(config_file, prefix):
  711. spec = importlib.util.spec_from_file_location(
  712. name="custom_config", location=config_file
  713. )
  714. custom_cfg = importlib.util.module_from_spec(spec)
  715. spec.loader.exec_module(custom_cfg)
  716. config = {}
  717. for key in dir(custom_cfg):
  718. if prefix + "_" in key:
  719. val = getattr(custom_cfg, key)
  720. config[key.replace(prefix + "_", "")] = val
  721. return config
  722. def get_features_config(config_file):
  723. spec = importlib.util.spec_from_file_location(
  724. name="custom_config", location=config_file
  725. )
  726. custom_cfg = importlib.util.module_from_spec(spec)
  727. spec.loader.exec_module(custom_cfg)
  728. features_config = {}
  729. for key in dir(custom_cfg):
  730. if "features_" in key:
  731. val = getattr(custom_cfg, key)
  732. features_config[key] = val
  733. return features_config
  734. def get_config_keyval(config_file, key, return_false_if_not_found=True):
  735. spec = importlib.util.spec_from_file_location(
  736. name="custom_config", location=config_file
  737. )
  738. custom_cfg = importlib.util.module_from_spec(spec)
  739. spec.loader.exec_module(custom_cfg)
  740. # Find the value of the key
  741. for k in dir(custom_cfg):
  742. if key in k:
  743. val = getattr(custom_cfg, k)
  744. return val
  745. if return_false_if_not_found:
  746. # If the key is not found, return False
  747. return False
  748. else:
  749. raise ValueError(f"Key {key} not found in {config_file}")
  750. def collect_preprocessing_stats(bids_path, pipeline_path, task):
  751. mne.set_log_level("ERROR")
  752. # Initialize the dataframe
  753. preprocessing_stats = pd.DataFrame(data=None, columns=["participant_id", "session"])
  754. preprocessing_stats.set_index(["participant_id", "session"], inplace=True)
  755. clean_epo_files = list(
  756. set(
  757. str(f)
  758. for f in BIDSPath(
  759. root=pipeline_path,
  760. task=task,
  761. session=None,
  762. suffix="epo",
  763. processing="clean",
  764. extension=".fif",
  765. check=False,
  766. ).match()
  767. )
  768. )
  769. # Number of bad channels
  770. for p in clean_epo_files:
  771. sess_num = get_entities_from_fname(p)["session"]
  772. # If there is no session, all sessions are 1
  773. if not sess_num:
  774. sess_num = "1"
  775. sub_num = "sub-" + get_entities_from_fname(p)["subject"]
  776. # Get the channel file
  777. chan_filename = p.replace("_proc-clean_epo.fif", "_bads.tsv")
  778. chan_file = pd.read_csv(chan_filename, sep="\t")
  779. msg = f"Collecting preprocessing stats."
  780. logger.info(
  781. **gen_log_kwargs(
  782. message=msg, subject=sub_num.replace("sub", ""), session=sess_num
  783. )
  784. )
  785. # Nubmer of bad channels
  786. preprocessing_stats.loc[(sub_num, sess_num), "n_bad_channels"] = len(chan_file)
  787. # Number of removed ICA components
  788. ica_frame = pd.read_csv(
  789. p.replace("_proc-clean_epo.fif", "_proc-ica_components.tsv"),
  790. sep="\t",
  791. )
  792. # Nubmer of bad icas
  793. preprocessing_stats.loc[(sub_num, sess_num), "n_bad_ica"] = len(
  794. ica_frame[ica_frame["status"] == "bad"]
  795. )
  796. # Number of total/removed expochs
  797. epochs = mne.read_epochs(p)
  798. # Get the events present in the epochs
  799. events_type = list(epochs.event_id.keys())
  800. preprocessing_stats.loc[(sub_num, sess_num), "total_clean_epochs"] = len(epochs)
  801. preprocessing_stats.loc[(sub_num, sess_num), "n_removed_epochs"] = len(
  802. epochs.drop_log
  803. ) - len(epochs)
  804. # Add number of epochs flagged because of boundary events
  805. n_boundary = len(
  806. [i for i, x in enumerate(epochs.drop_log) if "BAD boundary" in x]
  807. )
  808. preprocessing_stats.loc[(sub_num, sess_num), "boundary_n_removed_epochs"] = (
  809. n_boundary
  810. )
  811. # GEt epochs removed and remaining in each event type
  812. for event in events_type:
  813. epochs_cond = epochs[event]
  814. preprocessing_stats.loc[
  815. (sub_num, sess_num), event + "_total_clean_epochs"
  816. ] = len(epochs_cond)
  817. preprocessing_stats.sort_values(by=["participant_id", "session"], inplace=True)
  818. preprocessing_stats.reset_index(inplace=True)
  819. preprocessing_stats.to_csv(
  820. os.path.join(pipeline_path, f"task_{task}_preprocessing_stats.tsv"),
  821. sep="\t",
  822. index=False,
  823. )
  824. preprocessing_stats.describe().to_csv(
  825. os.path.join(pipeline_path, f"task_{task}_preprocessing_stats_desc.tsv"),
  826. sep="\t",
  827. index=False,
  828. )
  829. def compute_features(
  830. out_path,
  831. mne_bids_root,
  832. task,
  833. freq_bands,
  834. sourcecoords_file,
  835. freq_res=1,
  836. somato_chans=None,
  837. psd_freqmax=100,
  838. psd_freqmin=1,
  839. n_jobs=1,
  840. specificsubs="all",
  841. interpolate_bads=True,
  842. compute_sourcespace_features=True,
  843. ):
  844. logger.title("Custom step - Computing rest features")
  845. # Make sure the output path exists
  846. if not os.path.exists(out_path):
  847. os.makedirs(out_path)
  848. # FOR DEBUGGING
  849. # freq_res = 0.1
  850. # mne_bids_root = pipeline_path
  851. clean_epo_files = list(
  852. set(
  853. str(f)
  854. for f in BIDSPath(
  855. root=mne_bids_root,
  856. task=task,
  857. session=None,
  858. suffix="epo",
  859. processing="clean",
  860. extension=".fif",
  861. check=False,
  862. ).match()
  863. )
  864. )
  865. clean_epo_files.sort()
  866. # Get participants
  867. if specificsubs != "all":
  868. clean_epo_files = [
  869. f
  870. for f in clean_epo_files
  871. if get_entities_from_fname(f)["subject"] in specificsubs
  872. ]
  873. def _compute_features(
  874. p,
  875. task=task,
  876. out_path=out_path,
  877. freq_bands=freq_bands,
  878. freq_res=freq_res,
  879. somato_chans=somato_chans,
  880. psd_freqmax=psd_freqmax,
  881. psd_freqmin=psd_freqmin,
  882. sourcecoords_file=sourcecoords_file,
  883. interpolate_bads=interpolate_bads,
  884. compute_sourcespace_features=compute_sourcespace_features,
  885. ):
  886. # Data frame for features with a single row per subject
  887. features_frame = pd.DataFrame(
  888. columns=["participant_id", "session"],
  889. )
  890. features_frame.set_index(["participant_id", "session"], inplace=True)
  891. sub_num = "sub-" + get_entities_from_fname(p)["subject"]
  892. session = get_entities_from_fname(p)["session"]
  893. # If there is no session, all sessions are 1 and do not create a session folder
  894. if not get_entities_from_fname(p)["session"]:
  895. session = "1"
  896. session_out = ""
  897. session_file = ""
  898. else:
  899. session_out = "ses-" + session
  900. session_file = "ses-" + session + "_"
  901. # Create subject folder
  902. os.makedirs(os.path.join(out_path, sub_num, session_out), exist_ok=True)
  903. features_frame.loc[(sub_num, session), "participant_id"] = sub_num
  904. features_frame.loc[(sub_num, session), "session"] = session
  905. features_frame.loc[(sub_num, session), "task"] = task
  906. features_frame.loc[(sub_num, session), "file_name"] = p
  907. # Load preprocessed epochs
  908. clean_epo = mne.read_epochs(p)
  909. if interpolate_bads:
  910. clean_epo = clean_epo.interpolate_bads()
  911. # Keep only EEG channels
  912. clean_epo = clean_epo.pick_types(eeg=True)
  913. # logger
  914. msg = "Computing features - PSD."
  915. # Get subject from file
  916. logger.info(
  917. **gen_log_kwargs(
  918. message=msg, subject=sub_num.replace("sub-", ""), session=session
  919. )
  920. )
  921. # Zero pad the data to required duration to get the desired frequency resolution
  922. pad_s = (
  923. 1 / freq_res
  924. ) # Padding necessary to get the desired frequency resolution
  925. epo_dur = clean_epo.get_data().shape[2] / clean_epo.info["sfreq"]
  926. n_pad = int(int((pad_s - epo_dur) * clean_epo.info["sfreq"]) / 2)
  927. clean_epo_padded = mne.EpochsArray(
  928. np.pad(
  929. clean_epo.get_data(),
  930. pad_width=((0, 0), (0, 0), (n_pad, n_pad)),
  931. mode="constant",
  932. constant_values=0,
  933. ),
  934. info=clean_epo.info,
  935. tmin=-pad_s / 2,
  936. verbose=False,
  937. )
  938. # Compute power spectral density (and average across epochs)
  939. psd = clean_epo_padded.compute_psd(
  940. method="multitaper", n_jobs=-1, fmin=psd_freqmin, fmax=psd_freqmax
  941. ).average()
  942. # Save PSD
  943. psd.save(
  944. os.path.join(
  945. out_path,
  946. sub_num,
  947. session_out,
  948. f"{sub_num}_task-{task}_{session_file}psd.h5",
  949. ),
  950. overwrite=True,
  951. )
  952. #############################################################
  953. # Peak alpha power and COG (global, each channel, somato-sensory)
  954. #############################################################
  955. msg = "Computing features - Peak alpha."
  956. # Get subject from file
  957. logger.info(
  958. **gen_log_kwargs(
  959. message=msg, subject=sub_num.replace("sub-", ""), session=session
  960. )
  961. )
  962. # Peak alpha power and COG (global, each channel, somato-sensory)
  963. freqRange = (psd.freqs >= freq_bands["alpha"][0]) & (
  964. psd.freqs <= freq_bands["alpha"][1]
  965. )
  966. # Average spectrum (across channels)
  967. avgpow = psd.get_data().mean(axis=0)
  968. peak, prop = scipy.signal.find_peaks(avgpow[freqRange])
  969. # If more than one peak, keep the highest
  970. if len(peak) > 1:
  971. peak = peak[np.argmax(avgpow[freqRange][peak])]
  972. used_max_global = 0
  973. elif len(peak) == 0:
  974. # Use this workaround if no peak is found, use the max
  975. # and flag it as a potential problem
  976. peak = np.argmax(avgpow[freqRange])
  977. used_max_global = 1
  978. else:
  979. used_max_global = 0
  980. features_frame.loc[(sub_num, session), "alpha_peak_global"] = psd.freqs[
  981. freqRange
  982. ][peak]
  983. features_frame.loc[(sub_num, session), "alpha_peak_global_usedmax"] = (
  984. used_max_global
  985. )
  986. # Center of gravity
  987. features_frame.loc[(sub_num, session), "alpha_cog_global"] = np.sum(
  988. np.multiply(avgpow[freqRange], psd.freqs[freqRange])
  989. ) / np.sum(avgpow[freqRange])
  990. # Same for somatosenory channels
  991. if somato_chans:
  992. avgpowe_samato = psd.copy().pick(somato_chans).get_data().mean(axis=0)
  993. peak, prop = scipy.signal.find_peaks(avgpowe_samato[freqRange])
  994. if len(peak) > 1:
  995. peak = peak[np.argmax(avgpowe_samato[freqRange][peak])]
  996. used_max_somato = 0
  997. elif len(peak) == 0:
  998. # Use this workaround if no peak is found, use the max
  999. # and flag it as a potential problem
  1000. peak = np.argmax(avgpow[freqRange])
  1001. used_max_somato = 1
  1002. else:
  1003. used_max_somato = 0
  1004. features_frame.loc[(sub_num, session), "alpha_peak_somato"] = psd.freqs[
  1005. freqRange
  1006. ][peak]
  1007. features_frame.loc[(sub_num, session), "alpha_peak_somato_usedmax"] = (
  1008. used_max_somato
  1009. )
  1010. # Center of gravity
  1011. features_frame.loc[(sub_num, session), "alpha_cog_somato"] = np.sum(
  1012. np.multiply(avgpowe_samato[freqRange], psd.freqs[freqRange])
  1013. ) / np.sum(avgpowe_samato[freqRange])
  1014. # Same for all channels while we are at it
  1015. peak_alpha_all_chans = pd.DataFrame(
  1016. columns=["peak", "cog"], index=clean_epo.ch_names
  1017. )
  1018. for ch in clean_epo.ch_names:
  1019. avgpow = psd.copy().pick(ch).get_data().mean(axis=0)
  1020. peak, prop = scipy.signal.find_peaks(avgpow[freqRange])
  1021. if len(peak) > 1:
  1022. peak = peak[np.argmax(avgpow[freqRange][peak])]
  1023. used_max = 0
  1024. elif len(peak) == 0:
  1025. # Use this workaround if no peak is found, use the max
  1026. # and flag it as a potential problem
  1027. peak = np.argmax(avgpow[freqRange])
  1028. used_max = 1
  1029. else:
  1030. used_max = 0
  1031. peak_alpha_all_chans.loc[ch, "peak"] = psd.freqs[freqRange][peak]
  1032. peak_alpha_all_chans.loc[ch, "peak_usedmax"] = used_max
  1033. peak_alpha_all_chans.loc[ch, "cog"] = np.sum(
  1034. np.multiply(avgpow[freqRange], psd.freqs[freqRange])
  1035. ) / np.sum(avgpow[freqRange])
  1036. peak_alpha_all_chans.to_csv(
  1037. os.path.join(
  1038. out_path,
  1039. sub_num,
  1040. session_out,
  1041. f"{sub_num}_task-{task}_{session_file}peak_alpha_all_chans.tsv",
  1042. ),
  1043. sep="\t",
  1044. index=True,
  1045. )
  1046. #############################################################
  1047. # Average power in canonical bands
  1048. #############################################################
  1049. msg = "Computing features - Power and 1/f."
  1050. # Get subject from file
  1051. logger.info(
  1052. **gen_log_kwargs(
  1053. message=msg, subject=sub_num.replace("sub-", ""), session=session
  1054. )
  1055. )
  1056. # Average power in canonical bands
  1057. for band in freq_bands:
  1058. curr_freqRange = (psd.freqs >= freq_bands[band][0]) & (
  1059. psd.freqs <= freq_bands[band][1]
  1060. )
  1061. features_frame.loc[(sub_num, session), f"{band}_power_global"] = np.sum(
  1062. psd.get_data()[:, curr_freqRange].mean(axis=1)
  1063. )
  1064. # Same in somatosensory channels
  1065. if somato_chans:
  1066. features_frame.loc[(sub_num, session), f"{band}_power_somato"] = np.sum(
  1067. psd.copy()
  1068. .pick(somato_chans)
  1069. .get_data()[:, curr_freqRange]
  1070. .mean(axis=1)
  1071. )
  1072. #############################################################
  1073. # FOOF fit and metrics
  1074. #############################################################
  1075. # Initialize model object
  1076. fm = SpectralModel(verbose=False)
  1077. # Define frequency range across which to model the spectrum
  1078. freq_range = [psd_freqmin, psd_freqmax]
  1079. # Parameterize the power spectrum, and print out a report
  1080. fm.fit(psd.freqs, psd.get_data().mean(axis=0), freq_range)
  1081. (
  1082. features_frame.loc[(sub_num, session), "foof_offset"],
  1083. features_frame.loc[(sub_num, session), "foof_exponent"],
  1084. ) = fm.aperiodic_params_
  1085. # Save foof
  1086. # Disable logger here
  1087. # Ignore warnings for this:
  1088. with warnings.catch_warnings(): # Diable tight_layout warning
  1089. warnings.simplefilter("ignore")
  1090. fm.save(
  1091. os.path.join(
  1092. out_path,
  1093. sub_num,
  1094. session_out,
  1095. f"{sub_num}_task-{task}_{session_file}foof.h5",
  1096. )
  1097. )
  1098. fm.save_report(
  1099. os.path.join(
  1100. out_path,
  1101. sub_num,
  1102. session_out,
  1103. f"{sub_num}_task-{task}_{session_file}foof_report.jpg",
  1104. )
  1105. )
  1106. #############################################################
  1107. # PSD plots
  1108. #############################################################
  1109. fig = psd.plot(show=False)
  1110. fig.savefig(
  1111. os.path.join(
  1112. out_path,
  1113. sub_num,
  1114. session_out,
  1115. f"{sub_num}_task-{task}_{session_file}psd.jpg",
  1116. )
  1117. )
  1118. dat = psd.get_data().mean(axis=0)
  1119. plt.figure()
  1120. plt.plot(psd.freqs, dat, label="Average PSD")
  1121. plt.xlabel("Frequency (Hz)")
  1122. plt.ylabel("Power")
  1123. plt.title("Average PSD")
  1124. # Marker for peak alpha
  1125. plt.scatter(
  1126. features_frame.loc[(sub_num, session), "alpha_peak_global"],
  1127. dat[
  1128. psd.freqs == features_frame.loc[(sub_num, session), "alpha_peak_global"]
  1129. ],
  1130. color="r",
  1131. linestyle="--",
  1132. label="Peak alpha",
  1133. )
  1134. plt.scatter(
  1135. features_frame.loc[(sub_num, session), "alpha_cog_global"],
  1136. dat[
  1137. np.argmin(
  1138. np.abs(
  1139. psd.freqs
  1140. - features_frame.loc[(sub_num, session), "alpha_cog_global"]
  1141. )
  1142. )
  1143. ],
  1144. color="g",
  1145. linestyle="--",
  1146. label="COG alpha",
  1147. )
  1148. # Color the auc for each band
  1149. for band, color in zip(
  1150. freq_bands, ["red", "blue", "green", "yellow", "purple"]
  1151. ):
  1152. freqRange = (psd.freqs >= freq_bands[band][0]) & (
  1153. psd.freqs <= freq_bands[band][1]
  1154. )
  1155. plt.fill_between(
  1156. psd.freqs[freqRange], dat[freqRange], color=color, alpha=0.5, label=band
  1157. )
  1158. plt.legend()
  1159. plt.savefig(
  1160. os.path.join(
  1161. out_path,
  1162. sub_num,
  1163. session_out,
  1164. f"{sub_num}_task-{task}_{session_file}psd_bands.jpg",
  1165. )
  1166. )
  1167. plt.close("all")
  1168. if compute_sourcespace_features:
  1169. #############################################################
  1170. # Source space
  1171. #############################################################
  1172. msg = "Computing features - source space"
  1173. # Get subject from file
  1174. logger.info(
  1175. **gen_log_kwargs(
  1176. message=msg, subject=sub_num.replace("sub-", ""), session=session
  1177. )
  1178. )
  1179. # Reload epochs
  1180. clean_epo = mne.read_epochs(p)
  1181. # NOTE maybe should download instead
  1182. pos = pd.read_csv(sourcecoords_file)
  1183. pos_coord = dict()
  1184. # Divide to convert mm to m
  1185. pos_coord["rr"] = np.array([pos["R"], pos["A"], pos["S"]]).T / 1000
  1186. pos_coord["nn"] = np.array([pos["R"], pos["A"], pos["S"]]).T / 1000
  1187. labels = pos["ROI Name"]
  1188. # Setup the source space
  1189. src = mne.setup_volume_source_space(
  1190. "fsaverage", pos=pos_coord, verbose=False
  1191. )
  1192. # Get standard bem model
  1193. bem = os.path.join(
  1194. fetch_fsaverage(), "bem", "fsaverage-5120-5120-5120-bem-sol.fif"
  1195. )
  1196. # Make forward model
  1197. forward = mne.make_forward_solution(
  1198. clean_epo.info, src=src, trans="fsaverage", bem=bem, eeg=True
  1199. )
  1200. # Bandpass the data in the relevant frequency band
  1201. graph_measures = dict()
  1202. for band in freq_bands:
  1203. clean_epo_band = clean_epo.copy().filter(
  1204. freq_bands[band][0], freq_bands[band][1], n_jobs=-1
  1205. )
  1206. # Compute the covariance matrix from the data
  1207. cov = mne.compute_covariance(
  1208. clean_epo_band,
  1209. method="empirical",
  1210. keep_sample_mean=False,
  1211. verbose=False,
  1212. )
  1213. filters = make_lcmv(
  1214. clean_epo_band.info,
  1215. forward,
  1216. cov,
  1217. reg=0.05,
  1218. pick_ori="max-power",
  1219. rank=None,
  1220. )
  1221. # Apply the spatial filter
  1222. stcs = apply_lcmv_epochs(clean_epo_band, filters)
  1223. # Compute the connectivity via methods of Weighted PLI (wPLI) (only lower triangle)
  1224. con_wpli = spectral_connectivity_epochs(
  1225. data=stcs,
  1226. method="wpli",
  1227. mode="multitaper",
  1228. sfreq=clean_epo_band.info["sfreq"],
  1229. fmin=freq_bands[band][0],
  1230. fmax=freq_bands[band][1],
  1231. faverage=True,
  1232. n_jobs=-1,
  1233. verbose=False,
  1234. )
  1235. # Reshape the connectivity results back into a square matrix
  1236. con_wpli_matrix = (
  1237. con_wpli.get_data().squeeze().reshape(len(labels), len(labels))
  1238. )
  1239. # Mask the upper triangle
  1240. con_wpli_matrix[np.triu_indices(len(labels), 0)] = np.nan
  1241. # Save the connectivity matrix
  1242. np.save(
  1243. os.path.join(
  1244. out_path,
  1245. sub_num,
  1246. session_out,
  1247. f"{sub_num}_task-{task}_{session_file}connectivity_wpli_{band}.npy",
  1248. ),
  1249. con_wpli_matrix,
  1250. )
  1251. # Save the labels
  1252. np.save(
  1253. os.path.join(
  1254. out_path,
  1255. sub_num,
  1256. session_out,
  1257. f"{sub_num}_task-{task}_{session_file}connectivity_labels.npy",
  1258. ),
  1259. labels,
  1260. )
  1261. # Apply hilbert and stack data in 3D array with epochs x vertices x time
  1262. vtcs = []
  1263. for s in stcs:
  1264. vtcs.append(s.copy().apply_hilbert(envelope=True).data)
  1265. vtcs = np.stack(vtcs, axis=0)
  1266. # Compute the connectivity via methods of Amplitude Envelope Correlation (AEC)
  1267. con_aec = envelope_correlation(
  1268. data=vtcs, orthogonalize="pairwise", verbose=False
  1269. )
  1270. con_aec = (
  1271. con_aec.combine()
  1272. ) # Combine connectivity data over epochs based on the method of https://mne.tools/mne-connectivity/stable/auto_examples/mne_inverse_envelope_correlation.html#ex-envelope-correlation
  1273. # Retrieve the dense AEC matrix
  1274. con_aec_matrix = con_aec.get_data(output="dense")[
  1275. :, :, 0
  1276. ] # based on the method of https://mne.tools/mne-connectivity/stable/auto_examples/mne_inverse_envelope_correlation.html#ex-envelope-correlation
  1277. con_aec_matrix /= 0.577 # normalization because of under-estimation through orthogonalization, based on Discover-EEg pipeline and Hipp et al. 2012 Nature Neuroscience
  1278. con_aec_matrix[np.triu_indices(len(labels), 0)] = np.nan
  1279. # Save to file
  1280. np.save(
  1281. os.path.join(
  1282. out_path,
  1283. sub_num,
  1284. session_out,
  1285. f"{sub_num}_task-{task}_{session_file}connectivity_aec_{band}.npy",
  1286. ),
  1287. con_aec_matrix,
  1288. )
  1289. # Plot the connectivity matrix
  1290. plt.figure(figsize=(20, 20))
  1291. plt.imshow(con_wpli_matrix, cmap="viridis")
  1292. # Add labels to the matrix
  1293. plt.xticks(range(len(labels)), labels, rotation=90, fontsize=6)
  1294. plt.yticks(range(len(labels)), labels, fontsize=6)
  1295. plt.colorbar()
  1296. plt.savefig(
  1297. os.path.join(
  1298. out_path,
  1299. sub_num,
  1300. session_out,
  1301. f"{sub_num}_task-{task}_{session_file}connectivitymatrix_wpli_{band}.jpg",
  1302. )
  1303. )
  1304. plt.close("all")
  1305. # Plot the aec connectivity matrix
  1306. plt.figure(figsize=(20, 20))
  1307. plt.imshow(con_aec_matrix, cmap="viridis")
  1308. # Add labels to the matrix
  1309. plt.xticks(range(len(labels)), labels, rotation=90, fontsize=6)
  1310. plt.yticks(range(len(labels)), labels, fontsize=6)
  1311. plt.colorbar()
  1312. plt.savefig(
  1313. os.path.join(
  1314. out_path,
  1315. sub_num,
  1316. session_out,
  1317. f"{sub_num}_task-{task}_{session_file}connectivitymatrix_aec_{band}.jpg",
  1318. )
  1319. )
  1320. plt.close("all")
  1321. # Compute graph metrics for each connectivity matrix
  1322. graph_measures[band] = dict()
  1323. for conn_measure in ["aec", "wpli"]:
  1324. graph_measures[band][conn_measure] = dict()
  1325. if conn_measure == "wpli":
  1326. conn_matrix = con_wpli_matrix.copy()
  1327. else:
  1328. conn_matrix = con_aec_matrix.copy()
  1329. # # Fill nans with 0
  1330. # conn_matrix[np.isnan(conn_matrix)] = 0
  1331. # # Copy lower triangle to upper triangle
  1332. # conn_matrix = conn_matrix + conn_matrix.T - np.diag(np.diag(conn_matrix))
  1333. # Trehsold by the top 20% of the connectivity values
  1334. sortedValues = np.sort(np.abs(conn_matrix.flatten()))
  1335. # Remove NaNs
  1336. sortedValues = sortedValues[~np.isnan(sortedValues)]
  1337. # Get the threshold
  1338. threshold = sortedValues[int(0.8 * len(sortedValues))]
  1339. # Binarize the matrix
  1340. adjacency_matrix = np.abs(conn_matrix) >= threshold
  1341. # Plot the connectome
  1342. conn_matrix_plot = conn_matrix.copy()
  1343. # Remove nans
  1344. conn_matrix_plot[np.isnan(conn_matrix_plot)] = 0
  1345. # Make it symmetric
  1346. conn_matrix_plot = (
  1347. conn_matrix_plot
  1348. + conn_matrix_plot.T
  1349. - np.diag(np.diag(conn_matrix_plot))
  1350. )
  1351. plot_connectome(
  1352. conn_matrix_plot,
  1353. node_coords=pos_coord["rr"] * 1000,
  1354. edge_threshold="99%",
  1355. node_size=10,
  1356. title=f"{band} {conn_measure} thresholded at 99%",
  1357. colorbar=True,
  1358. edge_cmap="viridis",
  1359. edge_vmin=0,
  1360. edge_vmax=1,
  1361. )
  1362. plt.savefig(
  1363. os.path.join(
  1364. out_path,
  1365. sub_num,
  1366. session_out,
  1367. f"{sub_num}_task-{task}_{session_file}connectome_{conn_measure}_{band}.jpg",
  1368. )
  1369. )
  1370. # Compute graph measures
  1371. graph_measures[band][conn_measure]["threshold"] = threshold
  1372. # Degree - Number of connexions of each node
  1373. graph_measures[band][conn_measure]["degree"] = (
  1374. bct.degree.degrees_und(adjacency_matrix)
  1375. )
  1376. # Clustering coefficient - The percentage of existing triangles surrounding
  1377. graph_measures[band][conn_measure]["cc"] = (
  1378. bct.clustering.clustering_coef_bu(adjacency_matrix)
  1379. )
  1380. # Global clustering coefficient
  1381. graph_measures[band][conn_measure]["gcc"] = np.mean(
  1382. graph_measures[band][conn_measure]["cc"]
  1383. )
  1384. # Characteristic path length
  1385. distance = bct.distance.distance_bin(adjacency_matrix)
  1386. cpl = bct.distance.charpath(distance, 0, 0)[0]
  1387. # Global efficiency - The average of the inverse shortest path between two points of the network
  1388. graph_measures[band][conn_measure]["geff"] = (
  1389. bct.efficiency.efficiency_bin(adjacency_matrix)
  1390. )
  1391. # Small-worldness
  1392. randN = bct.makerandCIJ_und(
  1393. len(adjacency_matrix),
  1394. int(np.floor(np.sum(adjacency_matrix) / 2)),
  1395. )
  1396. gcc_rand = np.mean(bct.clustering.clustering_coef_bu(randN))
  1397. cpl_rand = bct.distance.charpath(
  1398. bct.distance.distance_bin(randN), 0, 0
  1399. )[0]
  1400. graph_measures[band][conn_measure]["smallworldness"] = (
  1401. graph_measures[band][conn_measure]["gcc"] / gcc_rand
  1402. ) / (cpl / cpl_rand)
  1403. # Plot degree at each node
  1404. plot_markers(
  1405. graph_measures[band][conn_measure]["degree"],
  1406. pos_coord["rr"] * 1000,
  1407. title="Degree - " + band + " - " + conn_measure,
  1408. )
  1409. plt.savefig(
  1410. os.path.join(
  1411. out_path,
  1412. sub_num,
  1413. session_out,
  1414. f"{sub_num}_task-{task}_{session_file}degree_{band}_{conn_measure}.jpg",
  1415. )
  1416. )
  1417. # Plot clustering coefficient at each node
  1418. plot_markers(
  1419. graph_measures[band][conn_measure]["cc"],
  1420. pos_coord["rr"] * 1000,
  1421. title="Clustering coefficient - " + band + " - " + conn_measure,
  1422. )
  1423. plt.savefig(
  1424. os.path.join(
  1425. out_path,
  1426. sub_num,
  1427. session_out,
  1428. f"{sub_num}_task-{task}_{session_file}cc_{band}_{conn_measure}.jpg",
  1429. )
  1430. )
  1431. # Add all global measures to features frame
  1432. for measure in ["gcc", "geff", "smallworldness"]:
  1433. features_frame.loc[
  1434. (sub_num, session), f"{band}_{conn_measure}_{measure}"
  1435. ] = graph_measures[band][conn_measure][measure]
  1436. # Save the graph measures
  1437. np.save(
  1438. os.path.join(
  1439. out_path,
  1440. sub_num,
  1441. session_out,
  1442. f"{sub_num}_task-{task}_graph_measures.npy",
  1443. ),
  1444. graph_measures,
  1445. )
  1446. # Save the features frame
  1447. features_frame.to_csv(
  1448. os.path.join(
  1449. out_path,
  1450. sub_num,
  1451. session_out,
  1452. f"{sub_num}_task-{task}_{session_file}features_frame.tsv",
  1453. ),
  1454. sep="\t",
  1455. index=False,
  1456. )
  1457. plt.close("all")
  1458. return features_frame
  1459. # Use joblib to parallelize the process
  1460. if n_jobs != 1:
  1461. frames = Parallel(n_jobs=n_jobs)(
  1462. delayed(_compute_features)(p) for p in clean_epo_files
  1463. )
  1464. else:
  1465. frames = []
  1466. for p in clean_epo_files:
  1467. frames.append(_compute_features(p))
  1468. # Concatenate all frames into a single dataframe if more than one subject
  1469. if len(frames) == 1:
  1470. all_frames = frames[0]
  1471. else:
  1472. # Concatenate all frames
  1473. all_frames = pd.concat(frames)
  1474. all_frames.to_csv(
  1475. os.path.join(out_path, f"task-{task}_features_frame.tsv"), sep="\t", index=True
  1476. )
  1477. # plots summarizing the features
  1478. def group_plot_features(out_path, task):
  1479. """
  1480. Plot the features for each subject.
  1481. """
  1482. # Initialize mne report
  1483. report = mne.Report()
  1484. report.add_section("Peak alpha power and COG", level=1)
  1485. # Load the features
  1486. features = pd.read_csv(
  1487. os.path.join(out_path, f"task-{task}_features_frame.tsv"), sep="\t"
  1488. )
  1489. # Distribution plot of alpha peak using a raincloud plot
  1490. fig = plt.figure(figsize=(10, 5))
  1491. sns.violinplot(
  1492. features["alpha_cog_global"],
  1493. label="Global",
  1494. )
  1495. sns.swarmplot(
  1496. features["alpha_cog_global"],
  1497. label="Global",
  1498. )
  1499. plt.figure(figsize=(10, 5))
  1500. sns.histplot(
  1501. features["alpha_peak_somato"],
  1502. bins=20,
  1503. kde=True,
  1504. color="red",
  1505. label="Somatosensory",
  1506. )
  1507. plt.figure(figsize=(10, 5))
  1508. sns.histplot(
  1509. features["alpha_peak_somato"],
  1510. bins=20,
  1511. kde=True,
  1512. color="red",
  1513. label="Somatosensory",
  1514. )
  1515. plt.title("Alpha peak frequency distribution at global and somatosensory channels")
  1516. plt.xlabel("Frequency (Hz)")
  1517. plt.ylabel("Density")
  1518. plt.legend()
  1519. def custom_tfr(
  1520. pipeline_path,
  1521. task,
  1522. freqs=np.arange(1, 100, 1),
  1523. n_cycles=None,
  1524. subjects="all",
  1525. decim=1,
  1526. n_jobs=1,
  1527. return_itc=True,
  1528. interpolate_bads=True,
  1529. average=True,
  1530. subtract_evoked=False,
  1531. return_average=True,
  1532. conditions=None,
  1533. crop=None,
  1534. ):
  1535. """Custom TFR function to compute time-frequency representations.
  1536. Parameters
  1537. ----------
  1538. freqs : array-like
  1539. Frequencies to compute the TFR for.
  1540. n_cycles : array-like
  1541. Number of cycles for each frequency.
  1542. decim : int
  1543. Decimation factor for the TFR.
  1544. n_jobs : int
  1545. Number of jobs to run in parallel. This is for the TFR computation. Partiicpants are not parallelized due to the computing demand of the TFR.
  1546. method : str
  1547. Method to compute the TFR. Can be 'multitaper', 'morlet', or 'cwt_morlet'.
  1548. return_itc : bool
  1549. Whether to return the inter-trial coherence (ITC) in addition to power.
  1550. interpolate_bads : bool
  1551. Whether to interpolate bad channels before computing the TFR.
  1552. average : bool
  1553. Whether to average the TFR across epochs.
  1554. return_average : bool
  1555. Whether to return the average TFR across epochs in addition to the single epoch TFRs (only if average is False).
  1556. crop : tuple of float
  1557. Time range to crop the TFR to. If None, no cropping is done.
  1558. """
  1559. # Get the clean epochs files
  1560. clean_epo_files = list(
  1561. set(
  1562. str(f)
  1563. for f in BIDSPath(
  1564. root=pipeline_path,
  1565. task=task,
  1566. session=None,
  1567. suffix="epo",
  1568. processing="clean",
  1569. extension=".fif",
  1570. check=False,
  1571. ).match()
  1572. )
  1573. )
  1574. clean_epo_files.sort()
  1575. # Keep only the subjects specified
  1576. if subjects != "all":
  1577. clean_epo_files = [
  1578. f
  1579. for f in clean_epo_files
  1580. if get_entities_from_fname(f)["subject"] in subjects
  1581. ]
  1582. # If n_cycles is not specified, use the frequencies/3 as n_cycles
  1583. if n_cycles is None:
  1584. n_cycles = freqs / 3.0
  1585. logger.title(f"Custom step - Computing TFR in {len(clean_epo_files)} files.")
  1586. def _compute_tfr(
  1587. p,
  1588. freqs=freqs,
  1589. n_cycles=n_cycles,
  1590. decim=decim,
  1591. n_jobs=n_jobs,
  1592. return_itc=return_itc,
  1593. interpolate_bads=interpolate_bads,
  1594. average=average,
  1595. crop=crop,
  1596. return_average=return_average,
  1597. subtract_evoked=False,
  1598. conditions=conditions,
  1599. ):
  1600. # turn off mne info logger to avoid clutter
  1601. mne.set_log_level("WARNING")
  1602. # Load the cleaned epochs
  1603. epo = mne.read_epochs(p)
  1604. if interpolate_bads:
  1605. epo = epo.interpolate_bads()
  1606. if conditions is not None:
  1607. # Select only the conditions specified
  1608. epo = epo[conditions]
  1609. sub_num = get_entities_from_fname(p)["subject"]
  1610. session = get_entities_from_fname(p)["session"]
  1611. msg = "Computing TFR"
  1612. # Get subject from file
  1613. logger.info(**gen_log_kwargs(message=msg, subject=sub_num, session=session))
  1614. # Compute the TFR
  1615. if not average:
  1616. # Compute the TFR for each epoch
  1617. if subtract_evoked:
  1618. # Subtract the evoked response from each epoch
  1619. epo = epo.subtract_evoked()
  1620. power = mne.time_frequency.tfr_morlet(
  1621. epo,
  1622. freqs=freqs,
  1623. n_cycles=n_cycles,
  1624. decim=decim,
  1625. n_jobs=n_jobs,
  1626. return_itc=return_itc, # Return ITC if specified
  1627. average=False, # Average across epochs
  1628. )
  1629. # If return_itc is True, power will be a tuple of (power, itc)
  1630. if return_itc:
  1631. power, itc = power
  1632. if crop:
  1633. itc.crop(tmin=crop[0], tmax=crop[1], include_tmax=True)
  1634. # Save ITC
  1635. itc.save(
  1636. p.replace("_proc-clean_epo.fif", "_itc_epo-tfr.h5"), overwrite=True
  1637. )
  1638. # Save the power object
  1639. if crop:
  1640. power.crop(tmin=crop[0], tmax=crop[1], include_tmax=True)
  1641. power.save(
  1642. p.replace("_proc-clean_epo.fif", "_power_epo-tfr.h5"),
  1643. overwrite=True,
  1644. )
  1645. if return_average:
  1646. for cond in epo.event_id:
  1647. # Sanitize condition name
  1648. cond_save = cond.replace(" ", "").replace("-", "").replace("/", "")
  1649. # Select epochs for the condition
  1650. power_cond = power[cond]
  1651. # Collect the number of epochs
  1652. n_epochs = len(epo[cond])
  1653. msg = f"Condition {cond_save} has {n_epochs} epochs."
  1654. logger.info(
  1655. **gen_log_kwargs(
  1656. message=msg, subject=sub_num, session=session, emoji="⚠️"
  1657. )
  1658. )
  1659. # Average the TFR across epochs
  1660. power_cond = power_cond.average()
  1661. # Save the average TFR object
  1662. power_cond.save(
  1663. p.replace(
  1664. "_proc-clean_epo.fif", "_power_" + cond_save + "_avg-tfr.h5"
  1665. ),
  1666. overwrite=True,
  1667. )
  1668. if return_itc:
  1669. itc_cond = itc[cond]
  1670. # Average the ITC across epochs
  1671. itc_cond.average()
  1672. # Save the average ITC object
  1673. itc_cond.save(
  1674. p.replace(
  1675. "_proc-clean_epo.fif",
  1676. "_itc_" + cond_save + "_avg-tfr.h5",
  1677. ),
  1678. overwrite=True,
  1679. )
  1680. else:
  1681. for cond in epo.event_id:
  1682. # Sanitize condition name
  1683. cond_save = cond.replace(" ", "").replace("-", "").replace("/", "")
  1684. # Select epochs for the condition
  1685. epochs_cond = epo[cond]
  1686. # Collect the number of epochs
  1687. n_epochs = len(epochs_cond)
  1688. msg = f"Condition {cond_save} has {n_epochs} epochs."
  1689. logger.info(
  1690. **gen_log_kwargs(
  1691. message=msg, subject=sub_num, session=session, emoji="⚠️"
  1692. )
  1693. )
  1694. if subtract_evoked:
  1695. # Subtract the evoked response from the epochs
  1696. epochs_cond = epochs_cond.subtract_evoked()
  1697. # Compute the TFR for the condition
  1698. power_cond = mne.time_frequency.tfr_morlet(
  1699. epochs_cond,
  1700. freqs=freqs,
  1701. n_cycles=n_cycles,
  1702. decim=decim,
  1703. n_jobs=n_jobs,
  1704. return_itc=return_itc, # Return ITC if specified
  1705. average=True, # Average across epochs
  1706. )
  1707. if return_itc:
  1708. power_cond, itc_cond = power_cond
  1709. if crop:
  1710. itc_cond.crop(tmin=crop[0], tmax=crop[1], include_tmax=True)
  1711. # Save ITC
  1712. itc_cond.save(
  1713. p.replace(
  1714. "_proc-clean_epo.fif", f"_itc+{cond_save}_avg-tfr.h5"
  1715. ),
  1716. overwrite=True,
  1717. )
  1718. # If crop is specified, crop the power object
  1719. if crop:
  1720. power_cond.crop(tmin=crop[0], tmax=crop[1], include_tmax=True)
  1721. # Save the TFR object
  1722. power_cond.save(
  1723. p.replace("_proc-clean_epo.fif", f"_power+{cond_save}_avg-tfr.h5"),
  1724. overwrite=True,
  1725. )
  1726. msg = "Done computing TFR"
  1727. # Get subject from file
  1728. logger.info(
  1729. **gen_log_kwargs(message=msg, subject=sub_num, session=session, emoji="✅")
  1730. )
  1731. for p in clean_epo_files:
  1732. _compute_tfr(
  1733. p,
  1734. freqs=freqs,
  1735. n_cycles=n_cycles,
  1736. decim=decim,
  1737. n_jobs=n_jobs,
  1738. return_itc=return_itc,
  1739. interpolate_bads=interpolate_bads,
  1740. average=average,
  1741. crop=crop,
  1742. return_average=return_average,
  1743. subtract_evoked=subtract_evoked,
  1744. conditions=conditions,
  1745. )
  1746. def run_pipeline_task(task, config_file):
  1747. #########################################################
  1748. # Update config file with task
  1749. #########################################################
  1750. update_config(config_file, {"task": task})
  1751. #########################################################
  1752. # Print some info
  1753. #########################################################
  1754. # Log the date
  1755. from datetime import datetime
  1756. logger.info(
  1757. msg=f"👍 Running preprocessing pipeline for task: {task} on {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}"
  1758. )
  1759. #########################################################
  1760. # Bad channels using pyprep (not yet implemented in mne-bids pipeline for EEG)
  1761. #########################################################
  1762. if get_config_keyval(config_file, "use_pyprep"):
  1763. run_bads_detection(**get_specific_config(config_file, "pyprep"))
  1764. #########################################################
  1765. # First pass mne-bids pipeline to get ICA components and create the components.tsv file
  1766. #########################################################
  1767. if get_config_keyval(config_file, "use_icalabel"):
  1768. os.system(
  1769. f"mne_bids_pipeline --config={config_file} --steps=init,preprocessing/_01_data_quality,preprocessing/_04_frequency_filter,preprocessing/_05_regress_artifact,preprocessing/_06a1_fit_ica"
  1770. )
  1771. else:
  1772. # If not using icalabel, use mne_bids_pipeline to find bad icas
  1773. os.system(
  1774. f"mne_bids_pipeline --config={config_file} --steps=init,preprocessing/_01_data_quality,preprocessing/_04_frequency_filter,preprocessing/_05_regress_artifact,preprocessing/_06a1_fit_ica,preprocessing/_06a2_find_ica_artifacts.py"
  1775. )
  1776. #########################################################
  1777. # Flag bad ICA components using mne_icalabel (not yet implemented in mne-bids pipeline for EEG)
  1778. #########################################################
  1779. if get_config_keyval(config_file, "use_icalabel"):
  1780. run_ica_label(**get_specific_config(config_file, "icalabel"))
  1781. #########################################################
  1782. # Run mne-bids pipeline ICA after componets are flagged for artifact, epochs and ptp rejection
  1783. #########################################################
  1784. os.system(
  1785. f"mne_bids_pipeline --config={config_file} --steps=preprocessing/_07_make_epochs,preprocessing/_08a_apply_ica,preprocessing/_09_ptp_reject"
  1786. )
  1787. #########################################################
  1788. # Collect preprocessing metrics
  1789. #########################################################
  1790. collect_preprocessing_stats(
  1791. bids_path=get_config_keyval(config_file, "bids_root"),
  1792. pipeline_path=get_config_keyval(config_file, "deriv_root"),
  1793. task=task,
  1794. )
  1795. #########################################################
  1796. # Extract rest features from the data and make some plots subject-wise
  1797. #########################################################
  1798. if get_config_keyval(config_file, "compute_rest_features"):
  1799. compute_features(**get_specific_config(config_file, "features"))
  1800. #########################################################
  1801. # Compute TFRs
  1802. #########################################################
  1803. if get_config_keyval(config_file, "custom_tfr"):
  1804. custom_tfr(**get_specific_config(config_file, "custom_tfr"))
  1805. # Done !
  1806. logger.info(f"✅ Pipeline completed successfully for task: {task}.")
  1807. from mne_bids_pipeline._config_template import create_template_config
  1808. from pathlib import Path
  1809. import textwrap
  1810. def generate_config(path="example_config.py", mne_bids_type="minimal"):
  1811. # Custom part of the config
  1812. with open(path, "w") as f:
  1813. f.write(
  1814. "########################################################\n# This is a custom config file for the labmp EEG pipeline.\n###########################################################"
  1815. )
  1816. f.write(
  1817. "\n\n# It includes the default config from mne-bids pipeline and some custom settings.\n\n"
  1818. )
  1819. f.write("# Imports")
  1820. f.write("\nimport os\nfrom mne_bids import BIDSPath, get_entities_from_fname\n")
  1821. f.write("# Global settings\n")
  1822. f.write("# bids_root: Path to the BIDS dataset\n")
  1823. f.write("bids_root = '/path/to/bids/dataset'\n")
  1824. f.write(
  1825. "# external_root: Path to the external data folder (e.g. Schaefer atlas)\n"
  1826. )
  1827. f.write("external_root = '/path/to/external/data'\n")
  1828. f.write(
  1829. "# log_type: how to log the messages from the pipeline. Use 'console' to print everything to the console or 'file' to log all to a file in the preprocessed folder (recommended)\n"
  1830. )
  1831. f.write("log_type = 'file'\n")
  1832. f.write(
  1833. "# use_pyprep: Set to True to use pyprep for bad channels detection. If false, no automated bad channels dection\n"
  1834. )
  1835. f.write("use_pyprep = True\n")
  1836. f.write(
  1837. "# use_icalabel: Set to True to use mne-icalabel for ICA component classification. If false, the mne-bids-pipeline default classification based on eog and ecg is used\n"
  1838. )
  1839. f.write("use_icalabel = True\n")
  1840. f.write(
  1841. "# use_custom_tfr: Set to True to use the custom TFR function to compute time-frequency representations. If false, no TFRs are computed\n"
  1842. )
  1843. f.write("custom_tfr = True\n")
  1844. f.write(
  1845. "# compute_rest_features: Set to True to compute rest features after preprocessing. If false, no features are computed\n"
  1846. )
  1847. f.write("compute_rest_features = True\n")
  1848. f.write("# tasks_to_process: List of tasks to process.\n")
  1849. f.write("tasks_to_process = []\n")
  1850. f.write(
  1851. "#config_validation: mne-bids-pipeline config validation. Leave to False because we use custom options.\n"
  1852. )
  1853. f.write("config_validation = False\n")
  1854. f.write(
  1855. "# subjects: List of subjects to process or 'all' to process all subjects.\n"
  1856. )
  1857. f.write("subjects = 'all'\n")
  1858. f.write(
  1859. "# sessions: List of sessions to process or leave empty to pcrocess all sessions.\n"
  1860. )
  1861. f.write("sessions = []\n")
  1862. f.write(
  1863. "#task: Task to process. This will be updated iteratively by the pipeline for each task. No need to change.\n"
  1864. )
  1865. f.write("task = ''\n")
  1866. f.write(
  1867. "# deriv_root: Path to the derivatives folder where the preprocessed data will be saved.\n"
  1868. )
  1869. f.write(
  1870. "deriv_root = os.path.join(bids_root, 'derivatives', 'task_' + task, 'preprocessed')\n"
  1871. )
  1872. f.write(
  1873. "# select_subjects: If True, only the subjects with a file for the current task will be processed. If False, pipeline will crash if missing task\n"
  1874. )
  1875. f.write("select_subjects = True\n")
  1876. f.write("if select_subjects:\n")
  1877. f.write(
  1878. " task_subs = list(set(str(f) for f in BIDSPath(root=bids_root, task=task, session=None, datatype='eeg', suffix='eeg', extension='vhdr').match()))\n"
  1879. )
  1880. f.write(
  1881. " task_subs = [get_entities_from_fname(f).get('subject') for f in task_subs]\n"
  1882. )
  1883. f.write(" if subjects != 'all':\n")
  1884. f.write(" # If subjects is not 'all', filter the task_subs list\n")
  1885. f.write(" subjects = [sub for sub in task_subs if sub in subjects]\n")
  1886. f.write(" else:\n")
  1887. f.write(" # If subjects is 'all', use all available subjects\n")
  1888. f.write(" subjects = task_subs\n")
  1889. f.write("\n")
  1890. f.write(
  1891. "########################################################\n# Options for pyprep bad channels detection\n###########################################################"
  1892. )
  1893. f.write(
  1894. "\n# pyprep_bids_path: Path to the BIDS dataset for pyprep, do not change unless you want a different path from the bids files\n"
  1895. )
  1896. f.write("pyprep_bids_path = bids_root\n")
  1897. f.write(
  1898. "# pyprep_pipeline_path: Path to the derivatives folder where the preprocessed data will be saved for pyprep, do not change\n"
  1899. )
  1900. f.write("pyprep_pipeline_path = deriv_root\n")
  1901. f.write("# pyprep_task: Task to process for pyprep, do not change\n")
  1902. f.write("pyprep_task = task\n")
  1903. f.write(
  1904. "# pyprep_ransac: Set to True to use RANSAC for bad channels detection, False to use only the other methods\n"
  1905. )
  1906. f.write("pyprep_ransac = False\n")
  1907. f.write(
  1908. "# pyprep_repeats: Number of repeats for the bad channel detection. This can improve detection by removing very bad channels and iterating again\n"
  1909. )
  1910. f.write("pyprep_repeats = 3\n")
  1911. f.write(
  1912. "# pyprep_average_reref: Set to True to average rereference the data before bad channels detection, False to use the original data\n"
  1913. )
  1914. f.write("pyprep_average_reref = False\n")
  1915. f.write(
  1916. "# pyprep_file_extension: File extension to use for the data files, default is .vhdr for BrainVision files\n"
  1917. )
  1918. f.write("pyprep_file_extension = '.vhdr'\n")
  1919. f.write(
  1920. "# pyprep_montage: Montage to use for the data, default is easycap-M1 for BrainVision files\n"
  1921. )
  1922. f.write("pyprep_montage = 'easycap-M1'\n")
  1923. f.write(
  1924. "# pyprep_l_pass: Low pass filter frequency for the data, default is 100.0 Hz\n"
  1925. )
  1926. f.write("pyprep_l_pass = 100.0\n")
  1927. f.write(
  1928. "# pyprep_notch: Notch filter frequency for the data, default is 60.0 Hz\n"
  1929. )
  1930. f.write("pyprep_notch = 60.0\n")
  1931. f.write(
  1932. "# pyprep_consider_previous_bads: Set to True to consider previous bad channels in the data (e.g. visually identified), False to ignore and clear them (e.g. when re-running the pipeline)\n"
  1933. )
  1934. f.write("pyprep_consider_previous_bads = False\n")
  1935. f.write(
  1936. "# pyprep_rename_anot_dict: Dictionary to rename the annotations to the format expected by MNE (e.g. BAD_)\n"
  1937. )
  1938. f.write("pyprep_rename_anot_dict = None\n")
  1939. f.write(
  1940. "# pyprep_overwrite_chans_tsv: Set to True to overwrite the channels.tsv file with the bad channels detected by pyprep, False to keep the original file and create a second file. mne-bids-pipeline will only use original file so not recommended to set to False\n"
  1941. )
  1942. f.write("pyprep_overwrite_chans_tsv = True\n")
  1943. f.write("# pyprep_n_jobs: Number of jobs to use for pyprep, default is 1\n")
  1944. f.write("pyprep_n_jobs = 1\n")
  1945. f.write(
  1946. "# pyprep_subjects: List of subjects to process for pyprep, default is same as the rest of the pipeline\n"
  1947. )
  1948. f.write("pyprep_subjects = subjects\n")
  1949. f.write(
  1950. "# pyprep_delete_breaks: Set to True to delete breaks in the data (only for this operation, the data file is not modified), False to keep them\n"
  1951. )
  1952. f.write("pyprep_delete_breaks = False\n")
  1953. f.write(
  1954. "pyprep_breaks_min_length = 20 # Minimum length of breaks in seconds to consider them as breaks\n"
  1955. )
  1956. f.write(
  1957. "pyprep_t_start_after_previous = 2 # Time in seconds to start after the last event\n"
  1958. )
  1959. f.write(
  1960. "pyprep_t_stop_before_next = 2 # Time in seconds to stop before the next event\n"
  1961. )
  1962. f.write(
  1963. "# pyprep_custom_bad_dict: Dictionary to specify custom bad channels for each subject. The format should be {taskname :{subject:[bad_chan_list]}} for example: {'eegtask': {'001': ['TP8']}} If not specified, the bad channels will only be detected automatically.\n"
  1964. )
  1965. f.write("pyprep_custom_bad_dict = None\n")
  1966. f.write("\n\n\n")
  1967. f.write(
  1968. "########################################################\n# Options for Icalabel \n###########################################################"
  1969. )
  1970. f.write(
  1971. "\n# icalabel_bids_path: Path to the BIDS dataset for icalabel, do not change unless you want a different path from the bids files\n"
  1972. )
  1973. f.write("icalabel_bids_path = bids_root\n")
  1974. f.write(
  1975. "# icalabel_pipeline_path: Path to the derivatives folder where the preprocessed data will be saved for icalabel, do not change\n"
  1976. )
  1977. f.write("icalabel_pipeline_path = deriv_root\n")
  1978. f.write("# icalabel_task: Task to process for icalabel, do not change\n")
  1979. f.write("icalabel_task = task\n")
  1980. f.write(
  1981. "# icalabel_prob_threshold: Probability threshold to use for icalabel, default is 0.8\n"
  1982. )
  1983. f.write("icalabel_prob_threshold = 0.8\n")
  1984. f.write(
  1985. "# icalabel_labels_to_keep: List of labels to keep for icalabel, default is ['brain', 'other']\n"
  1986. )
  1987. f.write("icalabel_labels_to_keep = ['brain', 'other']\n")
  1988. f.write("# icalabel_n_jobs: Number of jobs to use for icalabel, default is 1\n")
  1989. f.write("icalabel_n_jobs = 1\n")
  1990. f.write(
  1991. "# icalabel_subjects: List of subjects to process for icalabel, default is same as the rest of the pipeline\n"
  1992. )
  1993. f.write("icalabel_subjects = subjects\n")
  1994. f.write(
  1995. "# icalabel_keep_mnebids_bads: Set to True to keep the bad ica already flagged in the components.tsv file (e.g. visual inspection)\n"
  1996. )
  1997. f.write("icalabel_keep_mnebids_bads = False\n")
  1998. f.write("\n\n\n")
  1999. f.write(
  2000. "########################################################\n# Rest features extraction config\n###########################################################"
  2001. )
  2002. f.write(
  2003. "\n# features_bids_path: Path to the BIDS dataset for features extraction, do not change unless you want a different path from the bids files\n"
  2004. )
  2005. f.write("features_bids_path = bids_root\n")
  2006. f.write(
  2007. "# features_out_path: Path to the derivatives folder where the preprocessed data will be saved for features extraction\n"
  2008. )
  2009. f.write("features_out_path = deriv_root.replace('preprocessed', 'features')\n")
  2010. f.write(
  2011. "# features_task: Task to process for features extraction, do not change\n"
  2012. )
  2013. f.write("features_task = task\n")
  2014. f.write(
  2015. "# features_sourcecoords_file: Path to the source coordinates file for features extraction, default is Schaefer 2018 atlas\n"
  2016. )
  2017. f.write(
  2018. "features_sourcecoords_file = os.path.join(external_root, 'Schaefer2018_100Parcels_7Networks_order_FSLMNI152_1mm.Centroid_RAS.csv')\n"
  2019. )
  2020. f.write(
  2021. "# features_freq_res: Frequency resolution for the PSD, default is 0.1 Hz\n"
  2022. )
  2023. f.write("features_freq_res = 0.1\n")
  2024. f.write(
  2025. "# features_freq_bands: Frequency bands to use for the PSD, default is theta, alpha, beta, gamma\n"
  2026. )
  2027. f.write("features_freq_bands = {\n")
  2028. f.write(" 'theta': (4, 8 - features_freq_res),\n")
  2029. f.write(" 'alpha': (8, 13 - features_freq_res),\n")
  2030. f.write(" 'beta': (13, 30),\n")
  2031. f.write(" 'gamma': (30 + features_freq_res, 80),\n")
  2032. f.write("}\n")
  2033. f.write(
  2034. "# features_psd_freqmax: Maximum frequency for the PSD, default is 100 Hz\n"
  2035. )
  2036. f.write("features_psd_freqmax = 100\n")
  2037. f.write(
  2038. "# features_psd_freqmin: Minimum frequency for the PSD, default is 1 Hz\n"
  2039. )
  2040. f.write("features_psd_freqmin = 1\n")
  2041. f.write(
  2042. "# features_somato_chans: List of somatosensory channels to use for the features extraction, default is ['C3', 'C4', 'Cz']\n"
  2043. )
  2044. f.write("features_somato_chans = ['C3', 'C4', 'Cz']\n")
  2045. f.write(
  2046. "# features_subjects: List of specific subjects to compute features for, default is same as rest of pipeline)\n"
  2047. )
  2048. f.write("features_subjects = subjects\n")
  2049. f.write(
  2050. "# features_compute_sourcespace_features: Set to True to compute source space features, False to skip this step\n"
  2051. )
  2052. f.write("features_compute_sourcespace_features = False\n")
  2053. f.write(
  2054. "# features_n_jobs: Number of jobs to use for features extraction, default is 1\n"
  2055. )
  2056. f.write("features_n_jobs = 1\n")
  2057. f.write(
  2058. "# features_subjects: List of subjects to process for features extraction, default is same as the rest of the pipeline\n"
  2059. )
  2060. f.write("\n\n\n")
  2061. f.write(
  2062. "########################################################\n# custom tfr config\n###########################################################"
  2063. )
  2064. f.write(
  2065. "\n# custom_tfr_pipeline_path: Path to the preprocessed epochs, do not change unless you want a different path from the preprocessed epochs files\n"
  2066. )
  2067. f.write("custom_tfr_pipeline_path = deriv_root\n")
  2068. f.write("# custom_tfr_task: Task to process for TFR, do not change\n")
  2069. f.write("custom_tfr_task = task\n")
  2070. f.write(
  2071. "# custom_tfr_n_jobs: Number of jobs to use for TFR computation, default is 5\n"
  2072. )
  2073. f.write("custom_tfr_n_jobs = 5\n")
  2074. f.write(
  2075. "# custom_tfr_freqs: Frequencies to compute for TFR, default is np.arange(1, 100, 1)\n"
  2076. )
  2077. f.write("custom_tfr_freqs = np.arange(1, 100, 1)\n")
  2078. f.write(
  2079. "# custom_tfr_crop: Time interval to crop the TFR, default is None (no cropping)\n"
  2080. )
  2081. f.write("custom_tfr_crop = None\n")
  2082. f.write(
  2083. "# custom_tfr_n_cycles: Number of cycles for TFR computation, default is freqs/3.0\n"
  2084. )
  2085. f.write("custom_tfr_n_cycles = custom_tfr_freqs / 3.0\n")
  2086. f.write(
  2087. "# custom_tfr_decim: Decimation factor for TFR computation, default is 1\n"
  2088. )
  2089. f.write("custom_tfr_decim = 2\n")
  2090. f.write(
  2091. "# custom_tfr_return_itc: Whether to return inter-trial coherence, default is False\n"
  2092. )
  2093. f.write("custom_tfr_return_itc = False\n")
  2094. f.write(
  2095. "# custom_tfr_interpolate_bads: Whether to interpolate bad channels before computing the TFR, default is True\n"
  2096. )
  2097. f.write("custom_tfr_interpolate_bads = True\n")
  2098. f.write(
  2099. "# custom_tfr_average: Whether to average TFR across epochs, default is False\n"
  2100. )
  2101. f.write("custom_tfr_average = False\n")
  2102. f.write(
  2103. "# custom_tfr_return_average: Whether to return the average TFR in addition to the single trials, default is True\n"
  2104. )
  2105. f.write("custom_tfr_return_average = True\n")
  2106. f.write("\n\n\n")
  2107. if mne_bids_type == "minimal":
  2108. f.write("\n\n\n")
  2109. f.write(
  2110. "########################################################\n# Rest MNE-BIDS-PIPELINE OPTIONS (MININIMAL)\n###########################################################"
  2111. )
  2112. config_text = """\
  2113. # Number of jobs
  2114. n_jobs = 10
  2115. # The task to process.
  2116. # Whether the task should be treated as resting-state data.
  2117. task_is_rest = True
  2118. # The channel types to consider.
  2119. ch_types = ["eeg"]
  2120. # Specify EOG channels to use, or create virtual EOG channels.
  2121. eog_channels = ["Fp1", "Fp2"]
  2122. # The EEG reference to use. If `average`, will use the average reference,
  2123. eeg_reference = "average"
  2124. # eeg_template_montage
  2125. eeg_template_montage = "easycap-M1"
  2126. # You can specify the seed of the random number generator (RNG).
  2127. random_state = 42
  2128. # The low-frequency cut-off in the highpass filtering step.
  2129. l_freq = 0.5
  2130. # The high-frequency cut-off in the lowpass filtering step.
  2131. h_freq = 100.0
  2132. # Specifies frequency to remove using Zapline filtering. If None, zapline will not
  2133. # be used.
  2134. zapline_fline = 60.0
  2135. # Resampling
  2136. raw_resample_sfreq = 500
  2137. # ## Epoching
  2138. # Duration of epochs in seconds.
  2139. rest_epochs_duration = 5.0 # data are segmented into 5-second epochs
  2140. # Overlap between epochs in seconds
  2141. rest_epochs_overlap = 2.5 # with a 50% overlap
  2142. epochs_tmin = 0.0
  2143. epochs_tmax = 5.0
  2144. # if `None`, no baseline correction is applied.
  2145. baseline = None
  2146. # Whether to use a spatial filter to detect and remove artifacts. The BIDS
  2147. # Pipeline offers the use of signal-space projection (SSP) and independent
  2148. # component analysis (ICA).
  2149. spatial_filter = "ica"
  2150. # Peak-to-peak amplitude limits to exclude epochs from ICA fitting. This allows you to
  2151. # remove strong transient artifacts from the epochs used for fitting ICA, which could
  2152. # negatively affect ICA performance.
  2153. ica_reject = {"eeg": 300e-6}
  2154. # The ICA algorithm to use. `"picard-extended_infomax"` operates `picard` such that the
  2155. ica_algorithm = "extended_infomax" # extended infomax for mne icalabel
  2156. # ICA high pass filter for mne icalabel
  2157. ica_l_freq = 1.0
  2158. # Run source estimation or not
  2159. run_source_estimation = False
  2160. # How to handle bad epochs after ICA.
  2161. reject = "autoreject_local"
  2162. autoreject_n_interpolate = [4, 8, 16]"""
  2163. # Split the multi-line string into a list of individual lines
  2164. clean_text = textwrap.dedent(config_text)
  2165. f.write(clean_text)
  2166. f.close()
  2167. elif mne_bids_type == "full":
  2168. f.write(
  2169. "########################################################\n# Rest MNE-BIDS-PIPELINE OPTIONS (Full)\n###########################################################"
  2170. )
  2171. # generate mne conifg file
  2172. create_template_config(
  2173. target_path=Path(path.replace(".py", "_temp.py")), overwrite=True
  2174. )
  2175. # Read the generated config file and append all lines to f
  2176. with open(Path(path.replace(".py", "_temp.py")), "r") as temp_config:
  2177. for line in temp_config:
  2178. f.write(line)
  2179. f.close()
  2180. # Delete the temporary config file
  2181. os.remove(Path(path.replace(".py", "_temp.py")))
  2182. logger.info(f"✅ Config file generated at {path}.")
  2183. try:
  2184. if sys.argv[1] == "-r":
  2185. # Prepare to run the pipeline
  2186. run = True
  2187. config_file = sys.argv[2]
  2188. mne.set_log_level("ERROR")
  2189. plt.set_loglevel("ERROR")
  2190. plt.switch_backend("agg")
  2191. elif sys.argv[1] == "-gm":
  2192. # Generate the config file
  2193. config_file = sys.argv[2] if len(sys.argv) > 2 else "example_config_minimal.py"
  2194. run = False
  2195. generate_config(config_file, mne_bids_type="minimal")
  2196. elif sys.argv[1] == "-gf":
  2197. # Generate the config file
  2198. config_file = sys.argv[2] if len(sys.argv) > 2 else "example_config_full.py"
  2199. run = False
  2200. generate_config(config_file, mne_bids_type="full")
  2201. else:
  2202. run = False
  2203. if __name__ == "__main__":
  2204. print(
  2205. "ERROR! Usage: python coll_lab_eeg_pipeline.py -r <config_file> to run the pipeline or -gm <config_file> to generate minimal a config file or -gf <config_file> to generate a full config file."
  2206. )
  2207. sys.exit(1)
  2208. except:
  2209. run = False
  2210. if run:
  2211. """
  2212. Run the preprocessing pipeline for all tasks specified in the config file.
  2213. """
  2214. #########################################################
  2215. # Load tasks from the config file
  2216. #########################################################
  2217. tasks = get_config_keyval(config_file, "tasks_to_process")
  2218. # This file path
  2219. this_file_path = os.path.dirname(os.path.abspath(__file__))
  2220. # Print global message
  2221. msg = f"Welcome! 👋 The pipeline will be run sequentially for the following tasks: {', '.join(tasks)} using the data from the following BIDS folder:{get_config_keyval(config_file, 'bids_root')}."
  2222. logger.info(msg)
  2223. for idx, task in enumerate(tasks):
  2224. update_config(config_file, {"task": task})
  2225. deriv_root = get_config_keyval(config_file, "deriv_root")
  2226. if not os.path.exists(deriv_root):
  2227. os.makedirs(deriv_root)
  2228. log_path = os.path.join(deriv_root, task + "_pipeline.log")
  2229. py_command = (
  2230. f"import sys; "
  2231. f"sys.path.append(r'{this_file_path}'); "
  2232. f"import coll_lab_eeg_pipeline; "
  2233. f"coll_lab_eeg_pipeline.run_pipeline_task(r'{task}', r'{config_file}')"
  2234. )
  2235. # 2. Construct the final shell command, wrapping the Python command in double quotes
  2236. if get_config_keyval(config_file, "log_type") == "file":
  2237. cmd = f'python -c "{py_command}" > "{log_path}" 2>&1'
  2238. msg = f"👍 Running pipeline for task {idx+1} out of {len(tasks)}: {task} and logging to {log_path}"
  2239. logger.info(msg)
  2240. else:
  2241. cmd = f'python -c "{py_command}"'
  2242. print(
  2243. f"👍 Running pipeline for task {idx+1} out of {len(tasks)}: {task} and logging to console"
  2244. )
  2245. exit_code = os.system(cmd)
  2246. if exit_code != 0:
  2247. logger.error(
  2248. f"❌ Pipeline failed for task {task}. Check the log file for details."
  2249. )
  2250. sys.exit(1)
  2251. else:
  2252. logger.info(f"✅ Pipeline completed successfully for task {task}.")
  2253. #########################################
  2254. # Plotting functions
  2255. #########################################
  2256. def plot_tfr_nice_spectro(
  2257. tfr: AverageTFR,
  2258. chans: List[str],
  2259. cmap: str = "RdBu_r",
  2260. time_range: Tuple[Optional[float], Optional[float]] = (None, None),
  2261. freq_range: Optional[Tuple[float, float]] = None,
  2262. vlim: Tuple[float, float] = (None, None),
  2263. figsize: Tuple[float, float] = (2, 1),
  2264. nameout: str = "",
  2265. title: str = "",
  2266. cbar: bool = True,
  2267. fontsize: int = 8,
  2268. save_path: Optional[str] = None,
  2269. cbar_label: str = "Z-score",
  2270. cbar_ticks: Optional[List[float]] = None,
  2271. remove_cbar_ticks: bool = False,
  2272. combine: str = "mean",
  2273. y_ticks: Optional[List[float]] = None,
  2274. x_ticks: Optional[List[float]] = None,
  2275. hide_xlabel: bool = False,
  2276. hide_ylabel: bool = False,
  2277. remove_xticks: bool = False,
  2278. remove_yticks: bool = False,
  2279. remove_xticklabels: bool = False,
  2280. remove_yticklabels: bool = False,
  2281. extension: str = "svg",
  2282. dpi: int = 1200,
  2283. bbox_inches: str = "tight",
  2284. transparent: bool = True,
  2285. show=False,
  2286. mask: Optional[np.ndarray] = None,
  2287. custom_yticks=None,
  2288. custom_xticks=None,
  2289. mask_style: Optional[str] = None,
  2290. ):
  2291. """Plot a single channel time-frequency representation (TFR) as a spectrogram.
  2292. This function is a wrapper around the tfr.plot() method from MNE-Python,
  2293. providing extensive customization options for creating publication-quality figures.
  2294. Parameters
  2295. ----------
  2296. tfr : mne.time_frequency.AverageTFR
  2297. The time-frequency representation to plot.
  2298. chans : list of str
  2299. The list of channel names to plot. If multiple are provided, they are
  2300. averaged using the 'combine="mean"' parameter.
  2301. cmap : str, optional
  2302. The colormap to use for the plot, by default "RdBu_r".
  2303. time_range : tuple of float or None, optional
  2304. The time range (tmin, tmax) to plot in seconds. If an element is None,
  2305. the range is inferred from the TFR object, by default (None, None).
  2306. freq_range : tuple of float or None, optional
  2307. The frequency range (fmin, fmax) to plot in Hz. If None, the range is
  2308. inferred from the TFR object, by default None.
  2309. vlim : tuple of float, optional
  2310. The limits for the color scale (vmin, vmax), by default (-3, 3).
  2311. figsize : tuple of float, optional
  2312. The size of the figure (width, height) in inches, by default (2, 1).
  2313. nameout : str, optional
  2314. The base name of the output file (without extension), by default "".
  2315. title : str, optional
  2316. The title of the plot, by default "".
  2317. cbar : bool, optional
  2318. Whether to display the colorbar, by default True.
  2319. fontsize : int, optional
  2320. The base font size for title and labels, by default 8.
  2321. save_path : str, optional
  2322. The directory path to save the figure. If None, the figure is not
  2323. saved, by default None.
  2324. cbar_label : str, optional
  2325. The label for the colorbar, by default "Z-score".
  2326. y_ticks : list of float, optional
  2327. Custom y-axis tick locations, by default None.
  2328. x_ticks : list of float, optional
  2329. Custom x-axis tick locations, by default None.
  2330. hide_xlabel : bool, optional
  2331. If True, hides the x-axis label, by default False.
  2332. hide_ylabel : bool, optional
  2333. If True, hides the y-axis label, by default False.
  2334. remove_xticks : bool, optional
  2335. If True, removes the x-axis ticks entirely, by default False.
  2336. remove_yticks : bool, optional
  2337. If True, removes the y-axis ticks entirely, by default False.
  2338. remove_xticklabels : bool, optional
  2339. If True, removes the x-axis tick labels, by default False.
  2340. remove_yticklabels : bool, optional
  2341. If True, removes the y-axis tick labels, by default False.
  2342. extension : str, optional
  2343. The file extension for the saved figure, by default "svg".
  2344. dpi : int, optional
  2345. The resolution (dots per inch) for the saved figure, by default 1200.
  2346. bbox_inches : str, optional
  2347. Bounding box adjustment for saving, by default "tight".
  2348. transparent : bool, optional
  2349. If True, the saved figure will have a transparent background, by default True.
  2350. show : bool, optional
  2351. If True, displays the plot immediately. If False, the plot is not shown
  2352. but can be saved, by default False.
  2353. mask : np.ndarray, optional
  2354. A boolean mask to apply to the TFR data.
  2355. custom_yticks : list of float, optional
  2356. Custom y-axis tick locations, by default None.
  2357. custom_xticks : list of float, optional
  2358. Custom x-axis tick locations, by default None.
  2359. """
  2360. # Create the figure and axes objects
  2361. fig, ax = plt.subplots(figsize=figsize)
  2362. # Use the MNE plot function to draw the main spectrogram
  2363. tfr.plot(
  2364. chans,
  2365. baseline=None,
  2366. show=False,
  2367. vlim=vlim,
  2368. tmin=time_range[0],
  2369. tmax=time_range[1],
  2370. axes=ax,
  2371. cmap=cmap,
  2372. title=None, # Title is set manually later
  2373. colorbar=True,
  2374. combine=combine,
  2375. mask=mask,
  2376. mask_style=mask_style,
  2377. )
  2378. # --- Customize Plot Appearance ---
  2379. # Set font size for major ticks
  2380. ax.tick_params(axis="both", which="major", labelsize=fontsize - 2)
  2381. # Set labels and title with specified font size
  2382. ax.set_ylabel("Frequency (Hz)", fontsize=fontsize)
  2383. ax.set_xlabel("Time (s)", fontsize=fontsize)
  2384. ax.set_title(title, fontsize=fontsize)
  2385. # Set custom axis limits and ticks if provided
  2386. if time_range[0] is not None or time_range[1] is not None:
  2387. ax.set_xlim(time_range)
  2388. if freq_range is not None:
  2389. ax.set_ylim(freq_range)
  2390. if y_ticks is not None:
  2391. ax.set_yticks(y_ticks)
  2392. if x_ticks is not None:
  2393. ax.set_xticks(x_ticks)
  2394. # Customize the colorbar
  2395. if cbar:
  2396. # The colorbar is typically the second axes object created by tfr.plot
  2397. if len(fig.axes) > 1:
  2398. cbar_ax = fig.axes[1]
  2399. cbar_ax.set_ylabel(cbar_label, fontsize=fontsize)
  2400. cbar_ax.tick_params(labelsize=fontsize - 2)
  2401. if remove_cbar_ticks and len(fig.axes) > 1:
  2402. cbar_ax = fig.axes[1]
  2403. cbar_ax.set_yticks([])
  2404. else:
  2405. # If no colorbar is needed, remove the second axes if it exists
  2406. if len(fig.axes) > 1:
  2407. fig.delaxes(fig.axes[1])
  2408. # --- Hide or Remove Plot Elements ---
  2409. if hide_ylabel:
  2410. ax.set_ylabel("")
  2411. if remove_yticks:
  2412. ax.set_yticks([])
  2413. if custom_yticks is not None:
  2414. ax.set_yticks(custom_yticks)
  2415. if custom_xticks is not None:
  2416. ax.set_xticks(custom_xticks)
  2417. if hide_xlabel:
  2418. ax.set_xlabel("")
  2419. if remove_xticks:
  2420. ax.set_xticks([])
  2421. if remove_xticklabels:
  2422. ax.set_xticklabels([])
  2423. if remove_yticklabels:
  2424. ax.set_yticklabels([])
  2425. # --- Save Figure ---
  2426. if save_path is not None:
  2427. # Construct the full file path
  2428. full_path = os.path.join(
  2429. save_path,
  2430. f"{nameout}_spectro.{extension}",
  2431. )
  2432. # Save the figure
  2433. fig.savefig(
  2434. full_path,
  2435. dpi=dpi,
  2436. bbox_inches=bbox_inches,
  2437. transparent=transparent,
  2438. )
  2439. if not show:
  2440. plt.close(fig)
  2441. def plot_tfr_nice_topo(
  2442. tfr: AverageTFR,
  2443. fmin: float,
  2444. fmax: float,
  2445. tmin: float,
  2446. tmax: float,
  2447. title: str = "",
  2448. xlabel: str = "",
  2449. ylabel: str = "",
  2450. cmap: str = "RdBu_r",
  2451. vlim: Tuple[float, float] = (None, None),
  2452. figsize: Tuple[float, float] = (1, 1),
  2453. nameout: str = "",
  2454. cbar: bool = True,
  2455. fontsize: int = 10,
  2456. cbar_label: str = "Z-score",
  2457. save_path: Optional[str] = None,
  2458. extension: str = "png",
  2459. dpi: int = 800,
  2460. bbox_inches: str = "tight",
  2461. transparent: bool = True,
  2462. show: bool = False,
  2463. title_offset=6.0,
  2464. mask: Optional[np.ndarray] = None,
  2465. mask_params: Optional[dict] = None,
  2466. ):
  2467. """Plot a time-frequency representation (TFR) as a topographic map.
  2468. This function averages TFR data across a specified time and frequency
  2469. window and plots the result on a topographic map. It provides extensive
  2470. customization options for creating publication-quality figures.
  2471. Parameters
  2472. ----------
  2473. tfr : mne.time_frequency.AverageTFR
  2474. The time-frequency representation to plot.
  2475. fmin : float
  2476. The minimum frequency to include in the average (in Hz).
  2477. fmax : float
  2478. The maximum frequency to include in the average (in Hz).
  2479. tmin : float
  2480. The minimum time to include in the average (in seconds).
  2481. tmax : float
  2482. The maximum time to include in the average (in seconds).
  2483. title : str, optional
  2484. The title of the plot, by default "".
  2485. cmap : str, optional
  2486. The colormap to use for the plot, by default "RdBu_r".
  2487. vlim : tuple of float, optional
  2488. The limits for the color scale (vmin, vmax), by default (-1.5, 1.5).
  2489. figsize : tuple of float, optional
  2490. The size of the figure (width, height) in inches, by default (2, 2).
  2491. nameout : str, optional
  2492. The base name of the output file (without extension), by default "".
  2493. cbar : bool, optional
  2494. Whether to display the colorbar, by default True.
  2495. fontsize : int, optional
  2496. The base font size for the title, by default 10.
  2497. cbar_label : str, optional
  2498. The label for the colorbar, by default "Z-score".
  2499. save_path : str, optional
  2500. The directory path to save the figure. If None, the figure is not
  2501. saved, by default None.
  2502. extension : str, optional
  2503. The file extension for the saved figure, by default "png".
  2504. dpi : int, optional
  2505. The resolution (dots per inch) for the saved figure, by default 800.
  2506. bbox_inches : str, optional
  2507. Bounding box adjustment for saving, by default "tight".
  2508. transparent : bool, optional
  2509. If True, the saved figure will have a transparent background, by default True.
  2510. show : bool, optional
  2511. If True, displays the plot immediately. If False, the plot is not shown
  2512. but can be saved, by default False.
  2513. title_offset : float, optional
  2514. The offset for the title from the top of the plot, by default 6.0.
  2515. mask : np.ndarray, optional
  2516. A boolean mask to apply to the topographic data.
  2517. """
  2518. if mne is None:
  2519. raise ImportError("MNE-Python must be installed to use this function.")
  2520. # Average power over the specified time and frequency window
  2521. avg_power_topo = (
  2522. tfr.copy()
  2523. .crop(fmin=fmin, fmax=fmax, tmin=tmin, tmax=tmax)
  2524. .get_data()
  2525. .mean(axis=(1, 2))
  2526. )
  2527. # Create the figure and axes objects
  2528. fig, ax = plt.subplots(figsize=figsize)
  2529. # Plot the topomap
  2530. im, _ = mne.viz.plot_topomap(
  2531. avg_power_topo,
  2532. tfr.info,
  2533. vlim=vlim,
  2534. cmap=cmap,
  2535. contours=False,
  2536. axes=ax,
  2537. show=False,
  2538. mask=mask,
  2539. mask_params=mask_params,
  2540. )
  2541. # Set the title
  2542. ax.set_title(title, fontsize=fontsize, pad=title_offset)
  2543. if xlabel:
  2544. ax.set_xlabel(xlabel, fontsize=fontsize)
  2545. if ylabel:
  2546. ax.set_ylabel(ylabel, fontsize=fontsize)
  2547. # Add and customize the colorbar if requested
  2548. if cbar:
  2549. # Create an axis for the colorbar
  2550. cax = fig.add_axes(
  2551. [
  2552. ax.get_position().x1 + 0.01,
  2553. ax.get_position().y0,
  2554. 0.04,
  2555. ax.get_position().height,
  2556. ]
  2557. )
  2558. clb = fig.colorbar(im, cax=cax)
  2559. clb.set_label(cbar_label, fontsize=fontsize - 2)
  2560. clb.ax.tick_params(labelsize=fontsize - 3)
  2561. # Save the figure if a path is provided
  2562. if save_path is not None:
  2563. full_path = os.path.join(
  2564. save_path,
  2565. f"{nameout}_topo.{extension}",
  2566. )
  2567. fig.savefig(
  2568. full_path,
  2569. dpi=dpi,
  2570. bbox_inches=bbox_inches,
  2571. transparent=transparent,
  2572. )
  2573. # Close the figure to free up memory
  2574. if not show:
  2575. plt.close(fig)
  2576. def rm_ttest_tfr(
  2577. tfr_dict1,
  2578. tfr_dict2,
  2579. freq_range=None,
  2580. time_range=None,
  2581. channels=None,
  2582. n_permutations=1000,
  2583. clusters=None,
  2584. tail=0,
  2585. n_jobs=1,
  2586. seed=None,
  2587. decimate=None,
  2588. threshold=None,
  2589. adjacency_name=None,
  2590. adjacency_freqs=True,
  2591. alpha=0.05,
  2592. that_correction=False,
  2593. ):
  2594. # Ensure both dictionaries have the same keys
  2595. if set(tfr_dict1.keys()) != set(tfr_dict2.keys()):
  2596. raise ValueError("Both TFR dictionaries must have the same keys.")
  2597. # Assert all files in the dict have the same channels order
  2598. ch_names = None
  2599. for key in tfr_dict1.keys():
  2600. if ch_names is None:
  2601. ch_names = tfr_dict1[key].info["ch_names"]
  2602. else:
  2603. if tfr_dict1[key].info["ch_names"] != ch_names:
  2604. raise ValueError(
  2605. f"Channels for participant {key} does not match the expected channels order."
  2606. )
  2607. if tfr_dict2 is not None:
  2608. if tfr_dict2[key].info["ch_names"] != ch_names:
  2609. raise ValueError(
  2610. f"Channels for participant {key} does do not match the expected channels order."
  2611. )
  2612. # Drop non eeg channels if present
  2613. for key in tfr_dict1.keys():
  2614. tfr_dict1[key].pick(["eeg"])
  2615. if tfr_dict2 is not None:
  2616. tfr_dict2[key].pick(["eeg"])
  2617. if freq_range is not None:
  2618. for tfr in tfr_dict1.values():
  2619. tfr.crop(fmin=freq_range[0], fmax=freq_range[1])
  2620. if tfr_dict2 is not None:
  2621. for tfr in tfr_dict2.values():
  2622. tfr.crop(fmin=freq_range[0], fmax=freq_range[1])
  2623. if time_range is not None:
  2624. for tfr in tfr_dict1.values():
  2625. tfr.crop(tmin=time_range[0], tmax=time_range[1])
  2626. if tfr_dict2 is not None:
  2627. for tfr in tfr_dict2.values():
  2628. tfr.crop(tmin=time_range[0], tmax=time_range[1])
  2629. if channels is not None:
  2630. for tfr in tfr_dict1.values():
  2631. tfr.pick(channels)
  2632. if tfr_dict2 is not None:
  2633. for tfr in tfr_dict2.values():
  2634. tfr.pick(channels)
  2635. if decimate is not None:
  2636. for tfr in tfr_dict1.values():
  2637. tfr.decimate(decimate, verbose=False)
  2638. if tfr_dict2 is not None:
  2639. for tfr in tfr_dict2.values():
  2640. tfr.decimate(decimate, verbose=False)
  2641. # Extract data from both dictionaries
  2642. data1 = np.array([tfr.data for tfr in tfr_dict1.values()])
  2643. if tfr_dict2 is not None:
  2644. data2 = np.array([tfr.data for tfr in tfr_dict2.values()])
  2645. # Get the difference between the two datasets (test diff against zero)
  2646. data_diff = data1 - data2
  2647. else:
  2648. # If only one dictionary is provided, use it as the data_diff
  2649. data_diff = data1
  2650. # Dimensions check
  2651. tfr = list(tfr_dict1.values())[0] # Assuming all TFRs have the same shape
  2652. # MNE should always have participant, channel, frequency, time dimensions
  2653. assert np.argwhere(np.asarray(data_diff.shape) == len(tfr_dict1))[0][0] == 0
  2654. assert (
  2655. np.argwhere(np.asarray(data_diff.shape) == len(tfr.info["ch_names"]))[0][0] == 1
  2656. )
  2657. assert np.argwhere(np.asarray(data_diff.shape) == len(tfr.freqs))[0][0] == 2
  2658. assert np.argwhere(np.asarray(data_diff.shape) == len(tfr.times))[0][0] == 3
  2659. if that_correction:
  2660. # Apply that correction to the data_diff
  2661. stat_fun = partial(mne.stats.ttest_1samp_no_p, sigma=1e-3)
  2662. else:
  2663. stat_fun = mne.stats.ttest_1samp_no_p
  2664. # Perform the paired t-test across subjects
  2665. if clusters is None or clusters is False:
  2666. # Reshape data so its 2d with (OBS, Chans*Freqs*Times)
  2667. X = data_diff.reshape(data_diff.shape[0], -1)
  2668. # Perform the permutation t-test (t) NOTE (From mne documentation): When applying the test to multiple variables, the “tmax” method is used for adjusting the p-values of each variable for multiple comparisons.
  2669. t_stat, p_val = mne.stats.permutation_t_test(
  2670. X, n_permutations=n_permutations, tail=tail, n_jobs=n_jobs, seed=seed
  2671. )
  2672. # Reshape t_stat and p_val back to the original shape
  2673. t_stat = t_stat.reshape(data_diff.shape[1:])
  2674. p_val = p_val.reshape(data_diff.shape[1:])
  2675. else:
  2676. # If clusters are provided, use spatio-temporal cluster test
  2677. # Read the channel adjacency from the file if provided otherwise use the default adjacency
  2678. if adjacency_name is not None:
  2679. sensor_adjacency, _ = mne.channels.read_ch_adjacency(
  2680. adjacency_name, tfr.ch_names
  2681. )
  2682. else:
  2683. sensor_adjacency, _ = mne.channels.find_ch_adjacency(
  2684. tfr.info,
  2685. ch_type="eeg",
  2686. )
  2687. # Cmbine adjacency in time and frequency dimensions
  2688. if adjacency_freqs is True:
  2689. adjacency = mne.stats.combine_adjacency(
  2690. sensor_adjacency, len(tfr.freqs), len(tfr.times)
  2691. )
  2692. else:
  2693. adjacency = mne.stats.combine_adjacency(
  2694. sensor_adjacency,
  2695. np.zeros((len(tfr.freqs), len(tfr.freqs))),
  2696. len(tfr.times),
  2697. )
  2698. # If threshold is not provided, use the t-distribution to calculate the threshold
  2699. if threshold is None:
  2700. df = data_diff.shape[0] - 1 # degrees of freedom
  2701. if tail == 0:
  2702. t_thresh = scipy.stats.distributions.t.ppf(1 - alpha / 2, df=df)
  2703. elif tail == 1:
  2704. t_thresh = scipy.stats.distributions.t.ppf(1 - alpha, df=df)
  2705. elif threshold == "TFCE":
  2706. # If TFCE is used, set the threshold to None
  2707. t_thresh = dict(start=0, step=0.2)
  2708. t_stat, clusters, cluster_p_values, H0 = (
  2709. mne.stats.permutation_cluster_1samp_test(
  2710. data_diff,
  2711. n_permutations=n_permutations,
  2712. tail=tail,
  2713. n_jobs=n_jobs,
  2714. stat_fun=stat_fun,
  2715. seed=seed,
  2716. adjacency=adjacency,
  2717. threshold=t_thresh,
  2718. )
  2719. )
  2720. # Create a p-value array with the same shape as t_stat
  2721. p_val = np.ones_like(t_stat)
  2722. # Fill the p-values for the clusters
  2723. for cl, p in zip(clusters, cluster_p_values):
  2724. p_val[cl] = p
  2725. # --- Summarize the clusters in a pandas DataFrame ---
  2726. cluster_chans, cluster_freqs, cluster_times, cluster_tvals, cluster_pvals = (
  2727. [],
  2728. [],
  2729. [],
  2730. [],
  2731. [],
  2732. )
  2733. cluster_extent = []
  2734. for c_idx, c_pvals in zip(clusters, cluster_p_values):
  2735. cluster_chans.append(set([tfr.ch_names[c] for c in c_idx[0]]))
  2736. cluster_freqs.append(list(set([float(tfr.freqs[c]) for c in c_idx[1]])))
  2737. cluster_times.append(list(set([float(tfr.times[c]) for c in c_idx[2]])))
  2738. cluster_tvals.append(t_stat[c_idx])
  2739. cluster_pvals.append(c_pvals)
  2740. cluster_extent.append(
  2741. np.prod(np.asarray(c_idx).shape)
  2742. ) # Calculate the extent of the cluster
  2743. # Create a DataFrame to summarize the clusters
  2744. cluster_frame = pd.DataFrame(
  2745. {
  2746. "cluster_id": range(len(cluster_chans)),
  2747. "channels": cluster_chans,
  2748. "frequencies": cluster_freqs,
  2749. "times": cluster_times,
  2750. "t-values": cluster_tvals,
  2751. "p-values": cluster_pvals,
  2752. "extent": cluster_extent,
  2753. }
  2754. )
  2755. # Create an Evoked object for the t-statistics
  2756. cp = tfr.copy()
  2757. cp.data = t_stat
  2758. t_stat_evoked = cp.copy()
  2759. # Create an Evoked object for the p-values
  2760. cp_p = tfr.copy()
  2761. cp_p.data = p_val
  2762. p_val_evoked = cp_p.copy()
  2763. return {
  2764. "t_stat": t_stat,
  2765. "p_val": p_val,
  2766. "clusters": clusters,
  2767. "cluster_p_values": cluster_p_values,
  2768. "cluster_frame": cluster_frame,
  2769. "t_stat_evoked": t_stat_evoked,
  2770. "p_val_evoked": p_val_evoked,
  2771. "info": tfr.info,
  2772. }
  2773. def rm_anova_tfr(
  2774. tfr_data_list_of_dicts,
  2775. factor_levels=None, # Tuple of integers, e.g., (2, 3) for two factors with 2 and 3 levels
  2776. effects="A", # List of lists/tuples, e.g., [[0], [1], [0, 1]] for main effects and interaction
  2777. freq_range=None,
  2778. time_range=None,
  2779. channels=None,
  2780. n_permutations=1000,
  2781. n_jobs=1,
  2782. seed=None,
  2783. decimate=None,
  2784. p_threshold=0.05, # P-value threshold for F-statistic thresholding
  2785. adjacency_name="easycapM1",
  2786. ):
  2787. """
  2788. Performs an M-way repeated measures ANOVA using a cluster-based permutation test
  2789. on MNE-Python TFR (Time-Frequency Representation) data.
  2790. Parameters
  2791. ----------
  2792. tfr_data_list_of_dicts : list of dict
  2793. A list where each element is a dictionary representing a unique experimental
  2794. condition (cell) defined by the combination of factor levels. Each inner
  2795. dictionary should contain MNE TFR objects, keyed by subject ID.
  2796. Example:
  2797. [
  2798. {'sub01': tfr_A1B1_s1, 'sub02': tfr_A1B1_s2, ...}, # Cell (Level1_FactorA, Level1_FactorB)
  2799. {'sub01': tfr_A1B2_s1, 'sub02': tfr_A1B2_s2, ...}, # Cell (Level1_FactorA, Level2_FactorB)
  2800. ...
  2801. ]
  2802. All inner dictionaries must have the same subject IDs as keys.
  2803. factor_levels : tuple of int
  2804. A tuple where each integer represents the number of levels for a factor.
  2805. For example, (2, 3) means two factors, one with 2 levels and one with 3 levels.
  2806. effects : list of lists or tuples
  2807. A list specifying the effects to test. Each inner list/tuple contains
  2808. the 0-based indices of the factors involved in a main effect or interaction.
  2809. For example, [[0], [1], [0, 1]] tests the main effect of factor 0,
  2810. main effect of factor 1, and the interaction between factor 0 and 1.
  2811. freq_range : tuple of float | None
  2812. Frequency range (fmin, fmax) to crop the TFR data.
  2813. time_range : tuple of float | None
  2814. Time range (tmin, tmax) to crop the TFR data.
  2815. channels : list of str | None
  2816. List of channel names to pick.
  2817. n_permutations : int
  2818. The number of permutations to perform for the cluster test.
  2819. tail : int
  2820. The tail of the test. For F-tests, this should always be 1 (positive values).
  2821. n_jobs : int
  2822. Number of jobs to run in parallel. -1 uses all available CPU cores.
  2823. seed : int | None
  2824. Random seed for reproducibility of permutations.
  2825. decimate : int | None
  2826. Decimation factor for TFR data.
  2827. p_threshold : float
  2828. The p-value threshold used to determine the F-statistic threshold for
  2829. cluster formation.
  2830. adjacency_name : str | None
  2831. Name of the channel adjacency file (e.g., 'easycap-M1'). If None,
  2832. `mne.channels.find_ch_adjacency` will be used.
  2833. Returns
  2834. -------
  2835. dict
  2836. A dictionary containing the results:
  2837. 'f_stat' : np.ndarray
  2838. The observed F-statistics.
  2839. 'p_val' : np.ndarray
  2840. The p-values for each data point, adjusted for multiple comparisons
  2841. via cluster correction.
  2842. 'clusters' : list of tuples
  2843. A list of boolean masks, where each mask represents a significant cluster.
  2844. 'cluster_p_values' : np.ndarray
  2845. The p-values associated with each significant cluster.
  2846. 'f_stat_tfr' : mne.time_frequency.AverageTFR
  2847. An MNE TFR object containing the F-statistics.
  2848. 'p_val_tfr' : mne.time_frequency.AverageTFR
  2849. An MNE TFR object containing the cluster-corrected p-values.
  2850. 'info' : mne.Info
  2851. The MNE Info object from the TFR data.
  2852. """
  2853. # If effets and factors_levels are not provided, assume one way with number of levels equal to the number of dictionaries in tfr_data_list_of_dicts
  2854. if factor_levels is None:
  2855. factor_levels = (len(tfr_data_list_of_dicts),)
  2856. # --- Data Preprocessing ---
  2857. tfr_sample = None # To store a sample TFR for info and shape
  2858. if not tfr_data_list_of_dicts:
  2859. raise ValueError("tfr_data_list_of_dicts cannot be empty.")
  2860. # Get subject IDs from the first dictionary (assuming all dictionaries have the same subjects)
  2861. subject_ids = list(tfr_data_list_of_dicts[0].keys())
  2862. # Ensure all dictionaries have the same keys (subject IDs)
  2863. for i, tfr_dict_level in enumerate(tfr_data_list_of_dicts):
  2864. if set(tfr_dict_level.keys()) != set(subject_ids):
  2865. raise ValueError(
  2866. f"Dictionaries in tfr_data_list_of_dicts must have the same subject IDs. "
  2867. f"Mismatch found in dictionary at index {i}."
  2868. )
  2869. # Iterate through subjects first, then through conditions/levels to build the data array
  2870. # This ensures the data is ordered correctly for f_mway_rm:
  2871. # (subject1_condition1, subject1_condition2, ..., subject2_condition1, ...)
  2872. data_for_anova = [] # List to hold TFR data for all subjects
  2873. for tfr_dict_level in tfr_data_list_of_dicts:
  2874. all_cond_data = (
  2875. []
  2876. ) # List to hold TFR data for this condition across all subjects
  2877. for sub_id in subject_ids:
  2878. tfr = tfr_dict_level[sub_id]
  2879. # Pick only EEG channels
  2880. tfr.pick(["eeg"])
  2881. # Crop frequency range
  2882. if freq_range is not None:
  2883. tfr.crop(fmin=freq_range[0], fmax=freq_range[1])
  2884. # Crop time range
  2885. if time_range is not None:
  2886. tfr.crop(tmin=time_range[0], tmax=time_range[1])
  2887. # Pick specific channels
  2888. if channels is not None:
  2889. tfr.pick(channels)
  2890. # Decimate data
  2891. if decimate is not None:
  2892. tfr.decimate(decimate, verbose=False)
  2893. all_cond_data.append(tfr.data)
  2894. # Store a sample TFR object after all preprocessing for info and shape
  2895. if tfr_sample is None:
  2896. tfr_sample = tfr.copy()
  2897. # Stack the data for this subject
  2898. data_for_anova.append(np.stack(all_cond_data, axis=0))
  2899. # --- Dimension Checks ---
  2900. expected_observations = len(subject_ids) * np.prod(factor_levels)
  2901. if len(data_for_anova) * data_for_anova[0].shape[0] != expected_observations:
  2902. raise ValueError(
  2903. f"Mismatch in number of observations. Expected {expected_observations} "
  2904. f"based on n_replications and factor_levels, but got {data_for_anova.shape[0]}."
  2905. f"Ensure tfr_data_list_of_dicts contains the correct number of TFR objects "
  2906. f"for all subjects across all conditions."
  2907. )
  2908. if data_for_anova[0].shape[1] != len(tfr_sample.info["ch_names"]):
  2909. raise ValueError("Channel dimension mismatch after preprocessing.")
  2910. if data_for_anova[0].shape[2] != len(tfr_sample.freqs):
  2911. raise ValueError("Frequency dimension mismatch after preprocessing.")
  2912. if data_for_anova[0].shape[3] != len(tfr_sample.times):
  2913. raise ValueError("Time dimension mismatch after preprocessing.")
  2914. # --- Define the Statistical Function for ANOVA ---
  2915. # This function will be called by mne.stats.permutation_cluster_test
  2916. # It receives a subset of the data (e.g., a single time point across all observations)
  2917. # and should return the statistic (F-value in this case).
  2918. def stat_fun_anova(*args):
  2919. # args will be a tuple containing the data array for the current permutation.
  2920. # The data array will have shape (n_observations, n_features_at_this_point)
  2921. # where n_features_at_this_point is (n_channels * n_freqs * n_times) if flattened.
  2922. # f_mway_rm expects (n_observations, n_features)
  2923. data_reshaped_for_anova = np.swapaxes(args, 1, 0)
  2924. # Calculate F-values using MNE's f_mway_rm
  2925. # We only need the F-values, not the p-values for the permutation test's stat_fun
  2926. f_values = mne.stats.f_mway_rm(
  2927. data_reshaped_for_anova,
  2928. factor_levels=factor_levels,
  2929. effects=effects,
  2930. return_pvals=False,
  2931. )[
  2932. 0
  2933. ] # [0] to get the F-values array
  2934. return f_values
  2935. # --- Calculate F-value Threshold for Clustering ---
  2936. # This threshold is used to define what constitutes a "cluster" in the data.
  2937. f_thresh = mne.stats.f_threshold_mway_rm(
  2938. len(subject_ids), factor_levels, effects, p_threshold
  2939. )
  2940. # --- Prepare Adjacency Matrix for Spatio-Frequency-Temporal Clustering ---
  2941. # Adjacency defines which data points are considered "neighbors" for clustering.
  2942. # We need to combine adjacency for channels, frequencies, and time points.
  2943. # 1. Channel Adjacency: Based on sensor proximity.
  2944. if adjacency_name is not None:
  2945. sensor_adjacency, _ = mne.channels.read_ch_adjacency(
  2946. adjacency_name, tfr_sample.ch_names
  2947. )
  2948. else:
  2949. # Fallback if no specific adjacency name is provided
  2950. sensor_adjacency, _ = mne.channels.find_ch_adjacency(
  2951. tfr_sample.info, ch_type="eeg", verbose=False
  2952. )
  2953. # 2. Frequency Adjacency: Simple nearest-neighbor for frequencies.
  2954. # Combine all adjacencies into a single sparse matrix for 3D data (channels, freqs, times)
  2955. # The order here should match the order of dimensions in the data's feature space
  2956. # (i.e., after observations, which is channels, freqs, times).
  2957. full_adjacency = mne.stats.combine_adjacency(
  2958. sensor_adjacency,
  2959. len(tfr_sample.freqs), # Number of frequencies
  2960. len(tfr_sample.times), # Number of time points
  2961. )
  2962. # --- Perform Cluster-Based Permutation Test ---
  2963. # This is the main statistical test, which accounts for multiple comparisons.
  2964. F_obs, clusters_out, cluster_p_values_out, H0 = mne.stats.permutation_cluster_test(
  2965. data_for_anova, # Input data:list of conditions (n_observations, n_channels, n_freqs, n_times)
  2966. stat_fun=stat_fun_anova, # Our custom ANOVA statistic function
  2967. threshold=f_thresh, # F-value threshold for cluster formation
  2968. tail=1, # 1 for F-test (positive values)
  2969. n_jobs=n_jobs, # For parallel processing
  2970. n_permutations=n_permutations, # Number of permutations
  2971. buffer_size=10, # No buffering
  2972. seed=seed, # For reproducibility
  2973. adjacency=full_adjacency, # The combined spatio-frequency-temporal adjacency
  2974. )
  2975. # --- Create Cluster-Corrected P-value Array ---
  2976. # Initialize a p-value array with ones (non-significant by default)
  2977. p_val = np.ones_like(F_obs)
  2978. # Fill in the p-values for the significant clusters
  2979. if clusters_out: # Check if any clusters were found
  2980. for cl, p in zip(clusters_out, cluster_p_values_out):
  2981. p_val[cl] = p
  2982. # --- Summarize the clusters in a pandas DataFrame ---
  2983. cluster_chans, cluster_freqs, cluster_times, cluster_fvals, cluster_pvals = (
  2984. [],
  2985. [],
  2986. [],
  2987. [],
  2988. [],
  2989. )
  2990. cluster_extent = (
  2991. []
  2992. ) # To store the extent of each cluster (number of elements in the cluster)
  2993. for c_idx, c_pvals in zip(clusters_out, cluster_p_values_out):
  2994. cluster_chans.append(set([tfr_sample.ch_names[c] for c in c_idx[0]]))
  2995. cluster_freqs.append(list(set([float(tfr_sample.freqs[c]) for c in c_idx[1]])))
  2996. cluster_times.append(list(set([float(tfr_sample.times[c]) for c in c_idx[2]])))
  2997. cluster_fvals.append(F_obs[c_idx])
  2998. cluster_pvals.append(c_pvals)
  2999. cluster_extent.append(
  3000. np.prod(np.asarray(c_idx).shape)
  3001. ) # Calculate the extent of the cluster
  3002. cluster_frame = pd.DataFrame(
  3003. {
  3004. "cluster_id": range(len(cluster_chans)),
  3005. "channels": cluster_chans,
  3006. "frequencies": cluster_freqs,
  3007. "times": cluster_times,
  3008. "F-values": cluster_fvals,
  3009. "p-values": cluster_pvals,
  3010. "extent": cluster_extent,
  3011. }
  3012. )
  3013. # --- Create MNE TFR Objects for Results ---
  3014. # These TFR objects allow easy visualization of the F-statistics and p-values
  3015. # using MNE's plotting functions.
  3016. f_stat_tfr = tfr_sample.copy()
  3017. f_stat_tfr.data = F_obs # Assign the observed F-statistics
  3018. p_val_tfr = tfr_sample.copy()
  3019. p_val_tfr.data = p_val # Assign the cluster-corrected p-values
  3020. # --- Return Results ---
  3021. return {
  3022. "f_stat": F_obs,
  3023. "p_val": p_val,
  3024. "clusters": clusters_out,
  3025. "cluster_p_values": cluster_p_values_out,
  3026. "cluster_frame": cluster_frame,
  3027. "f_stat_tfr": f_stat_tfr,
  3028. "p_val_tfr": p_val_tfr,
  3029. "info": tfr_sample.info,
  3030. }
  3031. import pandas as pd
  3032. import mne
  3033. import numpy as np
  3034. def plot_clusters(
  3035. results_dict,
  3036. tvals,
  3037. figures_path,
  3038. test_name=None,
  3039. alpha=0.05,
  3040. fontsize=10,
  3041. x_ticks=None,
  3042. y_ticks=None,
  3043. figsize_spectro=(2, 1),
  3044. figsize_topo=(0.75, 0.75),
  3045. title_offset=3.0,
  3046. ):
  3047. """
  3048. Plots clusters from a summary DataFrame containing cluster information.
  3049. Parameters
  3050. ----------
  3051. summary : pd.DataFrame
  3052. DataFrame containing cluster information with columns:
  3053. 'channels', 'frequencies', 'times', 'extent', 'p-values'.
  3054. tvals : mne.time_frequency.AverageTFR
  3055. TFR object containing t-values for plotting. Must have the same info as used in the stats.
  3056. figures_path : str
  3057. Path to save the generated figures.
  3058. alpha : float
  3059. Significance threshold for p-values to filter clusters.
  3060. """
  3061. # If test_name is not provided, try to extract it from the summary DataFrame
  3062. if test_name is None:
  3063. test_name = "cluster_test"
  3064. # Make a folder for the figures if it doesn't exist
  3065. figures_path = os.path.join(figures_path, test_name + "_cluster_figures")
  3066. if not os.path.exists(figures_path):
  3067. os.makedirs(figures_path)
  3068. # Loop clusters and plot
  3069. idx = 0
  3070. tvals_summary = tvals.copy()
  3071. tvals_summary.data = np.ones(tvals.data.shape) * np.nan
  3072. channels_summary = []
  3073. for cluster, pval in zip(
  3074. results_dict["clusters"], results_dict["cluster_p_values"]
  3075. ):
  3076. if pval < alpha:
  3077. idx += 1
  3078. tvals_clusters = tvals.copy()
  3079. # Create a mask for the cluster
  3080. mask = np.zeros(tvals_clusters.data.shape, dtype=bool)
  3081. mask[cluster] = True
  3082. tvals_clusters.data[~mask] = np.nan
  3083. # Summarize the clusters in the summary tvals
  3084. tvals_summary.data[mask] = tvals_clusters.data[mask]
  3085. # Get the channels in the cluster
  3086. channels = [tvals_clusters.ch_names[i] for i in np.unique(cluster[0])]
  3087. channels_summary += channels
  3088. times = tvals_clusters.times[np.unique(cluster[2])]
  3089. freqs = tvals_clusters.freqs[np.unique(cluster[1])]
  3090. # Average across the cluster across times and frequencies
  3091. topo_cluster = np.nanmean(tvals_clusters.data, axis=(1, 2))
  3092. chan_mask = np.zeros(len(tvals_clusters.ch_names), dtype=bool)
  3093. chan_mask[np.unique(cluster[0])] = True
  3094. # Create the figure and axes objects
  3095. fig, ax = plt.subplots(figsize=figsize_topo)
  3096. # Plot the topomap
  3097. im, _ = mne.viz.plot_topomap(
  3098. np.where(np.isnan(topo_cluster), 0, topo_cluster),
  3099. tvals.info,
  3100. contours=False,
  3101. axes=ax,
  3102. show=False,
  3103. vlim=(-3, 3),
  3104. # mask=chan_mask,
  3105. # mask_params=dict(marker="o",
  3106. # markerfacecolor="w",
  3107. # markeredgecolor="k",
  3108. # linewidth=1,
  3109. # markersize=4
  3110. # )
  3111. )
  3112. ax.set_title(f"Cluster {idx}", fontsize=fontsize, pad=title_offset)
  3113. plt.savefig(
  3114. f"{figures_path}/{test_name}_cluster_{idx}_topo.svg",
  3115. bbox_inches="tight",
  3116. transparent=True,
  3117. )
  3118. fig, ax = plt.subplots(figsize=figsize_topo)
  3119. # Sum t-values across the cluster
  3120. topo_cluster = np.nansum(tvals_clusters.data, axis=(1, 2))
  3121. im, _ = mne.viz.plot_topomap(
  3122. np.where(np.isnan(topo_cluster), 0, topo_cluster),
  3123. tvals.info,
  3124. contours=False,
  3125. axes=ax,
  3126. show=False,
  3127. # vlim=(-3, 3),
  3128. # mask=chan_mask,
  3129. # mask_params=dict(marker="o",
  3130. # markerfacecolor="w",
  3131. # markeredgecolor="k",
  3132. # linewidth=1,
  3133. # markersize=4
  3134. # )
  3135. )
  3136. # Create an axis for the colorbar
  3137. cax = fig.add_axes(
  3138. [
  3139. ax.get_position().x1 + 0.01,
  3140. ax.get_position().y0,
  3141. 0.04,
  3142. ax.get_position().height,
  3143. ]
  3144. )
  3145. clb = fig.colorbar(im, cax=cax)
  3146. clb.set_label("Summed t-values", fontsize=fontsize - 2)
  3147. clb.set_ticks([])
  3148. ax.set_title(f"Cluster {idx}", fontsize=fontsize, pad=title_offset)
  3149. plt.savefig(
  3150. f"{figures_path}/{test_name}_cluster_{idx}_topo_sum.svg",
  3151. bbox_inches="tight",
  3152. transparent=True,
  3153. )
  3154. # Do a plot at the max time and frequency point in the cluster
  3155. # This craches with index 558 is out of bounds for axis 0 with size
  3156. # Find the maximum t-value in the cluster
  3157. # Channels is 0 , frequency 1, time 2
  3158. # Find the time point corresponding to the maximum t-value
  3159. max_time = tvals.times[
  3160. np.nanargmax(np.nanmax(np.abs(tvals_clusters.data), axis=(0, 1)))
  3161. ]
  3162. max_freq = tvals.freqs[
  3163. np.nanargmax(np.nanmax(np.abs(tvals_clusters.data), axis=(0, 2)))
  3164. ]
  3165. dat_max = np.where(
  3166. np.isnan(
  3167. tvals_clusters.data[
  3168. :,
  3169. np.argmin(np.abs(tvals.freqs - max_freq)),
  3170. np.argmin(np.abs(tvals.times - max_time)),
  3171. ]
  3172. ),
  3173. 0,
  3174. tvals_clusters.data[
  3175. :,
  3176. np.argmin(np.abs(tvals.freqs - max_freq)),
  3177. np.argmin(np.abs(tvals.times - max_time)),
  3178. ],
  3179. )
  3180. chan_mask = np.where(dat_max != 0, True, False)
  3181. fig, ax = plt.subplots(figsize=(1, 1))
  3182. im, _ = mne.viz.plot_topomap(
  3183. np.where(
  3184. np.isnan(
  3185. tvals_clusters.data[
  3186. :,
  3187. np.argmin(np.abs(tvals.freqs - max_freq)),
  3188. np.argmin(np.abs(tvals.times - max_time)),
  3189. ]
  3190. ),
  3191. 0,
  3192. tvals_clusters.data[
  3193. :,
  3194. np.argmin(np.abs(tvals.freqs - max_freq)),
  3195. np.argmin(np.abs(tvals.times - max_time)),
  3196. ],
  3197. ),
  3198. tvals.info,
  3199. contours=False,
  3200. axes=ax,
  3201. show=False,
  3202. vlim=(-3, 3),
  3203. mask=chan_mask,
  3204. )
  3205. ax.set_title(
  3206. f"Cluster {idx} at {max_freq:.0f} Hz, {max_time:.2f} s",
  3207. fontsize=10,
  3208. pad=3,
  3209. )
  3210. plt.savefig(
  3211. f"{figures_path}/{test_name}_cluster_{idx}_topo_MAX.svg",
  3212. bbox_inches="tight",
  3213. transparent=True,
  3214. )
  3215. # Plot spectrogram for the channels in the cluster
  3216. tvals_clusters.data = np.where(
  3217. np.isnan(tvals_clusters.data), 0, tvals_clusters.data
  3218. )
  3219. plot_tfr_nice_spectro(
  3220. tfr=tvals_clusters,
  3221. title=test_name + f"cluster {idx} (p={pval:.4f})",
  3222. chans=channels,
  3223. vlim=(-3, 3),
  3224. figsize=figsize_spectro,
  3225. nameout=test_name + f"cluster_{idx}_spectro",
  3226. cbar=True,
  3227. fontsize=fontsize,
  3228. cbar_label="T-value",
  3229. save_path=figures_path,
  3230. extension="svg",
  3231. bbox_inches="tight",
  3232. x_ticks=x_ticks,
  3233. y_ticks=y_ticks,
  3234. transparent=True,
  3235. show=True,
  3236. )
  3237. plot_tfr_nice_spectro(
  3238. tfr=tvals_clusters,
  3239. title=test_name + f"cluster {idx} (p={pval:.4f})",
  3240. chans="eeg",
  3241. # vlim=(-3, 3),
  3242. figsize=figsize_spectro,
  3243. nameout=test_name + f"cluster_{idx}_spectro_sum",
  3244. cbar=True,
  3245. fontsize=fontsize,
  3246. cbar_label="Summed T-values",
  3247. combine=lambda data: np.nansum(data, axis=0),
  3248. save_path=figures_path,
  3249. extension="svg",
  3250. bbox_inches="tight",
  3251. x_ticks=x_ticks,
  3252. y_ticks=y_ticks,
  3253. transparent=True,
  3254. show=True,
  3255. remove_cbar_ticks=True,
  3256. )
  3257. if idx != 0:
  3258. # Replace NaNs with zeros for better visualization
  3259. tvals_summary.data = np.where(
  3260. np.isnan(tvals_summary.data), 0, tvals_summary.data
  3261. )
  3262. # Plot spectrogram for the channels in the cluster
  3263. plot_tfr_nice_spectro(
  3264. tfr=tvals_summary,
  3265. title=test_name + f"Summary of significant clusters",
  3266. chans=channels_summary,
  3267. figsize=figsize_spectro,
  3268. nameout=test_name + f"cluster_summary_spectro",
  3269. cbar=True,
  3270. x_ticks=x_ticks,
  3271. y_ticks=y_ticks,
  3272. fontsize=fontsize,
  3273. cbar_label="T-value",
  3274. save_path=figures_path,
  3275. extension="svg",
  3276. bbox_inches="tight",
  3277. transparent=True,
  3278. show=True,
  3279. )
  3280. # Sum summary
  3281. plot_tfr_nice_spectro(
  3282. tfr=tvals_summary,
  3283. title=test_name + f"Summary of significant clusters",
  3284. chans="eeg",
  3285. figsize=figsize_spectro,
  3286. nameout=test_name + f"cluster_summary_spectro_sum",
  3287. cbar=True,
  3288. x_ticks=x_ticks,
  3289. y_ticks=y_ticks,
  3290. fontsize=fontsize,
  3291. cbar_label="Summed T-values",
  3292. remove_cbar_ticks=True,
  3293. # Do sum of t-values in the cluster
  3294. combine=lambda data: np.nansum(data, axis=0),
  3295. save_path=figures_path,
  3296. extension="svg",
  3297. bbox_inches="tight",
  3298. transparent=True,
  3299. show=True,
  3300. )

coll_lab_eeg_pipeline.py at commit c2e571c, no license · at the source

Overview

Authors: Nicolas Roy1,2, Coralie Deslauriers1, Thaliane Côté-Cazes1, Audrey Etcheverry1, Michel-Pierre Coll1,2
  1. École de Psychologie, Université Laval, Quebec City, QC, Canada
  2. Centre Interdisciplinaire de Recherche en Réadaptation et Intégration Sociale (Cirris), Quebec City, QC, Canada
Journal: Pain, volume 167, issue 9, pages e441-e452
Dates: received 10 February 2026; accepted 1 May 2026; published online 24 June 2026; in print September 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1097/j.pain.0000000000004044 · PMID 42350309 · PMCID PMC13528908 · OpenAlex W7165961531
Open access: hybrid, a free copy (OpenAlex)
Status: code verified
Categories: EEG (modality), human (organism), pain (population), cognitive (subfield)
Methods: Spectral & time-frequency, Connectivity, Statistics, Smoothing, state filtering, decompositions, Preprocessing, Evoked potentials, Physiology & signal measures
Keywords: Pain perception, Rhythmic visual entrainment, EEG, Neural oscillations
MeSH: Pain*, Pain Perception*, Periodicity*, Photic Stimulation*, Adolescent, Adult, Electric Stimulation, Electroencephalography, Female, Humans, Male, Pain Measurement, Young Adult (* major topic)
Topic: Pain Mechanisms and Treatments (Physiology, Medicine), according to OpenAlex
Citations: not cited yet (Europe PMC); 68 references in the paper

Abstract

The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.

Repository

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

mpcoll/2025_painrvs

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: c2e571c7b10ac1229569a839a23d24bef92ceb45, 9 April 2026
Languages: Python (9), Shell (1)
Size: 1,227 files, 10 scripts
Software Heritage: not archived
Found in: the acknowledgements
Holds: README, environment (experiment1_eeg/code/requirements.txt)
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (9 files), Matplotlib (8 files), MNE-Python (7 files), pandas (7 files), seaborn (6 files), Pingouin (5 files), SciPy (4 files), statsmodels (2 files), Brain Connectivity Toolbox (1 file), ICLabel (1 file), MNE-BIDS (1 file), MNE-Connectivity (1 file), Nilearn (1 file), PyPREP (1 file), specparam (formerly FOOOF) (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
11 files

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;
  • 10 scripts, each with its path and the digest of its content;
  • 8 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

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 2, 28 September 2026

  • Publisher: n/a → Lippincott Williams & Wilkins

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 4 keywords, 13 MeSH terms, 3 funders, 60 references.

Cite

This paper

Roy, N., Deslauriers, C., Côté-Cazes, T., Etcheverry, A., & Coll, M.-P. (2026). No effect of rhythmic visual stimulation on experimental pain perception. Pain, 167(9), e441-e452. https://doi.org/10.1097/j.pain.0000000000004044

BibTeX

@article{roy2026no,
author = {Roy, Nicolas and Deslauriers, Coralie and Côté-Cazes, Thaliane and Etcheverry, Audrey and Coll, Michel-Pierre},
title = {{No effect of rhythmic visual stimulation on experimental pain perception}},
journal = {Pain},
year = {2026},
month = jun,
volume = {167},
number = {9},
pages = {e441--e452},
publisher = {Lippincott Williams \& Wilkins},
issn = {0304-3959},
doi = {10.1097/j.pain.0000000000004044},
url = {https://doi.org/10.1097/j.pain.0000000000004044},
pmid = {42350309},
pmcid = {PMC13528908}
}

RIS

TY - JOUR
AU - Roy, Nicolas
AU - Deslauriers, Coralie
AU - Côté-Cazes, Thaliane
AU - Etcheverry, Audrey
AU - Coll, Michel-Pierre
TI - No effect of rhythmic visual stimulation on experimental pain perception
T2 - Pain
J2 - Pain
PY - 2026
DA - 2026/06/24
VL - 167
IS - 9
SP - e441
EP - e452
SN - 0304-3959
PB - Lippincott Williams & Wilkins
DO - 10.1097/j.pain.0000000000004044
UR - https://doi.org/10.1097/j.pain.0000000000004044
LA - en
ER -

CSL-JSON

{
"id": "10.1097/j.pain.0000000000004044",
"type": "article-journal",
"title": "No effect of rhythmic visual stimulation on experimental pain perception",
"container-title": "Pain",
"author": [
{
"family": "Roy",
"given": "Nicolas"
},
{
"family": "Deslauriers",
"given": "Coralie"
},
{
"family": "Côté-Cazes",
"given": "Thaliane"
},
{
"family": "Etcheverry",
"given": "Audrey"
},
{
"family": "Coll",
"given": "Michel-Pierre"
}
],
"container-title-short": "Pain",
"volume": "167",
"issue": "9",
"page": "e441-e452",
"DOI": "10.1097/j.pain.0000000000004044",
"PMID": "42350309",
"PMCID": "PMC13528908",
"ISSN": "0304-3959",
"publisher": "Lippincott Williams & Wilkins",
"URL": "https://doi.org/10.1097/j.pain.0000000000004044",
"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.1038/s41597-026-07350-9 [code]
An open multi-center MEG-EEG dataset for studying conscious visual perception.
Journal: Scientific data
In common: PyPREP, MNE-Connectivity, MNE-BIDS, 9 other tools, EEG, 2 references
[2] 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: PyPREP, MNE-Connectivity, MNE-BIDS, 9 other tools, cognitive, 1 reference
[3] doi:10.1038/s41597-026-07377-y [code]
An open-access multi-site fMRI dataset for investigating conscious visual perception.
Journal: Scientific data
In common: PyPREP, MNE-Connectivity, MNE-BIDS, 9 other tools
[4] doi:10.1093/cercor/bhag113 [code]
Long-term reliability and stability of parameterized resting state EEG: evidence from a five-year follow-up.
Journal: Cerebral cortex (New York, N.Y. : 1991)
In common: PyPREP, MNE-BIDS, specparam (formerly FOOOF), 7 other tools, EEG, 2 references
[5] doi:10.1371/journal.pbio.3003948 [code]
Neural encoding of pain is robust within but unstable between individuals.
Journal: PLoS biology
In common: pain, EEG, cognitive, 10 references
[6] doi:10.1162/imag.a.1269 [code]
From early to contemporary normative modeling: Mapping individual differences in neurophysiological signals.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: specparam (formerly FOOOF), ICLabel, Pingouin, 8 other tools, EEG, 2 references
[7] 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: PyPREP, MNE-BIDS, ICLabel, 6 other tools, EEG, 3 references
[8] doi:10.7554/elife.100605 [code]
Age-related changes in ‘cortical’ 1/f dynamics are linked to cardiac activity
Journal: n/a
In common: MNE-BIDS, specparam (formerly FOOOF), Pingouin, 7 other tools, 2 references
[9] 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, Pingouin, 8 other tools, EEG
[10] doi:10.3758/s13428-026-02997-z [code]
PyLossless: A non-destructive EEG processing pipeline.
Journal: Behavior research methods
In common: MNE-BIDS, ICLabel, MNE-Python, 3 other tools, EEG, 5 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.