OSCR

Rotation-based metric on the Riemannian manifold of SPD matrices with applications to source data selection for brain-computer interface transfer learning.

Code ↔ Paper

3 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 3 matches
  1. [1] § Results › Pole ratio for BCI data ↔ rotation_analysis.py, lines 1266–1342 · score 0.62 · circle represents, intra subject accuracy, south pole, pole ratio, north, class
  2. [2] § Materials and methods › BCI data ↔ rotation_analysis.py, lines 171–245 · score 0.61 · motor imagery, covariance matrix, filtered, channels, MOABB, 35 Hz
  3. [3] § Materials and methods › Data split for classification evaluation ↔ rotation_analysis.py, lines 436–585 · score 0.51 · target training, splits, MDM, transfer learning, classification, transformed

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,746 lines · 59 KB · MIT · 3 matches

  1. #!python
  2. # venv: rotation
  3. # Imports
  4. # General Libraries
  5. import os
  6. import sys
  7. import time
  8. import math
  9. import random
  10. import datetime
  11. import subprocess
  12. import copy
  13. import warnings
  14. warnings.filterwarnings("ignore")
  15. warnings.simplefilter(action="ignore", category=FutureWarning) # Ignore futurewarning warnings.
  16. import shutil
  17. from pathlib import Path
  18. # Numpy Library
  19. import numpy as np
  20. # Making numpy single threaded, so I can multithread this program myself:
  21. os.environ["OMP_NUM_THREADS"] = "1" # export OMP_NUM_THREADS=4
  22. os.environ["OPENBLAS_NUM_THREADS"] = "1" # export OPENBLAS_NUM_THREADS=4
  23. os.environ["MKL_NUM_THREADS"] = "1" # export MKL_NUM_THREADS=6
  24. os.environ["VECLIB_MAXIMUM_THREADS"] = "1" # export VECLIB_MAXIMUM_THREADS=4
  25. os.environ["NUMEXPR_NUM_THREADS"] = "1" # export NUMEXPR_NUM_THREADS=6
  26. # Pandas Library
  27. # import pandas as pd
  28. # Pyrieman library
  29. from pyriemann.classification import MDM
  30. from pyriemann.estimation import Covariances
  31. from pyriemann.transfer import TLCenter,TLScale,TLRotate
  32. from pyriemann.utils.distance import distance
  33. from pyriemann.utils.mean import mean_covariance
  34. from pyriemann.tangentspace import TangentSpace
  35. from pyriemann.clustering import Kmeans
  36. # Scikit-learn library
  37. from sklearn.metrics import accuracy_score
  38. from sklearn.model_selection import StratifiedShuffleSplit
  39. from sklearn.pipeline import make_pipeline
  40. from sklearn import svm
  41. # Matplotlib Library
  42. import matplotlib as mpl
  43. import matplotlib.pyplot as plt
  44. from matplotlib.patches import CirclePolygon
  45. # Moabb
  46. from moabb.datasets import BNCI2014_001,PhysionetMI, Dreyer2023,Stieger2021,Liu2024,Lee2019_MI,Cho2017
  47. from moabb.paradigms import MotorImagery
  48. # Multiprocessing
  49. from multiprocessing import Pool
  50. # Scipy
  51. from scipy.optimize import curve_fit
  52. from scipy.stats import ttest_ind
  53. # Other
  54. from PIL import Image
  55. mpl.rcParams['font.family'] = "serif"
  56. plt.rcParams.update({'font.size': 15})
  57. random_state = 42
  58. rng = np.random.default_rng(seed=random_state)
  59. # GLOBAL PARAMETERS
  60. NBR_THREADS = 1 # Nbr of threads for parallellisation
  61. DATASETS = [
  62. (BNCI2014_001(),['feet','tongue'],'BCI-IV_feet_tongue'),
  63. (PhysionetMI(),['feet','hands'],'Physio_feet_hands'),
  64. (Dreyer2023(),['right_hand','left_hand'],'Dreyer_left_right'),
  65. (BNCI2014_001(),['right_hand','left_hand'],'BCI-IV_left_right'),
  66. (PhysionetMI(),['right_hand','left_hand'],'Physio_left_right'),
  67. (Lee2019_MI(),['right_hand','left_hand'],'Lee_left_right'),
  68. (Cho2017(),['right_hand','left_hand'],'Cho_left_right'),
  69. ]
  70. PROCESSING_DICT={
  71. 'channels': ['FC3', 'FC4', 'C5', 'C3', 'C1', 'C2', 'C4', 'C6', 'CP3', 'CP4'],
  72. 'fmin':7,
  73. 'fmax':35,
  74. 'tmin':1,
  75. 'tmax':2,
  76. 'baseline':None,
  77. }
  78. print(__doc__)
  79. class Logger(object):
  80. """For saving everything printed to file as well."""
  81. def __init__(self, filename="_console_log.txt"):
  82. """Initialize the class."""
  83. self.terminal = sys.stdout
  84. self.log = open(filename, "a")
  85. def write(self, message):
  86. """Write both to stdout and to file."""
  87. self.terminal.write(message)
  88. self.log.write(message)
  89. def flush(self):
  90. """We don't do flush."""
  91. pass
  92. def print_git_version():
  93. """We want to know what version of the source code we're dealing with."""
  94. # Ideally, commit to git _every time_ before you run the code.
  95. try:
  96. git_output = subprocess.check_output("git log -n 1", shell=True, stderr=subprocess.STDOUT)
  97. print("\n\n### git log -n 1\n%s\n\n" % git_output.decode("utf-8"))
  98. git_output2 = subprocess.check_output("git status", shell=True, stderr=subprocess.STDOUT)
  99. print("### git status\n%s\n\n" % git_output2.decode("utf-8"))
  100. git_output3 = subprocess.check_output("git diff", shell=True, stderr=subprocess.STDOUT)
  101. print("### git diff\n%s\n\n\n" % git_output3.decode("utf-8"))
  102. except:
  103. print("\n\n### WARNING: The code you're running is NOT under GIT version control. You can do better. Behave.\n\n")
  104. def save_source_code(run_logdir):
  105. print("\n\n###\n### Saving the python source code used, from the file '%s':\n###" % __file__)
  106. shutil.copy2(__file__, f'{run_logdir}/_script.py')
  107. def write_description(file_name='_description'):
  108. with open(f'{run_logdir}/{file_name}.txt', 'a') as f:
  109. # Write some text to the file
  110. f.write("This is a description file.\n")
  111. f.write(f'NBR_THREADS: {NBR_THREADS}\n')
  112. for key in PROCESSING_DICT.keys():
  113. f.write(f'{key}: {PROCESSING_DICT[key]}\n')
  114. for (dataset,events,identifier) in DATASETS:
  115. f.write(f'\nDataset: {dataset.code}')
  116. f.write(f'\nClasses: {"_".join(events)}')
  117. f.write(f'\nidentifier: {"_".join(identifier)}\n')
  118. # -----------------------------------------------
  119. # -----------------------------------------------
  120. # LOADING DATA : MOABB
  121. # -----------------------------------------------
  122. def get_subject_dataset_MOABB(subject,dataset,events,processing_dict,run_logdir='.',data_dir='.'):
  123. '''
  124. subject: int
  125. subject number
  126. dataset: object with MOABB dataset
  127. Dataset object.
  128. event: list of str
  129. Event names in the dataset to be used.
  130. '''
  131. print(f'\nLoading data subj: {subject}')
  132. folder_name = f'{data_dir}/subj_{subject}'
  133. os.makedirs(folder_name, exist_ok=True)
  134. full_file_path_keep = os.path.join(folder_name, 'keep.npy')
  135. full_file_path_covs = os.path.join(folder_name, 'covariance_matrices.npy')
  136. full_file_path_str_labels = os.path.join(folder_name, 'str_labels.npy')
  137. if os.path.exists(full_file_path_keep) and os.path.exists(full_file_path_covs) and os.path.exists(full_file_path_str_labels):
  138. print(f'Loading stored data for subjects')
  139. # Load covs
  140. covs = np.load(full_file_path_covs,allow_pickle=False)
  141. print(f'covs loaded from {full_file_path_covs}')
  142. # Load str_labels
  143. str_labels_filtered = np.load(full_file_path_str_labels,allow_pickle=False)
  144. print(f'str_labels loaded from {full_file_path_str_labels}')
  145. keep = np.load(full_file_path_keep,allow_pickle=False)
  146. print(f'keep loaded from {full_file_path_keep}')
  147. if not keep: # Print diagnostics to file to know why excluded.
  148. for cla in np.unique(str_labels_filtered):
  149. check_spd_matrices(covs[str_labels_filtered==cla], name=f'Subject: {subject} class: {cla}',run_logdir=run_logdir)
  150. else:
  151. keep = True # placeholder
  152. paradigm = MotorImagery(
  153. fmin=processing_dict['fmin'],#7,
  154. fmax=processing_dict['fmax'],#35,
  155. tmin=processing_dict['tmin'],#1.,
  156. tmax=processing_dict['tmax'],#2.,
  157. baseline=processing_dict['baseline'],#None,
  158. channels=processing_dict['channels'])#CHANNELS) # Here is preprocessing settings to be added. Also select channels.
  159. # https://moabb.neurotechx.com/docs/generated/moabb.paradigms.MotorImagery.html#moabb.paradigms.MotorImagery
  160. X, str_labels, metadata = paradigm.get_data(dataset=dataset, subjects=[subject])
  161. # Filter out data with wrong labels.
  162. mask = np.isin(str_labels, events) # boolean mask
  163. X_filtered = X[mask]
  164. str_labels_filtered = str_labels[mask]
  165. # Compute covariance matrices on scaled data
  166. covs = Covariances().fit_transform(X_filtered)
  167. print(f"\n\n === Diagnostics for Subject: {subject} === ")
  168. for cla in np.unique(str_labels_filtered):
  169. keep_cla = check_spd_matrices(covs[str_labels_filtered==cla], name=f'Subject: {subject} class: {cla}',run_logdir=run_logdir)
  170. if not keep_cla:
  171. keep = False
  172. print()
  173. # Saving data
  174. if not os.path.exists(full_file_path_covs):
  175. np.save(full_file_path_covs,covs,allow_pickle=False)
  176. if not os.path.exists(full_file_path_keep):
  177. np.save(full_file_path_keep,keep,allow_pickle=False)
  178. if not os.path.exists(full_file_path_str_labels):
  179. np.save(full_file_path_str_labels,str_labels_filtered,allow_pickle=False)
  180. return covs, str_labels_filtered, keep
  181. def check_spd_matrices(covs, name="X",run_logdir='.'):
  182. """
  183. Check symmetry, SPD condition, and condition numbers for a stack of matrices.
  184. Parameters
  185. ----------
  186. covs : ndarray, shape (n_matrices, n, n)
  187. Input matrices to check.
  188. name : str, optional
  189. Label for reporting.
  190. """
  191. n_covs = covs.shape[0]
  192. max_asym = []
  193. min_eigs = []
  194. conds = []
  195. for i, A in enumerate(covs):
  196. # enforce symmetry check
  197. asym = np.max(np.abs(A - A.T))
  198. max_asym.append(asym)
  199. # eigenvalues
  200. w = np.linalg.eigvalsh(0.5 * (A + A.T)) # symmetrize for safety
  201. min_eigs.append(np.min(w))
  202. # condition number (ratio of largest to smallest eigenvalue)
  203. cond = np.max(w) / np.min(w) if np.min(w) > 0 else np.inf
  204. conds.append(cond)
  205. max_asym = np.array(max_asym)
  206. min_eigs = np.array(min_eigs)
  207. conds = np.array(conds)
  208. mean_cov = mean_covariance(covs,metric='riemann')
  209. print(f" --- Diagnostics for {name} --- ")
  210. print(f"Number of matrices: {n_covs}")
  211. print(f'ratio max/min/mean | max: {(mean_cov.max()-mean_cov.mean())/(mean_cov.max()-mean_cov.min()):.2e}, min: {(mean_cov.mean()-mean_cov.min())/(mean_cov.max()-mean_cov.min()):.2e}')
  212. print(f"Min eigenvalue | min: {min_eigs.min():.2e}, mean: {min_eigs.mean():.2e}")
  213. print(f"Condition number | max: {conds.max():.2e}, mean: {conds.mean():.2e}")
  214. with open(f'{run_logdir}/_diagnostics.txt', 'a') as f:
  215. f.write(f"\n\n --- Diagnostics for {name} --- ")
  216. f.write(f"\nNumber of matrices: {n_covs}")
  217. f.write(f'\nratio max/min/mean | max: {(mean_cov.max()-mean_cov.mean())/(mean_cov.max()-mean_cov.min()):.2e}, min: {(mean_cov.mean()-mean_cov.min())/(mean_cov.max()-mean_cov.min()):.2e}')
  218. f.write(f"\nMin eigenvalue | min: {min_eigs.min():.2e}, mean: {min_eigs.mean():.2e}")
  219. f.write(f"\nCondition number | max: {conds.max():.2e}, mean: {conds.mean():.2e}")
  220. ratio_in_matrix = (mean_cov.max()-mean_cov.mean())/(mean_cov.max()-mean_cov.min())
  221. # flag outliers
  222. bad_sym = np.where(max_asym > 1e-10)[0]
  223. bad_spd = np.where(min_eigs <= 0.01)[0]
  224. bad_cond = np.where(conds > 1e8)[0]
  225. keep = True
  226. print()
  227. if ratio_in_matrix > 0.8:
  228. print(f"[!] bad ratio {ratio_in_matrix}.")
  229. keep = False
  230. if bad_sym.size:
  231. print(f"[!] {len(bad_sym)} matrices not symmetric (idx: {bad_sym[:10]})")
  232. if bad_spd.size:
  233. print(f"[!] {len(bad_spd)} matrices bad eig, idx: {bad_spd[:10]})")
  234. keep = False
  235. if bad_cond.size:
  236. print(f"[!] {len(bad_cond)} matrices ill-conditioned (cond > 1e8, idx: {bad_cond[:10]})")
  237. if not (bad_sym.size or bad_spd.size or bad_cond.size or ratio_in_matrix > 0.8):
  238. print("All matrices look fine.")
  239. return keep
  240. # -----------------------------------------------
  241. # Multiprocessing helper function:
  242. # -----------------------------------------------
  243. def worker_function_loading_data(iteration_idx,dataset,subjects,events,processing_dict,run_logdir='.',data_dir='.'):
  244. # This function will be passed to the Pool.
  245. print(f"### worker_function ({subjects[iteration_idx]}/{len(subjects)}) for itr_{iteration_idx}")
  246. covs, str_labels, keep = get_subject_dataset_MOABB(subject=subjects[iteration_idx],dataset=dataset,events=events,processing_dict=processing_dict,run_logdir=run_logdir,data_dir=data_dir)
  247. subject_data = {
  248. 'X_cov' : covs,
  249. 'y': str_labels,
  250. 'y_enc' : np.array([f'subj_{subjects[iteration_idx]}/{str(item)}' for item in str_labels]),
  251. 'subject_nbr': subjects[iteration_idx],
  252. 'subj_id': f'subj_{subjects[iteration_idx]}'
  253. }
  254. return subject_data,keep
  255. def load_data_MOABB(subjects,dataset,events,processing_dict,run_logdir='.',data_dir_identifier='.'):
  256. all_subject_data = []
  257. excluded_subjects = []
  258. data_dir = f'data/raw_data/{data_dir_identifier}'
  259. if NBR_THREADS == 1:
  260. print('\n-------------------------------')
  261. print( '--> Single thread started < ---')
  262. print( '-------------------------------\n')
  263. for subj in subjects:
  264. covs, str_labels,keep = get_subject_dataset_MOABB(subj,dataset,events,processing_dict,run_logdir,data_dir)
  265. if keep:
  266. subject_data = {
  267. 'X_cov' : covs,
  268. 'y': str_labels,
  269. 'y_enc' : np.array([f'subj_{subj}/{str(item)}' for item in str_labels]),
  270. 'subject_nbr': subj,
  271. 'subj_id': f'subj_{subj}'
  272. }
  273. all_subject_data.append(subject_data)
  274. else:
  275. excluded_subjects.append(int(subj))
  276. else:
  277. # Multithreaded version:
  278. print('\n---------------------------------')
  279. print( '--> Multiprocessing started < ---')
  280. print( '---------------------------------\n')
  281. pool = Pool(processes=NBR_THREADS)
  282. results = []
  283. for iteration_idx in range(len(subjects)):
  284. res = pool.apply_async(worker_function_loading_data, args=(iteration_idx,dataset,subjects,events,processing_dict,run_logdir,data_dir))
  285. results.append(res)
  286. if res.get()[1]:
  287. all_subject_data.append(res.get()[0])
  288. else:
  289. excluded_subjects.append(int(subjects[iteration_idx]))
  290. print('---> Closing pools\n')
  291. pool.close()
  292. pool.join()
  293. print('---> Pools closed \n')
  294. print("!!!!!!!!!!!!!!!!!!!!!!!!!!!")
  295. print("!!!! EXCLUDED SUBJECTS !!!!")
  296. print(excluded_subjects)
  297. print("!!!!!!!!!!!!!!!!!!!!!!!!!!!")
  298. with open(f'{run_logdir}/_diagnostics.txt', 'a') as f:
  299. f.write(f"\n\n!!!!!!!!!!!!!!!!!!!!!!!!!!!")
  300. f.write(f"\n!!!! EXCLUDED SUBJECTS !!!!")
  301. f.write(f"\n{excluded_subjects}")
  302. f.write(f"\n!!!!!!!!!!!!!!!!!!!!!!!!!!!")
  303. return all_subject_data
  304. # ################################
  305. # ######## ANALYSIS #######
  306. # ################################
  307. def run_one_subject(data,iteration_idx,run_logdir='.',data_dir_identifier='.'):
  308. # get data
  309. X_covs = data[iteration_idx]['X_cov']
  310. y = data[iteration_idx]['y']
  311. subj_id = data[iteration_idx]['subj_id']
  312. subject_nbr = data[iteration_idx]['subject_nbr']
  313. print(f'\nSubj: {subject_nbr}')
  314. find_distance_between_classes_to_pole(data,iteration_idx,run_logdir,data_dir_identifier)
  315. find_transfer_learning_accuracies(data,iteration_idx,run_logdir,data_dir_identifier)
  316. return
  317. # ---- transfer learning
  318. def transfer_learning_one_target_vs_one_source(data, target_idx, source_idx,run_logdir='.',data_dir='.',data_dir_identifier='.'):
  319. X_covs_target = data[target_idx]['X_cov']
  320. y_target = data[target_idx]['y']
  321. y_enc_target = data[target_idx]['y_enc']
  322. subj_id_target = data[target_idx]['subj_id']
  323. subject_nbr_target = data[target_idx]['subject_nbr']
  324. X_covs_source = data[source_idx]['X_cov']
  325. y_source = data[source_idx]['y']
  326. y_enc_source = data[source_idx]['y_enc']
  327. subj_id_source = data[source_idx]['subj_id']
  328. subject_nbr_source = data[source_idx]['subject_nbr']
  329. full_path_accuracies = f"{data_dir}/{subj_id_source}_accuracies.npy"
  330. if os.path.exists(full_path_accuracies):
  331. # If it exists do nothing.
  332. print(f'Accuracies for {subj_id_target} with {subj_id_source} already exist, skipping...')
  333. scores = np.load(full_path_accuracies,allow_pickle=True).item()
  334. return scores
  335. clf = MDM(metric=dict(mean='riemann', distance='riemann'))
  336. splitter = StratifiedShuffleSplit(n_splits=4, train_size=0.75, random_state=42) # Shuffels data.
  337. nothing_score = np.zeros(splitter.get_n_splits())
  338. recenter_score = np.zeros(splitter.get_n_splits())
  339. scale_score = np.zeros(splitter.get_n_splits())
  340. rotate_score = np.zeros(splitter.get_n_splits())
  341. for i, (train_index,test_index) in enumerate(splitter.split(X_covs_target, y_target)):
  342. X_target_train = copy.deepcopy(X_covs_target[train_index])
  343. y_target_train = copy.deepcopy(y_target[train_index])
  344. y_enc_target_train = copy.deepcopy(y_enc_target[train_index])
  345. X_source_train = copy.deepcopy(X_covs_source)
  346. y_source_train = copy.deepcopy(y_source)
  347. y_enc_source_train = copy.deepcopy(y_enc_source)
  348. X_test = copy.deepcopy(X_covs_target[test_index])
  349. y_test = copy.deepcopy(y_target[test_index])
  350. y_enc_test = copy.deepcopy(y_enc_target[test_index])
  351. nbr_target_training_data = len(y_target_train)
  352. if target_idx == source_idx: # Intra subject
  353. clf.fit(X_target_train,y_target_train)
  354. y_pred = clf.predict(X_test)
  355. score = accuracy_score(y_test, y_pred)
  356. nothing_score[i] = score
  357. recenter_score[i] = score
  358. scale_score[i] = score
  359. rotate_score[i] = score
  360. continue # Continue with next split.
  361. # Format data in correct way for transfer learning
  362. X_train = np.concatenate((X_source_train,X_target_train))
  363. y_enc = np.concatenate((y_enc_source_train,y_enc_target_train))
  364. y_train = np.concatenate((y_source_train,y_target_train))
  365. # Before any transfer learning
  366. clf.fit(X_train[:(-nbr_target_training_data)],y_train[:(-nbr_target_training_data)])
  367. y_pred = clf.predict(X_test)
  368. nothing_score[i] = accuracy_score(y_test, y_pred)
  369. # Transfer learnig steps:
  370. tl_recenter = TLCenter(target_domain=subj_id_target,metric='riemann')
  371. tl_scale = TLScale(target_domain=subj_id_target, final_dispersion=1.0, centered_data=True, metric='riemann')
  372. tl_rotate = TLRotate(target_domain=subj_id_target, metric='riemann',tol_step=1e-9, maxiter=10000)
  373. # === RECENTER ===
  374. # Recenter the data.
  375. X_train = tl_recenter.fit_transform(X_train, y_enc)
  376. X_test = tl_recenter.transform(X_test)
  377. clf.fit(X_train[:(-nbr_target_training_data)],y_train[:(-nbr_target_training_data)])
  378. y_pred = clf.predict(X_test)
  379. recenter_score[i] = accuracy_score(y_test, y_pred)
  380. # === SCALE ===
  381. X_train = tl_scale.fit_transform(X_train, y_enc)
  382. X_test = tl_scale.transform(X_test)
  383. clf.fit(X_train[:(-nbr_target_training_data)],y_train[:(-nbr_target_training_data)])
  384. y_pred = clf.predict(X_test)
  385. scale_score[i] = accuracy_score(y_test, y_pred)
  386. # === ROTATE ===
  387. with warnings.catch_warnings(record=True) as w:
  388. warnings.simplefilter("always") # capture all warnings
  389. str_test_data = '_'.join(test_index.astype(str))
  390. rotations_path = f"data/rotations/{data_dir_identifier}/{subj_id_target}/{subj_id_source}_iter_{i}.npy"
  391. test_data_path = f"data/rotations/{data_dir_identifier}/{subj_id_target}/{subj_id_source}_iter_{i}_test_data.npy"
  392. convergece_file = f"data/rotations/{data_dir_identifier}/{subj_id_target}/convergence_info.txt"
  393. if os.path.exists(rotations_path) and os.path.exists(test_data_path):
  394. test_str = np.load(test_data_path,allow_pickle=False)
  395. if test_str == str_test_data:
  396. load_rotation = True
  397. else:
  398. bug # You have to remove rotational data since the test/target split are wrong.
  399. load_rotation = False
  400. else:
  401. load_rotation = False
  402. if load_rotation:
  403. Q_rot = np.load(rotations_path,allow_pickle=False)
  404. X_to_rotate = X_train[:-nbr_target_training_data] # The first samples in the training vector are source data. Only source data is rotated
  405. X_rotated = Q_rot @ X_to_rotate @ Q_rot.T
  406. X_train[:(-nbr_target_training_data)] = X_rotated
  407. print(f'Rotation loaded from {rotations_path}')
  408. else:
  409. X_train = tl_rotate.fit_transform(X_train, y_enc)
  410. # Save rotations
  411. if subj_id_target != subj_id_source:
  412. Q_rot = tl_rotate.rotations_[subj_id_source]
  413. os.makedirs(f'data/rotations/{data_dir_identifier}/{subj_id_target}', exist_ok=True)
  414. np.save(f"{rotations_path}", Q_rot,allow_pickle=False)
  415. np.save(f"{test_data_path}", str_test_data,allow_pickle=False)
  416. # print(f'Rotation saved to {rotations_path}')
  417. if w: # there were warnings
  418. for warn in w:
  419. print(f"Target: {subj_id_target} source: {subj_id_source} iter:{i}: Warning: {warn.message}")
  420. with open(convergece_file, 'a') as f:
  421. f.write(f"Target: {subj_id_target} source: {subj_id_source} iter:{i}: Warning: {warn.message}\n")
  422. else:
  423. print(f"Target: {subj_id_target} source: {subj_id_source} iter:{i}: --------------------")
  424. with open(convergece_file, 'a') as f:
  425. f.write(f"Target: {subj_id_target} source: {subj_id_source} iter:{i}: --------------------\n")
  426. clf.fit(X_train[:(-nbr_target_training_data)],y_train[:(-nbr_target_training_data)])
  427. y_pred = clf.predict(X_test)
  428. rotate_score[i] = accuracy_score(y_test, y_pred)
  429. scores = {
  430. 'nothing_score': nothing_score.mean(),
  431. 'recenter_score': recenter_score.mean(),
  432. 'scale_score': scale_score.mean(),
  433. 'rotate_score': rotate_score.mean(),
  434. }
  435. np.save(full_path_accuracies, scores,allow_pickle=True)
  436. return scores
  437. def find_transfer_learning_accuracies(data,iteration_idx,run_logdir='.',data_dir_identifier='.'):
  438. X_covs_target = data[iteration_idx]['X_cov']
  439. y_target = data[iteration_idx]['y']
  440. y_enc_target = data[iteration_idx]['y_enc']
  441. subj_id_target = data[iteration_idx]['subj_id']
  442. subject_nbr_target = data[iteration_idx]['subject_nbr']
  443. data_dir = f'data/transfer_learning/{data_dir_identifier}/{subj_id_target}'
  444. print(f'Creating data folder: {data_dir}')
  445. os.makedirs(f'{data_dir}', exist_ok=True)
  446. full_path_accuracies = f"{data_dir}/all_subj_accuracies.npy"
  447. if os.path.exists(full_path_accuracies):
  448. # If it exists do nothing.
  449. print(f'Accuracies for {subj_id_target} for all subj already exist, skipping...')
  450. return
  451. score_all_subj={}
  452. for source_idx in range(len(data)):
  453. subj_id_source = data[source_idx]['subj_id']
  454. scores = transfer_learning_one_target_vs_one_source(data=data, target_idx=iteration_idx, source_idx=source_idx,run_logdir=run_logdir,data_dir=data_dir,data_dir_identifier=data_dir_identifier)
  455. score_all_subj[subj_id_source] = scores
  456. np.save(full_path_accuracies,score_all_subj,allow_pickle=True)
  457. return
  458. # ---- distance data
  459. def find_distance_between_classes_to_pole(data,iteration_idx,run_logdir='.',data_dir_identifier='.'):
  460. X_covs = data[iteration_idx]['X_cov']
  461. y = data[iteration_idx]['y']
  462. y_enc = data[iteration_idx]['y_enc']
  463. subj_id = data[iteration_idx]['subj_id']
  464. subject_nbr = data[iteration_idx]['subject_nbr']
  465. data_dir = f'data/distance_pole/{data_dir_identifier}/{subj_id}'
  466. print(f'Creating data folder: {data_dir}')
  467. os.makedirs(f'{data_dir}', exist_ok=True)
  468. size_matrix = X_covs[0].shape[0]
  469. classes = np.unique(y)
  470. full_path_distances = f"{data_dir}/distances.npy"
  471. if os.path.exists(full_path_distances):
  472. # If it exists do nothing.
  473. print(f'Disntances for {subj_id} already exist, skipping...')
  474. return
  475. tl_recenter = TLCenter(target_domain=subj_id,metric='riemann')
  476. tl_scale = TLScale(target_domain=subj_id, final_dispersion=1.0, centered_data=True, metric='riemann')
  477. X = tl_recenter.fit_transform(X_covs, y_enc)
  478. X = tl_scale.fit_transform(X, y_enc)
  479. ref = mean_covariance(X)
  480. X_0 = X[y==classes[0]]
  481. X_1 = X[y==classes[1]]
  482. X_mean_0 = mean_covariance(X_0)
  483. X_mean_1 = mean_covariance(X_1)
  484. dist_0 = distance(ref,X_mean_0, metric='riemann', squared=True)
  485. dist_1 = distance(ref,X_mean_1, metric='riemann', squared=True)
  486. pole_north_0 = np.eye(size_matrix)* np.exp(+np.sqrt(dist_0/size_matrix))
  487. pole_south_0 = np.eye(size_matrix)* np.exp(-np.sqrt(dist_0/size_matrix))
  488. pole_north_1 = np.eye(size_matrix)* np.exp(+np.sqrt(dist_1/size_matrix))
  489. pole_south_1 = np.eye(size_matrix)* np.exp(-np.sqrt(dist_1/size_matrix))
  490. # Find intra subject accuracy
  491. clf = MDM(metric=dict(mean='riemann', distance='riemann'))
  492. splitter = StratifiedShuffleSplit(n_splits=4, train_size=0.75, random_state=42) # Shuffels data. Test size = training size in our case.
  493. all_score = np.zeros(splitter.get_n_splits())
  494. for i, (train_index,test_index) in enumerate(splitter.split(X_covs, y)):
  495. X_train = X_covs[train_index]
  496. y_train = y[train_index]
  497. X_test = X_covs[test_index]
  498. y_test = y[test_index]
  499. clf.fit(X_train,y_train)
  500. y_pred = clf.predict(X_test)
  501. all_score[i] = accuracy_score(y_test, y_pred)
  502. distances = {
  503. 'subj_id': subj_id,
  504. 'Intra-subject accuracy':all_score.mean(),
  505. 'ref-mean:0': distance(ref,X_mean_0, metric='riemann', squared=False), # Same distance as for poles
  506. 'mean-southpole:0':distance(X_mean_0,pole_south_0, metric='riemann', squared=False),
  507. 'mean-northpole:0':distance(X_mean_0,pole_north_0, metric='riemann', squared=False),
  508. 'ref-mean:1': distance(ref,X_mean_1, metric='riemann', squared=False), # Same distance as for poles
  509. 'mean-southpole:1':distance(X_mean_1,pole_south_1, metric='riemann', squared=False),
  510. 'mean-northpole:1':distance(X_mean_1,pole_north_1, metric='riemann', squared=False),
  511. }
  512. np.save(full_path_distances,distances,allow_pickle=True)
  513. return
  514. # -----------------------------------------------
  515. # Multiprocessing helper function:
  516. # -----------------------------------------------
  517. def worker_function(data,iteration_idx,run_logdir='.',data_dir_identifier='.'):
  518. # This function will be passed to the Pool.
  519. print(f"### worker_function ({iteration_idx+1}/{len(data)}) for itr_{iteration_idx}")
  520. # return run_one_iteration(n_matrices, mean,sigma, nbr_outliers,iteration_nbr,nbr_folds=nbr_folds,run_logdir=run_logdir,plotting=plotting)
  521. return run_one_subject(data,iteration_idx,run_logdir=run_logdir,data_dir_identifier=data_dir_identifier)
  522. # -----------------------------------------------
  523. # Run code single or multithreaded
  524. # -----------------------------------------------
  525. def run_code(data,run_logdir='.',identifier='',data_dir_identifier='.'):
  526. if NBR_THREADS == 1:
  527. print('\n-------------------------------')
  528. print( '--> Single thread started < ---')
  529. print( '-------------------------------\n')
  530. for iteration_idx in range(len(data)):
  531. run_one_subject(data,iteration_idx,run_logdir=run_logdir,data_dir_identifier=data_dir_identifier)
  532. else:
  533. # Multithreaded version:
  534. print('\n---------------------------------')
  535. print( '--> Multiprocessing started < ---')
  536. print( '---------------------------------\n')
  537. pool = Pool(processes=NBR_THREADS)
  538. results = []
  539. for iteration_idx in range(len(data)):
  540. res = pool.apply_async(worker_function, args=(data,iteration_idx,run_logdir,data_dir_identifier))
  541. results.append(res)
  542. print('---> Closing pools\n')
  543. pool.close()
  544. pool.join()
  545. print('---> Pools closed \n')
  546. all_res_dict = {}
  547. dist_dict_res = analyse_distance_data(data,data_dir_identifier,run_logdir)
  548. all_res_dict.update(dist_dict_res)
  549. acc_dict_res = analyse_accuracies_for_RPA_steps(data,data_dir_identifier,run_logdir)
  550. all_res_dict.update(acc_dict_res)
  551. return all_res_dict
  552. # ------- Accuracies
  553. def analyse_accuracies_for_RPA_steps(data, data_dir_identifier='.',run_logdir='.'):
  554. print('Data analysis for accuracies')
  555. accuracies_all_subj = []
  556. distances_all_subj = []
  557. list_of_subj = [item['subj_id'] for item in data]
  558. data_dir_acc = f'data/transfer_learning/{data_dir_identifier}'
  559. data_dir_dist = f'data/distance_pole/{data_dir_identifier}'
  560. print('Loading data from iterations...')
  561. for subj in list_of_subj:
  562. subdir = subj
  563. print(f'loading data from {subdir}')
  564. # == Accuracy data
  565. subdir_path = os.path.join(data_dir_acc, subdir)
  566. # Check if the path is indeed a directory
  567. if not os.path.isdir(subdir_path):
  568. print(f'Not a folder: {subdir}')
  569. continue # There might be a .DS_store.
  570. file_path = os.path.join(subdir_path, 'all_subj_accuracies.npy')
  571. # Check if the file exists
  572. if not os.path.exists(file_path):
  573. import pdb; pdb.set_trace()
  574. bug # There is something wrong! The file is missing.
  575. data_here = np.load(file_path,allow_pickle=True).item()
  576. accuracies_all_subj.append(data_here)
  577. # == Dist data
  578. subdir_path = os.path.join(data_dir_dist, subdir)
  579. # Check if the path is indeed a directory
  580. if not os.path.isdir(subdir_path):
  581. print(f'Not a folder: {subdir}')
  582. continue # There might be a .DS_store.
  583. file_path = os.path.join(subdir_path, 'distances.npy')
  584. # Check if the file exists
  585. if not os.path.exists(file_path):
  586. import pdb; pdb.set_trace()
  587. bug # There is something wrong! The file is missing.
  588. data_here = np.load(file_path,allow_pickle=True).item()
  589. distances_all_subj.append(data_here)
  590. print('...done')
  591. print('Plotting...')
  592. all_res_dict = {}
  593. plot_accuracies_data(accuracies_all_subj,list_of_subj,run_logdir=run_logdir,data_dir_identifier=data_dir_identifier)
  594. p_val_dict = plot_accuracy_vs_radius_and_latitude(accuracies_all_subj,distances_all_subj,list_of_subj,run_logdir=run_logdir,data_dir_identifier=data_dir_identifier)
  595. all_res_dict['p-val'] = p_val_dict
  596. print('... done')
  597. return all_res_dict
  598. def plot_accuracy_vs_radius_and_latitude(data_acc,data_dist,list_of_subj,run_logdir='.',data_dir_identifier=''):
  599. # --- accuracy matrix ---
  600. nbr_subjects = len(data_acc)
  601. # Placeholders
  602. nothing_matrix = np.zeros((nbr_subjects,nbr_subjects))
  603. recenter_matrix = np.zeros((nbr_subjects,nbr_subjects))
  604. scale_matrix = np.zeros((nbr_subjects,nbr_subjects))
  605. rotate_matrix = np.zeros((nbr_subjects,nbr_subjects))
  606. # Create matrices
  607. for row, row_dict in enumerate(data_acc):
  608. for col, source_id in enumerate(list_of_subj):
  609. nothing_matrix[row,col] = row_dict[source_id]['nothing_score']
  610. recenter_matrix[row,col] = row_dict[source_id]['recenter_score']
  611. scale_matrix[row,col] = row_dict[source_id]['scale_score']
  612. rotate_matrix[row,col] = row_dict[source_id]['rotate_score']
  613. if data_dir_identifier == 'BNCI2014-001_feet_tongue':
  614. title = 'BCI-IV 2a (feet/tongue)'
  615. extra_str_save = 'BCI_IV_feet_tongue'
  616. elif data_dir_identifier == 'BNCI2014-001_left_hand_right_hand':
  617. title = 'BCI-IV 2a (left/right)'
  618. extra_str_save = 'BCI_IV_left_right'
  619. elif data_dir_identifier == 'Cho2017_left_hand_right_hand':
  620. title = 'Cho (left/right)'
  621. extra_str_save = 'Cho'
  622. elif data_dir_identifier == 'Dreyer2023_left_hand_right_hand':
  623. title = 'Dreyer (left/right)'
  624. extra_str_save = 'Dreyer'
  625. elif data_dir_identifier == 'Lee2019-MI_left_hand_right_hand':
  626. title = 'Lee (left/right)'
  627. extra_str_save = 'Lee'
  628. elif data_dir_identifier == 'PhysionetMotorImagery_feet_hands':
  629. title = 'Physionet (feet/hands)'
  630. extra_str_save = 'physionet_feet_hands'
  631. elif data_dir_identifier == 'PhysionetMotorImagery_left_hand_right_hand':
  632. title = 'Physionet (left/right)'
  633. extra_str_save = 'physionet_left_right'
  634. else:
  635. title = ''
  636. extra_str_save = ''
  637. p_val_dict = plot_pole_ratio_to_acc_correlation_all_subj_for_paper(rotate_matrix,data_dist,list_of_subj,run_logdir=run_logdir,figname=f'lat_to_accuracy_correlation_{extra_str_save}',title=title)
  638. return p_val_dict
  639. def plot_pole_ratio_to_acc_correlation_all_subj_for_paper(accuracy_matrix,data_dist,list_of_subj,run_logdir='.',figname='fig',title=''):
  640. # ----- All subject in one plot -----
  641. nbr_subjects = len(list_of_subj)
  642. distance_vector = np.zeros(nbr_subjects)
  643. for i, subj_id in enumerate(list_of_subj):
  644. distance_vector[i] = data_dist[i]['mean-southpole:0'] / (data_dist[i]['mean-northpole:0'] + data_dist[i]['mean-southpole:0'] )
  645. distance_matrix = np.tile(distance_vector, (nbr_subjects, 1))
  646. diff_to_target_matrix = distance_matrix - distance_vector[:,None] # Source- target
  647. accuracy_matrix_adjusted = accuracy_matrix - np.diag(accuracy_matrix)[:,None]
  648. mask = (np.diag(accuracy_matrix) > 0.55) & (np.diag(accuracy_matrix) < 0.95)
  649. masked_accuracy_matrix_adjusted = accuracy_matrix_adjusted[mask][:,mask]
  650. masked_diff_to_target_matrix = diff_to_target_matrix[mask][:,mask]
  651. fig, axs = plt.subplots(nrows=1,ncols=1,figsize=(5,4),layout="constrained")
  652. # Highlight areas for statistic significance.
  653. axs.axvspan(
  654. xmin=-0.2,
  655. xmax=0.2,
  656. color='lightgray',
  657. alpha=0.5,
  658. edgecolor=None,
  659. )
  660. axs.axvspan(
  661. xmin=-0.1,
  662. xmax=0.1,
  663. color='darkgray',
  664. alpha=0.5,
  665. edgecolor=None,
  666. )
  667. axs.axvspan(
  668. xmin=-0.05,
  669. xmax=0.05,
  670. color='dimgrey',
  671. alpha=0.5,
  672. edgecolor=None,
  673. )
  674. # Plot data
  675. x_to_plot = masked_diff_to_target_matrix.ravel()
  676. y_to_plot = masked_accuracy_matrix_adjusted.ravel() * 100 # * 100 to get %
  677. noise = (rng.random(y_to_plot.shape) - 0.5)*0.05 * 100 # * 100 to get %
  678. print(f'Included subjects in plot_pole_ratio_to_acc_correlation_one_subj() plot: {np.sum(mask)}')
  679. print(np.arange(1,len(mask)+1)[mask])
  680. plot_pole_ratio_to_acc_correlation_one_subj(x_to_plot, y_to_plot+noise, title='',ax=axs)
  681. # Fit line to data.
  682. def model(x, a,b,c):
  683. return a*x**2 + b*x + c
  684. popt, pcov = curve_fit(
  685. model,
  686. x_to_plot,
  687. y_to_plot,
  688. p0=[-1.0,0.0,0.0],
  689. )
  690. a_hat = popt[0]
  691. b_hat = popt[1]
  692. c_hat = popt[2]
  693. # Plot fitted line
  694. x_smooth = np.linspace(x_to_plot.min(), x_to_plot.max(), 300)
  695. y_smooth = model(x_smooth, *popt)
  696. axs.plot(x_smooth, y_smooth, color='black', label=f'{a_hat:.3f}x**2 + {b_hat:.3f}x + {c_hat:.3f}')
  697. plt.title(title,fontsize=12)
  698. plt.ylabel(r'$\Delta$Acc (Improvement vs Intra-subject accuracy)',fontsize=10)
  699. plt.xlabel(r'$\Delta \rho$ (Difference in pole ratio)',fontsize=10)
  700. save_string = f'{figname}' # Filename
  701. plt.savefig(f'{run_logdir}/png_{save_string}.png',dpi=600) # Save the figure in the specified directory
  702. plt.savefig(f'{run_logdir}/pdf_{save_string}.pdf') # Save the figure in the specified directory
  703. plt.close(fig)
  704. p_test_res = {}
  705. # hypotestestning
  706. mask_1 = (np.abs(x_to_plot) <= 0.2)
  707. mask_2 = (np.abs(x_to_plot) > 0.2)
  708. y_bin_0 = y_to_plot[mask_1]
  709. y_bin_1 = y_to_plot[mask_2]
  710. _, p_ttest = ttest_ind(y_bin_0, y_bin_1, equal_var=False,alternative='greater')
  711. p_test_res['x<=0.2|0.2<x'] = p_ttest
  712. p_test_res['Val-diff: x<=0.2|0.2<x'] = np.mean(y_bin_0) - np.mean(y_bin_1)
  713. mask_1 = (np.abs(x_to_plot) <= 0.1)
  714. mask_2 = (np.abs(x_to_plot) > 0.1)
  715. y_bin_0 = y_to_plot[mask_1]
  716. y_bin_1 = y_to_plot[mask_2]
  717. _, p_ttest = ttest_ind(y_bin_0, y_bin_1, equal_var=False,alternative='greater')
  718. p_test_res['x<=0.1|0.1<x'] = p_ttest
  719. p_test_res['Val-diff: x<=0.1|0.1<x'] = np.mean(y_bin_0) - np.mean(y_bin_1)
  720. mask_1 = (np.abs(x_to_plot) <= 0.05)
  721. mask_2 = (np.abs(x_to_plot) > 0.05)
  722. y_bin_0 = y_to_plot[mask_1]
  723. y_bin_1 = y_to_plot[mask_2]
  724. _, p_ttest = ttest_ind(y_bin_0, y_bin_1, equal_var=False,alternative='greater')
  725. p_test_res['x<=0.05|0.05<x'] = p_ttest
  726. p_test_res['Val-diff: x<=0.05|0.05<x'] = np.mean(y_bin_0) - np.mean(y_bin_1)
  727. mask_1 = (np.abs(x_to_plot) <= 0.1)
  728. mask_2 = (np.abs(x_to_plot) > 0.1) & (np.abs(x_to_plot) <= 0.2)
  729. y_bin_0 = y_to_plot[mask_1]
  730. y_bin_1 = y_to_plot[mask_2]
  731. _, p_ttest = ttest_ind(y_bin_0, y_bin_1, equal_var=False,alternative='greater')
  732. p_test_res['x<=0.1|0.1<x<=0.2'] = p_ttest
  733. p_test_res['Val-diff: x<=0.1|0.1<x<=0.2'] = np.mean(y_bin_0) - np.mean(y_bin_1)
  734. mask_1 = (np.abs(x_to_plot) <= 0.05)
  735. mask_2 = (np.abs(x_to_plot) > 0.05) & (np.abs(x_to_plot) <= 0.2)
  736. y_bin_0 = y_to_plot[mask_1]
  737. y_bin_1 = y_to_plot[mask_2]
  738. _, p_ttest = ttest_ind(y_bin_0, y_bin_1, equal_var=False,alternative='greater')
  739. p_test_res['x<=0.05|0.05<x<=0.2'] = p_ttest
  740. p_test_res['Val-diff: x<=0.05|0.05<x<=0.2'] = np.mean(y_bin_0) - np.mean(y_bin_1)
  741. mask_1 = (np.abs(x_to_plot) <= 0.05)
  742. mask_2 = (np.abs(x_to_plot) > 0.05) & (np.abs(x_to_plot) <= 0.1)
  743. y_bin_0 = y_to_plot[mask_1]
  744. y_bin_1 = y_to_plot[mask_2]
  745. _, p_ttest = ttest_ind(y_bin_0, y_bin_1, equal_var=False,alternative='greater')
  746. p_test_res['x<=0.05|0.05<x<=0.1'] = p_ttest
  747. p_test_res['Val-diff: x<=0.05|0.05<x<=0.1'] = np.mean(y_bin_0) - np.mean(y_bin_1)
  748. return p_test_res
  749. def plot_pole_ratio_to_acc_correlation_one_subj(diff_to_target,accuracy_vector,title,ax):
  750. print(f'Plotting {title} lattitude to accuracy correlation')
  751. ax.scatter(diff_to_target,accuracy_vector,alpha=0.25,color='tab:red')
  752. ax.set_title(title)
  753. ax.axvline(0,color='black',linestyle=':')
  754. ax.axhline(0,color='black',linestyle=':')
  755. return
  756. def plot_accuracies_data(data,list_of_subj,run_logdir='.',data_dir_identifier=''):
  757. nbr_subjects = len(data)
  758. # Placeholders
  759. nothing_matrix = np.zeros((nbr_subjects,nbr_subjects))
  760. recenter_matrix = np.zeros((nbr_subjects,nbr_subjects))
  761. scale_matrix = np.zeros((nbr_subjects,nbr_subjects))
  762. rotate_matrix = np.zeros((nbr_subjects,nbr_subjects))
  763. # Create matrices
  764. for row, row_dict in enumerate(data):
  765. for col, source_id in enumerate(list_of_subj):
  766. nothing_matrix[row,col] = row_dict[source_id]['nothing_score']
  767. recenter_matrix[row,col] = row_dict[source_id]['recenter_score']
  768. scale_matrix[row,col] = row_dict[source_id]['scale_score']
  769. rotate_matrix[row,col] = row_dict[source_id]['rotate_score']
  770. mask = (np.diag(nothing_matrix) > 0.55) & (np.diag(nothing_matrix) < 0.95)
  771. list_of_matrices = [nothing_matrix[mask][:,mask],recenter_matrix[mask][:,mask],scale_matrix[mask][:,mask],rotate_matrix[mask][:,mask]]
  772. sub_titles = ['No transfer learning','Recentering','Scale','Rotate (full RPA)']
  773. plot_four_matrices_for_paper(list_of_matrices,sub_titles,run_logdir=run_logdir,filename='four_accuracies',vmin=0.5,vmax=1,title_identifier=data_dir_identifier)
  774. return
  775. def plot_four_matrices_for_paper(list_of_matrices,sub_titles,run_logdir='.',filename='fig',vmin=0.5,vmax=1,title_identifier=''):
  776. figsize = (10,3.25)
  777. fig, axs = plt.subplots(figsize=figsize,ncols=4,nrows=1,layout='constrained') #
  778. for i, ax in enumerate(axs.flatten()):
  779. matrix_to_plot = list_of_matrices[i]
  780. row_sums = np.sum(matrix_to_plot, axis=1) # Sum of rows
  781. col_sums = np.sum(matrix_to_plot, axis=0) # Sum of columns
  782. row_sorted_indices = np.argsort(-row_sums) # Indices to sort rows in descending order
  783. col_sorted_indices = np.argsort(-col_sums)
  784. matrix_to_plot = matrix_to_plot[row_sorted_indices, :]
  785. matrix_to_plot = matrix_to_plot[:, col_sorted_indices]
  786. cax_0 = ax.matshow(matrix_to_plot, cmap='binary',vmin=vmin,vmax=vmax)
  787. # ax.set_title(sub_titles[i],fontsize=12)
  788. ax.set_xticks([])
  789. ax.set_yticks([])
  790. if i ==0:
  791. ax.set_ylabel('Target users',fontsize=10)
  792. ax.set_xlabel('Source users',fontsize=10)
  793. ax.grid(False)
  794. ax.set_title(sub_titles[i],fontsize=12)
  795. if title_identifier == 'BNCI2014-001_feet_tongue':
  796. title = 'BCI-IV 2a (feet/tongue)'
  797. extra_str_save = 'BCI_IV_feet_tongue'
  798. elif title_identifier == 'BNCI2014-001_left_hand_right_hand':
  799. title = 'BCI-IV 2a (left/right)'
  800. extra_str_save = 'BCI_IV_left_right'
  801. elif title_identifier == 'Cho2017_left_hand_right_hand':
  802. title = 'Cho (left/right)'
  803. extra_str_save = 'Cho'
  804. elif title_identifier == 'Dreyer2023_left_hand_right_hand':
  805. title = 'Dreyer (left/right)'
  806. extra_str_save = 'Dreyer'
  807. elif title_identifier == 'Lee2019-MI_left_hand_right_hand':
  808. title = 'Lee (left/right)'
  809. extra_str_save = 'Lee'
  810. elif title_identifier == 'PhysionetMotorImagery_feet_hands':
  811. title = 'Physionet (feet/hands)'
  812. extra_str_save = 'physionet_feet_hands'
  813. elif title_identifier == 'PhysionetMotorImagery_left_hand_right_hand':
  814. title = 'Physionet (left/right)'
  815. extra_str_save = 'physionet_left_right'
  816. else:
  817. title = ''
  818. extra_str_save = ''
  819. plt.suptitle(title,fontsize=14)
  820. # Colorbar
  821. norm = mpl.colors.Normalize(vmin=50, vmax=100)
  822. cmap = mpl.colormaps['binary']
  823. sm = mpl.cm.ScalarMappable(norm=norm, cmap=cmap)
  824. sm.set_array([]) # required but the content is irrelevant
  825. # one colorbar for the whole figure
  826. cbar = fig.colorbar(sm, ax=ax, shrink = 0.75)
  827. cbar.set_label("Accuracy [%]",fontsize=10)
  828. # optional: custom ticks
  829. cbar.set_ticks([50, 60, 70, 80, 90, 100])
  830. save_string = f'{filename}_{extra_str_save}' # Filename
  831. plt.savefig(f'{run_logdir}/png_{save_string}.jpg',dpi=600) # Save the figure in the specified directory
  832. plt.close(fig)
  833. return
  834. # ------ Distance data
  835. def analyse_distance_data(data, data_dir_identifier='.',run_logdir='.'):
  836. print('Data analysis for distances')
  837. distances_all_subj =[]
  838. list_of_subj = [item['subj_id'] for item in data]
  839. data_dir = f'data/distance_pole/{data_dir_identifier}'
  840. print('Loading data from iterations...')
  841. # for subdir in os.listdir(data_folder_path):
  842. for subj in list_of_subj:
  843. subdir = subj
  844. print(f'loading data from {subdir}')
  845. subdir_path = os.path.join(data_dir, subdir)
  846. # Check if the path is indeed a directory
  847. if not os.path.isdir(subdir_path):
  848. print(f'Not a folder: {subdir}')
  849. continue # There might be a .DS_store.
  850. # == Distance data
  851. file_path = os.path.join(subdir_path, 'distances.npy')
  852. # Check if the file exists
  853. if not os.path.exists(file_path):
  854. import pdb; pdb.set_trace()
  855. bug # There is something wrong! The file is missing.
  856. data_here = np.load(file_path,allow_pickle=True).item()
  857. distances_all_subj.append(data_here)
  858. print('...done')
  859. print('Plotting...')
  860. dict_data = plot_pole_ratio_for_paper(distances_all_subj,run_logdir=run_logdir,data_dir_identifier=data_dir_identifier)
  861. print('... done')
  862. return dict_data
  863. def plot_pole_ratio_for_paper(distances_all_subj,run_logdir='.',data_dir_identifier=None):
  864. if data_dir_identifier == 'BNCI2014-001_feet_tongue':
  865. idx_to_highlight = [8]
  866. title = 'BCI-IV 2a (feet/tongue)'
  867. extra_str_save = 'BCI_IV_feet_tongue'
  868. elif data_dir_identifier == 'BNCI2014-001_left_hand_right_hand':
  869. idx_to_highlight = [3]
  870. title = 'BCI-IV 2a (left/right)'
  871. extra_str_save = 'BCI_IV_left_right'
  872. elif data_dir_identifier == 'Cho2017_left_hand_right_hand':
  873. idx_to_highlight = [14]
  874. title = 'Cho (left/right)'
  875. extra_str_save = 'Cho'
  876. elif data_dir_identifier == 'Dreyer2023_left_hand_right_hand':
  877. idx_to_highlight = [81]
  878. title = 'Dreyer (left/right)'
  879. extra_str_save = 'Dreyer'
  880. elif data_dir_identifier == 'Lee2019-MI_left_hand_right_hand':
  881. idx_to_highlight = [48]
  882. title = 'Lee (left/right)'
  883. extra_str_save = 'Lee'
  884. elif data_dir_identifier == 'PhysionetMotorImagery_feet_hands':
  885. idx_to_highlight = [3]
  886. title = 'Physionet (feet/hands)'
  887. extra_str_save = 'physionet_feet_hands'
  888. elif data_dir_identifier == 'PhysionetMotorImagery_left_hand_right_hand':
  889. idx_to_highlight = [76]
  890. title = 'Physionet (left/right)'
  891. extra_str_save = 'physionet_left_right'
  892. else:
  893. idx_to_highlight = [0]
  894. title = ''
  895. extra_str_save = ''
  896. fig, ax = plt.subplots(nrows=1,ncols=1,figsize=(5,4),layout="constrained")
  897. all_pole_ratios = np.zeros(len(distances_all_subj))
  898. for i, dist_dict in enumerate(distances_all_subj):
  899. if i in idx_to_highlight:
  900. pole_ratio = plot_pole_ratio_one_subj_for_paper(dist_dict,ax,highlight=True)
  901. else:
  902. pole_ratio = plot_pole_ratio_one_subj_for_paper(dist_dict,ax,highlight=False)
  903. all_pole_ratios[i] = pole_ratio
  904. # Colorbar
  905. norm = mpl.colors.Normalize(vmin=50, vmax=100)
  906. cmap = mpl.colormaps['viridis']
  907. sm = mpl.cm.ScalarMappable(norm=norm, cmap=cmap)
  908. sm.set_array([]) # required but the content is irrelevant
  909. # one colorbar for the whole figure
  910. cbar = fig.colorbar(sm, ax=ax, shrink = 0.75)
  911. cbar.set_label("Intra-subject accuracy [%]",fontsize=10)
  912. # optional: custom ticks
  913. cbar.set_ticks([50, 60, 70, 80, 90, 100])
  914. ax.set_xlabel('Distance',fontsize=10)
  915. ax.set_ylabel('Distance',fontsize=10)
  916. plt.title(title,fontsize=12)
  917. ax.set_aspect("equal")
  918. save_string = f'distances_to_poles_for_paper_{extra_str_save}' # Filename
  919. plt.savefig(f'{run_logdir}/png_{save_string}.png',dpi=600) # Save the figure in the specified directory
  920. plt.savefig(f'{run_logdir}/pdf_{save_string}.pdf') # Save the figure in the specified directory
  921. plt.close(fig)
  922. dict_to_return = {
  923. 'mean_pole_ratio':np.mean(all_pole_ratios),
  924. 'std_pole_ratio':np.std(all_pole_ratios),
  925. }
  926. return dict_to_return
  927. def plot_pole_ratio_one_subj_for_paper(dist_dict,ax,highlight=False):
  928. print(f'Plotting: {dist_dict['subj_id']} earth/circle representation.')
  929. # Class 0 circle
  930. r_0 = dist_dict['ref-mean:0']
  931. r_1 = dist_dict['ref-mean:1']
  932. r = (r_0 + r_1) / 2
  933. color_0 = 'tab:orange'
  934. color_1 = 'tab:blue'
  935. acc_score = dist_dict['Intra-subject accuracy']
  936. norm = mpl.colors.Normalize(vmin=0.5, vmax=1.0)
  937. cmap = mpl.colormaps['viridis']
  938. circle_color = cmap(norm(acc_score))
  939. # Circle
  940. circle_poly = CirclePolygon(
  941. xy=(0,0),
  942. radius=r,
  943. resolution=40,
  944. edgecolor=circle_color,
  945. facecolor="none",
  946. linewidth=2,
  947. alpha=0.5,
  948. zorder=1,
  949. )
  950. ax.add_patch(circle_poly)
  951. # Class 0 point,
  952. ratio_away_from_northpole = dist_dict['mean-northpole:0'] / (dist_dict['mean-northpole:0'] + dist_dict['mean-southpole:0'])
  953. angle = ratio_away_from_northpole * np.pi + np.pi/2
  954. ax.scatter(
  955. x = np.cos(angle) * r ,
  956. y = np.sin(angle) * r,
  957. color=color_0,
  958. edgecolor='black' if highlight else color_0,
  959. linewidth = 2 if highlight else 1,
  960. marker='o',
  961. s=100,
  962. alpha=0.75 if highlight else 0.5,
  963. zorder=5 if highlight else 4,
  964. )
  965. # Class 1 point
  966. ratio_away_from_northpole = dist_dict['mean-northpole:1'] / (dist_dict['mean-northpole:1'] + dist_dict['mean-southpole:1'])
  967. angle = ratio_away_from_northpole * np.pi - np.pi/2
  968. ax.scatter(
  969. x = np.cos(-angle) * r,
  970. y = np.sin(-angle) * r,
  971. color=color_1,
  972. edgecolor='black' if highlight else color_1,
  973. linewidth = 2 if highlight else 1,
  974. marker='o',
  975. s=100,
  976. alpha=0.75 if highlight else 0.5,
  977. zorder=5 if highlight else 4,
  978. )
  979. # north and south pole
  980. ax.scatter(
  981. y = [r, -r],
  982. x = [0, 0],
  983. color='tab:gray',
  984. marker='x',
  985. alpha=0.7,
  986. zorder=3,
  987. #s=20,
  988. )
  989. return ratio_away_from_northpole
  990. def highlight_cell(row, col, ax=None, fill=False, **kwargs):
  991. """
  992. Highlights a specific cell in a matrix plot by drawing a rectangle around it.
  993. This function adds a rectangular patch to highlight a specific cell in a plot (typically a heatmap or matrix plot).
  994. The cell is defined by its row and column indices, and the rectangle can optionally be filled with a color.
  995. Parameters
  996. ----------
  997. row : int
  998. The row index of the cell to highlight.
  999. col : int
  1000. The column index of the cell to highlight.
  1001. ax : matplotlib.axes.Axes, optional
  1002. The axes object on which to draw the rectangle. If None, the current axes will be used. Default is None.
  1003. fill : bool, optional
  1004. Whether to fill the rectangle with color. Default is False.
  1005. **kwargs : dict, optional
  1006. Additional keyword arguments to pass to `matplotlib.patches.Rectangle` (e.g., edgecolor, facecolor, linewidth).
  1007. Returns
  1008. -------
  1009. rect : matplotlib.patches.Rectangle
  1010. The rectangle patch added to the plot.
  1011. Notes
  1012. -----
  1013. This function is typically used to highlight specific cells in heatmaps or matrix plots, where each cell represents
  1014. a data point. The rectangle can be customized using the `**kwargs` to control the appearance of the highlight.
  1015. """
  1016. y=row
  1017. x=col
  1018. rect = plt.Rectangle((x-.5, y-.5), 1,1, fill=fill, **kwargs)
  1019. ax = ax or plt.gca()
  1020. ax.add_patch(rect)
  1021. return rect
  1022. def visualize_synthetic_2x2_SPD_matrices(run_logdir='.'):
  1023. # Sample SPD(2) by sampling a,c>0 and |b|<sqrt(ac)
  1024. # matrix = [[a,b],[b,c]]
  1025. n = 2000
  1026. a = np.random.exponential(1, n)
  1027. c = np.random.exponential(1, n)
  1028. b_max = np.sqrt(a*c)
  1029. b = (2*np.random.rand(n) - 1) * b_max * 0.9 # stay inside
  1030. fig = plt.figure(figsize=(5,5))
  1031. ax = fig.add_subplot(projection="3d")
  1032. ax.scatter(1, 0, 1, s=10,color='black')
  1033. # Build a grid for the boundary surface ac - b^2 = 0 -> c = b^2 / a
  1034. a_max = a.max()
  1035. b_max = b.max()
  1036. A, C = np.meshgrid(
  1037. np.linspace(1e-6, a_max, 60),
  1038. np.linspace(1e-6, a_max, 60)
  1039. )
  1040. # Boundary: b = ±sqrt(a c)
  1041. B_pos = np.sqrt(A * C)
  1042. B_neg = -B_pos
  1043. # "Radius" in SPD metric
  1044. x = 1 # choose the distance from I
  1045. # Parameters: φ for eigenvalues, θ for rotation
  1046. n_phi = 80
  1047. n_theta = 80
  1048. phi = np.linspace(0, 2*np.pi, n_phi)
  1049. theta = np.linspace(0, np.pi, n_theta) # π is enough due to symmetry
  1050. Phi, Theta = np.meshgrid(phi, theta)
  1051. # Eigenvalues on the circle in log-space
  1052. lam1 = np.exp(x * np.cos(Phi))
  1053. lam2 = np.exp(x * np.sin(Phi))
  1054. cosT = np.cos(Theta)
  1055. sinT = np.sin(Theta)
  1056. # Map (φ, θ) -> (a, b, c)
  1057. A11 = lam1 * cosT**2 + lam2 * sinT**2
  1058. A12 = (lam1 - lam2) * cosT * sinT
  1059. A22 = lam1 * sinT**2 + lam2 * cosT**2
  1060. # Plot the "sphere" around I as a surface in (a,b,c)-space
  1061. ax.plot_wireframe(
  1062. A11, A12, A22,
  1063. rcount=15, ccount=15,
  1064. linewidth=0.4,
  1065. alpha=0.1,
  1066. edgecolor="black",
  1067. )
  1068. A_flat = np.stack(
  1069. [
  1070. A11.ravel(),
  1071. A12.ravel(),
  1072. A12.ravel(),
  1073. A22.ravel(),
  1074. ],
  1075. axis=-1, # shape (N, 4)
  1076. ).reshape(-1, 2, 2)
  1077. Q = random_orthogonal(2)
  1078. for i in range(10):
  1079. idx = rng.choice(len(A_flat))
  1080. A = A_flat[idx]
  1081. A_rotated = Q @ A @ Q.T
  1082. color = f"C{i % 10}"
  1083. ax.scatter([A[0,0]], [A[0,1]], [A[1,1]], s=60,marker='*',color=color)
  1084. ax.scatter([A_rotated[0,0]], [A_rotated[0,1]], [A_rotated[1,1]], s=40,color=color)
  1085. ax.plot([A[0,0], A_rotated[0,0]], [A[0,1], A_rotated[0,1]], [A[1,1], A_rotated[1,1]],color=color,linewidth=1)
  1086. lim = 1.5
  1087. ax.set_xlim([0,lim*2])
  1088. ax.set_ylim([-lim,lim])
  1089. ax.set_zlim([0,lim*2])
  1090. ax.axis('equal')
  1091. ax.set_xlabel(r"$c_{11}$",fontsize=10)
  1092. ax.set_ylabel(r"$c_{12}$",fontsize=10)
  1093. ax.set_zlabel(r"$c_{22}$",fontsize=10)
  1094. elevations = 30 - np.abs(np.linspace(-15,15,11))
  1095. aziums = np.linspace(-90,90,11)
  1096. for i,(elev,azim) in enumerate(zip(elevations,aziums)):
  1097. ax.view_init(elev=elev, azim=azim) # elevation and azimuth in degrees
  1098. save_string = f'SPD_synthetic_{i}' # Filename
  1099. plt.tight_layout()
  1100. plt.savefig(f'{run_logdir}/png_{save_string}.png',dpi=600) # Save the figure in the specified directory
  1101. plt.savefig(f'{run_logdir}/pdf_{save_string}.pdf') # Save the figure in the specified directory
  1102. plt.close(fig)
  1103. return
  1104. def random_orthogonal(n):
  1105. A = rng.normal(size=(n, n))
  1106. Q, R = np.linalg.qr(A)
  1107. # Fix sign so det(Q)=±1 stays consistent
  1108. d = np.diag(R)
  1109. Q *= np.sign(d)
  1110. return Q
  1111. def combine_images_to_figure(
  1112. paths,
  1113. out_pdf="combined.pdf",
  1114. out_png="combined.png",
  1115. ncols=None,
  1116. figsize=None,
  1117. dpi_png=600,
  1118. facecolor="white",
  1119. ):
  1120. """
  1121. Combine raster images into a single Matplotlib figure and save as PDF + PNG.
  1122. Parameters
  1123. ----------
  1124. paths : list[str|Path]
  1125. Paths to input images.
  1126. out_pdf : str|Path
  1127. Output PDF path (set to None to skip).
  1128. out_png : str|Path
  1129. Output PNG path (set to None to skip).
  1130. ncols : int|None
  1131. Number of columns in the grid. Default: all images in one row.
  1132. figsize : tuple[float,float]|None
  1133. Figure size in inches. If None, choose a reasonable default.
  1134. dpi_png : int
  1135. DPI used for PNG output (controls pixel size of the saved PNG).
  1136. facecolor : str
  1137. Background color for the figure.
  1138. Returns
  1139. -------
  1140. (Path|None, Path|None)
  1141. The saved PDF and PNG paths (or None if skipped).
  1142. """
  1143. paths = [Path(p) for p in paths]
  1144. if not paths:
  1145. raise ValueError("paths is empty")
  1146. # Load images
  1147. imgs = [Image.open(p) for p in paths]
  1148. # Layout
  1149. n = len(imgs)
  1150. if ncols is None:
  1151. ncols = n
  1152. if ncols <= 0:
  1153. raise ValueError("ncols must be >= 1")
  1154. nrows = (n + ncols - 1) // ncols
  1155. if figsize is None:
  1156. # Simple heuristic: ~4 inches per column, ~3 inches per row
  1157. figsize = (4 * ncols, 3 * nrows)
  1158. fig, axes = plt.subplots(nrows, ncols, figsize=figsize, facecolor=facecolor)
  1159. # Make axes always iterable
  1160. if nrows == 1 and ncols == 1:
  1161. axes = [axes]
  1162. else:
  1163. axes = axes.ravel()
  1164. # Plot images
  1165. for ax, im in zip(axes, imgs):
  1166. ax.imshow(im)
  1167. ax.axis("off")
  1168. # Hide unused axes
  1169. for ax in axes[len(imgs):]:
  1170. ax.axis("off")
  1171. plt.tight_layout(pad=0.5, w_pad=0.5, h_pad=0.5)
  1172. saved_pdf = None
  1173. saved_png = None
  1174. if out_pdf is not None:
  1175. out_pdf = Path(out_pdf)
  1176. fig.savefig(out_pdf, bbox_inches="tight")
  1177. saved_pdf = out_pdf
  1178. if out_png is not None:
  1179. out_png = Path(out_png)
  1180. fig.savefig(out_png, dpi=dpi_png, bbox_inches="tight")
  1181. saved_png = out_png
  1182. plt.close(fig)
  1183. return saved_pdf, saved_png
  1184. # -----------------------------------------------
  1185. # Main
  1186. # -----------------------------------------------
  1187. # execute only if run as a script
  1188. if __name__ == "__main__":
  1189. begin_time = datetime.datetime.now()
  1190. log_folder = "my_logs"
  1191. if not os.path.exists(log_folder):
  1192. os.makedirs(log_folder)
  1193. # Create run directory to store models and logs in:
  1194. root_logdir = os.path.join(os.curdir, log_folder)
  1195. run_id = "%s_on_%s" % (time.strftime("run_%Y_%m_%d-%H_%M_%S"), os.uname()[1])
  1196. run_logdir = os.path.join(root_logdir, run_id)
  1197. output_folder = run_logdir
  1198. os.mkdir(run_logdir)
  1199. sys.stderr = Logger("{}/_console_err.txt".format(run_logdir))
  1200. sys.stdout = Logger("{}/_console_log.txt".format(run_logdir))
  1201. print(
  1202. "\n### main.py was started on %s in %s with logs in %s"
  1203. % (os.uname()[1], os.getcwd(), run_logdir)
  1204. )
  1205. print(sys.argv)
  1206. print_git_version()
  1207. save_source_code(run_logdir)
  1208. now = begin_time
  1209. now_string = str(now).replace(" ", "_").replace(":","_")
  1210. write_description()
  1211. print('\nRUNNING MAIN CODE\n')
  1212. run_logdir_here = f'{run_logdir}/synthetic_data'
  1213. os.mkdir(run_logdir_here)
  1214. visualize_synthetic_2x2_SPD_matrices(run_logdir=run_logdir_here)
  1215. paths = [
  1216. f'./{run_logdir}/synthetic_data/png_SPD_synthetic_0.png',
  1217. f'./{run_logdir}/synthetic_data/png_SPD_synthetic_2.png',
  1218. f'./{run_logdir}/synthetic_data/png_SPD_synthetic_5.png'
  1219. ]
  1220. combine_images_to_figure(
  1221. paths,
  1222. out_pdf=None,
  1223. out_png=f"{run_logdir}/png_synthetic_combined.jpg",
  1224. ncols=None,
  1225. figsize=(7,5),
  1226. dpi_png=600,
  1227. facecolor="white",
  1228. )
  1229. all_res = []
  1230. # -----------------------------------------------
  1231. # Run the code.
  1232. # -----------------------------------------------
  1233. for (dataset,events,identifier) in DATASETS:
  1234. print()
  1235. print()
  1236. print()
  1237. print('=======================')
  1238. print('=======================')
  1239. print('=======================')
  1240. print(f'Running {identifier}')
  1241. print('=======================')
  1242. print('=======================')
  1243. print('=======================')
  1244. print()
  1245. events.sort()
  1246. run_logdir_dataset = f'{run_logdir}/{identifier}'
  1247. data_dir_identifier = f'{dataset.code}_{"_".join(events)}'
  1248. os.mkdir(run_logdir_dataset)
  1249. all_subject_data = load_data_MOABB(subjects=dataset.subject_list,dataset=dataset,events=events,processing_dict=PROCESSING_DICT,run_logdir=run_logdir_dataset,data_dir_identifier=data_dir_identifier)
  1250. res = run_code(data=all_subject_data,run_logdir=run_logdir_dataset,identifier=identifier,data_dir_identifier=data_dir_identifier)
  1251. all_res.append(res)
  1252. print()
  1253. print()
  1254. print()
  1255. print()
  1256. print('RESULTS: ')
  1257. for i, (dataset,events,identifier) in enumerate(DATASETS):
  1258. print()
  1259. print(identifier)
  1260. res = all_res[i]
  1261. for key in res.keys():
  1262. if key == 'p-val':
  1263. print(f'p-val: z | q')
  1264. for key_2 in res[key]:
  1265. print(f'{key_2:20s}: p<0.05: {str(res[key][key_2]<0.05):6s}, p<0.01: {str(res[key][key_2]<0.01):6s}, pval={res[key][key_2]}')
  1266. else:
  1267. print(f'{key}: {res[key]:.3f} ({res[key]})')
  1268. paths = [
  1269. f'{run_logdir}/Physio_feet_hands/png_distances_to_poles_for_paper_physionet_feet_hands.png',
  1270. f'{run_logdir}/Dreyer_left_right/png_distances_to_poles_for_paper_Dreyer.png',
  1271. f'{run_logdir}/BCI-IV_feet_tongue/png_distances_to_poles_for_paper_BCI_IV_feet_tongue.png'
  1272. ]
  1273. combine_images_to_figure(
  1274. paths,
  1275. out_pdf=None,
  1276. out_png=f"{run_logdir}/png_pole_ratio_combined.jpg",
  1277. ncols=None,
  1278. figsize=(7,5),
  1279. dpi_png=600,
  1280. facecolor="white",
  1281. )
  1282. paths = [
  1283. f'{run_logdir}/Physio_feet_hands/png_lat_to_accuracy_correlation_physionet_feet_hands.png',
  1284. f'{run_logdir}/Dreyer_left_right/png_lat_to_accuracy_correlation_Dreyer.png',
  1285. f'{run_logdir}/BCI-IV_feet_tongue/png_lat_to_accuracy_correlation_BCI_IV_feet_tongue.png'
  1286. ]
  1287. combine_images_to_figure(
  1288. paths,
  1289. out_pdf=None,
  1290. out_png=f"{run_logdir}/png_lat_to_accuracy_correlation_combined.jpg",
  1291. ncols=None,
  1292. figsize=(7,5),
  1293. dpi_png=600,
  1294. facecolor="white",
  1295. )
  1296. plt.close('all')
  1297. print("\n...work complete!")
  1298. print("### Total run time: %s" % (datetime.datetime.now() - begin_time))

rotation_analysis.py at commit f8522b3, under MIT · at the source

Overview

Authors: Frida Heskebeck1, Bo Bernhardsson1, Carolina Bergeling2
  1. Department of Automatic Control, Lund University, Lund, Sweden
  2. Department of Mathematics and Natural Sciences, Blekinge Institute of Technology, Karlskrona, Sweden
Institutions: Lund University (Sweden); Blekinge Institute of Technology (Sweden)
Journal: Frontiers in human neuroscience, volume 20, article 1824613
Dates: received 6 March 2026; accepted 30 April 2026; published online 21 May 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.3389/fnhum.2026.1824613 · PMID 42253789 · PMCID PMC13233444 · OpenAlex W7161984432
Open access: gold, a free copy (OpenAlex)
Status: code verified
Methods: Spectral & time-frequency, Statistics, Machine learning
Keywords: BCI, pole ratio, Riemann geometry, source data selection, transfer learning
Topic: EEG and Brain-Computer Interfaces (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Citations: not cited yet (Europe PMC); 27 references in the paper

Abstract

This paper introduces the pole ratio metric and presents a sphere-based view of symmetric positive-definite matrix rotations on the Riemannian manifold of symmetric positive-definite matrices equipped with the affine-invariant Riemannian metric. The pole ratio quantifies whether data from different users lie on this Riemannian manifold in a way that enables effective transfer learning. The sphere-based view provides insight into the rotational step of transfer learning using the Riemannian Procrustes analysis method and highlights the limitations of rotation. For effective transfer learning, selecting appropriate source data is essential for good performance. The pole ratio is shown to be an effective metric for selecting source data. The main contribution of the paper is the insight into the limitations of rotations on a Riemannian manifold; the usefulness of the pole ratio as a source selection metric is a natural extension of this insight. This paper focuses on Brain-Computer Interfaces (BCIs), but the sphere-based view of rotations of symmetric positive-definite matrix data and the pole ratio are applicable to any field that models two-class data using symmetric positive-definite matrices.

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 3 matches between paragraphs and lines of code.

gitlab.control.lth.se/fridah/pole_ratio_public

License: MIT
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: f8522b32f7375ffefc79c31ed2e02a2e69e6b2d1, 10 March 2026
Languages: Python (1)
Size: 9 files, 1 script
Software Heritage: not archived
Found in: the text, “Materials and methods”
Holds: README, license file, environment (requirements.txt)
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: Matplotlib (1 file), MOABB (1 file), NumPy (1 file), Pillow (1 file), pyRiemann (1 file), scikit-learn (1 file), SciPy (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
3 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;
  • 1 script, each with its path and the digest of its content;
  • 3 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

Publicly available datasets were analyzed in this study. This data can be found here: https://moabb.neurotechx.com/docs/index.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 2, 28 September 2026

  • Funding: added Knut och Alice Wallenbergs Stiftelse

Version 1, 28 September 2026: the first record

Recorded: type, language, journal, volume, pages, dates, 3 authors, 5 keywords, 25 references.

Cite

This paper

Heskebeck, F., Bernhardsson, B., & Bergeling, C. (2026). Rotation-based metric on the Riemannian manifold of SPD matrices with applications to source data selection for brain-computer interface transfer learning. Frontiers in human neuroscience, 20, 1824613. https://doi.org/10.3389/fnhum.2026.1824613

BibTeX

@article{heskebeck2026rotation,
author = {Heskebeck, Frida and Bernhardsson, Bo and Bergeling, Carolina},
title = {{Rotation-based metric on the Riemannian manifold of SPD matrices with applications to source data selection for brain-computer interface transfer learning}},
journal = {Frontiers in human neuroscience},
year = {2026},
month = may,
volume = {20},
pages = {1824613},
publisher = {Frontiers Media SA},
issn = {1662-5161},
doi = {10.3389/fnhum.2026.1824613},
url = {https://doi.org/10.3389/fnhum.2026.1824613},
pmid = {42253789},
pmcid = {PMC13233444}
}

RIS

TY - JOUR
AU - Heskebeck, Frida
AU - Bernhardsson, Bo
AU - Bergeling, Carolina
TI - Rotation-based metric on the Riemannian manifold of SPD matrices with applications to source data selection for brain-computer interface transfer learning
T2 - Frontiers in human neuroscience
J2 - Front Hum Neurosci
PY - 2026
DA - 2026/05/21
VL - 20
SP - 1824613
SN - 1662-5161
PB - Frontiers Media SA
DO - 10.3389/fnhum.2026.1824613
UR - https://doi.org/10.3389/fnhum.2026.1824613
LA - en
ER -

CSL-JSON

{
"id": "10.3389/fnhum.2026.1824613",
"type": "article-journal",
"title": "Rotation-based metric on the Riemannian manifold of SPD matrices with applications to source data selection for brain-computer interface transfer learning",
"container-title": "Frontiers in human neuroscience",
"author": [
{
"family": "Heskebeck",
"given": "Frida"
},
{
"family": "Bernhardsson",
"given": "Bo"
},
{
"family": "Bergeling",
"given": "Carolina"
}
],
"container-title-short": "Front Hum Neurosci",
"volume": "20",
"page": "1824613",
"DOI": "10.3389/fnhum.2026.1824613",
"PMID": "42253789",
"PMCID": "PMC13233444",
"ISSN": "1662-5161",
"publisher": "Frontiers Media SA",
"URL": "https://doi.org/10.3389/fnhum.2026.1824613",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
21
]
]
}
}

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/fnhum.2026.1895016 [code]
A confidence-gated source selection strategy for cross-session transfer in brain-computer interfaces.
Journal: Frontiers in human neuroscience
In common: MOABB, pyRiemann, scikit-learn, 1 other tool, 7 references
[2] doi:10.1002/hbm.70528 [code]
Explainable AI Insights Into EEG Classification and Its Alignment to Neural Correlates.
Journal: Human brain mapping
In common: MOABB, Pillow, scikit-learn, 3 other tools, 1 reference
[3] 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, scikit-learn, SciPy, 2 other tools, 1 reference
[4] doi:10.1371/journal.pone.0347671 [code]
RMETNet: A cross-subject motor imagery EEG signal classification model based on TSLANet and riemannian geometry features.
Journal: PloS one
In common: pyRiemann, scikit-learn, SciPy, 2 other tools, 1 reference
[5] doi:10.7554/elife.100605 [code]
Age-related changes in ‘cortical’ 1/f dynamics are linked to cardiac activity
Journal: n/a
In common: pyRiemann, scikit-learn, SciPy, 2 other tools, 1 reference
[6] doi:10.1038/s41597-026-07807-x [code]
A Multimodal fNIRS-EEG Dataset for Unilateral Limb Motor Imagery.
Journal: Scientific data
In common: SciPy, Matplotlib, NumPy, 3 references
[7] doi:10.3390/s26175327 [code]
Subject Identity Confounds qEEG Emotion Recognition on DEAP and DREAMER.
Journal: Sensors (Basel, Switzerland)
In common: pyRiemann, scikit-learn, SciPy, 2 other tools
[8] doi:10.3389/fnhum.2026.1869918 [code]
Single-subject auditory ERP-BCI performance enhancement in ALS via an AI coding assistant prompt.
Journal: Frontiers in human neuroscience
In common: pyRiemann, scikit-learn, SciPy, 2 other tools
[9] doi:10.3389/fpsyg.2026.1774068 [code]
Analysis of cognitive mechanisms in phoneme perception and pronunciation errors among Korean language learners.
Journal: Frontiers in psychology
In common: pyRiemann, scikit-learn, SciPy, 2 other tools
[10] doi:10.1038/s43856-026-01817-x [code]
Visual prompt engineering for multimodal and irregularly sampled medical data.
Journal: Communications medicine
In common: Pillow, scikit-learn, SciPy, 2 other tools, 1 reference

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.