OSCR

Explainable AI Insights Into EEG Classification and Its Alignment to Neural Correlates.

Code ↔ Paper

12 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 12 matches · 2 of them tie a paragraph to a whole file, not to given lines: weak matches, whose lines are not tinted
  1. [1] § Making Neuroscientific Sense of Relevance Distributions › Signal Representations ↔ relevance/results.py, lines 1066–1099 · score 0.78 · left parietal, left frontal, right frontal, right temporal, right occipital, bands
  2. [2] § Making Neuroscientific Sense of Relevance Distributions › Signal Representations ↔ main_relevance.ipynb, lines 802–909 · score 0.78 · left parietal, left frontal, right frontal, right temporal, right occipital, bands
  3. [3] § Making Neuroscientific Sense of Relevance Distributions › Functional Grouping ↔ main_relevance.ipynb, lines 802–909 · score 0.73 · right frontal, right temporal, right occipital, brain region, channels
  4. [4] § Making Neuroscientific Sense of Relevance Distributions › Independent Component Analysis ↔ relevance/results.py, lines 1066–1099 · score 0.69 · left frontal, right parietal, right temporal, occipital, brain, channel
  5. [5] § Classification › Model ↔ training/train.py, lines 82–211 · score 0.67 · cross entropy loss, Adam, dropout, optimizer, validation, batch
  6. [6] § Materials and Methods › Data ↔ main_classification.ipynb, lines 367–458 · score 0.62 · motor imagery, external attention, auditory attention, internal, classification
  7. [7] § Materials and Methods › Preprocessing ↔ data/data_handler_kul.py, the whole file · a weak match · score 0.57 · 1–60 Hz, segmented, resampled, preprocessing, windows, epochs
  8. [8] § Signal Representations › Topographic Maps ↔ relevance/results.py, lines 418–500 · score 0.57 · power spectral density, frequency bands, PSD, 30 Hz, 12 Hz, 60 Hz
  9. [9] § Results ↔ main_classification.ipynb, lines 271–364 · score 0.56 · motor imagery, external attention, auditory attention, internal, classification
  10. [10] § Classification › Model ↔ crp/DFT_LRP/synthetic_example.py, lines 109–161 · score 0.56 · cross entropy loss, optimizer, batch, classification, training, model
  11. [11] § Classification › Leave‐One‐Out Cross‐Validation ↔ training/train.py, lines 25–78 · score 0.50 · cross validation, trained models, seed, CV
  12. [12] § Materials and Methods › Preprocessing ↔ data/data_handler_cho.py, the whole file · a weak match · score 0.50 · 1–60 Hz, resampled, preprocessing, windows, epochs, channels

Paper

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

The paper is loaded when this pane is shown.

The authors' code

Python · 2,111 lines · 77 KB · no license · 3 matches

  1. import os
  2. import mne
  3. import torch
  4. import joblib
  5. import matplotlib
  6. import numpy as np
  7. import matplotlib.pyplot as plt
  8. import relevance.utils as rel_utils
  9. import training.utils as train_utils
  10. import matplotlib.patches as mpatches
  11. from matplotlib.colors import ListedColormap
  12. from copy import copy
  13. from scipy.spatial.distance import pdist
  14. from mne_icalabel import label_components
  15. from matplotlib.ticker import MaxNLocator
  16. import seaborn as sns
  17. import relevance.utils as utils
  18. import matplotlib.pyplot as plt
  19. from training import utils as train_utils
  20. from matplotlib.colors import LinearSegmentedColormap
  21. from matplotlib import gridspec
  22. from mne.filter import filter_data
  23. from matplotlib.colors import LinearSegmentedColormap
  24. class Results:
  25. def __init__(self, base_result_path, CV_params, num_folds, sampling_rate, class_labels, layer, num_filter, select_samples_by_class,
  26. use_which_data, num_selected_samples, ds_name, reversed_classes, select_correct, channel_names,
  27. use_what_for_similarity, testing=False):
  28. self.base_result_path = base_result_path
  29. self.use_what_for_similarity = use_what_for_similarity
  30. self.sampling_rate = sampling_rate
  31. self.class_labels = class_labels
  32. self.layer = layer
  33. self.num_filter = num_filter
  34. self.num_CVs = len(CV_params)
  35. self.CV_params = CV_params
  36. self.num_folds = num_folds
  37. self.select_samples_by_class = select_samples_by_class
  38. self.use_which_data = use_which_data
  39. self.num_selected_samples = num_selected_samples
  40. self.ds_name = ds_name
  41. self.reversed_classes = reversed_classes
  42. self.select_correct = select_correct
  43. self.channel_names = channel_names
  44. self.testing = testing
  45. self.correlation_dict_c0_R = None
  46. self.correlation_dict_c1_R = None
  47. self.correlation_dict_c0_X = None
  48. self.correlation_dict_c1_X = None
  49. self.data = {}
  50. self.cluster_data = {}
  51. self.ica_components = {}
  52. self.cluster_keywords = ["embedding", "indices", "labels", "unique_labels", "unique_labels_plotting", "indices_plotting"]
  53. self.data_keywords = ["X_sim", "X_time_files", "X_freq_files", "R_time_files", "R_freq_files", "y_files"]
  54. for label in self.class_labels:
  55. self.ica_components[label] = {}
  56. for cv_idx in range(self.num_CVs):
  57. self.ica_components[label][cv_idx] = {}
  58. for kw in self.cluster_keywords:
  59. self.cluster_data[kw] = {}
  60. for label in self.class_labels:
  61. self.cluster_data[kw][label] = {}
  62. for cv_idx in range(self.num_CVs):
  63. self.cluster_data[kw][label][cv_idx] = []
  64. for kw in self.data_keywords:
  65. self.data[kw] = {}
  66. for label in self.class_labels:
  67. self.data[kw][label] = {}
  68. for cv_idx in range(self.num_CVs):
  69. self.data[kw][label][cv_idx] = {}
  70. for fold_idx in range(self.num_folds):
  71. self.data[kw][label][cv_idx][fold_idx] = []
  72. self._prepare_paths()
  73. self._collect_data()
  74. # ---------------------------------------------------------------------
  75. def set_data_class_cv_fold(self, data, data_identifier, class_label, cv_idx, fold_idx):
  76. self.data[data_identifier][class_label][cv_idx][fold_idx] = data
  77. # ---------------------------------------------------------------------
  78. def get_data_class_cv_fold(self, data_identifier, class_label, cv_idx, fold_idx):
  79. return self.data[data_identifier][class_label][cv_idx][fold_idx]
  80. # ---------------------------------------------------------------------
  81. def set_data_class_cv(self, data, data_identifier, class_label, cv_idx):
  82. if data_identifier in self.cluster_keywords:
  83. self.cluster_data[data_identifier][class_label][cv_idx] = data
  84. else:
  85. self.data[data_identifier][class_label][cv_idx] = data
  86. # ---------------------------------------------------------------------
  87. def get_data_class_cv(self, data_identifier, class_label, cv_idx):
  88. if data_identifier in self.cluster_keywords:
  89. return self.cluster_data[data_identifier][class_label][cv_idx]
  90. else:
  91. folds = []
  92. for fold_idx in range(self.num_folds):
  93. data = self.data[data_identifier][class_label][cv_idx][fold_idx]
  94. folds += data
  95. return folds
  96. # ---------------------------------------------------------------------
  97. def _get_data_of_cluster(self, data_identifier, class_label, cv_idx, cluster_idx):
  98. data = self.get_data_class_cv(data_identifier, class_label, cv_idx)
  99. indices = self.get_data_class_cv("indices", class_label, cv_idx)
  100. unique_labels = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  101. if cluster_idx in unique_labels:
  102. index = unique_labels.index(cluster_idx)
  103. else:
  104. print("not a valid cluster index:", cluster_idx, ", valid: ", unique_labels)
  105. return None
  106. indices = indices[index]
  107. data = [data[i] for i in indices]
  108. return data
  109. # ---------------------------------------------------------------------
  110. def _get_num_filter_per_subject_of_cluster(self, class_label, cv_idx, cluster_idx):
  111. indices = self.get_data_class_cv("indices", class_label, cv_idx)
  112. unique_labels = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  113. if cluster_idx in unique_labels:
  114. index = unique_labels.index(cluster_idx)
  115. else:
  116. print("not a valid cluster index:", cluster_idx, ", valid: ", unique_labels)
  117. return None
  118. indices = indices[index]
  119. filter_person_relation = []
  120. for fold_idx in range(self.num_folds):
  121. for _ in range(self.num_filter):
  122. filter_person_relation.append(fold_idx)
  123. filter_person_relation = [filter_person_relation[i] for i in indices]
  124. num_filter_per_subject = [0 for _ in range(self.num_folds)]
  125. for p in filter_person_relation:
  126. num_filter_per_subject[p] += 1
  127. return num_filter_per_subject, filter_person_relation
  128. # ---------------------------------------------------------------------
  129. def get_samples_of_cluster(self, class_label, cv_idx, cluster_idx, time_domain=True, relevance=False):
  130. data_identifier = "_time_files" if time_domain else "_freq_files"
  131. data_identifier = "R" + data_identifier if relevance else "X" + data_identifier
  132. filenames_X = self._get_data_of_cluster(data_identifier, class_label, cv_idx, cluster_idx)
  133. filenames_y = self._get_data_of_cluster("y_files", class_label, cv_idx, cluster_idx)
  134. all_samples = []
  135. all_labels = []
  136. for fn_X, fn_y in zip(filenames_X, filenames_y):
  137. samples = joblib.load(fn_X)
  138. labels = joblib.load(fn_y)
  139. all_samples.append(samples)
  140. all_labels += list(labels)
  141. all_samples = np.concatenate(all_samples, axis=0)
  142. return all_samples, all_labels
  143. # ---------------------------------------------------------------------
  144. def compute_correlations(self, CV_params, calc_correlations_new=True, discard_negative=True, use_relevance=True, use_X=True, plot=False):
  145. if use_relevance:
  146. if self.correlation_dict_c0_R is None or calc_correlations_new:
  147. print("Compute correlations for class 0")
  148. self.correlation_dict_c0_R = self._calc_correlations(class_label=0, CV_params=CV_params, plot=plot, discard_negative=discard_negative, use_relevance=True)
  149. if self.correlation_dict_c1_R is None or calc_correlations_new:
  150. print("Compute correlations for class 1")
  151. self.correlation_dict_c1_R = self._calc_correlations(class_label=1, CV_params=CV_params, plot=plot, discard_negative=discard_negative, use_relevance=True)
  152. if use_X:
  153. if self.correlation_dict_c0_X is None or calc_correlations_new:
  154. print("Compute correlations for class 0")
  155. self.correlation_dict_c0_X = self._calc_correlations(class_label=0, CV_params=CV_params, plot=plot, discard_negative=discard_negative, use_relevance=False)
  156. if self.correlation_dict_c1_X is None or calc_correlations_new:
  157. print("Compute correlations for class 1")
  158. self.correlation_dict_c1_X = self._calc_correlations(class_label=1, CV_params=CV_params, plot=plot, discard_negative=discard_negative, use_relevance=False)
  159. # ---------------------------------------------------------------------
  160. @staticmethod
  161. def assemble_path(CV_params, base_result_path, num_samples_to_select, ds_name, layer, class_label,
  162. select_samples_by_class, reversed_classes, select_correct,
  163. filter_sim_sample_based, use_which_data, testing=False):
  164. nr = CV_params['nr']
  165. if testing:
  166. testing_str = "_testing"
  167. else:
  168. testing_str = ""
  169. if select_samples_by_class:
  170. sample_select_str = "_samplesOfClass"
  171. else:
  172. sample_select_str = ""
  173. use_which_data = "_" + use_which_data
  174. if filter_sim_sample_based:
  175. filter_sim_sample_based_str = "_sampleBased"
  176. else:
  177. filter_sim_sample_based_str = ""
  178. if select_correct == True:
  179. select_correct_str = "_correct"
  180. elif select_correct == False:
  181. select_correct_str = "_incorrect"
  182. else:
  183. select_correct_str = ""
  184. if reversed_classes:
  185. sample_select_str += "Reversed"
  186. if class_label == "both":
  187. sample_select_str = ""
  188. if "seed" in CV_params:
  189. seed_str = f"_seed{CV_params['seed']}"
  190. else:
  191. seed_str = ""
  192. path_base = os.path.join(base_result_path, f"results_relevance_{num_samples_to_select}{use_which_data}{testing_str}", ds_name, layer)
  193. path_for_cv = os.path.join(path_base, f"class_{class_label}{sample_select_str}" +
  194. f"{select_correct_str}{filter_sim_sample_based_str}{seed_str}_{nr}")
  195. return path_base, path_for_cv
  196. # ---------------------------------------------------------------------
  197. def plot_distribution_of_selected_samples(self):
  198. for cv_idx in range(self.num_CVs):
  199. for class_label in self.class_labels:
  200. suptitle = f"Class {class_label} - CV iteration {cv_idx}"
  201. path = self.paths_to_CV_iteration[class_label][cv_idx]
  202. # path_get_data_dict[class_label][cv_idx]
  203. selected_samples_per_subject = joblib.load(os.path.join(path, "selected_samples_per_subject"))
  204. x_labels = []
  205. num_samples_per_subj = []
  206. num_unique_per_subj = []
  207. for key, value in selected_samples_per_subject.items():
  208. x_labels.append(key)
  209. num_samples_per_subj.append(len(value))
  210. num_unique_per_subj.append(len(set(value)))
  211. width_per_person = 6/10
  212. width = self.num_folds * width_per_person
  213. print(width_per_person)
  214. _, axes = plt.subplots(1,2, figsize=(width,2.5), sharey="row")
  215. for i in range(0, 2):
  216. ax = axes[i]
  217. if i == 0:
  218. values = num_samples_per_subj
  219. title = "Num selected samples"
  220. else:
  221. values = num_unique_per_subj
  222. title = "Num unique samples selected"
  223. ax.bar(x_labels, values)
  224. ax.set_xticks(x_labels)
  225. ax.set_xlabel("Test Subjects")
  226. ax.set_ylabel("Amount of samples")
  227. ax.set_title(title)
  228. plt.suptitle(suptitle, y=1.1)
  229. plt.show()
  230. # ---------------------------------------------------------------------
  231. def plot_single_filter(self):
  232. tmp = joblib.load(os.path.join(self.paths_to_CV_iteration[0][0], "total_rel_per_filter_final"))
  233. num_total = len(tmp)
  234. num_filters_per_model = num_total // self.num_folds
  235. del tmp
  236. folds = [i for i in range(self.num_folds)]
  237. include_cvs = [0]
  238. include_folds = [2]
  239. for class_label in self.class_labels:
  240. print("-"*20, class_label, "-"*20)
  241. for cv in self.CV_params:
  242. cv_nr = cv['nr']
  243. if cv_nr not in include_cvs:
  244. continue
  245. for fold in folds:
  246. if fold not in include_folds:
  247. continue
  248. path_store_data = self.paths_to_CV_iteration[class_label][cv_nr]
  249. titles = []
  250. X_freq_list = []
  251. R_freq_list = []
  252. rel_per_filter = []
  253. for f in range(num_filters_per_model):
  254. X_freq = joblib.load(os.path.join(path_store_data, f"X_freq_select_{fold}_{f}"))
  255. R_freq = joblib.load(os.path.join(path_store_data, f"R_freq_select_{fold}_{f}"))
  256. X_freq = np.abs(X_freq)**2
  257. discard_negative_rel = True
  258. if discard_negative_rel:
  259. R_freq[R_freq < 0] = 0
  260. title = f"CV: {cv['nr']}, Fold: {fold}, Filter {f}"
  261. titles.append(title)
  262. X_freq_list.append(X_freq)
  263. R_freq_list.append(R_freq)
  264. rel_per_filter.append(np.sum(R_freq))
  265. rel_per_filter = np.array(rel_per_filter)
  266. indices = np.flip(np.argsort(rel_per_filter))
  267. rel_per_filter = rel_per_filter[indices]
  268. titles = [titles[i] for i in indices]
  269. x_labels = [str(i) for i in indices]
  270. plt.bar(x_labels, rel_per_filter)
  271. plt.xlabel('Filters')
  272. plt.ylabel('Relevance')
  273. plt.title(f"Class {class_label}, CV {cv['nr']}, Fold {fold}")
  274. plt.show()
  275. for f in range(0, 3): # num_filters_per_model):
  276. X_freq = X_freq_list[f]
  277. R_freq = R_freq_list[f]
  278. title = titles[f]
  279. bands = [(0,4,"delta"), (4,8,"theta"), (8,12,"alpha"), (12,30,"beta"), (30,60,"gamma")]
  280. fig, axes = plt.subplots(2, len(bands), figsize=(6, 4))
  281. _, X_freq = self._plot_freq_bands_topo(X_freq, self.channel_names, bands, axes[0], title=title, global_scale=True, norm=False, log_scale=False, scale_bands=False, cmap="Reds",
  282. label="", colorbar_offset=(0.0,0.5), vlim=None, mean=False, negative_lim=False)
  283. _, R_freq = self._plot_freq_bands_topo(R_freq, self.channel_names, bands, axes[1], title="", global_scale=True, norm=False, log_scale=False, scale_bands=False, cmap="Reds",
  284. label="", colorbar_offset=(0.0,0.2), vlim=None, mean=False, negative_lim=False)
  285. # self._plot_freq_bands_topo_combined(X_freq, R_freq, self.channel_names, bands, axes[2], title="", global_scale=True, norm=False, log_scale=False, scale_bands=False, cmap="Reds",
  286. # label="", colorbar_offset=(0.0,0.2), vlim=None, mean=False, negative_lim=False)
  287. plt.show()
  288. # -------------------------------------------------------------------------------------------------------
  289. def plot_cluster_signal(self, class_label, cv_idx,
  290. title=None, exclude_cluster=[], sig_log_scale=False,
  291. sig_global_scale=True, rel_global_scale=True, psd=True, sig_scale_bands=False, rel_scale_bands=False, sig_norm=True, rel_norm=True,
  292. cluster_order=None, rel_common_scale=False, sig_common_scale=False, sig_mean=True, rel_mean=False, rel_vlim=None, sig_vlim=None, negative_lim=False,
  293. rel_cmap="bwr", remove_neg=False, scale_per_band=False, bands=None, scale_1_div_f=True, plot_ideal_topo=False):
  294. title_was_none = title is None
  295. bands = [(0,4,"delta"), (4,8,"theta"), (8,12,"alpha"), (12,30,"beta"), (30,60,"gamma")]
  296. scale_factors = [] # 2, 6, 10, 21, 45]
  297. first_scalar = (bands[0][0]+bands[0][1])/2
  298. for band in bands:
  299. val = (band[0]+band[1])/(2*first_scalar)
  300. scale_factors.append(val)
  301. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  302. if exclude_cluster is None:
  303. exclude_cluster = []
  304. for clstr in exclude_cluster:
  305. if clstr in cidx_keys:
  306. cidx_keys.remove(clstr)
  307. if cluster_order is not None:
  308. cidx_keys = np.array(cidx_keys)[cluster_order]
  309. for cidx_key in cidx_keys:
  310. print(cidx_key)
  311. X_freq, y = self.get_samples_of_cluster(class_label, cv_idx, cidx_key, time_domain=False)
  312. R_freq, _ = self.get_samples_of_cluster(class_label, cv_idx, cidx_key, time_domain=False, relevance=True)
  313. if psd:
  314. X_freq = X_freq**2
  315. if title_was_none:
  316. if cidx_key == "NC":
  317. title_ = cidx_key
  318. elif cidx_key == "all":
  319. title_ = "All Filters"
  320. else:
  321. title_ = f"Cluster {cidx_key}"
  322. else:
  323. title_ = title
  324. # fig, axes = plt.subplots(3, len(bands), figsize=(12, 6))
  325. if plot_ideal_topo:
  326. fig, axes = plt.subplots(3, len(bands), figsize=(9, 5))
  327. else:
  328. fig, axes = plt.subplots(2, len(bands), figsize=(6, 5))
  329. if sig_log_scale:
  330. label = "log_10(PSD)"
  331. else:
  332. label = "PSD"
  333. colorbar_offset = (0.1, 0)
  334. s_vlim, X_freq = self._plot_freq_bands_topo(X_freq, self.channel_names, ax=axes[0], global_scale=sig_global_scale, log_scale=sig_log_scale,
  335. scale_bands=sig_scale_bands, bands=bands, norm=sig_norm, title=title_, label=label, colorbar_offset=(0.0,0.5),
  336. vlim=sig_vlim, mean=sig_mean, negative_lim=False, cmap="Reds", scale_per_band=scale_per_band, scale_1_div_f=scale_1_div_f,
  337. scale_factors=scale_factors)
  338. colorbar_offset = (0, 0.25)
  339. r_vlim, R_freq = self._plot_freq_bands_topo(R_freq, self.channel_names, ax=axes[1], global_scale=rel_global_scale, log_scale=False,
  340. scale_bands=rel_scale_bands, bands=bands, norm=rel_norm, title=title_, label="Relevance", colorbar_offset=colorbar_offset,
  341. vlim=rel_vlim, mean=rel_mean, negative_lim=negative_lim, cmap=rel_cmap, scale_per_band=False, scale_1_div_f=scale_1_div_f,
  342. scale_factors=scale_factors, remove_neg=remove_neg)
  343. if plot_ideal_topo:
  344. self._plot_freq_bands_topo_combined(X_freq, R_freq, self.channel_names, ax=axes[2], bands=bands, title=title_, label="Ideal", colorbar_offset=(0.0, 0.2))
  345. plt.show()
  346. plt.tight_layout()
  347. # ---------------------------------------------------------------------
  348. def _plot_freq_bands_topo_combined(self, X_freq, R_freq, ch_names, bands, ax, title="", cmap=None, label="", colorbar_offset=[0,0]):
  349. X_all = []
  350. for band_idx in range(len(X_freq)):
  351. ax_band = ax[band_idx]
  352. X_band = np.empty_like(X_freq[band_idx])
  353. for chn in range(X_freq[band_idx].shape[0]):
  354. X_val = X_freq[band_idx][chn]
  355. R_val = R_freq[band_idx][chn]
  356. if R_val > 0:
  357. res = X_val * R_val
  358. else:
  359. res = (1-X_val) * np.abs(R_val)
  360. X_band[chn] = res
  361. X_all.append(X_band)
  362. self._plot_topography(X_band, 120, ch_names, title="", label="", plot=True, sp_size=2, show_ch_names=False, vlim=(0,1), ax=ax_band, cmap=cmap)
  363. ax_band.set_title("")
  364. # if i == len(bands)-1 and global_scale:
  365. # if suptitle:
  366. # cax = plt.axes([1.0 + colorbar_offset[0], 0.0 + colorbar_offset[1], 0.01, 0.2])
  367. # else:
  368. # cax = plt.axes([1.0 + colorbar_offset[0], 0.0 + colorbar_offset[1], 0.01, 0.2])
  369. # plt.colorbar(im, cax=cax)
  370. # ---------------------------------------------------------------------
  371. def plot_cluster_relevance_subplot(self, class_label, cv_idx, ax, exclude_cluster=[], colors=None, show_xlabel=True, bar_width=0.5):
  372. rels = []
  373. abs_rels = []
  374. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  375. for clstr in exclude_cluster:
  376. if clstr in cidx_keys:
  377. cidx_keys.remove(clstr)
  378. for cidx_key in cidx_keys:
  379. rel = self._calc_cluster_relevance(class_label, cv_idx, cidx_key)
  380. rels.append(rel)
  381. abs_rels.append(np.abs(rel))
  382. labels = cidx_keys
  383. x = np.arange(len(labels))
  384. rels = np.array(rels)
  385. abs_rels = np.array(abs_rels)
  386. data = abs_rels
  387. for i in range(len(data)):
  388. color, alpha = colors[i]
  389. ax.bar(x[i], data[i], color=color, alpha=alpha, width=bar_width)
  390. ax.set_xlim(left=-0.5, right=len(x)-0.5)
  391. if show_xlabel:
  392. ax.set_xlabel('Cluster')
  393. ax.set_ylabel('Mean Relevance')
  394. ax.set_xticks(x, labels)
  395. # ---------------------------------------------------------------------
  396. def plot_cluster_coherence_subplot(self, class_label, cv_idx, ax, colors, exclude_cluster=[], show_xlabel=True, bar_width=0.5):
  397. coherence_per_cluster = []
  398. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  399. for clstr in exclude_cluster:
  400. print(clstr, exclude_cluster)
  401. if clstr in cidx_keys:
  402. cidx_keys.remove(clstr)
  403. for cidx_key in cidx_keys:
  404. coh = self._calc_cluster_coherence(class_label, cv_idx, cidx_key)
  405. coherence_per_cluster.append(coh)
  406. labels = cidx_keys
  407. coherence_list = []
  408. for coherence in coherence_per_cluster:
  409. coherence_list.append(coherence)
  410. coherence_list = np.array(coherence_list)
  411. labels = np.array(labels)
  412. x = np.arange(len(labels))
  413. ax.set_xlim(left=-0.5, right=len(x)-0.5)
  414. for idx, (xi, cl) in enumerate(zip(x, coherence_list)):
  415. color, alpha = colors[idx]
  416. ax.bar(xi, cl, width=bar_width, color=color, alpha=alpha)
  417. # ax.bar(x, coherence_list, width=0.3, color='blue', alpha=0.5)
  418. ax.set_ylabel("Coherence")
  419. if show_xlabel:
  420. ax.set_xlabel("Cluster")
  421. ax.set_xticks(x, labels)
  422. # ---------------------------------------------------------------------
  423. def plot_cluster_amount_of_filter_subplot(self, class_label, cv_idx, ax, colors, exclude_cluster=[], show_xlabel=True, bar_width=0.5):
  424. num_filter = []
  425. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  426. for clstr in exclude_cluster:
  427. if clstr in cidx_keys:
  428. cidx_keys.remove(clstr)
  429. rotation_angles = []
  430. for cidx_key in cidx_keys:
  431. rotation_angles.append(0)
  432. nf, _ = self._get_num_filter_per_subject_of_cluster(class_label, cv_idx, cidx_key)
  433. nf = np.sum(nf)
  434. num_filter.append(nf)
  435. num_filter = np.array(num_filter)
  436. labels = cidx_keys
  437. x = np.arange(len(labels))
  438. for i in range(len(num_filter)):
  439. color, alpha = colors[i]
  440. ax.bar(x[i], num_filter[i], color=color, alpha=alpha, width=bar_width)
  441. if show_xlabel:
  442. ax.set_xlabel("Cluster")
  443. ax.set_ylabel("Amount of filters")
  444. ax.set_xticks(x, labels)
  445. # ---------------------------------------------------------------------
  446. def plot_cluster_class_label_distribution(self, class_label, cv_idx, title="", normalize=True, exclude_cluster=[]):
  447. data = []
  448. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  449. for clstr in exclude_cluster:
  450. if clstr in cidx_keys:
  451. cidx_keys.remove(clstr)
  452. for idx in range(len(cidx_keys)):
  453. cidx_key = cidx_keys[idx]
  454. _, y = self.get_samples_of_cluster(class_label, cv_idx, cidx_key, time_domain=True, relevance=False)
  455. if normalize:
  456. norm = len(y)
  457. else:
  458. norm = 1
  459. num_0 = (len(y)-np.sum(y)) / norm
  460. num_1 = np.sum(y) / norm
  461. data.append([num_0, num_1])
  462. data_array = np.array(data)
  463. num_bars = len(data[0])
  464. bar_width = 0.30
  465. bar_distance = 0.25
  466. group_positions = np.arange(len(data))*1.5
  467. fig, ax = plt.subplots(figsize=(3,3))
  468. colors = ["blue", "red"]
  469. for i in range(num_bars):
  470. bar_positions = group_positions + i * (bar_width + bar_distance)
  471. ax.bar(bar_positions, data_array[:, i], bar_width, label=f'Class {i}', color=colors[i])
  472. ax.set_xticks(group_positions + ((num_bars - 1) * (bar_width + bar_distance)) / 2)
  473. cluster_labels = cidx_keys
  474. ax.set_xticklabels(cluster_labels, rotation=0)
  475. # Set labels and title
  476. ax.set_xlabel('Clusters')
  477. if normalize:
  478. ax.set_ylabel('Proportion of class labels')
  479. ax.set_ylim((0, 1))
  480. else:
  481. ax.set_ylabel('Number of labels')
  482. ax.set_title('Distribution of class labels in clusters')
  483. ax.legend()
  484. plt.show()
  485. # ---------------------------------------------------------------------
  486. def plot_cluster_distribution_stacked_subplot(self, class_label, cv_idx, ax, cluster_color_map=None, exclude_cluster=[], show_xticks=True):
  487. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  488. for clstr in exclude_cluster:
  489. if clstr in cidx_keys:
  490. cidx_keys.remove(clstr)
  491. colors = []
  492. list_of_lists = []
  493. legend_labels = []
  494. label_colors = {}
  495. for idx, cluster_idx in enumerate(cidx_keys):
  496. colors.append(cluster_color_map[cluster_idx])
  497. filenames = self._get_data_of_cluster("X_time_files", class_label, cv_idx, cluster_idx)
  498. num_filter_per_fold = [0]*self.num_folds
  499. for fn in filenames:
  500. fold = fn.split("_")[-2]
  501. num_filter_per_fold[int(fold)] += 1
  502. list_of_lists.append(num_filter_per_fold)
  503. if cluster_idx == "NC":
  504. legend_label = cluster_idx
  505. else:
  506. legend_label = f"C{cluster_idx}"
  507. legend_labels.append(legend_label)
  508. label_colors[legend_label] = colors[idx]
  509. x = np.arange(self.num_folds)
  510. labels = [str(xi+1) for xi in x]
  511. bottom = np.zeros(self.num_folds) # Bottom position for each bar, initialized to zeros
  512. for i, list_ in enumerate(list_of_lists):
  513. ax.bar(x, list_, bottom=bottom, color=colors[i], alpha=1, label=legend_labels[i])
  514. bottom += np.array(list_) # Update bottom positions for the next stacked bar
  515. ax.set_xlim(left=-1, right=self.num_folds)
  516. ax.legend(title='')
  517. ax.yaxis.set_major_locator(MaxNLocator(integer=True))
  518. if not show_xticks:
  519. ax.set_xticks([])
  520. ax.set_xlabel("")
  521. else:
  522. ax.set_xticks(x, labels)
  523. ax.set_xlabel("Models")
  524. ax.set_ylabel('#Filters')
  525. # ---------------------------------------------------------------------
  526. def plot_cluster_distribution_stacked(self, class_label, cv_idx, title="", ylim=None, exclude_cluster=[], figsize=(7, 1), cluster_color_map=None, ax=None):
  527. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  528. for clstr in exclude_cluster:
  529. if clstr in cidx_keys:
  530. cidx_keys.remove(clstr)
  531. colors = []
  532. list_of_lists = []
  533. legend_labels = []
  534. label_colors = {}
  535. for idx, cluster_idx in enumerate(cidx_keys):
  536. colors.append(cluster_color_map[cluster_idx])
  537. filenames = self._get_data_of_cluster("X_time_files", class_label, cv_idx, cluster_idx)
  538. num_filter_per_fold = [0]*self.num_folds
  539. for fn in filenames:
  540. fold = fn.split("_")[-2]
  541. num_filter_per_fold[int(fold)] += 1
  542. list_of_lists.append(num_filter_per_fold)
  543. legend_labels.append(f"Cluster {cluster_idx}")
  544. label_colors[f"Cluster {cluster_idx}"] = colors[idx]
  545. x = np.arange(self.num_folds)
  546. labels = [str(xi) for xi in x]
  547. _, ax = plt.subplots(1,1, figsize=figsize)
  548. bottom = np.zeros(self.num_folds) # Bottom position for each bar, initialized to zeros
  549. for i, list_ in enumerate(list_of_lists):
  550. ax.bar(x, list_, bottom=bottom, color=colors[i], alpha=0.5, label=legend_labels[i])
  551. bottom += np.array(list_) # Update bottom positions for the next stacked bar
  552. ax.legend(title='')
  553. ax.yaxis.set_major_locator(MaxNLocator(integer=True))
  554. ax.set_xlabel("Models")
  555. ax.set_xticks(x, labels)
  556. ax.set_title(title)
  557. ax.set_ylabel('Amount of filters')
  558. if ylim is not None:
  559. ax.set_ylim(*ylim)
  560. plt.tight_layout()
  561. plt.show()
  562. # ---------------------------------------------------------------------
  563. def plot_cluster_distribution(self, class_label, cv_idx, title="", ylim=None, exclude_cluster=[], figsize=(7, 3), cluster_color_map=None):
  564. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  565. for clstr in exclude_cluster:
  566. if clstr in cidx_keys:
  567. cidx_keys.remove(clstr)
  568. list_of_lists = []
  569. legend_labels = []
  570. label_colors = {}
  571. colors = []
  572. for idx, cluster_idx in enumerate(cidx_keys):
  573. filenames = self._get_data_of_cluster("X_time_files", class_label, cv_idx, cluster_idx)
  574. num_filter_per_fold = [0]*self.num_folds
  575. color = cluster_color_map[cluster_idx]
  576. for fn in filenames:
  577. fold = fn.split("_")[-2]
  578. num_filter_per_fold[int(fold)] += 1
  579. list_of_lists.append(num_filter_per_fold)
  580. colors.append(color)
  581. if cluster_idx != "NC":
  582. legend_labels.append(f"Cluster {cluster_idx}")
  583. else:
  584. legend_labels.append(f"{cluster_idx}")
  585. label_colors[f"Cluster {cluster_idx}"] = color
  586. # -----------------------------------------------------
  587. def create_colormap(color):
  588. cmap = LinearSegmentedColormap.from_list('custom_cmap', [(1, 1, 1), color])
  589. return cmap
  590. values = np.array(list_of_lists)
  591. # Create the colormaps for each row
  592. cmaps = [create_colormap(color) for color in colors]
  593. # Plot the values for each row using the corresponding colormap
  594. fig, axs = plt.subplots(len(values), 1, figsize=(12, len(values)*0.5), gridspec_kw = {'wspace':0, 'hspace':0})
  595. for i, (row_values, cmap) in enumerate(zip(values, cmaps)):
  596. heatmap = axs[i].imshow([row_values], cmap=cmap, aspect='auto')
  597. if i == len(values)-1:
  598. custom_labels = [str(v+1) for v in range(len(row_values))]
  599. axs[i].set_xticks(np.arange(len(row_values)))
  600. axs[i].set_xticklabels(custom_labels)
  601. else:
  602. axs[i].set_xticks([])
  603. axs[i].set_yticks([])
  604. axs[i].grid(False)
  605. # axs[i].set_title(f'Row {i+1}')
  606. # axs[i].set_yticks([])
  607. # axs[i].set_xticks(np.arange(len(row_values)))
  608. axs[i].set_ylabel(cidx_keys[i], rotation=0, labelpad=20, y=0.4)
  609. for y in range(row_values.shape[0]):
  610. val = row_values[y]
  611. if val == 0:
  612. axs[i].text(y, 0, '0', color='red', ha='center', va='center')
  613. else:
  614. if val <= 3:
  615. axs[i].text(y, 0, f'{val}', color='gray', ha='center', va='center')
  616. plt.subplots_adjust(wspace=None, hspace=None)
  617. plt.tight_layout()
  618. plt.show()
  619. # ---------------------------------------------------------------------
  620. def compute_ica(self, class_label, cv_idx, exclude_cluster=[], overwrite_if_exists=False, band=False):
  621. cluster_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  622. for clstr in exclude_cluster:
  623. if clstr in cluster_keys:
  624. cluster_keys.remove(clstr)
  625. for cluster_idx in cluster_keys:
  626. compute_components = True
  627. cluster_idx_str = cluster_idx
  628. if band != False:
  629. cluster_idx_str = f"{cluster_idx}_{band}"
  630. if cluster_idx_str in self.ica_components[class_label][cv_idx].keys():
  631. if not overwrite_if_exists:
  632. compute_components = False
  633. if compute_components:
  634. ica, labels, proba = self._calc_ica_on_cluster(class_label, cv_idx, cluster_idx, time_domain=True, band=band)
  635. self.ica_components[class_label][cv_idx][cluster_idx_str] = (ica, labels, proba)
  636. # ---------------------------------------------------------------------
  637. def plot_ica(self, class_label, cv_idx, exclude_cluster=[], figsize=(6,6), title="", time=True, band=False, ica_title_fontsize=None):
  638. info = mne.create_info(self.channel_names, self.sampling_rate, ch_types='eeg')
  639. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  640. for clstr in exclude_cluster:
  641. if clstr in cidx_keys:
  642. cidx_keys.remove(clstr)
  643. for cluster_idx in cidx_keys:
  644. print(" "*5, "Cluster:", cluster_idx)
  645. cluster_idx_str = cluster_idx
  646. if not time:
  647. cluster_idx_str = f"{cluster_idx_str}_freq"
  648. if band is not None and band is not False:
  649. cluster_idx_str = f"{cluster_idx_str}_{band}"
  650. try:
  651. ica, labels, proba = self.ica_components[class_label][cv_idx][cluster_idx_str]
  652. except Exception as e:
  653. print(f"ICA components not computed for class label {class_label}, CV nr. {cv_idx}, Cluster {cluster_idx_str}")
  654. return
  655. R_time, _ = self.get_samples_of_cluster(class_label, cv_idx, cluster_idx, time_domain=False, relevance=True)
  656. info = mne.create_info(self.channel_names, self.sampling_rate, ch_types='eeg')
  657. R_epochs = mne.EpochsArray(R_time, info)
  658. R_epochs.set_montage('standard_1020')
  659. R_ica = ica.get_sources(R_epochs).get_data()
  660. R_ica = np.abs(R_ica)
  661. R_ica = np.sum(R_ica, axis=0)
  662. R_ica_plot = np.sum(R_ica, axis=1)
  663. R_ica_rank = np.abs(R_ica_plot)
  664. indices = np.flip(np.argsort(R_ica_rank))
  665. labels_sorted = [(i, labels[i]) for i in indices]
  666. rels_sorted_rank = R_ica_rank[indices]
  667. rels_sorted_plot = R_ica_plot[indices]
  668. rels_sorted_rank /= np.sum(rels_sorted_rank)
  669. count = 4
  670. rels_sorted_rank = rels_sorted_rank[:count]
  671. rels_sorted_rank *= 100
  672. # --------------------------------
  673. colors = []
  674. flip = False
  675. for rel in rels_sorted_plot:
  676. if flip:
  677. if rel < 0:
  678. colors.append('blue')
  679. else:
  680. colors.append('red')
  681. else:
  682. colors.append("teal")
  683. def add_leading_zeros(id):
  684. id = str(id)
  685. if len(id) < 3:
  686. id = (3-len(id))*"0" + id
  687. return id
  688. labels = [f"{add_leading_zeros(id)}_{name}" for (id, name) in labels_sorted[:count]]
  689. # --------------------------------
  690. fig = plt.figure(figsize=figsize)
  691. gs = gridspec.GridSpec(1, 5) # , height_ratios=[3.5, 3.5])
  692. if band is None or band is False:
  693. title = f"Cluster {cluster_idx}"
  694. else:
  695. title = f"Cluster {cluster_idx} ({band})"
  696. plt.suptitle(title, y=1.55)
  697. ax1 = plt.subplot(gs[0, 0])
  698. ax2 = plt.subplot(gs[0, 1])
  699. ax3 = plt.subplot(gs[0, 2])
  700. ax4 = plt.subplot(gs[0, 3])
  701. ax5 = plt.subplot(gs[0, 4])
  702. axes = [ax2, ax3, ax4, ax5]
  703. ax = ax1
  704. y = np.arange(len(labels))
  705. ax.xaxis.grid(False)
  706. ax.yaxis.grid(True, alpha=0.5)
  707. y = np.flip(y)
  708. rels_sorted_rank = np.flip(rels_sorted_rank)
  709. ax.barh(y, rels_sorted_rank, height=0.5, label='', align='center', color=colors, alpha=0.75)
  710. ax.set_yticks(y)
  711. ax.set_yticklabels(labels)
  712. ax.set_xlabel('Relevance (%)')
  713. # ax.set_title("ICA Components")
  714. ica.plot_components(picks=indices[:count], plot_std=True, sensors=False, axes=axes)
  715. plt.show()
  716. # ---------------------------------------------------------------------
  717. def compute_functional_groups(self, freq_signal, brain_regions=None, bands=None, only_positive=False, mean=True):
  718. # print(freq_signal.shape)
  719. # (32, 193)
  720. if brain_regions is None:
  721. brain_regions = ["left temporal", "right temporal", "left parietal", "right parietal",
  722. "left occipital", "right occipital", "left central", "right central",
  723. "left frontal", "right frontal"]
  724. if bands is None:
  725. bands = [(0,4,"delta"), (4,8,"theta"), (8,12,"alpha"), (12,30,"beta"), (30,60,"gamma")]
  726. if only_positive:
  727. freq_signal[freq_signal < 0] = 0
  728. f = 3
  729. values = []
  730. for low, high, band in bands:
  731. low, high = int(low*f), int(high*f)
  732. for br in brain_regions:
  733. _, indices = rel_utils.get_channels_by_brain_region(self.channel_names, brain_region=br)
  734. rel_br = freq_signal[indices,:]
  735. rel_fb = rel_br[:, low:high] # select freq bins
  736. if mean:
  737. rel_fb = np.mean(rel_fb)
  738. values.append(rel_fb)
  739. return values
  740. # ---------------------------------------------------------------------
  741. def plot_cluster_functional_grouping2(self, class_label, cv_idx, brain_regions=None, bands=None, cluster_idx=0,
  742. aggregation="mean", only_positive=False, abs=True, title="", axes=None, from_ideal_topo=False, use_X=False):
  743. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  744. if cluster_idx not in cidx_keys:
  745. print(f"the specified cluster {cluster_idx} does not exist")
  746. return
  747. R_freq1, _ = self.get_samples_of_cluster(0, cv_idx, cluster_idx, time_domain=False, relevance=True)
  748. R_freq2, _ = self.get_samples_of_cluster(1, cv_idx, cluster_idx, time_domain=False, relevance=True)
  749. rels = {}
  750. if only_positive:
  751. R_freq1[R_freq1 < 0] = 0
  752. R_freq2[R_freq2 < 0] = 0
  753. for low, high, band in bands:
  754. low, high = int(low*3), int(high*3)
  755. rels[band] = []
  756. colors = []
  757. for (_, channels, color) in brain_regions:
  758. colors.append(color)
  759. indices = [i for i, ch_name in enumerate(self.channel_names) if ch_name in channels]
  760. rel_br1 = R_freq1[:,indices,:]
  761. rel_fb1 = rel_br1[:, :, low:high] # select freq bins
  762. rel_br2 = R_freq2[:,indices,:]
  763. rel_fb2 = rel_br2[:, :, low:high] # select freq bins
  764. if aggregation == "mean":
  765. rel_fb1 = np.mean(rel_fb1)
  766. rel_fb2 = np.mean(rel_fb2)
  767. elif aggregation == "max":
  768. rel_fb1 = np.max(rel_fb1)
  769. rel_fb2 = np.max(rel_fb2)
  770. elif aggregation == "sum":
  771. rel_fb1 = np.sum(rel_fb1)
  772. rel_fb2 = np.sum(rel_fb2)
  773. elif aggregation == "median":
  774. rel_fb1 = np.median(rel_fb1)
  775. rel_fb2 = np.median(rel_fb2)
  776. rel_fb = np.abs(rel_fb1 - rel_fb2)
  777. rels[band].append(rel_fb)
  778. for i, band in enumerate(bands):
  779. try:
  780. ax = axes[i]
  781. except:
  782. ax = axes
  783. for bidx, (region, _, color) in enumerate(brain_regions):
  784. ax.bar(region, rels[band[2]][bidx], color=color)
  785. ax.set_title(f"$\\{band[2]}$")
  786. if i == 0:
  787. ax.set_ylabel(f"{aggregation} relevance")
  788. ax.tick_params(axis='x', rotation=90) # Adjust the rotation angle as needed
  789. ax.xaxis.grid(False)
  790. # -------------------------------------------------------------------------------------------------
  791. def plot_cluster_functional_grouping(self, class_label, cv_idx, brain_regions=None, bands=None, cluster_idx=0,
  792. aggregation="mean", only_positive=False, abs=True, title="", axes=None, from_ideal_topo=False, use_X=False):
  793. cidx_keys = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  794. if cluster_idx not in cidx_keys:
  795. print(f"the specified cluster {cluster_idx} does not exist")
  796. return
  797. R_freq, _ = self.get_samples_of_cluster(class_label, cv_idx, cluster_idx, time_domain=False, relevance=True)
  798. X_freq, _ = self.get_samples_of_cluster(class_label, cv_idx, cluster_idx, time_domain=False, relevance=False)
  799. if from_ideal_topo:
  800. R_freq = R_freq / np.max(np.abs(R_freq))
  801. X_freq = X_freq / np.max(np.abs(X_freq))
  802. mask = R_freq > 0
  803. R_freq = np.where(mask, R_freq*X_freq, (1-X_freq)*np.abs(R_freq))
  804. print(np.min(R_freq), np.max(R_freq))
  805. if use_X:
  806. R_freq = X_freq
  807. if abs:
  808. R_freq = np.abs(R_freq)
  809. if only_positive:
  810. R_freq[R_freq < 0] = 0
  811. total_relevance = np.mean((np.sum(R_freq)))
  812. rels = {}
  813. rels_unaggregated = {}
  814. for low, high, band in bands:
  815. low, high = int(low*3), int(high*3)
  816. rels[band] = []
  817. rels_unaggregated[band] = []
  818. colors = []
  819. for (_, channels, color) in brain_regions:
  820. colors.append(color)
  821. indices = [i for i, ch_name in enumerate(self.channel_names) if ch_name in channels]
  822. rel_br = R_freq[:,indices,:]
  823. rel_fb = rel_br[:, :, low:high] # select freq bins
  824. if aggregation == "mean":
  825. rel_fb = np.mean(rel_fb)
  826. elif aggregation == "max":
  827. rel_fb = np.max(rel_fb)
  828. elif aggregation == "sum":
  829. rel_fb = np.sum(rel_fb)
  830. elif aggregation == "median":
  831. rel_fb = np.median(rel_fb)
  832. # if abs:
  833. # rel_fb = np.abs(rel_fb)
  834. # rel_fb /= total_relevance
  835. # rel_fb *= 100
  836. rels[band].append(rel_fb)
  837. # ------------
  838. rel_fb = rel_br[:, :, low:high] # select freq bins
  839. rels_unaggregated[band].append(rel_fb)
  840. for i, band in enumerate(bands):
  841. try:
  842. ax = axes[i]
  843. except:
  844. ax = axes
  845. for bidx, (region, _, color) in enumerate(brain_regions):
  846. ax.bar(region, rels[band[2]][bidx], color=color)
  847. ax.set_title(f"$\\{band[2]}$")
  848. if i == 0:
  849. ax.set_ylabel(f"Mean Relevance")
  850. ax.tick_params(axis='x', rotation=90) # Adjust the rotation angle as needed
  851. ax.xaxis.grid(False)
  852. return rels_unaggregated
  853. # ---------------------------------------------------------------------
  854. def plot_correlation_barchart(self, ax=None, colors=None, xticks=None, use_relevance=True):
  855. if use_relevance:
  856. if self.correlation_dict_c0_R is None or self.correlation_dict_c1_R is None:
  857. print("Correlations on relevance have not been computed yet. Compute correlations by calling compute_correlation(..)")
  858. return
  859. else:
  860. if self.correlation_dict_c0_X is None or self.correlation_dict_c1_X is None:
  861. print("Correlations on relevance have not been computed yet. Compute correlations by calling compute_correlation(..)")
  862. return
  863. within_group_labels = ['Same', 'Different']
  864. if use_relevance:
  865. corr_dicts = [self.correlation_dict_c0_R, self.correlation_dict_c1_R]
  866. else:
  867. corr_dicts = [self.correlation_dict_c0_X, self.correlation_dict_c1_X]
  868. if xticks is None:
  869. outer_group_labels = []
  870. else:
  871. outer_group_labels = xticks
  872. data = []
  873. for corr_dict in corr_dicts:
  874. corrs = []
  875. corrs.append(np.mean(corr_dict['CV_x_CV_corr_list']))
  876. corrs.append(np.mean(corr_dict['Fold_x_Fold_corr_list']))
  877. data.append(corrs)
  878. # ---------------------
  879. data = np.array(data)
  880. num_bars = len(data[0])
  881. bar_width = 0.30
  882. bar_distance = 0.25
  883. group_positions = np.arange(len(data))*1.5
  884. if ax is None:
  885. fig, ax = plt.subplots(figsize=(3,3))
  886. if colors is None:
  887. colors = ["blue", "red"]
  888. for i in range(num_bars):
  889. label = within_group_labels[i]
  890. bar_positions = group_positions + i * (bar_width + bar_distance)
  891. bars = ax.bar(bar_positions, data[:, i], bar_width, label=label, color=colors[i])
  892. for bar in bars:
  893. height = bar.get_height()
  894. ax.text(bar.get_x() + bar.get_width() / 2, height + 0.0, f'{height:.2f}', ha='center', va='bottom', fontsize=8)
  895. ax.set_ylim(0, 1)
  896. ax.set_xticks(group_positions + ((num_bars - 1) * (bar_width + bar_distance)) / 2)
  897. # cluster_labels = cidx_keys
  898. ax.set_xticklabels(outer_group_labels, rotation=0)
  899. ax.set_xlabel('Classes')
  900. ax.set_ylabel('Mean Correlation')
  901. # ax.set_title('Correlation: across CVs vs. across Folds', pad=10)
  902. ax.legend(loc='lower right')
  903. return data
  904. # ---------------------------------------------------------------------
  905. def plot_correlation_matrices(self, CV_params, class_label, axes=None, use_relevance=True, over_folds=True, tick_size=9):
  906. if class_label == 0:
  907. if use_relevance:
  908. corr_dict = self.correlation_dict_c0_R
  909. else:
  910. corr_dict = self.correlation_dict_c1_X
  911. else:
  912. if use_relevance:
  913. corr_dict = self.correlation_dict_c1_R
  914. else:
  915. corr_dict = self.correlation_dict_c1_X
  916. if corr_dict is None:
  917. print("Correlations on relevance have not been computed yet. Compute correlations by calling compute_correlation(..)")
  918. return
  919. # -----------------------------------------------------
  920. if over_folds:
  921. label = "Fold"
  922. title = "CV"
  923. key_mat = "Fold_x_Fold_matrices"
  924. else:
  925. label = "CV"
  926. title = "Fold"
  927. key_mat = "CV_x_CV_matrices"
  928. corr_mats = corr_dict[key_mat]
  929. heatmaps = []
  930. for idx, corr_mat in enumerate(corr_mats):
  931. title_ = f"{title} {idx}"
  932. if not over_folds:
  933. ticklabels = [f"{params['nr']}" for params in CV_params]
  934. else:
  935. ticklabels = [subj for subj in range(self.num_folds)]
  936. if axes is None:
  937. _, ax = plt.subplots(1, 1, figsize=(2.75, 2.5))
  938. else:
  939. ax = axes[idx]
  940. sns.set(font_scale=0.75)
  941. ax.set_title(f"{title_}")
  942. hm = sns.heatmap(corr_mat, annot=False, cmap="coolwarm", linewidths=0.5, xticklabels=ticklabels,
  943. yticklabels=ticklabels, ax=ax, vmin=-1, vmax=1, cbar=False)
  944. hm.set_xticklabels(hm.get_xticklabels(), rotation=0, ha="right", fontsize=tick_size)
  945. hm.set_yticklabels(hm.get_yticklabels(), rotation=0, ha="right", fontsize=tick_size)
  946. ax.set_xlabel(label)
  947. ax.set_ylabel(label)
  948. heatmaps.append(hm)
  949. return heatmaps
  950. # -------------------------------------------------------------------------------------
  951. def plot_LRP_baseline(self, dh, params, class_label, title="", proportion_of_samples=1.0, discard_negative_rel=False, reversed_classes=False, model_id_of_subj=None, subject="all",
  952. select_correct=None, select_samples_by_class=True, composite=None, use_which_data="train-test"):
  953. X_freq, R_freq = utils.load_data_and_compute_relevance(dh, params, class_label, title, proportion_of_samples=proportion_of_samples,
  954. discard_negative_rel=discard_negative_rel, reversed_classes=reversed_classes, model_id_of_subj=model_id_of_subj, subject=subject,
  955. select_correct=select_correct, select_samples_by_class=select_samples_by_class, composite=composite, freq_domain=True,
  956. use_which_data=use_which_data)
  957. bands = [(0,4,"delta"), (4,8,"theta"), (8,12,"alpha"), (12,30,"beta"), (30,60,"gamma")]
  958. fig, axes = plt.subplots(2, len(bands), figsize=(6, 4))
  959. self._plot_freq_bands_topo(X_freq, self.channel_names, bands, axes[0], title=title, global_scale=True, norm=False, log_scale=False, scale_bands=True, cmap="Reds",
  960. label="", colorbar_offset=(0.0,0.5), vlim=None, mean=False, negative_lim=False)
  961. self._plot_freq_bands_topo(R_freq, self.channel_names, bands, axes[1], title="", global_scale=True, norm=False, log_scale=False, scale_bands=False, cmap="Reds",
  962. label="", colorbar_offset=(0.0,0.2), vlim=None, mean=False, negative_lim=False)
  963. plt.show()
  964. # -------------------------------------------------------------------------------------
  965. def plot_sensors_of_brain_regions(self, brain_regions, cmap=None, figsize=(5,5)):
  966. plt.rc('font', size=10)
  967. info = mne.create_info(self.channel_names, sfreq=128, ch_types='eeg')
  968. info.set_montage('standard_1020')
  969. fig, ax = plt.subplots(1,1, figsize=figsize)
  970. all_indices = []
  971. colors = []
  972. for _, (_, channels, color) in enumerate(brain_regions):
  973. indices = [i for i, ch_name in enumerate(self.channel_names) if ch_name in channels]
  974. all_indices.append(indices)
  975. colors.append(color)
  976. cmap = LinearSegmentedColormap.from_list('custom_cmap', colors)
  977. fig = mne.viz.plot_sensors(info, show_names=True, ch_type='eeg', to_sphere=True, axes=ax, ch_groups=all_indices,
  978. pointsize=125, linewidth=0, sphere="auto", cmap=cmap)
  979. ax.set_facecolor("white")
  980. plt.show()
  981. return fig
  982. # -------------------------------------------------------------------------------------
  983. def _plot_freq_bands_topo_original(self, signal, ch_names, bands, ax, title="", global_scale=True, norm=True, log_scale=False, scale_bands=False, cmap=None,
  984. label="", colorbar_offset=[0,0], vlim=None, mean=False, negative_lim=True, scale_per_band=False, scale_1_div_f=False, scale_factors=None):
  985. suptitle = title
  986. # if cmap is None:
  987. # if np.min(signal) < 0 or negative_lim:
  988. # cmap = "bwr"
  989. # elif not negative_lim:
  990. # cmap = "Reds"
  991. # else:
  992. # cmap = "Reds"
  993. if len(signal.shape) == 3:
  994. if mean:
  995. signal = np.mean(signal, axis=0)
  996. else:
  997. signal = np.sum(signal, axis=0)
  998. signal_bands = []
  999. mult = signal.shape[-1] // 60
  1000. min_per_band = []
  1001. max_per_band = []
  1002. # -------------------------------
  1003. # sum over bins, then find min and max per band
  1004. for (a, b, _) in bands:
  1005. freq = np.copy(signal[:, mult*a:mult*b])
  1006. if mean:
  1007. freq = np.mean(freq, axis=1)
  1008. else:
  1009. freq = np.sum(freq, axis=1)
  1010. signal_bands.append(freq)
  1011. min_per_band.append(np.min(freq))
  1012. max_per_band.append(np.max(freq))
  1013. freq_max = np.max(max_per_band)
  1014. freq_min = np.min(min_per_band)
  1015. # total_sum = np.sum(signal_bands)
  1016. # print("total_sum: ", total_sum)
  1017. # -------------------------------
  1018. # divide by global max (max value considering all bands)
  1019. if global_scale and norm:
  1020. for i in range(len(signal_bands)):
  1021. # freq_max = total_sum
  1022. signal_bands[i] /= freq_max
  1023. max_per_band[i] /= freq_max
  1024. min_per_band[i] /= freq_max
  1025. # -------------------------------
  1026. # scale bands: bring each band to the same amplitude
  1027. if scale_bands:
  1028. # scalars = [math.floor(1/val) for val in max_per_band]
  1029. max_total = np.max(max_per_band)
  1030. scalar_strs = []
  1031. if scale_1_div_f and scale_factors is not None:
  1032. scalars = scale_factors
  1033. for scalar in scalars:
  1034. scalar_strs.append(f"{scalar:.1f}")
  1035. for i in range(len(signal_bands)):
  1036. # signal_bands[i] /= max_total
  1037. signal_bands[i] *= scalars[i]
  1038. max_val = -np.inf
  1039. for sb in signal_bands:
  1040. if np.max(np.abs(sb)) > max_val:
  1041. max_val = np.max(np.abs(sb))
  1042. for i in range(len(signal_bands)):
  1043. # signal_bands[i] /= max_total
  1044. signal_bands[i] /= max_val
  1045. else:
  1046. scalars = []
  1047. for val in max_per_band:
  1048. scalar = max_total / val
  1049. scalars.append(scalar)
  1050. scalar_strs.append(f"{scalar:.1f}")
  1051. for i in range(len(signal_bands)):
  1052. signal_bands[i] *= scalars[i]
  1053. for i in range(len(bands)):
  1054. ax_band = ax[i]
  1055. band = bands[i]
  1056. freq = signal_bands[i]
  1057. if not global_scale and norm:
  1058. freq_max = np.max(np.abs(freq))
  1059. if i == 0:
  1060. label_ = label
  1061. else:
  1062. label_ = ""
  1063. title = utils.freq_band_name_to_latex(band[2])
  1064. if vlim is None:
  1065. if global_scale:
  1066. if negative_lim:
  1067. vlim = (-np.max(signal_bands), np.max(signal_bands))
  1068. else:
  1069. vlim = (0, np.max(signal_bands))
  1070. else:
  1071. if negative_lim:
  1072. vlim = (-np.max(freq), np.max(freq))
  1073. else:
  1074. vlim = (0, np.max(freq))
  1075. if log_scale:
  1076. vlim = (vlim[0]+1, vlim[1])
  1077. if scale_bands and not scale_per_band:
  1078. title += f" ($\\times$ {scalar_strs[i]})"
  1079. if scale_per_band:
  1080. vlim=(None, None)
  1081. im = self._plot_topography(freq, 120, ch_names, title=title, label=label_, plot=True, sp_size=2, show_ch_names=False, vlim=vlim, ax=ax_band, cmap=cmap)
  1082. if not scale_per_band:
  1083. if i == len(bands)-1 and global_scale:
  1084. if suptitle:
  1085. cax = plt.axes([1.0 + colorbar_offset[0], 0.0 + colorbar_offset[1], 0.01, 0.2])
  1086. else:
  1087. cax = plt.axes([1.0 + colorbar_offset[0], 0.0 + colorbar_offset[1], 0.01, 0.2])
  1088. plt.colorbar(im, cax=cax)
  1089. else:
  1090. step = 0.22
  1091. pos = step + i * step
  1092. cax = plt.axes([pos + colorbar_offset[0], 0.0 + colorbar_offset[1], 0.01, 0.2])
  1093. plt.colorbar(im, cax=cax)
  1094. if suptitle:
  1095. plt.suptitle(suptitle, y=0.9)
  1096. return vlim, signal
  1097. # -----------------------------------------------------------------------------------------
  1098. def _plot_freq_bands_topo(self, signal, ch_names, bands, ax, title="", global_scale=True, norm=True, log_scale=False, scale_bands=False, cmap=None,
  1099. label="", colorbar_offset=[0,0], vlim=None, mean=False, negative_lim=True, scale_per_band=False, scale_1_div_f=False, scale_factors=None, remove_neg=False):
  1100. suptitle = title
  1101. if len(signal.shape) == 3:
  1102. if mean:
  1103. signal = np.mean(signal, axis=0)
  1104. else:
  1105. signal = np.sum(signal, axis=0)
  1106. signal_bands = []
  1107. mult = signal.shape[-1] // 60
  1108. min_per_band = []
  1109. max_per_band = []
  1110. # -------------------------------
  1111. # sum over bins, then find min and max per band
  1112. for (a, b, _) in bands:
  1113. freq = np.copy(signal[:, mult*a:mult*b])
  1114. if remove_neg:
  1115. freq[freq < 0] = 0
  1116. if mean:
  1117. freq = np.mean(freq, axis=1)
  1118. else:
  1119. freq = np.sum(freq, axis=1)
  1120. signal_bands.append(freq)
  1121. min_per_band.append(np.min(freq))
  1122. max_per_band.append(np.max(freq))
  1123. freq_max = np.max(max_per_band)
  1124. freq_min = np.min(min_per_band)
  1125. # -------------------------------
  1126. # divide by global max (max value considering all bands)
  1127. if global_scale and norm:
  1128. for i in range(len(signal_bands)):
  1129. # freq_max = total_sum
  1130. signal_bands[i] /= freq_max
  1131. max_per_band[i] /= freq_max
  1132. min_per_band[i] /= freq_max
  1133. # -------------------------------
  1134. # scale bands: bring each band to the same amplitude
  1135. if scale_bands:
  1136. max_total = np.max(max_per_band)
  1137. scalar_strs = []
  1138. if scale_1_div_f and scale_factors is not None:
  1139. scalars = scale_factors
  1140. for scalar in scalars:
  1141. scalar_strs.append(f"{scalar:.1f}")
  1142. for i in range(len(signal_bands)):
  1143. signal_bands[i] *= scalars[i]
  1144. max_val = -np.inf
  1145. for sb in signal_bands:
  1146. if np.max(np.abs(sb)) > max_val:
  1147. max_val = np.max(np.abs(sb))
  1148. for i in range(len(signal_bands)):
  1149. signal_bands[i] /= max_val
  1150. else:
  1151. scalars = []
  1152. for val in max_per_band:
  1153. scalar = max_total / val
  1154. scalars.append(scalar)
  1155. scalar_strs.append(f"{scalar:.1f}")
  1156. for i in range(len(signal_bands)):
  1157. signal_bands[i] *= scalars[i]
  1158. # -------------------------------
  1159. freqs = []
  1160. for i in range(len(bands)):
  1161. ax_band = ax[i]
  1162. band = bands[i]
  1163. freq = signal_bands[i]
  1164. if not global_scale and norm:
  1165. freq_max = np.max(np.abs(freq))
  1166. if i == 0:
  1167. label_ = label
  1168. else:
  1169. label_ = ""
  1170. title = utils.freq_band_name_to_latex(band[2])
  1171. if vlim is None:
  1172. if global_scale:
  1173. if negative_lim:
  1174. vlim = (-np.max(signal_bands), np.max(signal_bands))
  1175. else:
  1176. vlim = (0, np.max(signal_bands))
  1177. else:
  1178. if negative_lim:
  1179. vlim = (-np.max(freq), np.max(freq))
  1180. else:
  1181. vlim = (0, np.max(freq))
  1182. if scale_bands:
  1183. title += f" ($\\times$ {scalar_strs[i]})"
  1184. im = self._plot_topography(freq, 120, ch_names, title=title, label=label_, plot=True, sp_size=2, show_ch_names=False, vlim=vlim, ax=ax_band, cmap=cmap)
  1185. freqs.append(freq)
  1186. if i == len(bands)-1 and global_scale:
  1187. if suptitle:
  1188. cax = plt.axes([1.0 + colorbar_offset[0], 0.0 + colorbar_offset[1], 0.01, 0.2])
  1189. else:
  1190. cax = plt.axes([1.0 + colorbar_offset[0], 0.0 + colorbar_offset[1], 0.01, 0.2])
  1191. plt.colorbar(im, cax=cax)
  1192. if suptitle:
  1193. plt.suptitle(suptitle, y=0.8)
  1194. return vlim, freqs
  1195. # -----------------------------------------------------------------------
  1196. def _plot_topography(self, signal, sampling_rate, ch_names=None, show_ch_names=True, title="", label="", plot=True, vlim=(-1,1), sp_size=3, ax=None, cmap="RdBu_r"):
  1197. if ch_names is None:
  1198. ch_names = []
  1199. ax.set_title(title)
  1200. if label:
  1201. ax.set_ylabel(label, rotation=90, labelpad=20)
  1202. info = mne.create_info(ch_names, sfreq=sampling_rate, ch_types='eeg')
  1203. info.set_montage('standard_1020')
  1204. if not show_ch_names:
  1205. ch_names = None
  1206. im, cm = mne.viz.plot_topomap(signal, info, names=ch_names, show=False, axes=ax, vlim=vlim, cmap=cmap,
  1207. image_interp="linear")
  1208. plt.tight_layout()
  1209. if plot:
  1210. plt.plot()
  1211. else:
  1212. plt.close()
  1213. return im
  1214. # -----------------------------------------------------------------------
  1215. def _calc_ica_on_cluster(self, class_label, cv_idx, cluster_idx, time_domain=True, band=False):
  1216. X, y = self.get_samples_of_cluster(class_label, cv_idx, cluster_idx, time_domain=time_domain)
  1217. # # 8-12 Hz
  1218. # low = 3*8
  1219. # high = 3*12
  1220. if band == "alpha":
  1221. low = 8
  1222. high = 12
  1223. elif band == "beta":
  1224. low = 12
  1225. high = 30
  1226. elif band == "gamma":
  1227. low = 30
  1228. high = 60
  1229. else:
  1230. low = 0
  1231. high = 60
  1232. info = mne.create_info(self.channel_names, self.sampling_rate, ch_types='eeg')
  1233. epochs = mne.EpochsArray(X, info)
  1234. epochs.set_montage('standard_1020')
  1235. if time_domain and band != False:
  1236. print("filter time")
  1237. epochs = epochs.copy().filter(l_freq=low, h_freq=high)
  1238. ica = mne.preprocessing.ICA(
  1239. n_components=None,
  1240. max_iter="auto",
  1241. method="infomax", # Use the "extended Infomax" algorithm as specified by ICLabel
  1242. random_state=0,
  1243. )
  1244. picks = np.arange(len(self.channel_names)-1)
  1245. ica.fit(epochs)
  1246. ic_labels = label_components(epochs, ica, method="iclabel")
  1247. labels = ic_labels["labels"]
  1248. proba = ic_labels["y_pred_proba"]
  1249. return ica, labels, proba
  1250. # ---------------------------------------------------------------------
  1251. def _collect_data(self):
  1252. tmp = joblib.load(os.path.join(self.paths_to_CV_iteration[0][0], "total_rel_per_filter_final"))
  1253. num_total = len(tmp)
  1254. num_filters_per_model = num_total // self.num_folds
  1255. del tmp
  1256. c=0
  1257. fp_dict = {}
  1258. for s in range(self.num_folds):
  1259. for _ in range(num_filters_per_model):
  1260. fp_dict[c] = s
  1261. c += 1
  1262. for class_label in self.class_labels:
  1263. for cv_idx in range(self.num_CVs):
  1264. for fold_idx in range(self.num_folds):
  1265. path_store_data = self.paths_to_CV_iteration[class_label][cv_idx]
  1266. for f in range(num_total):
  1267. if fold_idx != fp_dict[f]:
  1268. continue
  1269. try:
  1270. if self.layer != "b3_flatten":
  1271. fn = os.path.join(path_store_data, f"filter_activation_map_{fold_idx}_{f % num_filters_per_model}")
  1272. X_sim = joblib.load(fn)
  1273. fn = os.path.join(path_store_data, f"filter_activation_map_rel_{fold_idx}_{f % num_filters_per_model}")
  1274. X_sim_rel = joblib.load(fn)
  1275. else:
  1276. X_sim = X_sim_rel = None
  1277. # samples_freq
  1278. fn = os.path.join(path_store_data, f"X_freq_select_{fold_idx}_{f % num_filters_per_model}")
  1279. X_sim_samples = joblib.load(fn)
  1280. if self.use_what_for_similarity in ["samples_psd", "samples_rel_combined"]:
  1281. X_sim_samples = np.abs(X_sim_samples)**2
  1282. X_sim_samples = np.mean(X_sim_samples, axis=0)
  1283. a, b = X_sim_samples.shape
  1284. X_sim_samples = X_sim_samples[:,:192].reshape(a, b // 12, 12)
  1285. X_sim_samples = X_sim_samples.mean(axis=-1)
  1286. X_sim_samples = X_sim_samples.flatten()
  1287. X_sim_samples /= np.max(X_sim_samples)
  1288. # samples_rel
  1289. fn = os.path.join(path_store_data, f"R_freq_select_{fold_idx}_{f % num_filters_per_model}")
  1290. R_sim_samples = joblib.load(fn)
  1291. if self.use_what_for_similarity == "samples_rel_pos":
  1292. R_sim_samples[R_sim_samples < 0] = 0
  1293. # R_sim_samples = R_sim_samples**2
  1294. R_sim_samples = np.mean(R_sim_samples, axis=0)
  1295. a, b = R_sim_samples.shape
  1296. # functional_groups = self.compute_functional_groups(R_sim_samples)
  1297. R_sim_samples = R_sim_samples[:,:192].reshape(a, b // 12, 12)
  1298. R_sim_samples = R_sim_samples.mean(axis=-1)
  1299. R_sim_samples = R_sim_samples.flatten()
  1300. R_sim_samples /= np.max(np.abs(R_sim_samples))
  1301. except Exception as e:
  1302. print(">>> skip filter, no data:", f, e)
  1303. continue
  1304. if self.use_what_for_similarity == "combined":
  1305. # X_sim = torch.cat([X_sim, X_sim_rel])
  1306. X_sim = torch.mul(X_sim, X_sim_rel)
  1307. elif self.use_what_for_similarity == "activation_map_relevance":
  1308. X_sim = X_sim_rel
  1309. elif self.use_what_for_similarity == "samples_freq":
  1310. X_sim = X_sim_samples
  1311. elif self.use_what_for_similarity == "samples_rel":
  1312. X_sim = R_sim_samples
  1313. elif self.use_what_for_similarity == "samples_rel_pos":
  1314. X_sim = R_sim_samples
  1315. elif self.use_what_for_similarity == "samples_psd":
  1316. X_sim = X_sim_samples
  1317. elif self.use_what_for_similarity == "samples_rel_combined":
  1318. X_sim = np.concatenate([X_sim_samples, R_sim_samples])
  1319. # elif self.use_what_for_similarity == "samples_rel_functional_groups":
  1320. # X_sim = functional_groups
  1321. else:
  1322. pass
  1323. if torch.is_tensor(X_sim):
  1324. X_sim = X_sim.cpu().detach().numpy()
  1325. self.data["X_sim"][class_label][cv_idx][fold_idx].append(X_sim)
  1326. self.data["X_freq_files"][class_label][cv_idx][fold_idx].append(os.path.join(path_store_data, f"X_freq_select_{fold_idx}_{f % num_filters_per_model}"))
  1327. self.data["X_time_files"][class_label][cv_idx][fold_idx].append(os.path.join(path_store_data, f"X_time_select_{fold_idx}_{f % num_filters_per_model}"))
  1328. self.data["R_freq_files"][class_label][cv_idx][fold_idx].append(os.path.join(path_store_data, f"R_freq_select_{fold_idx}_{f % num_filters_per_model}"))
  1329. self.data["R_time_files"][class_label][cv_idx][fold_idx].append(os.path.join(path_store_data, f"R_time_select_{fold_idx}_{f % num_filters_per_model}"))
  1330. self.data["y_files"][class_label][cv_idx][fold_idx].append(os.path.join(path_store_data, f"y_select_{fold_idx}_{f % num_filters_per_model}"))
  1331. # ---------------------------------------------------------------------
  1332. def _prepare_paths(self):
  1333. self.paths_to_CV_iteration = {}
  1334. for class_label in self.class_labels:
  1335. self.paths_to_CV_iteration[class_label] = {}
  1336. for i in range(self.num_CVs):
  1337. _, path_for_cv = Results.assemble_path(self.CV_params[i], self.base_result_path,
  1338. self.num_selected_samples, self.ds_name, self.layer, class_label,
  1339. self.select_samples_by_class, self.reversed_classes, self.select_correct,
  1340. False, self.use_which_data, self.testing)
  1341. self.paths_to_CV_iteration[class_label][i] = path_for_cv
  1342. # ---------------------------------------------------------------------
  1343. def _calc_cluster_coherence(self, class_label, cv_idx, cluster_idx):
  1344. embedding = self.get_data_class_cv("embedding", class_label, cv_idx)
  1345. labels = self.get_data_class_cv("labels", class_label, cv_idx)
  1346. indices = self.get_data_class_cv("indices", class_label, cv_idx)
  1347. unique_labels = copy(self.get_data_class_cv("unique_labels", class_label, cv_idx))
  1348. index = unique_labels.index(cluster_idx)
  1349. cluster_points = embedding[indices[index]]
  1350. pairwise_distances = pdist(cluster_points, metric='euclidean')
  1351. average_distance = np.mean(pairwise_distances)
  1352. if np.isnan(average_distance):
  1353. average_distance = -1
  1354. inverse = 1.0/average_distance
  1355. return inverse
  1356. # ---------------------------------------------------------------------
  1357. def _calc_cluster_relevance(self, class_label, cv_idx, cluster_idx):
  1358. R_time, _ = self.get_samples_of_cluster(class_label, cv_idx, cluster_idx, time_domain=True, relevance=True)
  1359. return np.mean(R_time)
  1360. # ---------------------------------------------------------------------
  1361. def _calc_correlations(self, class_label, CV_params, plot=True, discard_negative=False, take_abs=False, use_relevance=True):
  1362. correlations_over_elements_dict = {
  1363. "Fold_x_Fold_matrices": [], "CV_x_CV_matrices": [],
  1364. "Fold_x_Fold_corr_list": [], "CV_x_CV_corr_list": []
  1365. }
  1366. for_CV_x_CV, for_fold_x_fold = self._gather_data_for_correlation(CV_params, class_label, use_relevance=use_relevance)
  1367. for over_folds in [False, True]:
  1368. if over_folds:
  1369. corr_dict = for_fold_x_fold
  1370. label = "Fold"
  1371. key_list = "Fold_x_Fold_corr_list"
  1372. key_mat = "Fold_x_Fold_matrices"
  1373. else:
  1374. corr_dict = for_CV_x_CV
  1375. label = "CV"
  1376. key_list = "CV_x_CV_corr_list"
  1377. key_mat = "CV_x_CV_matrices"
  1378. for key in corr_dict.keys():
  1379. ticklabels = [f"{params['nr']}" for params in CV_params]
  1380. R_sim_list = corr_dict[key]
  1381. # get max value (R_sim_list is ragged, so np.max doesn't work)
  1382. max_val = -np.inf
  1383. for R_sim in R_sim_list:
  1384. max_ = np.max(R_sim)
  1385. if max_ > max_val:
  1386. max_val = max_
  1387. R_sim_list_ = []
  1388. for R_sim in R_sim_list:
  1389. R_sim_ = R_sim
  1390. if discard_negative:
  1391. R_sim_[R_sim_ < 0] = 0
  1392. elif take_abs:
  1393. R_sim_ = np.abs(R_sim_)
  1394. R_sim_ = np.sum(R_sim_, axis=0)
  1395. R_sim_ = R_sim_.flatten()
  1396. R_sim_list_.append(R_sim_)
  1397. R_sim_list = np.array(R_sim_list_)
  1398. corr_mat = np.corrcoef(R_sim_list)
  1399. if plot:
  1400. _, ax = plt.subplots(1, 1, figsize=(2.5, 2.5))
  1401. if np.min(corr_mat) < 0:
  1402. vmin = -1
  1403. else:
  1404. vmin = 0
  1405. sns.set(font_scale=0.5)
  1406. ax.set_title(f"Correlation on {label}s - R - {key}")
  1407. hm1 = sns.heatmap(corr_mat, annot=False, cmap="coolwarm", linewidths=0.5, xticklabels=ticklabels,
  1408. yticklabels=ticklabels, ax=ax, vmin=vmin, vmax=1)
  1409. hm1.set_xticklabels(hm1.get_xticklabels(), rotation=90, ha="right")
  1410. ax.set_xlabel(label)
  1411. ax.set_ylabel(label)
  1412. plt.show()
  1413. lower_triangle_indices = np.tril_indices(corr_mat.shape[0], k=-1)
  1414. lower_triangle_elements = corr_mat[lower_triangle_indices]
  1415. correlations_over_elements_dict[key_list] += list(lower_triangle_elements)
  1416. correlations_over_elements_dict[key_mat].append(corr_mat)
  1417. if plot:
  1418. plt.tight_layout()
  1419. plt.show()
  1420. return correlations_over_elements_dict
  1421. # -------------------------------------------------------------------------------------
  1422. def _gather_data_for_correlation(self, CV_params, class_label, use_relevance=True):
  1423. for_CV_x_CV = {}
  1424. folds = [i for i in range(self.num_folds)]
  1425. for fold_idx, fold in enumerate(folds):
  1426. for_CV_x_CV[fold_idx] = []
  1427. for cv_idx, cv in enumerate(CV_params):
  1428. if use_relevance:
  1429. R_freq_fns = self.get_data_class_cv_fold("R_freq_files", class_label, cv_idx, fold_idx)
  1430. else:
  1431. R_freq_fns = self.get_data_class_cv_fold("X_freq_files", class_label, cv_idx, fold_idx)
  1432. R_freq = []
  1433. for fn in R_freq_fns:
  1434. R_freq.append(joblib.load(fn))
  1435. R_freq = np.concatenate(R_freq, axis=0)
  1436. for_CV_x_CV[fold_idx].append(R_freq)
  1437. for_fold_x_fold = {}
  1438. for cv_idx, cv in enumerate(CV_params):
  1439. for_fold_x_fold[cv_idx] = []
  1440. for fold_idx, fold in enumerate(folds):
  1441. if use_relevance:
  1442. R_freq_fns = self.get_data_class_cv_fold("R_freq_files", class_label, cv_idx, fold_idx)
  1443. else:
  1444. R_freq_fns = self.get_data_class_cv_fold("X_freq_files", class_label, cv_idx, fold_idx)
  1445. R_freq = []
  1446. for fn in R_freq_fns:
  1447. R_freq.append(joblib.load(fn))
  1448. R_freq = np.concatenate(R_freq, axis=0)
  1449. for_fold_x_fold[cv_idx].append(R_freq)
  1450. return for_CV_x_CV, for_fold_x_fold
  1451. # -------------------------------------------------------------------------------------

results.py at commit 75a0576, no license · at the source

Overview

Authors: Hendrik Eilts1, Gabriel Ivucic1, Niklas Koenen2, Marvin N Wright2, Tanja Schultz1, Felix Putze1
  1. Cognitive Systems Lab, Universität Bremen, Bremen, Germany
  2. Statistical Methods in Epidemiology, Leibniz Institute for Prevention Research and Epidemiology—BIPS, Bremen, Germany
Journal: Human brain mapping, volume 47, issue 6, article e70528
Dates: received 22 September 2025; accepted 7 April 2026; published online 26 April 2026; in print April 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1002/hbm.70528 · PMID 42037083 · PMCID PMC13111923 · OpenAlex W7156255653
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: EEG (modality), human (organism), cognitive (subfield)
Methods: Spectral & time-frequency, Connectivity, Smoothing, state filtering, decompositions, Preprocessing, Physiology & signal measures
Keywords: BCI, CRP, EEG, XAI
MeSH: Attention*, Auditory Perception*, Deep Learning*, Electroencephalography*, Convolutional Neural Networks, Humans (* major topic)
Topic: EEG and Brain-Computer Interfaces (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: Deutsche Forschungsgemeinschaft (447089431, 459360854)
Citations: not cited yet (Europe PMC); 41 references in the paper

Abstract

While deep learning has drastically improved the performance of electroencephalography (EEG) analysis, it remains unclear what these models, such as EEGNet, learn from the data and how their learned features relate to neuroscientific concepts. In this work, we introduce a comprehensive interpretability framework for deep learning models of neural data based on Concept Relevance Propagation (CRP), an extension of layer‐wise relevance propagation that enables the analysis of abstract concepts encoded by individual neurons and filters. We apply CRP to individual filters of convolutional neural networks (EEGNet) trained using leave‐one‐out cross‐validation. To identify common classification strategies across models, we guide the selection of representative data for individual filters using relevance maximization, reduce dimensionality via UMAP, and identify clusters of filters encoding similar concepts through density‐based clustering. To gain insight into the neural correlates of these tasks, we analyze the learned features across multiple data domains without requiring model retraining. We integrate a virtual inspection layer to project explanations into the frequency domain, enabling the simultaneous analysis of spatial, temporal, and spectral aspects using topographic maps, functional grouping, and independent component analysis (ICA). Using three EEG classification tasks—auditory attention, internal/external attention, and motor imagery—we demonstrate that our approach reveals interpretable, task‐relevant neural patterns that generalize across participants. Overall, this framework provides a step toward understanding the models itself and gaining insights into the tasks in terms of neuroscience.

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

Repository

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

hendrik-eilts/XAI-EEGNet

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 75a0576615a155571c2fe71edea2b66f4e9d171e, 28 April 2026
Languages: Python (33), Jupyter (4)
Size: 42 files, 37 scripts
Software Heritage: not archived
Found in: the references
Holds: README, environment (requirements.txt), 4 notebooks
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (27 files), PyTorch (23 files), Matplotlib (11 files), MNE-Python (9 files), pandas (3 files), SciPy (3 files), seaborn (3 files), ICLabel (2 files), Pillow (2 files), scikit-learn (2 files), MNE-Connectivity (1 file), MOABB (1 file), UMAP (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
38 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;
  • 37 scripts, each with its path and the digest of its content;
  • 12 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

No dataset and no data link were found in the paper.

Data Availability Statement

Data sharing not applicable to this article as no datasets were generated or analysed during the current study.

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, 6 authors, 4 keywords, 6 MeSH terms, 1 funder, 21 references.

Cite

This paper

Eilts, H., Ivucic, G., Koenen, N., Wright, M. N., Schultz, T., & Putze, F. (2026). Explainable AI Insights Into EEG Classification and Its Alignment to Neural Correlates. Human brain mapping, 47(6), e70528. https://doi.org/10.1002/hbm.70528

BibTeX

@article{eilts2026explainable,
author = {Eilts, Hendrik and Ivucic, Gabriel and Koenen, Niklas and Wright, Marvin N and Schultz, Tanja and Putze, Felix},
title = {{Explainable AI Insights Into EEG Classification and Its Alignment to Neural Correlates}},
journal = {Human brain mapping},
year = {2026},
month = apr,
volume = {47},
number = {6},
pages = {e70528},
publisher = {Wiley},
issn = {1065-9471},
doi = {10.1002/hbm.70528},
url = {https://doi.org/10.1002/hbm.70528},
pmid = {42037083},
pmcid = {PMC13111923}
}

RIS

TY - JOUR
AU - Eilts, Hendrik
AU - Ivucic, Gabriel
AU - Koenen, Niklas
AU - Wright, Marvin N
AU - Schultz, Tanja
AU - Putze, Felix
TI - Explainable AI Insights Into EEG Classification and Its Alignment to Neural Correlates
T2 - Human brain mapping
J2 - Hum Brain Mapp
PY - 2026
DA - 2026/04/01
VL - 47
IS - 6
SP - e70528
SN - 1065-9471
PB - Wiley
DO - 10.1002/hbm.70528
UR - https://doi.org/10.1002/hbm.70528
LA - en
ER -

CSL-JSON

{
"id": "10.1002/hbm.70528",
"type": "article-journal",
"title": "Explainable AI Insights Into EEG Classification and Its Alignment to Neural Correlates",
"container-title": "Human brain mapping",
"author": [
{
"family": "Eilts",
"given": "Hendrik"
},
{
"family": "Ivucic",
"given": "Gabriel"
},
{
"family": "Koenen",
"given": "Niklas"
},
{
"family": "Wright",
"given": "Marvin N"
},
{
"family": "Schultz",
"given": "Tanja"
},
{
"family": "Putze",
"given": "Felix"
}
],
"container-title-short": "Hum Brain Mapp",
"volume": "47",
"issue": "6",
"page": "e70528",
"DOI": "10.1002/hbm.70528",
"PMID": "42037083",
"PMCID": "PMC13111923",
"ISSN": "1065-9471",
"publisher": "Wiley",
"URL": "https://doi.org/10.1002/hbm.70528",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
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.3389/fncom.2026.1786996 [code]
Schumann-anchored golden ratio organization of human neural oscillations.
Journal: Frontiers in computational neuroscience
In common: MNE-Connectivity, UMAP, MNE-Python, 7 other tools, EEG
[2] doi:10.1097/j.pain.0000000000004044 [code]
No effect of rhythmic visual stimulation on experimental pain perception.
Journal: Pain
In common: MNE-Connectivity, ICLabel, MNE-Python, 5 other tools, EEG, cognitive, 1 reference
[3] doi:10.1002/hbm.70628 [code]
EEG Biomarkers for Affective Disorders Diagnosis: An Evaluation and Validation Study.
Journal: Human brain mapping
In common: MNE-Connectivity, ICLabel, MNE-Python, 6 other tools, EEG
[4] doi:10.1038/s41598-026-56070-y [code]
SSDLabeler: realistic semi-synthetic data generation for multi-label artifact classification in EEG.
Journal: Scientific reports
In common: ICLabel, PyTorch, seaborn, 4 other tools, EEG, 3 references
[5] doi:10.1038/s41597-025-05174-7 [code]
A large-scale MEG and EEG dataset for object recognition in naturalistic scenes
Journal: n/a
In common: ICLabel, MNE-Python, Pillow, 7 other tools, EEG
[6] doi:10.1038/s41746-026-02778-0 [code]
Trust-gated synthetic EEG augmentation reduces performance drops when generalizing to new patients.
Journal: NPJ digital medicine
In common: MOABB, MNE-Python, PyTorch, 5 other tools, EEG, 1 reference
[7] doi:10.1523/eneuro.0041-26.2026 [code]
Ocular Speech Tracking Persists in Blindness, but Its Dynamics and Oculo-Cerebral Connectivity Depend on Visual Status.
Journal: eNeuro
In common: MNE-Connectivity, MNE-Python, Pillow, 6 other tools, cognitive
[8] doi:10.1002/mds.70348 [code]
Electroencephalography-Based Clustering Reveals Robust Neurophysiological Subtypes in Parkinson's Disease.
Journal: Movement disorders : official journal of the Movement Disorder Society
In common: ICLabel, UMAP, MNE-Python, 6 other tools, EEG
[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, MNE-Python, 5 other tools, EEG
[10] doi:10.1371/journal.pone.0351872 [code]
Decoding visual object recognition from EEG signals.
Journal: PloS one
In common: MNE-Python, Pillow, PyTorch, 5 other tools, EEG, 2 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.