OSCR

Bidirectional Mamba-2 boosts EEG super-resolution via regression and diffusion.

Code ↔ Paper

18 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 18 matches
  1. [1] § 2 Materials and methods › 2.3 Datasets and preprocessing ↔ utils.py, lines 169–198 · score 0.75 · notch filter, bad channels, MNE, FIR, head, 40 Hz
  2. [2] § 3 Results ↔ models.py, lines 236–371 · score 0.74 · n_layers, d_model, d_state, Bidirectional Mamba blocks, diffusion modeling, electrode
  3. [3] § 2 Materials and methods › 2.1 Architecture of the proposed system ↔ diffusion_conditioning_ablations.py, lines 28–57 · score 0.74 · squaredcos_cap_v2, DDPM scheduler, beta schedule, clipping, timesteps, linearly
  4. [4] § 2 Materials and methods › 2.1 Architecture of the proposed system ↔ ablations.py, lines 28–53 · score 0.72 · squaredcos_cap_v2, DDPM scheduler, beta schedule, timesteps, linearly, prediction
  5. [5] § 3 Results › 3.2 Results for spatio-temporal super-resolution ↔ compare_maser_topomap.py, lines 213–307 · score 0.67 · PSD topomap comparison, MASER x4, DiBiMa, x8, signal
  6. [6] § 3 Results › 3.2 Results for spatio-temporal super-resolution ↔ utils.py, lines 956–1075 · score 0.66 · spectral fidelity, frequency domains, temporal super resolution, waveform, EEG
  7. [7] § 2 Materials and methods › 2.3 Datasets and preprocessing ↔ visualize_input.ipynb, lines 268–407 · score 0.64 · notch filter, MNE, FIR, scoring, 40 Hz, 50 Hz
  8. [8] § 2 Materials and methods › 2.1 Architecture of the proposed system ↔ compare_to_maser.py, lines 235–315 · score 0.63 · MASER models, FLOPs, DiBiMa, inference, devices, position
  9. [9] § 3 Results › 3.3 Super-resolution explainability and downstream classification task ↔ eeg_fid.py, lines 175–205 · score 0.61 · EEG FID scores, DiBiMa, MASER, x4, models
  10. [10] § 2 Materials and methods › 2.4 Metrics ↔ metrics.py, lines 14–27 · score 0.60 · Peak Signal, Noise Ratio, PSNR, Metrics
  11. [11] § 2 Materials and methods › 2.2 Mamba blocks and the bidirectional mamba layer ↔ models.py, lines 667–809 · score 0.58 · bidirectional Mamba, Mamba blocks, fused, BiMamba, kernel, backward
  12. [12] § 3 Results › 3.2 Results for spatio-temporal super-resolution ↔ app.py, lines 157–184 · score 0.58 · PSD topomap, Qualitative comparison, BiMa, super resolution, x4, x8
  13. [13] § 3 Results › 3.3 Super-resolution explainability and downstream classification task ↔ downstream_task/metrics_class.py, lines 17–61 · score 0.58 · ROC AUC, F1 score, binary, accuracy, metrics, downstream
  14. [14] § 3 Results › 3.3 Super-resolution explainability and downstream classification task ↔ eeg_fid.py, lines 25–69 · score 0.57 · EEG FID scores, ResNet50, pretrained
  15. [15] § 3 Results › 3.3 Super-resolution explainability and downstream classification task ↔ downstream_task/metrics.py, lines 94–144 · score 0.57 · F1 score, ROC AUC, curves, accuracy, classification, metrics
  16. [16] § 2 Materials and methods › 2.4 Metrics ↔ metrics.py, lines 99–118 · score 0.53 · Pearson Correlation Coefficient, PCC, Metrics, channel
  17. [17] § 2 Materials and methods › 2.4 Metrics ↔ utils.py, lines 956–1075 · score 0.52 · Pearson Correlation Coefficient, PCC
  18. [18] § 2 Materials and methods › 2.1 Architecture of the proposed system ↔ diffusion_conditioning_ablations.py, lines 214–318 · score 0.52 · internal residuals, Electrode position, mamba blocks, embeddings, layer, resolution

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 · 1,120 lines · 44 KB · MIT · 3 matches

  1. import math
  2. import os
  3. import numpy as np
  4. from sklearn.preprocessing import StandardScaler
  5. import torch
  6. from torch.utils.data import Dataset
  7. import mne
  8. from mne.datasets import eegbci
  9. from scipy import signal
  10. import shutil
  11. import matplotlib.pyplot as plt
  12. from umap import UMAP
  13. import random
  14. import gc
  15. import torch
  16. import scipy
  17. from sklearn.model_selection import train_test_split
  18. from mne.preprocessing import find_bad_channels_lof, find_bad_channels_maxwell
  19. # Set random seeds for reproducibility
  20. seed = 2
  21. np.random.seed(seed)
  22. torch.manual_seed(seed)
  23. #case 1 of https://ieeexplore.ieee.org/document/9796118
  24. seed_channels = ["Fp1", "Fpz", "Fp2", "AF3", "AF4", "F7", "F5", "F3", "F1", "Fz", "F2", "F4", "F6", "F8", "FT7", "FC5", "FC3", "FC1", "FCz", "FC2", "FC4", "FC6", "FT8", "T7", "C5", "C3", "C1", "Cz", "C2", "C4", "C6", "T8", "TP7", "CP5", "CP3", "CP1", "CPz", "CP2", "CP4", "CP6", "TP8", "P7", "P5", "P3", "P1", "Pz", "P2", "P4", "P6", "P8", "PO7", "PO5", "PO3", "POz", "PO4", "PO6", "PO8", "CB1", "O1", "Oz", "O2", "CB2"]
  25. mmi_channels = ['FC5', 'FC3', 'FC1', 'FCz', 'FC2', 'FC4', 'FC6', 'C5', 'C3', 'C1', 'Cz', 'C2', 'C4', 'C6', 'CP5', 'CP3', 'CP1', 'CPz', 'CP2', 'CP4', 'CP6', 'Fp1', 'Fpz', 'Fp2', 'AF7', 'AF3', 'AFz', 'AF4', 'AF8', 'F7', 'F5', 'F3', 'F1', 'Fz', "F2", "F4", "F6", "F8", "FT7", "FT8", "T7", "T8", "T9", "T10", "TP7", "TP8", "P7", "P5", "P3", "P1", "Pz", "P2", "P4", "P6", "P8", "PO7", "PO3", "POz", "PO4", "PO8", "O1", "Oz", "O2", "Iz"]
  26. map_seed_channels = {seed_channels[i]: i for i in range(len(seed_channels))}
  27. map_mmi_channels = {mmi_channels[i]: i for i in range(len(mmi_channels))}
  28. case1_seed = {
  29. 'x2': ['AF3', 'AF4', 'FT7', 'FC5', 'FC3', 'FC1', 'FCz', 'FC2', 'FC4', 'FC6', 'FT8', 'TP7', 'CP5', 'CP3', 'CP1', 'CPz', 'CP2', 'CP4', 'CP6', 'TP8', 'CB1', 'CB2', 'PO3', 'PO4', 'PO5', 'PO7', 'PO6', 'PO8', 'POz', 'O1', 'Oz', 'O2'],
  30. 'x4': ['Fp1', 'Fp2', 'F5', 'Fz', 'F6', 'C3', 'Cz', 'C4', 'T7', 'T8', 'P5', 'Pz', 'P6', 'O1', 'O2'],
  31. 'x8': ['AF3', 'AF4', 'FC5', 'FC6', 'CP5', 'CP6', 'PO5', 'PO6']
  32. }
  33. case1_mmi = {
  34. "x2": ["Fpz", "AF7", "AF3", "AFz", "AF4", "AF8", "F7", "F3", "Fz", "F4", "F8", "FT7", "FC3", "FCz", "FC4", "FT8", "T7", "C3", "Cz", "C4", "T8", "TP7", "CP3", "CPz", "CP4", "TP8", "P7", "P3", "Pz", "P4", "P8", "PO7", "PO3", "POz", "PO4", "PO8", "Oz", "Iz"],
  35. "x4": ["Fp1", "Fp2", "F5", "Fz", "F6", "T7", "C3", "Cz", "C4", "T8", "P5", "Pz", "P6", "O1", "Oz", "O2"],
  36. "x8": ["AF3", "AF4", "FC5", "FC6", "CP5", "CP6", "O1", "O2"]
  37. }
  38. case2_mmi = {
  39. "x2": ["Fp1", "Fp2", "F5", "F1", "F2", "F6", "FC5", "FC1", "FC2", "FC6", "T9", "C5", "C1", "C2", "C6", "T10", "CP5", "CP1", "CP2", "CP6", "P5", "P1", "P2", "P6", "O1", "O2"],
  40. "x4": ["Fpz", "AF3", "AF4", "FC5", "FC1", "FC2", "FC6", "T9", "CP5", "CP1", "CP2", "CP6", "T10", "PO3", "PO4", "Oz"],
  41. "x8": ["Fp1", "Fp2", "Fz", "T7", "T8", "Pz", "O1", "O2"]
  42. }
  43. case2_seed = {
  44. "x2": ["Fpz", "AF3", "AF4", "F7", "F3", "Fz", "F4", "F8", "FT7", "FC3", "FCz", "FC4", "FT8", "T7", "C3", "Cz", "C4", "T8", "TP7", "CP3", "CPz", "CP4", "TP8", "P7", "P3", "Pz", "P4", "P8", "PO7", "PO3", "POz", "PO4", "PO8", "Oz"],
  45. "x4": ["Fpz", "AF3", "AF4", "FC5", "FC1", "FC2", "FC6", "CP5", "CP1", "CP2", "CP6", "CB1", "PO3", "PO4", "CB2", "Oz"],
  46. "x8": ["Fp1", "Fp2", "Fz", "T7", "T8", "Pz", "O1", "O2"]
  47. }
  48. unmask_channels = {
  49. "mmi":{
  50. "x2": [map_mmi_channels[i] for i in case1_mmi['x2']],
  51. "x4": [map_mmi_channels[i] for i in case1_mmi['x4']], # Your 16ch (frontal/central/parietal) #[0, 2, 4, 6, 14, 16, 18, 20, 22, 25, 27, 42, 43, 56, 58, 61]
  52. "x8": [map_mmi_channels[i] for i in case1_mmi['x8']] # Even 10-10 coverage
  53. },
  54. "seed":{
  55. "x2": [map_seed_channels[i] for i in case1_seed['x2']],
  56. "x4": [map_seed_channels[i] for i in case1_seed['x4']], # Frontal/central/parietal
  57. "x8": [map_seed_channels[i] for i in case1_seed['x8']] # Even 10-10 coverage
  58. }
  59. }
  60. #reverse ch_name:i to i:ch_name
  61. unmask_channels = {
  62. dataset: {
  63. key: sorted(value) for key, value in channels.items()
  64. } for dataset, channels in unmask_channels.items()
  65. }
  66. map_runs_dataset = {
  67. "mmi": range(1, 15),
  68. "seed": None # All files in folder
  69. }
  70. map_tasks = {
  71. 'task1': 'open-close right or left fist',
  72. 'task3': 'open-close both fists or both feet',
  73. 'task2': 'imagine right or left fist',
  74. 'task4': 'imagine both fists or both feet'
  75. }
  76. map_labels_seed = {
  77. -1: 0, # negative
  78. 0: 1, # neutral
  79. 1: 2 # positive
  80. }
  81. labels_subject_mapping = {
  82. 1: 1,
  83. 2: 0,
  84. 3: -1,
  85. 4: -1,
  86. 5: 0,
  87. 6: 1,
  88. 7: -1,
  89. 8: 0,
  90. 9: 1,
  91. 10: 1,
  92. 11: 0,
  93. 12: -1,
  94. 13: 0,
  95. 14: 1,
  96. 15: -1
  97. }
  98. map_runs_mmi = {
  99. 1: 'eyes_open',
  100. 2: 'eyes_closed',
  101. 3: map_tasks['task1'],
  102. 4: map_tasks['task2'],
  103. 5: map_tasks['task3'],
  104. 6: map_tasks['task4'],
  105. 7: map_tasks['task1'],
  106. 8: map_tasks['task2'],
  107. 9: map_tasks['task3'],
  108. 10: map_tasks['task4'],
  109. 11: map_tasks['task1'],
  110. 12: map_tasks['task2'],
  111. 13: map_tasks['task3'],
  112. 14: map_tasks['task4']
  113. }
  114. map_labels_mmi = {
  115. 'eyes_open': 0,
  116. 'eyes_closed': 1,
  117. 'open-close right or left fist': 2,
  118. 'open-close both fists or both feet': 3,
  119. 'imagine right or left fist': 4,
  120. 'imagine both fists or both feet': 5
  121. }
  122. map_labels_mmi_rev = {v: k for k, v in map_labels_mmi.items()}
  123. map_annotations = {
  124. "T0": "rest",
  125. "T1": "right or left fist",
  126. "T2": "both fists or both feet"
  127. }
  128. def get_lr_data_temporal(eeg_hr, factor=2):
  129. """Downsample high-res EEG temporally by factor."""
  130. if eeg_hr.ndim == 2:
  131. _, num_samples = eeg_hr.shape
  132. elif eeg_hr.ndim == 3:
  133. _, _, num_samples = eeg_hr.shape
  134. downsampled_length = num_samples // factor
  135. eeg_lr = signal.resample(eeg_hr, downsampled_length, axis=-1)
  136. return eeg_lr
  137. def get_lr_data_spatial(eeg_64, dataset_name, sr_ratio):
  138. """Extract low-res from 64ch EEG using unmask_channels"""
  139. indices = unmask_channels[dataset_name][f"x{sr_ratio}"]
  140. if eeg_64.ndim == 2:
  141. return eeg_64[indices, :]
  142. elif eeg_64.ndim == 3:
  143. return eeg_64[:, indices, :]
  144. # -------------------------------
  145. # Preprocessing MNE
  146. # -------------------------------
  147. def _preprocess_raw(raw, dataset_name, verbose=False):
  148. fs_cut = 40 if dataset_name == "mmi" else 50
  149. raw.pick_types(eeg=True, eog=False, stim=False)
  150. if dataset_name == "seed":
  151. # SKIP dev_head_t - not needed for EEG preprocessing [web:42]
  152. # Bad channel detection & cleaning (your function)
  153. raw = detect_and_clean_seed_trial(raw)
  154. if raw is None:
  155. return None
  156. # SEED filters
  157. raw.notch_filter(50.0, fir_design='firwin', verbose=verbose)
  158. raw.filter(1.0, fs_cut, fir_design='firwin', verbose=verbose)
  159. # Scale (confirm units with raw.plot() first)
  160. data = raw.get_data()
  161. #data = data * 1e-4
  162. data = (data - data.mean(1, keepdims=True)) / (data.std(1, keepdims=True) + 1e-8)
  163. raw._data = data
  164. else:
  165. raw.notch_filter(50.0, fir_design='firwin', verbose=verbose)
  166. raw.filter(0.5, fs_cut, fir_design='firwin', verbose=verbose)
  167. return raw
  168. # Extract valid EEG positions (C, 3), nan-mask if needed
  169. def get_electrode_positions(raw, channel_order=None):
  170. """
  171. Extract electrode positions from raw EEG data.
  172. Parameters
  173. ----------
  174. raw : mne.io.Raw
  175. Raw EEG data with montage set
  176. channel_order : list, optional
  177. Desired channel order for reordering
  178. Returns
  179. -------
  180. locs : torch.Tensor
  181. Electrode positions (C_eeg, 3)
  182. """
  183. picks = mne.pick_types(raw.info, eeg=True)
  184. locs = np.array([raw.info['chs'][p]['loc'][:3] for p in picks]) # (C_eeg, 3)
  185. # Check for NaNs and report which channels are problematic
  186. if np.isnan(locs).any():
  187. nan_channels = []
  188. for i, p in enumerate(picks):
  189. if np.isnan(locs[i]).any():
  190. ch_name = raw.info['chs'][p]['ch_name']
  191. nan_channels.append(ch_name)
  192. # Get available montage channels for debugging
  193. montage = raw.get_montage()
  194. if montage is not None:
  195. available_channels = list(montage.get_positions()['ch_pos'].keys())
  196. print(f"⚠️ Available montage channels: {available_channels[:10]}...") # Show first 10
  197. raise ValueError(
  198. f"NaNs persist after montage; check channel names. "
  199. f"Problematic channels: {nan_channels}. "
  200. f"Raw channels: {[raw.info['chs'][p]['ch_name'] for p in picks]}"
  201. )
  202. if channel_order: # Reorder to model input (e.g., 64 chs)
  203. reorder_idx = [channel_order.index(ch) for ch in raw.ch_names if ch in channel_order]
  204. locs = locs[reorder_idx]
  205. locs = torch.tensor(locs, dtype=torch.float32) # (C_eeg, 3)
  206. return locs
  207. def load_seed_channel_positions(filepath):
  208. channels = []
  209. positions = []
  210. with open(filepath, 'r') as f:
  211. for line in f:
  212. parts = line.strip().split()
  213. if len(parts) >= 4:
  214. ch_name = parts[-1] # Or parts[-1] if label last
  215. theta_deg = float(parts[1]) # Azimuth, negative=left
  216. phi_frac = float(parts[2]) # Elevation fraction (~0-0.6)
  217. phi_rad = np.radians(phi_frac * 90) # Adjust scale to ~90° max
  218. theta_rad = np.radians(theta_deg)
  219. x = np.sin(phi_rad) * np.cos(theta_rad)
  220. y = np.sin(phi_rad) * np.sin(theta_rad)
  221. z = np.cos(phi_rad)
  222. channels.append(ch_name)
  223. positions.append([x, y, z])
  224. return channels, np.array(positions)
  225. def set_montage(signal, dataset_name, pos=None, channel_names=None, fs=160):
  226. """
  227. Set montage for raw EEG data.
  228. For MMI: pos=None, uses standard_1020
  229. For SEED: pos=array of positions from .pos file
  230. """
  231. #print(channel_names)
  232. # Clean channel names
  233. cleaned_names = []
  234. for ch in channel_names:
  235. clean = ch.strip().rstrip('.')
  236. cleaned_names.append(clean)
  237. # Create MNE RawArray
  238. info = mne.create_info(ch_names=cleaned_names, sfreq=fs, ch_types='eeg')
  239. raw = mne.io.RawArray(signal, info)
  240. if dataset_name == "mmi":
  241. # MMI dataset: Use standard montage
  242. try:
  243. montage = mne.channels.make_standard_montage('standard_1020')
  244. #montage_channels = list(montage.get_positions()['ch_pos'].keys())
  245. #print(f"Montage channels available: {montage_channels[:10]}...") # Show first 10
  246. #print(f"Raw channels: {cleaned_names}")
  247. raw.set_montage(montage, on_missing='warn')
  248. except Exception as e:
  249. print(f"⚠️ Error setting standard montage: {e}")
  250. else:
  251. # SEED dataset: Use custom positions from .pos file
  252. try:
  253. ch_pos = {channel_names[i]: pos[i] for i in range(len(channel_names))}
  254. montage = mne.channels.make_dig_montage(
  255. ch_pos=ch_pos,
  256. coord_frame='head'
  257. )
  258. raw.set_montage(montage, on_missing='warn')
  259. except Exception as e:
  260. print(f"⚠️ Error setting custom montage: {e}")
  261. raise
  262. return raw
  263. def download_mmi_data(subject_ids, runs, project_path, demo = False, verbose=False, is_classification=False):
  264. datas = []
  265. labels = []
  266. positions = []
  267. channel_names = None
  268. for i, subject in enumerate(subject_ids):
  269. #print(f"Processing subject {subject}")
  270. print(f"⬇️ Downloading/Reading data for subject: {i+1}/{len(subject_ids)}", end='\r')
  271. if demo:
  272. if i == 1:
  273. break # For demo, process only first subject
  274. for run in runs:
  275. if is_classification and run not in [1, 2]:
  276. continue # Skip non-classification runs
  277. local_path = os.path.join(project_path, f'S{subject:03d}R{run:02d}.edf')
  278. #print(f"Checking local path: {local_path}")
  279. if not os.path.exists(local_path):
  280. print(f"Downloading S{subject:03d}R{run:02d}.edf...")
  281. try:
  282. eegbci.load_data(subject, [run], path=project_path, update_path=True, force_update=False)
  283. downloaded = os.path.join(
  284. os.path.dirname(project_path), "mmi",
  285. 'MNE-eegbci-data', 'files', 'eegmmidb', '1.0.0',
  286. f'S{subject:03d}', f'S{subject:03d}R{run:02d}.edf'
  287. )
  288. os.makedirs(os.path.dirname(local_path), exist_ok=True)
  289. shutil.move(downloaded, local_path)
  290. print(f"Saved to: {local_path}")
  291. except Exception as e:
  292. print(f"Download error: {e}")
  293. continue
  294. try:
  295. raw = mne.io.read_raw_edf(local_path, preload=True, verbose=verbose)
  296. mne.datasets.eegbci.standardize(raw)
  297. signal = raw.get_data()
  298. #print(raw.ch_names)
  299. if channel_names is None:
  300. channel_names = raw.ch_names
  301. else:
  302. if channel_names != raw.ch_names:
  303. raise ValueError("Inconsistent channel names across recordings.")
  304. #we don't have the positions for mmi, we wait to set the montage and then extract them
  305. raw = set_montage(signal, "mmi", pos=None, channel_names=channel_names)
  306. if raw is None:
  307. print(f"⚠️ set_montage returned None for {local_path}")
  308. continue
  309. raw = _preprocess_raw(raw, dataset_name="mmi", verbose=verbose)
  310. if raw is None:
  311. print(f"⚠️ _preprocess_raw returned None for {local_path}")
  312. continue
  313. position = get_electrode_positions(raw, channel_order=None)
  314. if position is None:
  315. print(f"⚠️ No valid electrode positions for {local_path}")
  316. continue
  317. positions.append(position)
  318. data = raw.get_data()
  319. data = data*1e3 # Convert to µV
  320. data = data.astype(np.float32)
  321. datas.append(data)
  322. label = map_runs_mmi[run]
  323. labels.append(map_labels_mmi[label])
  324. except Exception as e:
  325. print(f"⚠️ Exception during processing {local_path}: {e}")
  326. import traceback
  327. traceback.print_exc() # This will show the exact line causing the error
  328. continue
  329. # Clean up downloaded files
  330. path_to_remove = os.path.join(
  331. os.path.dirname(project_path), "mmi",
  332. 'MNE-eegbci-data', 'files', 'eegmmidb'
  333. )
  334. if os.path.exists(path_to_remove):
  335. shutil.rmtree(path_to_remove)
  336. labels = torch.tensor(np.array(labels), dtype=torch.int32)
  337. positions = np.array(positions)
  338. positions = torch.tensor(positions, dtype=torch.float32)
  339. return datas, labels, positions, channel_names
  340. def load_mat(filepath):
  341. mat = scipy.io.loadmat(filepath)
  342. return mat
  343. def load_seed_data(subject_ids, project_path, demo=False, verbose=False):
  344. seed_datapath = os.path.join(project_path, "Preprocessed_EEG")
  345. files = os.listdir(seed_datapath)
  346. if len(files) == 0:
  347. raise ValueError("No SEED data found in the specified path.")
  348. datas = []
  349. labels = []
  350. positions = []
  351. channels, position = load_seed_channel_positions(os.path.join(project_path, "channel_62_pos.locs"))
  352. #print(position)
  353. i = 0
  354. for file in files:
  355. for subject in subject_ids:
  356. print(f"Processing subject {i+1}/{len(subject_ids)}", end='\r')
  357. if demo:
  358. if i == 1:
  359. break # For demo, process only first subject
  360. if file.startswith(f'{int(subject)}_'):
  361. filepath = os.path.join(seed_datapath, file)
  362. try:
  363. data_hr = load_mat(filepath)
  364. sfreq = 200 # Original SEED sampling rate
  365. eeg_keys = [key for key in data_hr.keys() if "eeg" in key.lower()]
  366. if len(eeg_keys) == 0:
  367. print(f"No EEG data found in {filepath}.")
  368. continue
  369. else:
  370. for key in eeg_keys:
  371. data = data_hr[key] # Shape: (channels, samples)
  372. #raw = mne.io.RawArray(data, mne.create_info(ch_names=channels, sfreq=sfreq, ch_types='eeg'))
  373. raw = set_montage(data, "seed", pos=position, channel_names=channels, fs=sfreq)
  374. raw = _preprocess_raw(raw, dataset_name="seed", verbose=verbose)
  375. if raw is None:
  376. continue
  377. data = raw.get_data()
  378. data = data.astype(np.float32)
  379. datas.append(data)
  380. label = map_labels_seed[labels_subject_mapping[subject]]
  381. labels.append(label)
  382. positions.append(torch.tensor(position, dtype=torch.float32))
  383. i += 1
  384. except Exception as e:
  385. print(f"⚠️ Exception during processing {filepath}: {e}")
  386. positions = np.array(positions)
  387. positions = torch.tensor(positions, dtype=torch.float32)
  388. labels = np.array(labels, dtype=np.int32)
  389. labels = torch.tensor(labels)
  390. return datas, labels, positions, channels
  391. def download_eegbci_data(subject_ids, runs, project_path, demo=False, dataset_name="mmi", is_classification=False, verbose=False):
  392. """
  393. Scarica i dati EEG BCI per i soggetti e le sessioni specificate.
  394. Salva i file EDF localmente in project_path.
  395. """
  396. if dataset_name.lower() not in ["mmi", "seed"]:
  397. raise ValueError("Dataset not supported. Use 'mmi' or 'seed'.")
  398. if dataset_name.lower() == "mmi":
  399. return download_mmi_data(subject_ids, runs, project_path, demo=demo, verbose=verbose, is_classification=is_classification)
  400. else:
  401. return load_seed_data(subject_ids, project_path, demo=demo, verbose=verbose)
  402. def clear_memory():
  403. gc.collect()
  404. torch.cuda.empty_cache()
  405. def set_seed(seed):
  406. torch.manual_seed(seed)
  407. np.random.seed(seed)
  408. if torch.cuda.is_available():
  409. torch.cuda.manual_seed_all(seed)
  410. class EEGWindowsDataset(Dataset):
  411. """
  412. Minimal dataset for pre-processed, pre-windowed EEG data.
  413. Assumes windows are already normalized and ready to use.
  414. Only generates LR on-the-fly during __getitem__.
  415. """
  416. def __init__(self, windows, labels, positions, sr_type="temporal", dataset_name="mmi",
  417. target_channels=64, multiplier=2, channel_names=None, fs_hr=160):
  418. """
  419. Args:
  420. windows: Tensor/array (N, C, T) - preprocessed HR windows
  421. labels: Tensor/array (N,) - class labels
  422. positions: Tensor/array (N, C, 3) or (C, 3) - electrode positions
  423. sr_type: 'temporal' or 'spatial'
  424. num_channels: int - for spatial SR (number of LR channels)
  425. multiplier: int - temporal downsampling factor
  426. channel_names: list - optional channel names
  427. """
  428. self.datas_hr = torch.tensor(windows, dtype=torch.float32) if not isinstance(windows, torch.Tensor) else windows.float()
  429. self.labels = torch.tensor(labels, dtype=torch.long) if not isinstance(labels, torch.Tensor) else labels.long()
  430. # Handle positions
  431. if isinstance(positions, torch.Tensor):
  432. self.positions = positions.float()
  433. else:
  434. self.positions = torch.tensor(positions, dtype=torch.float32)
  435. # Broadcast if single position (C, 3) -> (N, C, 3)
  436. if self.positions.dim() == 2:
  437. self.positions = self.positions.unsqueeze(0).expand(len(self.datas_hr), -1, -1)
  438. self.sr_type = sr_type
  439. self.target_channels = target_channels
  440. self.multiplier = multiplier
  441. self.channel_names = channel_names
  442. self.dataset_name = dataset_name
  443. self.num_classes = len(torch.unique(self.labels))
  444. self.ref_position = self.positions[0]
  445. self.fs_hr = fs_hr
  446. self.fs_lr = int(fs_hr//self.multiplier) if sr_type == "temporal" else fs_hr
  447. self.num_channels = math.ceil(self.target_channels / self.multiplier) if sr_type == "spatial" else self.datas_hr.shape[1]
  448. print(f"✅ EEGWindowsDataset: {len(self.datas_hr)} windows, shape {self.datas_hr.shape}")
  449. def _downsample(self, data, factor):
  450. """Temporal downsampling via signal.resample"""
  451. downsampled_length = data.shape[1] // factor
  452. return signal.resample(data, downsampled_length, axis=1)
  453. def __len__(self):
  454. return len(self.datas_hr)
  455. def __getitem__(self, idx):
  456. hr_data = self.datas_hr[idx] # (C, T)
  457. pos = self.positions[idx] # (C, 3)
  458. label = self.labels[idx] # scalar
  459. # Generate LR on-the-fly
  460. if self.sr_type == "temporal":
  461. lr_data = self._downsample(hr_data.numpy(), self.multiplier)
  462. lr_data = torch.tensor(lr_data, dtype=torch.float32)
  463. else: # spatial
  464. lr_data = get_lr_data_spatial(hr_data.clone(), dataset_name=self.dataset_name, sr_ratio=self.multiplier)
  465. return lr_data, hr_data, pos, label
  466. # -------------------------------
  467. # EEGDataset for Super-Resolution EEG
  468. # -------------------------------
  469. class EEGDataset(Dataset):
  470. """
  471. Dataset EEG preprocessato per Super-Resolution EEG (EEG BCI dataset).
  472. Genera segmenti LR (160/sr_factor Hz) e HR (160 Hz) sincronizzati.
  473. """
  474. def __init__(self, subject_ids, data_folder, dataset_name = "mmi", sr_type="temporal", seconds=10, verbose=False, demo=False, num_channels=64, multiplier=2, normalize=False, is_classification=False):
  475. self.project_path = data_folder
  476. self.verbose = verbose
  477. self.normalize = normalize
  478. self.demo = demo
  479. self.seconds = seconds
  480. self.num_channels = num_channels
  481. self.dataset_name = dataset_name
  482. self.runs = map_runs_dataset[self.dataset_name]
  483. self.sr_type = sr_type # 'temporal' or 'spatial'
  484. self.multiplier = multiplier # Downsampling factor for temporal or spatial SR
  485. self.fs_hr = 160 if self.dataset_name == "mmi" else 200
  486. self.hr_window_length = self.fs_hr * self.seconds
  487. self.is_classification = is_classification
  488. self.scaler = StandardScaler() #MinMaxScaler(feature_range=(0, 1))
  489. self.datas, self.labels, self.positions, self.channel_names = download_eegbci_data(
  490. subject_ids, runs=self.runs, project_path=self.project_path,
  491. demo=self.demo, dataset_name=self.dataset_name, verbose=self.verbose, is_classification=self.is_classification
  492. )
  493. self.positions = self.positions.cpu()
  494. self.labels = self.labels.cpu()
  495. self.num_classes = 6 if self.dataset_name == "mmi" else 3 #len(torch.unique(self.labels))
  496. self.ref_position = self.positions[0] # Assuming all raws have same channel positions
  497. self.datas_hr, self.positions = self._split_windows(self.hr_window_length)
  498. print(f"\n✅ Number of segments created: {len(self.datas_hr)}")
  499. #self.datas_hr = self._zscore_normalization(torch.tensor(self.datas_hr)).numpy()
  500. #print(f"\nData z-score normalization complete: {self.datas_hr.shape}")
  501. if self.normalize:
  502. self.datas_hr = self._normalize_data()
  503. print(f"\nData normalization complete: {self.datas_hr.shape}")
  504. def _zscore_normalization(self, data):
  505. # Z-score per channel (common for EEG DL)
  506. mean = data.mean(dim=-1, keepdim=True) # [B,C,1]
  507. std = data.std(dim=-1, keepdim=True)
  508. std = torch.clamp(std, min=1e-5) # Avoid div0
  509. normalized = (data - mean) / std # [-3,3] typical
  510. return normalized
  511. def _normalize_data(self):
  512. # self.datas_hr shape: (N_windows, C, T)
  513. N, C, T = self.datas_hr.shape
  514. # Flatten to (N*C, T) - treat each channel independently
  515. data_reshaped = self.datas_hr.reshape(N * C, T) # (N*C, T)
  516. # Fit scaler per channel
  517. data_norm = self.scaler.fit_transform(data_reshaped)
  518. # Reshape back to (N, C, T)
  519. data_normalized = torch.tensor(data_norm, dtype=torch.float32).reshape(N, C, T)
  520. return data_normalized
  521. def _downsample(self, data, factor):
  522. _, num_samples = data.shape
  523. downsampled_length = num_samples // factor
  524. downsampled_data = signal.resample(data, downsampled_length, axis=1)
  525. return downsampled_data
  526. def _split_windows(self, window_length, stride=None): # stride=window_length for non-overlap
  527. if stride is None:
  528. stride = window_length # Non-overlapping by default
  529. datas_hr = []
  530. positions = []
  531. labels = []
  532. for i, data in enumerate(self.datas): # self.datas = list of (C, T_i) arrays
  533. print(f"Splitting data {i+1}/{len(self.datas)} (len={data.shape[1]})", end='\r')
  534. position = self.positions[i] # (C, 3)
  535. T = data.shape[1]
  536. num_windows = (T - window_length) // stride + 1 # Floor div for valid windows
  537. for j in range(num_windows):
  538. start = j * stride
  539. end = start + window_length
  540. window = data[:, start:end].astype(np.float32) # (C, window_length)
  541. datas_hr.append(window)
  542. positions.append(position) # Same pos for all windows from this trial
  543. labels.append(self.labels[i])
  544. datas_hr = torch.tensor(np.array(datas_hr), dtype=torch.float32) # (N_windows_total, C, W)
  545. positions = torch.stack(positions) # (N_windows_total, C, 3)
  546. self.labels = torch.tensor(np.array(labels), dtype=torch.int32).cpu()
  547. return datas_hr, positions
  548. def __len__(self):
  549. return len(self.datas_hr)
  550. def __getitem__(self, idx):
  551. device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
  552. hr_data = self.datas_hr[idx]
  553. pos = self.positions[idx]
  554. label = self.labels[idx]
  555. if self.sr_type == "temporal":
  556. lr_data = self._downsample(hr_data, self.multiplier)
  557. else: # spatial
  558. if isinstance(hr_data, torch.Tensor):
  559. lr_data = hr_data.clone()
  560. else:
  561. lr_data = hr_data.copy() # numpy array
  562. lr_data = get_lr_data_spatial(lr_data, dataset_name=self.dataset_name, sr_ratio=self.multiplier)
  563. lr_data = torch.tensor(lr_data, dtype=torch.float32)
  564. return lr_data.to(device), hr_data.to(device), pos.to(device), label.to(device)
  565. def train_test_val_split_patients(patients, test_size=0.2, val_size=0.1, random_state=42):
  566. train_patients, test_patients = train_test_split(patients, test_size=test_size, random_state=random_state)
  567. relative_val_size = val_size / (1 - test_size)
  568. train_patients, val_patients = train_test_split(train_patients, test_size=relative_val_size, random_state=random_state)
  569. return train_patients, val_patients, test_patients
  570. def train_test_val_split(windows, labels, positions, test_size=0.2, val_size=0.1, random_state=42):
  571. train_data, test_data, train_labels, test_labels, train_positions, test_positions = train_test_split(
  572. windows, labels, positions, test_size=test_size, random_state=random_state, stratify=labels
  573. )
  574. relative_val_size = val_size / (1 - test_size)
  575. train_data, val_data, train_labels, val_labels, train_positions, val_positions = train_test_split(
  576. train_data, train_labels, train_positions, test_size=relative_val_size, random_state=random_state, stratify=train_labels
  577. )
  578. return (train_data, train_labels, train_positions), (val_data, val_labels, val_positions), (test_data, test_labels, test_positions)
  579. def plot_mean_timeseries(timeseries, save_path=None):
  580. fig = plt.figure(figsize=(12, 4))
  581. for key, value in timeseries.items():
  582. if value.ndim == 3:
  583. value = value[0] # Take the first sample in the batch
  584. print(f"Plotting {key} with shape {value.shape}, min {value.min()}, max {value.max()}")
  585. mean_signal = np.mean(value, axis=0) # Mean across channels
  586. plt.plot(mean_signal, label=key)
  587. plt.title(f'Mean Timeseries Across Channels')
  588. plt.xlabel('Time (samples)')
  589. plt.ylabel('Amplitude')
  590. plt.grid()
  591. plt.legend()
  592. if save_path:
  593. fig.savefig(save_path, dpi=300)
  594. print(f"Saved figure to {save_path}")
  595. plt.close(fig)
  596. def clear_directory(path, ignore=[]):
  597. """Remove all files in the specified directory."""
  598. if os.path.exists(path):
  599. for filename in os.listdir(path):
  600. if filename == '.gitignore' or filename in ignore:
  601. continue # Skip .gitignore files and ignored files/folders
  602. file_path = os.path.join(path, filename)
  603. try:
  604. if os.path.isfile(file_path) or os.path.islink(file_path):
  605. os.unlink(file_path)
  606. elif os.path.isdir(file_path):
  607. shutil.rmtree(file_path)
  608. except Exception as e:
  609. print(f'Failed to delete {file_path}. Reason: {e}')
  610. def add_zero_channels(input_tensor, target_channels=64, dataset_name="mmi", multiplier=2):
  611. """Add zero channels to input_tensor to match target_channels."""
  612. if input_tensor.ndim == 2:
  613. batch_size = None
  614. nchs = input_tensor.size(0)
  615. length = input_tensor.size(1)
  616. else:
  617. batch_size = input_tensor.size(0)
  618. nchs = input_tensor.size(1)
  619. length = input_tensor.size(2)
  620. if batch_size is None:
  621. input_target = torch.zeros((target_channels, length), device=input_tensor.device)
  622. else:
  623. input_target = torch.zeros((batch_size, target_channels, length), device=input_tensor.device)
  624. channels_to_use = unmask_channels[dataset_name][f"x{multiplier}"]
  625. for i, ch in enumerate(channels_to_use):
  626. if batch_size is None:
  627. input_target[ch, :] = input_tensor[i, :]
  628. else:
  629. input_target[:, ch, :] = input_tensor[:, i, :]
  630. return input_target
  631. def compute_latents(model, dataloader, device, split = "train", map_labels=None):
  632. """
  633. Computes latent vectors from model for all data in dataloader.
  634. Args:
  635. model: PyTorch model with return_latent=True support
  636. dataloader: yields (eeg_lr, eeg_hr, pos, label)
  637. device: torch device
  638. """
  639. latent_vectors = []
  640. labels = []
  641. print("Computing latents for training set", end='\r')
  642. for i, (eeg_lr, eeg_hr, pos, label) in enumerate(dataloader):
  643. print(f"Processing batch {i+1}/{len(dataloader)}", end='\r')
  644. eeg_lr = eeg_lr.to(device)
  645. eeg_hr = eeg_hr.to(device)
  646. pos = pos.to(device)
  647. label = label.to(device) # In case needed
  648. with torch.no_grad():
  649. if model.__class__.__name__ == "DiBiMa_Diff":
  650. batch_size = eeg_lr.size(0)
  651. t = torch.full((batch_size,), model.train_scheduler.num_train_timesteps - 1,
  652. device=device, dtype=torch.long)
  653. x_t_hr = torch.randn_like(eeg_hr).to(device)
  654. latent = model(x_t_hr, t, lr=eeg_lr, pos=pos, label=label, return_latent=True)[-1]
  655. else:
  656. latent = model(eeg_lr, return_latent=True)[-1]
  657. # Flatten to (B, D)
  658. latent = latent.reshape(latent.size(0), -1).cpu().numpy()
  659. latent_vectors.append(latent)
  660. # Collect labels per sample
  661. for l in label.flatten():
  662. labels.append(map_labels[l.item()] if map_labels else l.item())
  663. # Memory management
  664. del eeg_lr, eeg_hr, pos, label, latent
  665. torch.cuda.empty_cache()
  666. gc.collect()
  667. return np.vstack(latent_vectors), np.array(labels)
  668. def plot_umap_latent_space(model, dataloader_train, dataloader_test, save_path=None, map_labels=None, seed=42):
  669. """
  670. Plots UMAP projection of model latent space colored by labels.
  671. Args:
  672. model: PyTorch model with return_latent=True support
  673. dataloader_train: yields (eeg_lr, eeg_hr, pos, label)
  674. dataloader_test: yields (eeg_lr, eeg_hr, pos, label)
  675. save_path: Optional save path for figure
  676. map_labels: Optional dict {int: str} for label mapping
  677. seed: Random seed for reproducibility
  678. """
  679. # Collect all latent vectors and labels
  680. latent_vectors_train = []
  681. latent_vectors_test = []
  682. labels_train = []
  683. labels_test = []
  684. model.eval()
  685. device = next(model.parameters()).device
  686. print("Computing latents for training set...")
  687. latent_vectors_train, labels_train = compute_latents(model, dataloader_train, device, split="train", map_labels=map_labels)
  688. print("\nComputing latents for test set...")
  689. latent_vectors_test, labels_test = compute_latents(model, dataloader_test, device, split="test", map_labels=map_labels)
  690. print(f"UMAP on {latent_vectors_train.shape[0]} samples, {latent_vectors_train.shape[1]} dims")
  691. # Fit UMAP ONCE on full dataset
  692. reducer = UMAP(n_neighbors=15, min_dist=0.1, metric='euclidean',
  693. random_state=seed, n_jobs=-1)
  694. embedding_train = reducer.fit_transform(latent_vectors_train)
  695. embedding_test = reducer.transform(latent_vectors_test)
  696. # Plot by unique labels
  697. plt.figure(figsize=(12, 8))
  698. u_labels = np.unique(labels_test)
  699. print("Generating scatter plot...")
  700. for ul in u_labels:
  701. mask = labels_test == ul
  702. if np.sum(mask) > 0:
  703. plt.scatter(embedding_test[mask, 0], embedding_test[mask, 1],
  704. label=str(ul), alpha=0.6, s=20)
  705. print(f" {ul}: {np.sum(mask)} points")
  706. plt.title('UMAP Latent Space Projection', fontsize=14)
  707. plt.xlabel('UMAP 1', fontsize=12)
  708. plt.ylabel('UMAP 2', fontsize=12)
  709. plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
  710. plt.grid(True, alpha=0.3)
  711. plt.tight_layout()
  712. if save_path:
  713. plt.savefig(save_path, dpi=300, bbox_inches='tight')
  714. print(f"Saved: {save_path}")
  715. plt.show()
  716. plt.close()
  717. print("UMAP visualization complete!")
  718. def tensor2raw(eeg_tensor, info):
  719. """
  720. Convert a PyTorch tensor back to an MNE Raw object.
  721. """
  722. eeg_data = eeg_tensor.cpu().numpy()
  723. info = mne.create_info(ch_names=info['ch_names'], sfreq=info['sfreq'], ch_types='eeg')
  724. raw = mne.io.RawArray(eeg_data, info)
  725. return raw
  726. import torch
  727. import torch.nn as nn
  728. import torch.nn.functional as F
  729. class ReconstructionLoss(nn.Module):
  730. def __init__(
  731. self,
  732. lambda_spectral: float = 0.1,
  733. lambda_l2: float = 1e-6, # Reduced
  734. freq_bands: dict = None
  735. ):
  736. super().__init__()
  737. self.lambda_spectral = lambda_spectral
  738. self.lambda_l2 = lambda_l2
  739. # EEG frequency bands (Hz)
  740. self.freq_bands = freq_bands or {
  741. 'delta': (0.5, 4),
  742. 'theta': (4, 8),
  743. 'alpha': (8, 13),
  744. 'beta': (13, 30),
  745. 'gamma': (30, 45)
  746. }
  747. def spectral_loss(self, pred, target, sample_rate=250):
  748. """Preserve spectral power in EEG bands"""
  749. pred_fft = torch.fft.rfft(pred, dim=-1)
  750. target_fft = torch.fft.rfft(target, dim=-1)
  751. # Power spectral density
  752. pred_psd = torch.abs(pred_fft) ** 2
  753. target_psd = torch.abs(target_fft) ** 2
  754. return F.mse_loss(pred_psd, target_psd)
  755. def forward(self, pred, target, model: nn.Module):
  756. # Base reconstruction
  757. mse_loss = F.mse_loss(pred, target)
  758. # Spectral preservation
  759. spectral_loss = self.spectral_loss(pred, target)
  760. # L2 only on convolutional/linear weights, exclude biases/norms
  761. l2_reg = sum(
  762. torch.sum(p ** 2)
  763. for name, p in model.named_parameters()
  764. if 'weight' in name and len(p.shape) >= 2
  765. )
  766. total_loss = (
  767. mse_loss
  768. + self.lambda_spectral * spectral_loss
  769. + self.lambda_l2 * l2_reg
  770. )
  771. return total_loss
  772. def random_matplotlib_color():
  773. """Generate single random RGB color (0-1 range)."""
  774. return tuple(random.random() for _ in range(3))
  775. def generate_colors(n_colors: int = 1, method: str = 'hsv_uniform'):
  776. """
  777. Generate visually distinct colors for 64-channel signals (ECG/PPG).
  778. Args:
  779. n_colors: Number of colors (default 1)
  780. method: 'hsv_uniform' (recommended), 'random', or 'tab20'
  781. Returns:
  782. List of (R,G,B) tuples, matplotlib-ready
  783. """
  784. if method == 'tab20':
  785. # Matplotlib's built-in qualitative colormap (20 colors, repeats)
  786. cmap = plt.cm.tab20(np.linspace(0, 1, n_colors))
  787. return [tuple(color[:3]) for color in cmap]
  788. elif method == 'random':
  789. # Pure random (may have clashes)
  790. return [random_matplotlib_color() for _ in range(n_colors)]
  791. else: # 'hsv_uniform' - best perceptual separation
  792. colors = []
  793. for i in range(n_colors):
  794. hue = i / n_colors
  795. sat = np.clip(0.7 + 0.2 * np.sin(i * np.pi / 8), 0.6, 1.0)
  796. val = np.clip(0.9 + 0.05 * np.sin(i * np.pi / 4), 0.85, 1.0)
  797. color = plt.cm.hsv(hue)[:3]
  798. color = tuple(np.array(color) * np.array([1, sat, val]))
  799. colors.append(color)
  800. return colors
  801. class EEGSuperResolutionLoss(nn.Module):
  802. def __init__(self,
  803. lambda_grad=1.0,
  804. lambda_corr=0.5,
  805. lambda_freq=0.3,
  806. use_freq_loss=True):
  807. """
  808. Composite loss for EEG temporal super-resolution
  809. Args:
  810. lambda_grad: Weight for gradient loss (controls sharp transitions)
  811. lambda_corr: Weight for Pearson correlation loss (temporal coherence)
  812. lambda_freq: Weight for frequency domain loss (spectral fidelity)
  813. use_freq_loss: Whether to include frequency domain loss
  814. """
  815. super().__init__()
  816. self.lambda_grad = lambda_grad
  817. self.lambda_corr = lambda_corr
  818. self.lambda_freq = lambda_freq
  819. self.use_freq_loss = use_freq_loss
  820. self.mse = nn.MSELoss()
  821. self.__name__ = "EEGSuperResolutionLoss"
  822. def gradient_loss(self, pred, target):
  823. """
  824. Gradient difference loss - captures rate of change
  825. Helps preserve sharp transitions
  826. """
  827. # Temporal gradient (first derivative along time axis)
  828. pred_grad = pred[:, :, 1:] - pred[:, :, :-1]
  829. target_grad = target[:, :, 1:] - target[:, :, :-1]
  830. return F.mse_loss(pred_grad, target_grad)
  831. def pearson_correlation_loss(self, pred, target):
  832. """
  833. Pearson correlation coefficient loss
  834. Preserves temporal waveform patterns
  835. """
  836. # Flatten spatial dimensions, keep batch and time
  837. pred_flat = pred.reshape(pred.size(0), -1)
  838. target_flat = target.reshape(target.size(0), -1)
  839. # Center the data
  840. pred_centered = pred_flat - pred_flat.mean(dim=1, keepdim=True)
  841. target_centered = target_flat - target_flat.mean(dim=1, keepdim=True)
  842. # Covariance
  843. cov = (pred_centered * target_centered).sum(dim=1)
  844. # Standard deviations
  845. pred_std = torch.sqrt((pred_centered ** 2).sum(dim=1) + 1e-8)
  846. target_std = torch.sqrt((target_centered ** 2).sum(dim=1) + 1e-8)
  847. # Pearson correlation coefficient
  848. pcc = cov / (pred_std * target_std + 1e-8)
  849. # Loss is 1 - PCC (want to maximize correlation)
  850. return (1 - pcc).mean()
  851. def frequency_domain_loss(self, pred, target):
  852. """
  853. FFT-based frequency domain loss
  854. Preserves spectral characteristics
  855. """
  856. # Apply FFT along time dimension
  857. pred_fft = torch.fft.rfft(pred, dim=-1)
  858. target_fft = torch.fft.rfft(target, dim=-1)
  859. # Compute magnitude spectrum loss
  860. pred_mag = torch.abs(pred_fft)
  861. target_mag = torch.abs(target_fft)
  862. mag_loss = F.mse_loss(pred_mag, target_mag)
  863. # Compute phase loss (optional, helps with temporal alignment)
  864. pred_phase = torch.angle(pred_fft)
  865. target_phase = torch.angle(target_fft)
  866. phase_loss = F.mse_loss(pred_phase, target_phase)
  867. return mag_loss + 0.1 * phase_loss
  868. def forward(self, pred, target):
  869. """
  870. Combined loss function
  871. Args:
  872. pred: Predicted HR EEG (B, C, T)
  873. target: Ground truth HR EEG (B, C, T)
  874. """
  875. # Base MSE loss
  876. loss_mse = self.mse(pred, target)
  877. # Gradient loss for sharp transitions
  878. loss_grad = self.gradient_loss(pred, target)
  879. # Pearson correlation for temporal coherence
  880. loss_corr = self.pearson_correlation_loss(pred, target)
  881. # Frequency domain loss
  882. if self.use_freq_loss:
  883. loss_freq = self.frequency_domain_loss(pred, target)
  884. else:
  885. loss_freq = 0.0
  886. # Combined loss
  887. total_loss = (loss_mse +
  888. self.lambda_grad * loss_grad +
  889. self.lambda_corr * loss_corr)
  890. if self.use_freq_loss:
  891. total_loss += self.lambda_freq * loss_freq
  892. # Return total and individual losses for monitoring
  893. return total_loss, {
  894. 'mse': loss_mse.item(),
  895. 'gradient': loss_grad.item(),
  896. 'correlation': loss_corr.item(),
  897. 'frequency': loss_freq.item() if self.use_freq_loss else 0.0,
  898. 'total': total_loss.item()
  899. }
  900. def detect_and_clean_seed_trial(raw, reject_threshold=0.16):
  901. """Version-safe bad channel detection for SEED"""
  902. # Basic stats without montage dependency
  903. data = raw.get_data()
  904. # Flat channels: variance near zero
  905. variances = data.var(axis=1)
  906. flat_idx = np.where(variances < 1e-12)[0]
  907. # Noisy outliers: z-score variance >5
  908. var_z = (variances - variances.mean()) / (variances.std() + 1e-8)
  909. noisy_idx = np.where(var_z > 5)[0]
  910. # Saturated: low first differences (repeats)
  911. diffs = np.abs(np.diff(data, axis=1)).mean(axis=1)
  912. sat_idx = np.where(diffs < 1e-10)[0]
  913. all_bads = list(set(list(flat_idx) + list(noisy_idx) + list(sat_idx)))
  914. raw.info['bads'] = [raw.ch_names[i] for i in all_bads]
  915. n_bad_frac = len(all_bads) / raw.info['nchan']
  916. #print(f"Trial bad fraction: {n_bad_frac:.1%} ({len(all_bads)}/{raw.info['nchan']})")
  917. if n_bad_frac > reject_threshold:
  918. return None
  919. # Interpolate (works without montage if no spatial ops needed)
  920. try:
  921. raw.interpolate_bads(reset_bads=False)
  922. except Exception as e:
  923. print(f"Interpolation failed (ok for DL): {e}")
  924. # Still return - bads marked for your model to ignore
  925. return raw
  926. if __name__ == "__main__":
  927. for dataset_name in ["mmi", "seed"]:
  928. sr_type = "spatial"
  929. for sr_ratio in ["x2", "x4", "x8"]:
  930. channels = unmask_channels[dataset_name][sr_ratio]
  931. print(f"{dataset_name} {sr_type} {sr_ratio}: {len(channels)} channels -> {channels}")

utils.py at commit 8a666bd, under MIT · at the source

Overview

Authors: Ugo Lomoio1,2, Pietro Lió3, Pietro Hiram Guzzi1, Pierangelo Veltri2
  1. Department of Surgical and Medical Sciences, Magna Graecia University, Catanzaro 88100, Italy
  2. DIMES, University of Calabria, Rende 87036, Italy
  3. Computer Science and Technology, University of Cambridge, Cambridge CB2 1TN, United Kingdom
Institutions: Magna Graecia University (Italy); University of Calabria (Italy); University of Cambridge (United Kingdom)
Journal: Bioinformatics (Oxford, England), volume 42, issue 5, article btag169
Dates: received 15 February 2026; accepted 25 March 2026; published online 15 April 2026; in print May 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1093/bioinformatics/btag169 · PMID 41984820 · PMCID PMC13143424 · OpenAlex W7154494588
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: EEG (modality), human (organism), methods / tools (subfield)
Methods: Spectral & time-frequency, Preprocessing, Connectivity, Machine learning
MeSH: Electroencephalography*, Signal Processing, Computer-Assisted*, Algorithms, Brain, Humans (* major topic)
Topic: Functional Brain Connectivity Studies (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Citations: cited by 1 paper (Europe PMC); 22 references in the paper

Abstract

Motivations: Electroencephalography (EEG) is a non-invasive method that records brain electrical activity from scalp electrodes, offering millisecond temporal resolution but limited spatial detail due to sparse sensor layouts.

Results: We present DiBiMa-EEGSR, a bidirectional Mamba-2 diffusion framework for spatio-temporal EEG super-resolution that reconstructs high-resolution signals from standard low-density recordings without additional hardware. The method formulates super-resolution as conditional generative inference and integrates a diffusion process with a bidirectional state-space backbone to model long-range temporal dependencies with linear complexity. Conditioning on low-resolution inputs, electrode positions and task labels enables anatomically coherent and context-aware reconstruction. A one-step sampling strategy substantially reduces inference time while preserving fidelity. Across two public benchmarks, the approach improves reconstruction accuracy, spatial coherence and spectral preservation over convolutional, transformer-based and prior diffusion models in both spatial and temporal upsampling tasks, providing a scalable pathway toward high-resolution electrophysiological imaging.

Availability and implementation: Code to reproduce ablation experiments, training and evaluation of the proposed BiMa and DiBiMa EEGSR models are available at https://github.com/UgoLomoio/DiBiMa-EEGSR.git. Model weights are available at https://huggingface.co/Ugo96/DiBiMa-EEGSR while an interactive demo for EEG spatial super-resolution using our models can be found at https://huggingface.co/spaces/Ugo96/DiBiMa-EEGSR-Demo.

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

Repositories

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

UgoLomoio/DiBiMa-EEGSR

License: MIT
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 8a666bdf5d8af48d8cf7b657f2e7afa2d03f9093, 18 February 2026
Languages: Python (26), Jupyter (1)
Size: 140 files, 27 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, environment (requirements.txt), 1 notebook
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (25 files), NumPy (14 files), scikit-learn (13 files), Matplotlib (12 files), MNE-Python (12 files), pandas (10 files), PyTorch Lightning (9 files), SciPy (8 files), Plotly (2 files), UMAP (2 files)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
29 files

huggingface.co/ugo96/dibima-eegsr

License: MIT
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 59fb893083e774ce0ab7c13a943710520313526e, 10 February 2026
Size: 26 files, 0 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
1 file

huggingface.co/spaces/ugo96/dibima-eegsr-demo

License: MIT
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 1be1ea39e18f068e49b362bb283857d6b994684f, 11 February 2026
Languages: Python (92), CUDA (9), C/C++ (3), C++ (1)
Size: 163 files, 105 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, environment (Dockerfile, requirements.txt, mamba_install/setup.py, maser/requirements.txt), tests
Not found: license file, CITATION.cff, continuous integration, documentation
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers

Availability and implementation

Code to reproduce ablation experiments, training and evaluation of the proposed BiMa and DiBiMa EEGSR models are available at https://github.com/UgoLomoio/DiBiMa-EEGSR.git. Model weights are available at https://huggingface.co/Ugo96/DiBiMa-EEGSR while an interactive demo for EEG spatial super-resolution using our models can be found at https://huggingface.co/spaces/Ugo96/DiBiMa-EEGSR-Demo.

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

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:

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

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

Data

Datasets cited

Data availability

Physionet M/MI is an open-source dataset freely available at: https://physionet.org/content/eegmmidb/1.0.0/. SEED dataset is available upon request at: https://bcmi.sjtu.edu.cn/home/seed/downloads.html

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

Versions

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

Version 1, 28 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 4 authors, 5 MeSH terms, 17 references.

Cite

This paper

Lomoio, U., Lió, P., Guzzi, P. H., & Veltri, P. (2026). Bidirectional Mamba-2 boosts EEG super-resolution via regression and diffusion. Bioinformatics (Oxford, England), 42(5), btag169. https://doi.org/10.1093/bioinformatics/btag169

BibTeX

@article{lomoio2026bidirectional,
author = {Lomoio, Ugo and Lió, Pietro and Guzzi, Pietro Hiram and Veltri, Pierangelo},
title = {{Bidirectional Mamba-2 boosts EEG super-resolution via regression and diffusion}},
journal = {Bioinformatics (Oxford, England)},
year = {2026},
month = may,
volume = {42},
number = {5},
pages = {btag169},
publisher = {Oxford University Press},
issn = {1367-4803},
doi = {10.1093/bioinformatics/btag169},
url = {https://doi.org/10.1093/bioinformatics/btag169},
pmid = {41984820},
pmcid = {PMC13143424}
}

RIS

TY - JOUR
AU - Lomoio, Ugo
AU - Lió, Pietro
AU - Guzzi, Pietro Hiram
AU - Veltri, Pierangelo
TI - Bidirectional Mamba-2 boosts EEG super-resolution via regression and diffusion
T2 - Bioinformatics (Oxford, England)
J2 - Bioinformatics
PY - 2026
DA - 2026/05/01
VL - 42
IS - 5
SP - btag169
SN - 1367-4803
PB - Oxford University Press
DO - 10.1093/bioinformatics/btag169
UR - https://doi.org/10.1093/bioinformatics/btag169
LA - en
ER -

CSL-JSON

{
"id": "10.1093/bioinformatics/btag169",
"type": "article-journal",
"title": "Bidirectional Mamba-2 boosts EEG super-resolution via regression and diffusion",
"container-title": "Bioinformatics (Oxford, England)",
"author": [
{
"family": "Lomoio",
"given": "Ugo"
},
{
"family": "Lió",
"given": "Pietro"
},
{
"family": "Guzzi",
"given": "Pietro Hiram"
},
{
"family": "Veltri",
"given": "Pierangelo"
}
],
"container-title-short": "Bioinformatics",
"volume": "42",
"issue": "5",
"page": "btag169",
"DOI": "10.1093/bioinformatics/btag169",
"PMID": "41984820",
"PMCID": "PMC13143424",
"ISSN": "1367-4803",
"publisher": "Oxford University Press",
"URL": "https://doi.org/10.1093/bioinformatics/btag169",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
1
]
]
}
}

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.1007/s10916-026-02374-5 [code]
Attention-Enhanced U-Net for Sensor-Efficient High-Density EEG Reconstruction in Wearable Brain Monitoring Systems.
Journal: Journal of medical systems
In common: MNE-Python, PyTorch, scikit-learn, 3 other tools, methods / tools, EEG, 5 references
[2] doi:10.1038/s41598-026-56070-y [code]
SSDLabeler: realistic semi-synthetic data generation for multi-label artifact classification in EEG.
Journal: Scientific reports
In common: PyTorch Lightning, PyTorch, scikit-learn, 3 other tools, physionet.org/content/eegmmidb, EEG
[3] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: PyTorch Lightning, UMAP, Plotly, 6 other tools
[4] doi:10.1038/s41467-026-73996-z [code]
Genetic architecture of white matter microstructure captured by unsupervised deep representation learning of fractional anisotropy maps.
Journal: Nature communications
In common: PyTorch Lightning, UMAP, Plotly, 6 other tools
[5] doi:10.1523/eneuro.0023-26.2026 [code]
Real-Time Segmentation and Classification of Birdsong Syllables for Learning Experiments.
Journal: eNeuro
In common: PyTorch Lightning, UMAP, Plotly, 6 other tools
[6] doi:10.1007/s12021-026-09803-3 [code]
NeuroFusion: A Unified Framework for Generalized Visual Stimulus Decoding from fMRI Across Datasets and Subjects.
Journal: Neuroinformatics
In common: PyTorch Lightning, MNE-Python, PyTorch, 5 other tools, methods / tools
[7] doi:10.1162/imag.a.1207 [code]
Investigating the temporal dynamics and modeling of mid-level feature representations in humans.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: PyTorch Lightning, MNE-Python, PyTorch, 5 other tools, EEG
[8] doi:10.3390/bioengineering13080924 [code]
Deep Learning-Based Temporal Gait Analysis Using a Smartphone IMU in Older Adults with and Without Non-Specific Low Back Pain.
Journal: Bioengineering (Basel, Switzerland)
In common: PyTorch Lightning, Plotly, PyTorch, 5 other tools, methods / tools
[9] doi:10.1038/s41597-025-05174-7 [code]
A large-scale MEG and EEG dataset for object recognition in naturalistic scenes
Journal: n/a
In common: MNE-Python, Plotly, PyTorch, 5 other tools, methods / tools, EEG
[10] doi:10.1038/s41467-026-76939-w [code]
HIPPIE: a generative model for electrophysiological analysis across species, technologies, and modalities.
Journal: Nature communications
In common: PyTorch Lightning, UMAP, PyTorch, 5 other tools, methods / tools

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.