OSCR

The representational geometry of emotional states in basolateral amygdala.

Code ↔ Paper

5 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 5 matches
  1. [1] § Methods › Multiselectivity analysis ↔ decodanda/classes.py, lines 1473–1603 · score 0.59 · neural activity space, model iteration, break, classified, geometry, population
  2. [2] § Results › Specialized readouts for tremble and valence ↔ notebooks/CCGP.ipynb, lines 23–51 · score 0.53 · high dimensional, mixed selectivity, high CCGP, high decoding, representational geometry, disentangled
  3. [3] § Methods › Visualization of neural patterns in a reduced-dimensionality space ↔ decodanda/classes.py, lines 1473–1603 · score 0.52 · neural activity space, binary variables, geometry, population, vectors, decoding
  4. [4] § Methods › Neural decoding analysis › Balanced sampling ↔ decodanda/classes.py, lines 1821–1919 · score 0.52 · vice versa, neural activity, vectors, training, variables, decoding
  5. [5] § Methods › CCGP ↔ decodanda/classes.py, lines 1107–1253 · score 0.52 · decoding variable, decoding dichotomy, classes, classification, geometry, training

Paper

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

The paper is loaded when this pane is shown.

The authors' code

Python · 3,035 lines · 138 KB · GPL-3.0 · 4 matches

  1. # Copyright (C) 2023 Lorenzo Posani
  2. #
  3. # This program is free software: you can redistribute it and/or modify
  4. # it under the terms of the GNU General Public License as published by
  5. # the Free Software Foundation, either version 3 of the License, or
  6. # (at your option) any later version.
  7. #
  8. # This program is distributed in the hope that it will be useful,
  9. # but WITHOUT ANY WARRANTY; without even the implied warranty of
  10. # MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
  11. # GNU General Public License for more details.
  12. #
  13. import copy
  14. from typing import Tuple, Union
  15. import matplotlib.pyplot as plt
  16. import numpy as np
  17. import scipy.stats.stats
  18. from numpy import ndarray
  19. from .imports import *
  20. from .utilities import generate_binary_words, string_digits, sample_training_testing_from_rasters, CrossValidator, \
  21. log_dichotomy, hamming, sample_from_rasters, generate_dichotomies, semantic_score, z_pval, DictSession, \
  22. contiguous_chunking, non_contiguous_mask, cosine, generate_words, plot_confusion_matrix, enforce_min_time_separation
  23. from .visualize import corr_scatter, visualize_decoding, plot_perfs_null_model, visualize_PCA
  24. # Main class
  25. class Decodanda:
  26. def __init__(self,
  27. data: Union[list, dict],
  28. conditions: dict,
  29. classifier: any = 'svc',
  30. neural_attr: str = 'raster',
  31. trial_attr: str = 'trial',
  32. squeeze_trials: bool = False,
  33. min_data_per_condition: int = 2,
  34. min_trials_per_condition: int = 2,
  35. min_activations_per_cell: int = 1,
  36. min_time_separation: Optional[float] = None,
  37. time_attr: Optional[str] = None,
  38. trial_chunk: Optional[int] = None,
  39. exclude_silent: bool = False,
  40. verbose: bool = False,
  41. zscore: bool = False,
  42. fault_tolerance: bool = False,
  43. debug: bool = False,
  44. **kwargs
  45. ):
  46. """
  47. Main class that implements the decoding pipelines with built-in best practices.
  48. It works by separating the input data into all possible conditions - defined as specific
  49. combinations of variable values - and sampling data points from these conditions
  50. according to the specific decoding problem.
  51. Parameters
  52. ----------
  53. data
  54. A dictionary or a list of dictionaries each containing
  55. (1) the neural data (2) a set of variables that we want to decode from the neural data
  56. (3) a trial number. See the ``Data Structure`` section for more details.
  57. If a list is passed, the analyses will be performed on the pseudo-population built by pooling
  58. all the data sets in the list.
  59. conditions
  60. A dictionary that specifies which values for which variables of `data` we want to decode.
  61. See the ``Data Structure`` section for more details.
  62. classifier
  63. The classifier used for all decoding analyses. Default: ``sklearn.svm.LinearSVC``.
  64. neural_attr
  65. The key under which the neural features are stored in the ``data`` dictionary.
  66. trial_attr
  67. The key under which the trial numbers are stored in the ``data`` dictionary.
  68. Each different trial is considered as an independent sample to be used in
  69. during cross validation.
  70. If ``None``, trials are defined as consecutive bouts of data in time
  71. where all the variables have a constant value.
  72. squeeze_trials
  73. If True, all population vectors corresponding to the same trial number for the same
  74. condition will be squeezed into a single average activity vector.
  75. min_data_per_condition
  76. The minimum number of data points per each condition, defined as a specific
  77. combination of values of all variables in the ``conditions`` dictionary,
  78. that a data set needs to have to be included in the analysis.
  79. min_trials_per_condition
  80. The minimum number of unique trial numbers per each condition, defined as a specific
  81. combination of values of all variables in the ``conditions`` dictionary,
  82. that a data set needs to have to be included in the analysis.
  83. min_activations_per_cell
  84. The minimum number of non-zero bins that single neurons / features need to have to be
  85. included into the analysis.
  86. min_time_separation
  87. The minimum time difference, computed using ``time_attr``, between data assigned to different trial numbers.
  88. This prevents signals with long autocorrelations (e.g., calcium activity) from spilling over between
  89. training and testing trials.
  90. time_attr
  91. Name of the session field/attribute containing the time vector (one value per sample/bin),
  92. used to compute ``min_time_separation``.
  93. trial_chunk
  94. Only used when ``trial_attr=None``. The maximum number of consecutive data points
  95. within the same bout. Bouts longer than ``trial_chunk`` data points are split into
  96. different trials.
  97. exclude_silent
  98. If ``True``, all silent population vectors (only zeros) are excluded from the analysis.
  99. verbose
  100. If ``True``, most operations and analysis results are logged in standard output.
  101. zscore
  102. If ``True``, neural features are z-scored before being separated into conditions.
  103. fault_tolerance
  104. If ``True``, the constructor raises a warning instead of an error if no data set
  105. passes the inclusion criteria specified by ``min_data_per_condition`` and ``min_trials_per_condition``.
  106. debug
  107. If ``True``, operations are super verbose. Do not use unless you are developing.
  108. Data structure
  109. --------------
  110. Decodanda works with datasets organized into Python dictionaries.
  111. For ``N`` recorded neurons and ``T`` trials (or time bins), the data dictionary must contain:
  112. 1. a ``TxN`` array, under the ``raster`` key
  113. This is the set of features we use to decode. Can be continuous (e.g., calcium fluorescence) or discrete (e.g., spikes) values.
  114. 2. a ``Tx1`` array specifying a ``trial`` number
  115. This array will define the subdivisions for cross validation: trials (or time bins) that share the
  116. same ```trial``` value will always go together in either training or testing samples.
  117. 3. a ``Tx1`` array for each variable we want to decode
  118. Each value will be used as a label for the ``raster`` feature. Make sure these arrays are
  119. synchronized with the ``raster`` array.
  120. Say we have a data set with N=50 neurons, T=800 time bins divided into 80 trials, where two experimental
  121. variables are specified ``stimulus`` and ``action``.
  122. A properly-formatted data set would look like this:
  123. >>> data = {
  124. >>> 'raster': [[0, 1, ..., 0], ..., [0, 2, ..., 1]], # <800x50 array>, neural activations
  125. >>> 'stimulus': ['A', 'A', 'B', ..., 'B'], # <800x1 array>, values of the stimulus variable
  126. >>> 'action': ['left', 'left', 'none', ..., 'left'], # <800x1 array>, values of the action variable
  127. >>> 'trial': [1, 1, 1, ..., 2, 2, 2, ..., 80, 80, 80], # <800x1 array>, trial number, 80 unique numbers
  128. >>> }
  129. The ``conditions`` dictionary is used to specify which variables - out of
  130. all the keywords in the ``data`` dictionary, and which and values - out of
  131. all possible values of each specified variable - we want to decode.
  132. It has to be in the form ``{key: [value1, value2]}``:
  133. >>> conditions = {
  134. >>> 'stimulus': ['A', 'B'],
  135. >>> 'action': ['left', 'right']
  136. >>> }
  137. If more than one variable is specified, `Decodanda` will balance all
  138. conditions during each decoding analysis to disentangle
  139. the variables and avoid confounding correlations.
  140. Examples
  141. --------
  142. Using the data set defined above:
  143. >>> from decodanda import Decodanda
  144. >>>
  145. >>> dec = Decodanda(
  146. >>> data=data,
  147. >>> conditions=conditions
  148. >>> verbose=True)
  149. >>>
  150. [Decodanda] building conditioned rasters for session 0
  151. (stimulus = A, action = left): Selected 150 time bin out of 800, divided into 15 trials
  152. (stimulus = A, action = right): Selected 210 time bin out of 800, divided into 21 trials
  153. (stimulus = B, action = left): Selected 210 time bin out of 800, divided into 21 trials
  154. (stimulus = B, action = right): Selected 230 time bin out of 800, divided into 23 trials
  155. The constructor divides the data into conditions using the ``stimulus`` and ``action`` values
  156. and stores them in the ``self.conditioned_rasters`` object.
  157. This condition structure is the basis for all the balanced decoding analyses.
  158. """
  159. # casting single session to a list so that it is compatible with all loops below
  160. if type(data) != list:
  161. data = [data]
  162. # handling dictionaries as sessions
  163. # TODO: change default behavior with dictionaries instead of data structures
  164. if type(data[0]) == dict:
  165. dict_sessions = []
  166. for session in data:
  167. dict_sessions.append(DictSession(session))
  168. data = dict_sessions
  169. # check whether conditions are binary
  170. self._is_multiclass = False
  171. for key in conditions:
  172. if len(conditions[key]) != 2:
  173. self._is_multiclass = True
  174. # handling discrete dict conditions
  175. if type(list(conditions.values())[0]) == list:
  176. conditions = _generate_conditions_from_dic(conditions)
  177. # setting input parameters
  178. self.data = data
  179. self.conditions = conditions
  180. if classifier == 'svc':
  181. if self._is_multiclass:
  182. # One-vs-one linear SVM for multiclass
  183. classifier = SVC(kernel='linear', class_weight='balanced', C=1.0, max_iter=5000)
  184. else:
  185. # Faster solver for binary case
  186. classifier = LinearSVC(dual=False, C=1.0, class_weight='balanced', max_iter=5000)
  187. self.classifier = classifier
  188. # private params
  189. self._min_data_per_condition = min_data_per_condition
  190. self._min_trials_per_condition = min_trials_per_condition
  191. self._min_activations_per_cell = min_activations_per_cell
  192. self._verbose = verbose
  193. self._debug = debug
  194. self._zscore = zscore
  195. self._exclude_silent = exclude_silent
  196. self._neural_attr = neural_attr
  197. self._trial_attr = trial_attr
  198. self._trial_chunk = trial_chunk
  199. self._min_time_separation = min_time_separation
  200. self._time_attr = time_attr
  201. self._trial_average = squeeze_trials
  202. # setting session(s) data
  203. self.n_sessions = len(data)
  204. self.n_conditions = len(conditions)
  205. self._max_conditioned_data = 0
  206. self._min_conditioned_data = 10 ** 6
  207. self.n_neurons = 0
  208. self.n_brains = 0
  209. self.which_brain = []
  210. self._time_separation_masks = []
  211. self._session_trial_vectors = []
  212. # keys and stuff
  213. self._condition_vectors = generate_words(self.conditions)
  214. self._semantic_keys = list(self.conditions.keys())
  215. self._generate_semantic_vectors()
  216. # decoding weights
  217. self.decoding_weights = {}
  218. self.decoding_weights_null = {}
  219. # creating conditioned array with the following structure:
  220. # define a condition_vector with boolean values for each semantic condition, es. 100
  221. # use this vector as the key for a dictionary
  222. # as a value, create a list of neural data for each session conditioned as per key
  223. # >>> main object: neural rasters conditioned to semantic vector <<<
  224. self.conditioned_rasters = {string_digits(w): [] for w in self._condition_vectors}
  225. # conditioned null model index is the chunk division used for null model shuffles
  226. self.conditioned_trial_index = {string_digits(w): [] for w in self._condition_vectors}
  227. # >>> main part: create conditioned arrays <<< ---------------------------------
  228. self._divide_data_into_conditions(data)
  229. # \ >>> main part: create conditioned arrays <<< --------------------------------
  230. # raising exceptions
  231. if self.n_brains == 0:
  232. if not fault_tolerance:
  233. raise RuntimeError(
  234. "\n[Decodanda] No session passed the minimum data threshold for conditioned arrays.\n\t\t"
  235. "Check for mutually-exclusive conditions or try using less restrictive thresholds.")
  236. else:
  237. # derived attributes
  238. self._compute_centroids()
  239. # null model variables
  240. self.random_translations = {string_digits(w): [] for w in self._condition_vectors}
  241. self.subset = np.arange(self.n_neurons)
  242. self.ordered_conditioned_rasters = {}
  243. self.ordered_conditioned_trial_index = {}
  244. for w in self.conditioned_rasters.keys():
  245. self.ordered_conditioned_rasters[w] = self.conditioned_rasters[w].copy()
  246. self.ordered_conditioned_trial_index[w] = self.conditioned_trial_index[w].copy()
  247. # basic decoding functions
  248. def _train(self, training_raster_A, training_raster_B, label_A, label_B):
  249. training_labels_A = np.repeat(label_A, training_raster_A.shape[0]).astype(object)
  250. training_labels_B = np.repeat(label_B, training_raster_B.shape[0]).astype(object)
  251. training_raster = np.vstack([training_raster_A, training_raster_B])
  252. training_labels = np.hstack([training_labels_A, training_labels_B])
  253. self.classifier = sklearn.base.clone(self.classifier)
  254. training_raster = training_raster[:, self.subset]
  255. self.classifier.fit(training_raster, training_labels)
  256. def _test(self, testing_raster_A, testing_raster_B, label_A, label_B):
  257. testing_labels_A = np.repeat(label_A, testing_raster_A.shape[0]).astype(object)
  258. testing_labels_B = np.repeat(label_B, testing_raster_B.shape[0]).astype(object)
  259. testing_raster = np.vstack([testing_raster_A, testing_raster_B])
  260. testing_labels = np.hstack([testing_labels_A, testing_labels_B])
  261. testing_raster = testing_raster[:, self.subset]
  262. if self._debug:
  263. print("Real labels")
  264. print(testing_labels)
  265. print("Predicted labels")
  266. print(self.classifier.predict(testing_raster))
  267. performance = self.classifier.score(testing_raster, testing_labels)
  268. return performance
  269. def _one_cv_step(self, dic, training_fraction, ndata, shuffled=False, testing_trials=None, dic_key=None):
  270. if dic_key is None:
  271. dic_key = self._dic_key(dic)
  272. set_A = dic[0]
  273. label_A = ''
  274. for d in set_A:
  275. label_A += (self._semantic_vectors[d] + ' ')
  276. label_A = label_A[:-1]
  277. set_B = dic[1]
  278. label_B = ''
  279. for d in set_B:
  280. label_B += (self._semantic_vectors[d] + ' ')
  281. label_B = label_B[:-1]
  282. training_array_A = []
  283. training_array_B = []
  284. testing_array_A = []
  285. testing_array_B = []
  286. # allow for unbalanced dichotomies
  287. n_conditions_A = float(len(dic[0]))
  288. n_conditions_B = float(len(dic[1]))
  289. fraction = n_conditions_A / n_conditions_B
  290. for d in set_A:
  291. training, testing = sample_training_testing_from_rasters(self.conditioned_rasters[d],
  292. int(ndata / fraction),
  293. training_fraction,
  294. self.conditioned_trial_index[d],
  295. debug=self._debug,
  296. testing_trials=testing_trials)
  297. if self._debug:
  298. plt.title('Condition A')
  299. print("Sampling for condition A, d=%s" % d)
  300. print("Conditioned raster mean:")
  301. print(np.nanmean(self.conditioned_rasters[d][0], 0))
  302. training_array_A.append(training)
  303. testing_array_A.append(testing)
  304. for d in set_B:
  305. training, testing = sample_training_testing_from_rasters(self.conditioned_rasters[d],
  306. int(ndata),
  307. training_fraction,
  308. self.conditioned_trial_index[d],
  309. debug=self._debug,
  310. testing_trials=testing_trials)
  311. training_array_B.append(training)
  312. testing_array_B.append(testing)
  313. if self._debug:
  314. plt.title('Condition B')
  315. print("Sampling for condition B, d=%s" % d)
  316. print("Conditioned raster mean:")
  317. print(np.nanmean(self.conditioned_rasters[d][0], 0))
  318. training_array_A = np.vstack(training_array_A)
  319. training_array_B = np.vstack(training_array_B)
  320. testing_array_A = np.vstack(testing_array_A)
  321. testing_array_B = np.vstack(testing_array_B)
  322. if self._debug:
  323. selectivity_training = np.nanmean(training_array_A, 0) - np.nanmean(training_array_B, 0)
  324. selectivity_testing = np.nanmean(testing_array_A, 0) - np.nanmean(testing_array_B, 0)
  325. corr_scatter(selectivity_training, selectivity_testing, 'Selectivity (training)', 'Selectivity (testing)')
  326. if self._zscore:
  327. big_raster = np.vstack([training_array_A, training_array_B]) # z-scoring using the training data
  328. big_mean = np.nanmean(big_raster, 0)
  329. big_std = np.nanstd(big_raster, 0)
  330. big_std[big_std == 0] = np.inf
  331. training_array_A = (training_array_A - big_mean) / big_std
  332. training_array_B = (training_array_B - big_mean) / big_std
  333. testing_array_A = (testing_array_A - big_mean) / big_std
  334. testing_array_B = (testing_array_B - big_mean) / big_std
  335. self._train(training_array_A, training_array_B, label_A, label_B)
  336. if hasattr(self.classifier, 'coef_'):
  337. if dic_key and not shuffled:
  338. if dic_key not in self.decoding_weights.keys():
  339. self.decoding_weights[dic_key] = []
  340. self.decoding_weights[dic_key].append(self.classifier.coef_)
  341. if dic_key and shuffled:
  342. if dic_key not in self.decoding_weights_null.keys():
  343. self.decoding_weights_null[dic_key] = []
  344. self.decoding_weights_null[dic_key].append(self.classifier.coef_)
  345. performance = self._test(testing_array_A, testing_array_B, label_A, label_B)
  346. return performance
  347. def _one_X_cv_step(self, dic1, dic2, training_fraction, ndata, shuffled=False):
  348. # Training rasters
  349. training_set_A = dic1[0]
  350. training_set_B = dic1[1]
  351. testing_set_A = dic2[0]
  352. testing_set_B = dic2[1]
  353. training_array_A = []
  354. training_array_B = []
  355. testing_array_A = []
  356. testing_array_B = []
  357. # allow for unbalanced dichotomies
  358. n_conditions_A = float(len(dic1[0]))
  359. n_conditions_B = float(len(dic1[1]))
  360. fraction = n_conditions_A / n_conditions_B
  361. for d in training_set_A:
  362. training, testing = sample_training_testing_from_rasters(self.conditioned_rasters[d],
  363. int(ndata / fraction),
  364. training_fraction,
  365. self.conditioned_trial_index[d],
  366. debug=self._debug)
  367. if self._debug:
  368. plt.title('Condition A')
  369. print("Sampling for condition A, d=%s" % d)
  370. print("Conditioned raster mean:")
  371. print(np.nanmean(self.conditioned_rasters[d][0], 0))
  372. training_array_A.append(training)
  373. if d in testing_set_A:
  374. testing_array_A.append(testing)
  375. elif d in testing_set_B:
  376. testing_array_B.append(testing)
  377. for d in training_set_B:
  378. training, testing = sample_training_testing_from_rasters(self.conditioned_rasters[d],
  379. int(ndata),
  380. training_fraction,
  381. self.conditioned_trial_index[d],
  382. debug=self._debug)
  383. training_array_B.append(training)
  384. if d in testing_set_A:
  385. testing_array_A.append(testing)
  386. elif d in testing_set_B:
  387. testing_array_B.append(testing)
  388. training_array_A = np.vstack(training_array_A)
  389. training_array_B = np.vstack(training_array_B)
  390. testing_array_A = np.vstack(testing_array_A)
  391. testing_array_B = np.vstack(testing_array_B)
  392. if self._debug:
  393. selectivity_training = np.nanmean(training_array_A, 0) - np.nanmean(training_array_B, 0)
  394. selectivity_testing = np.nanmean(testing_array_A, 0) - np.nanmean(testing_array_B, 0)
  395. corr_scatter(selectivity_training, selectivity_testing, 'Selectivity (training)', 'Selectivity (testing)')
  396. if self._zscore:
  397. big_raster = np.vstack([training_array_A, training_array_B]) # z-scoring using the training data
  398. big_mean = np.nanmean(big_raster, 0)
  399. big_std = np.nanstd(big_raster, 0)
  400. big_std[big_std == 0] = np.inf
  401. training_array_A = (training_array_A - big_mean) / big_std
  402. training_array_B = (training_array_B - big_mean) / big_std
  403. testing_array_A = (testing_array_A - big_mean) / big_std
  404. testing_array_B = (testing_array_B - big_mean) / big_std
  405. self._train(training_array_A, training_array_B, 'A', 'B')
  406. performance = self._test(testing_array_A, testing_array_B, 'A', 'B')
  407. return performance
  408. # Sampling functions
  409. def balanced_resample(self, condition_names=False, ndata=None, z_score=None, min_ar=0):
  410. """
  411. Parameters
  412. ----------
  413. condition_names: if True, verbose names for conditions are used, otherwise a binary notation is used. Default: False.
  414. ndata: optional, number of resampled activity vectors per condition. If not specified,
  415. the maximum number of activity vectors across all conditions is used.
  416. z_score: if True, the resampled rasters are z-scored with respect to all the conditions.
  417. min_ar: neurons below a minimum activity rate (fraction of bins with non-zero activity) threshold specified
  418. by the ``min_ar`` parameter will be excluded from the sampled data.
  419. Returns
  420. -------
  421. balanced resampled rasters
  422. """
  423. if ndata is None:
  424. ndata = self._max_conditioned_data
  425. if z_score is None:
  426. z_score = self._zscore
  427. resampled_rasters = {}
  428. for key in self.conditioned_rasters:
  429. if condition_names:
  430. condition_key = self._semantic_vectors[key]
  431. else:
  432. condition_key = key
  433. resampled_rasters[condition_key] = []
  434. for n in range(self.n_brains):
  435. x = self.conditioned_rasters[key][n]
  436. sampling_index = np.random.randint(0, x.shape[0], ndata)
  437. resampled_rasters[condition_key].append(x[sampling_index])
  438. resampled_rasters[condition_key] = np.hstack(resampled_rasters[condition_key])
  439. if z_score:
  440. for i in range(self.n_neurons):
  441. big_x = np.hstack([r[:, i] for r in resampled_rasters.values()])
  442. bigmean = np.nanmean(big_x)
  443. bigstd = np.nanstd(big_x)
  444. for key in resampled_rasters:
  445. resampled_rasters[key][:, i] = (resampled_rasters[key][:, i] - bigmean) / bigstd
  446. if min_ar:
  447. X = np.vstack([resampled_rasters[key] for key in resampled_rasters])
  448. activityrate = np.nanmean(X > 0, 0)
  449. for key in resampled_rasters:
  450. resampled_rasters[key] = resampled_rasters[key][:, activityrate > min_ar]
  451. return resampled_rasters
  452. def split_resample(self, fraction=0.5, condition_names=False, ndata=None, z_score=None, min_ar=0):
  453. """
  454. Parameters ----------
  455. fraction: the fraction of trials used to sample from to fill the first data set (
  456. raster_A). The remaining fraction (1-``fraction``) is used to sample the second data set (raster_B)
  457. condition_names: if True, verbose names for conditions are used, otherwise a binary notation is used. Default: False.
  458. ndata: optional, number of resampled activity vectors per condition. If not specified,
  459. the maximum number of activity vectors across all conditions is used.
  460. z_score: if True, the resampled rasters are z-scored with respect to all the conditions.
  461. min_ar: neurons below a minimum activity rate (fraction of bins with non-zero activity) threshold specified
  462. by the ``min_ar`` parameter will be excluded from the sampled data.
  463. Returns
  464. -------
  465. rasters_A, rasters_B - dictionaries with resampled data for all conditions from different trials
  466. """
  467. if ndata is None:
  468. ndata = self._max_conditioned_data
  469. if z_score is None:
  470. z_score = self._zscore
  471. resampled_rasters_A = {}
  472. resampled_rasters_B = {}
  473. for key in self.conditioned_rasters:
  474. if condition_names:
  475. condition_key = self._semantic_vectors[key]
  476. else:
  477. condition_key = key
  478. training, testing = sample_training_testing_from_rasters(self.conditioned_rasters[key],
  479. ndata=ndata,
  480. training_fraction=fraction,
  481. trials=self.conditioned_trial_index[key])
  482. resampled_rasters_A[condition_key] = training
  483. resampled_rasters_B[condition_key] = testing
  484. if z_score:
  485. for i in range(self.n_neurons):
  486. big_x = np.hstack(
  487. [r[:, i] for r in resampled_rasters_A.values()] + [r[:, i] for r in resampled_rasters_B.values()])
  488. bigmean = np.nanmean(big_x)
  489. bigstd = np.nanstd(big_x)
  490. for key in resampled_rasters_A:
  491. resampled_rasters_A[key][:, i] = (resampled_rasters_A[key][:, i] - bigmean) / bigstd
  492. resampled_rasters_B[key][:, i] = (resampled_rasters_B[key][:, i] - bigmean) / bigstd
  493. if min_ar:
  494. X = np.vstack(
  495. [resampled_rasters_A[key] for key in resampled_rasters_A] + [resampled_rasters_B[key] for key in
  496. resampled_rasters_B])
  497. activityrate = np.nanmean(X > 0, 0)
  498. for key in resampled_rasters_A:
  499. resampled_rasters_A[key] = resampled_rasters_A[key][:, activityrate > min_ar]
  500. resampled_rasters_B[key] = resampled_rasters_B[key][:, activityrate > min_ar]
  501. return resampled_rasters_A, resampled_rasters_B
  502. # Dichotomy analysis functions
  503. def decode_dichotomy(self,
  504. dichotomy: Union[str, list],
  505. training_fraction: float,
  506. cross_validations: int = 10,
  507. ndata: Optional[int] = None,
  508. shuffled: bool = False,
  509. parallel: bool = False,
  510. testing_trials: Optional[list] = None,
  511. dic_key: Optional[str] = None,
  512. subsample: Optional[float] = 0,
  513. **kwargs) -> ndarray:
  514. """
  515. Function that performs cross-validated decoding of a specific dichotomy.
  516. Decoding is performed by sampling a balanced amount of data points from each condition in each class of the
  517. dichotomy, so to ensure that only the desired variable is analyzed by balancing confounds.
  518. Before sampling, each condition is individually divided into training and testing bins
  519. by using the ``self.trial`` array specified in the data structure when constructing the ``Decodanda`` object.
  520. Parameters
  521. ----------
  522. dichotomy : str || list
  523. The dichotomy to be decoded, expressed in a double-list binary format, e.g. [['10', '11'], ['01', '00']], or as a variable name.
  524. training_fraction:
  525. the fraction of trials used for training in each cross-validation fold.
  526. cross_validations:
  527. the number of cross-validations.
  528. ndata:
  529. the number of data points (population vectors) sampled for training and for testing for each condition.
  530. shuffled:
  531. if True, population vectors for each condition are sampled in a random way compatibly with a null model for decoding performance.
  532. parallel:
  533. if True, each cross-validation is performed by a dedicated thread (experimental, use with caution).
  534. testing_trials:
  535. if specified, data sampled from the specified trial numbers will be used for testing, and the remaining ones for training.
  536. dic_key:
  537. if specified, weights of the decoding analysis will be saved in self.decoding_weights using dic_key as the dictionary key.
  538. subsample:
  539. if >0, a random subsample of neurons of size=subsample will be used at each cross-validation
  540. Returns
  541. -------
  542. performances: list of decoding performance values for each cross-validation.
  543. Note
  544. ----
  545. ``dichotomy`` can be passed as a string or as a list.
  546. If a string is passed, it has to be a name of one of the variables specified in the conditions dictionary.
  547. If a list is passed, it needs to contain two lists in the shape [[...], [...]].
  548. Each sub list contains the conditions used to define one of the two decoded classes
  549. in binary notation.
  550. For example, if the data set has two variables
  551. ``stimulus`` :math:`\\in` {-1, 1} and ``action`` :math:`\\in` {-1, 1}, the condition
  552. ``stimulus=-1`` & ``action=-1`` will correspond to the binary notation ``'00'``,
  553. the condition ``stimulus=+1`` & ``action=-1`` will correspond to ``10`` and so on.
  554. Therefore, the notation:
  555. >>> dic = 'stimulus'
  556. is equivalent to
  557. >>> dic = [['00', '01'], ['10', '11']]
  558. and
  559. >>> dic = 'action'
  560. is equivalent to
  561. >>> dic = [['00', '10'], ['01', '11']]
  562. However, not all dichotomies have names (are semantic). For example, the dichotomy
  563. >>> [['01','10'], ['00', '11']]
  564. can only be defined using the binary notation.
  565. Note that this function gives you the flexibility to use sub-sets of conditions, for example
  566. >>> dic = [['10'], ['01']]
  567. will decode stimulus=1 & action=-1 vs. stimulus=-1 & action=1
  568. Example
  569. -------
  570. >>> data = generate_synthetic_data(keyA='stimulus', keyB='action')
  571. >>> dec = Decodanda(data=data, conditions={'stimulus': [-1, 1], 'action': [-1, 1]})
  572. >>> perfs = dec.decode_dichotomy('stimulus', training_fraction=0.75, cross_validations=10)
  573. >>> perfs
  574. [0.82, 0.87, 0.75, ..., 0.77] # 10 values
  575. """
  576. if type(dichotomy) == str:
  577. dic = self._dichotomy_from_key(dichotomy)
  578. else:
  579. dic = dichotomy
  580. if ndata is None and self.n_brains == 1:
  581. ndata = self._max_conditioned_data
  582. if ndata is None and self.n_brains > 1:
  583. ndata = max(self._max_conditioned_data, 2 * self.n_neurons)
  584. if subsample:
  585. self._generate_random_subset(subsample)
  586. if shuffled:
  587. self._shuffle_conditioned_arrays(dic)
  588. if self._verbose and not shuffled:
  589. print(dic, ndata)
  590. log_dichotomy(self, dic, ndata, 'Decoding')
  591. count = tqdm(range(cross_validations))
  592. else:
  593. count = range(cross_validations)
  594. if parallel:
  595. # TODO: add subsample to the parallel routine
  596. pool = Pool()
  597. res = pool.map(CrossValidator(classifier=self.classifier,
  598. conditioned_rasters=self.conditioned_rasters,
  599. conditioned_trial_index=self.conditioned_trial_index,
  600. dic=dic,
  601. training_fraction=training_fraction,
  602. ndata=ndata,
  603. subset=self.subset,
  604. semantic_vectors=self._semantic_vectors,
  605. z_score=self._zscore,
  606. dic_key=dic_key),
  607. range(cross_validations))
  608. performances = np.asarray([r[0] for r in res])
  609. if len(res[0][1]):
  610. key = list(res[0][1].keys())[0]
  611. weights = {key: [r[1][key] for r in res]}
  612. print(performances, weights)
  613. else:
  614. performances = np.zeros(cross_validations)
  615. if self._verbose and not shuffled:
  616. print('\nLooping over decoding cross validation folds:')
  617. for i in count:
  618. if subsample:
  619. self._generate_random_subset(subsample)
  620. performances[i] = self._one_cv_step(dic=dic, training_fraction=training_fraction, ndata=ndata,
  621. shuffled=shuffled, testing_trials=testing_trials, dic_key=dic_key)
  622. if subsample:
  623. self._generate_random_subset(self.n_neurons)
  624. if shuffled:
  625. self._order_conditioned_rasters()
  626. return np.asarray(performances)
  627. def CCGP_dichotomy(self, dichotomy: Union[str, list],
  628. resamplings: int = 3,
  629. ndata: Optional[int] = None,
  630. max_semantic_dist: int = 1,
  631. split_rule='OneOut',
  632. shuffled: bool = False,
  633. **kwargs):
  634. """
  635. Function that performs the cross-condition generalization performance analysis (CCGP, Bernardi et al. 2020, Cell)
  636. for a given variable, specified through its corresponding dichotomy. This function tests how well a given
  637. coding strategy for the given variable generalizes when the other variables are changed.
  638. Parameters
  639. ----------
  640. dichotomy : str || list
  641. The dichotomy corresponding to the variable to be tested, expressed in a double-list binary format, e.g. [['10', '11'], ['01', '00']], or as a variable name.
  642. resamplings:
  643. The number of iterations for each decoding analysis. The returned performance value is the average over these resamplings.
  644. ndata:
  645. The number of data points (population vectors) sampled for training and for testing for each condition.
  646. max_semantic_dist:
  647. The maximum semantic distance (number of variables that change value) between conditions in the held-out pair used to test the classifier.
  648. split_rule:
  649. The way conditions are split in training and testing. OneOut (default), name of a variable, or dichotomy in the double-list binary format. If OneOut is used, one pair of conditions is held out and the rest is used to train the classifier; if a variable is specified, then CCGP is computed specifically across that variable, balancing any third (or further) variables during sampling.
  650. shuffled:
  651. If True, the data is sampled according to geometrical null model for CCGP that keeps variables decodable but breaks the generalization. See Bernardi et al 2020 & Boyle, Posani et al. 2023.
  652. Returns
  653. -------
  654. performances: list of performance values for each cross-condition training-testing split.
  655. Note
  656. ----
  657. This function trains the ``self._classifier`` to decode the given variable in a sub-set
  658. of conditions, and tests it on the held-out set.
  659. The split of training and testing conditions is decided by the ``max_semantic_dist`` parameter: if set to 1,
  660. only pairs of conditions that have all variables in common except the specified one are held out to test the
  661. classifier.
  662. For example, if the data set has two variables
  663. ``stimulus`` :math:`\\in` {-1, 1} and ``action`` :math:`\\in` {-1, 1}, to compute CCGP for ``stimulus``
  664. with ``max_semantic_dist=1`` this function will train the classifier on
  665. ``(stimulus = -1, action = -1)`` vs. ``(stimulus = 1, action = -1)``
  666. And test it on
  667. ``(stimulus = -1, action = 1)`` vs. ``(stimulus = 1, action = 1)``
  668. note that action is kept fixed within the training and testing conditions.
  669. If instead we use ``max_semantic_dist=2``, all possible combinations are used, including training on
  670. ``(stimulus = -1, action = -1)`` vs. ``(stimulus = 1, action = 1)``
  671. and testing on
  672. ``(stimulus = -1, action = 1)`` vs. ``(stimulus = 1, action = -1)``
  673. ``dichotomy`` can be passed as a string or as a list.
  674. If a string is passed, it has to be a name of one of the variables specified in the conditions dictionary.
  675. If a list is passed, it needs to contain two lists in the shape [[...], [...]].
  676. Each sub list contains the conditions used to define one of the two decoded classes
  677. in binary notation.
  678. For example, if the data set has two variables
  679. ``stimulus`` :math:`\\in` {-1, 1} and ``action`` :math:`\\in` {-1, 1}, the condition
  680. ``stimulus=-1`` & ``action=-1`` will correspond to the binary notation ``'00'``,
  681. the condition ``stimulus=+1`` & ``action=-1`` will correspond to ``10`` and so on.
  682. Therefore, if ``stimulus`` is the first variable in the conditions dictionary, its corresponding dichotomy is
  683. >>> stimulus = [['00', '01'], ['10', '11']]
  684. Example
  685. -------
  686. >>> data = generate_synthetic_data(keyA='stimulus', keyB='action')
  687. >>> dec = Decodanda(data=data, conditions={'stimulus': [-1, 1], 'action': [-1, 1]})
  688. >>> perfs = dec.CCGP_dichotomy('stimulus')
  689. >>> perfs
  690. [0.82, 0.87] # 2 values
  691. """
  692. if self._is_multiclass:
  693. raise RuntimeError(f"{self.__class__.__name__}: this analysis is only defined for binary variables. ")
  694. if type(dichotomy) == str:
  695. dic = self._dichotomy_from_key(dichotomy)
  696. else:
  697. dic = dichotomy
  698. if ndata is None and self.n_brains == 1:
  699. ndata = self._max_conditioned_data
  700. if ndata is None and self.n_brains > 1:
  701. ndata = max(self._max_conditioned_data, 2 * self.n_neurons)
  702. all_performances = []
  703. if not shuffled and self._verbose:
  704. log_dichotomy(self, dic, ndata, 'Cross-condition decoding')
  705. iterable = tqdm(range(resamplings))
  706. elif not shuffled:
  707. iterable = range(resamplings)
  708. else:
  709. iterable = range(1)
  710. for n in iterable:
  711. performances = []
  712. set_A = dic[0]
  713. set_B = dic[1]
  714. if split_rule == 'OneOut':
  715. # loop over all possible held-out pairs
  716. for i in range(len(set_A)):
  717. for j in range(len(set_B)):
  718. test_condition_A = set_A[i]
  719. test_condition_B = set_B[j]
  720. if hamming(string_digits(test_condition_A),
  721. string_digits(test_condition_B)) <= max_semantic_dist:
  722. training_conditions_A = [x for iA, x in enumerate(set_A) if iA != i]
  723. training_conditions_B = [x for iB, x in enumerate(set_B) if iB != j]
  724. training_array_A = []
  725. training_array_B = []
  726. label_A = ''
  727. label_B = ''
  728. for ck in training_conditions_A:
  729. arr = sample_from_rasters(self.conditioned_rasters[ck], ndata=ndata)
  730. training_array_A.append(arr)
  731. label_A += (self._semantic_vectors[ck] + ' ')
  732. for ck in training_conditions_B:
  733. arr = sample_from_rasters(self.conditioned_rasters[ck], ndata=ndata)
  734. training_array_B.append(arr)
  735. label_B += (self._semantic_vectors[ck] + ' ')
  736. training_array_A = np.vstack(training_array_A)
  737. training_array_B = np.vstack(training_array_B)
  738. testing_array_A = sample_from_rasters(self.conditioned_rasters[test_condition_A],
  739. ndata=ndata)
  740. testing_array_B = sample_from_rasters(self.conditioned_rasters[test_condition_B],
  741. ndata=ndata)
  742. if shuffled:
  743. rotation_A = np.arange(testing_array_A.shape[1]).astype(int)
  744. rotation_B = np.arange(testing_array_B.shape[1]).astype(int)
  745. np.random.shuffle(rotation_A)
  746. np.random.shuffle(rotation_B)
  747. testing_array_A = testing_array_A[:, rotation_A]
  748. testing_array_B = testing_array_B[:, rotation_A]
  749. if self._zscore:
  750. big_raster = np.vstack(
  751. [training_array_A, training_array_B]) # z-scoring using the training data
  752. big_mean = np.nanmean(big_raster, 0)
  753. big_std = np.nanstd(big_raster, 0)
  754. big_std[big_std == 0] = np.inf
  755. training_array_A = (training_array_A - big_mean) / big_std
  756. training_array_B = (training_array_B - big_mean) / big_std
  757. testing_array_A = (testing_array_A - big_mean) / big_std
  758. testing_array_B = (testing_array_B - big_mean) / big_std
  759. self._train(training_array_A, training_array_B, label_A, label_B)
  760. performance = self._test(testing_array_A, testing_array_B, label_A, label_B)
  761. performances.append(performance)
  762. elif type(split_rule) == str:
  763. split_dichotomy = self._dichotomy_from_key(split_rule)
  764. training_conditions_A = [c for c in dic[0] if c in split_dichotomy[0]]
  765. training_conditions_B = [c for c in dic[1] if c in split_dichotomy[0]]
  766. testing_conditions_A = [c for c in dic[0] if c in split_dichotomy[1]]
  767. testing_conditions_B = [c for c in dic[1] if c in split_dichotomy[1]]
  768. training_array_A = []
  769. training_array_B = []
  770. label_A = ''
  771. label_B = ''
  772. for ck in training_conditions_A:
  773. arr = sample_from_rasters(self.conditioned_rasters[ck], ndata=ndata)
  774. training_array_A.append(arr)
  775. label_A += (self._semantic_vectors[ck] + ' ')
  776. for ck in training_conditions_B:
  777. arr = sample_from_rasters(self.conditioned_rasters[ck], ndata=ndata)
  778. training_array_B.append(arr)
  779. label_B += (self._semantic_vectors[ck] + ' ')
  780. training_array_A = np.vstack(training_array_A)
  781. training_array_B = np.vstack(training_array_B)
  782. if self._verbose:
  783. print(f'\nCCGP: Training on {label_A} vs {label_B}')
  784. testing_array_A = []
  785. testing_array_B = []
  786. label_A_test = ''
  787. label_B_test = ''
  788. for ck in testing_conditions_A:
  789. arr = sample_from_rasters(self.conditioned_rasters[ck], ndata=ndata)
  790. testing_array_A.append(arr)
  791. label_A_test += (self._semantic_vectors[ck] + ' ')
  792. for ck in testing_conditions_B:
  793. arr = sample_from_rasters(self.conditioned_rasters[ck], ndata=ndata)
  794. testing_array_B.append(arr)
  795. label_B_test += (self._semantic_vectors[ck] + ' ')
  796. testing_array_A = np.vstack(testing_array_A)
  797. testing_array_B = np.vstack(testing_array_B)
  798. if self._verbose:
  799. print(f'CCGP: Testing on {label_A_test} vs {label_B_test}')
  800. if shuffled:
  801. rotation_A = np.arange(testing_array_A.shape[1]).astype(int)
  802. rotation_B = np.arange(testing_array_B.shape[1]).astype(int)
  803. np.random.shuffle(rotation_A)
  804. np.random.shuffle(rotation_B)
  805. testing_array_A = testing_array_A[:, rotation_A]
  806. testing_array_B = testing_array_B[:, rotation_A]
  807. if self._zscore:
  808. big_raster = np.vstack(
  809. [training_array_A, training_array_B]) # z-scoring using the training data
  810. big_mean = np.nanmean(big_raster, 0)
  811. big_std = np.nanstd(big_raster, 0)
  812. big_std[big_std == 0] = np.inf
  813. training_array_A = (training_array_A - big_mean) / big_std
  814. training_array_B = (training_array_B - big_mean) / big_std
  815. testing_array_A = (testing_array_A - big_mean) / big_std
  816. testing_array_B = (testing_array_B - big_mean) / big_std
  817. self._train(training_array_A, training_array_B, label_A, label_B)
  818. performance1 = self._test(testing_array_A, testing_array_B, label_A, label_B)
  819. self._train(testing_array_A, testing_array_B, label_A, label_B)
  820. performance2 = self._test(training_array_A, training_array_B, label_A, label_B)
  821. performances = [performance1, performance2]
  822. all_performances.append(performances)
  823. return np.nanmean(all_performances, 0)
  824. def parallelism_score_dichotomy(self, dichotomy: Union[str, list],
  825. max_semantic_dist: int = 1,
  826. shuffled: bool = False,
  827. method: str = 'pearson',
  828. return_combinations: bool = False):
  829. if self._is_multiclass:
  830. raise RuntimeError(f"{self.__class__.__name__}: this analysis is only defined for binary variables. ")
  831. if type(dichotomy) == str:
  832. dic = self._dichotomy_from_key(dichotomy)
  833. else:
  834. dic = dichotomy
  835. ndata = 2 * self._max_conditioned_data
  836. coding_directions = []
  837. set_A = dic[0]
  838. set_B = dic[1]
  839. for i in range(len(set_A)):
  840. for j in range(len(set_B)):
  841. test_condition_A = set_A[i]
  842. test_condition_B = set_B[j]
  843. if hamming(string_digits(test_condition_A), string_digits(test_condition_B)) <= max_semantic_dist:
  844. testing_array_A = sample_from_rasters(self.conditioned_rasters[test_condition_A], ndata=ndata)
  845. testing_array_B = sample_from_rasters(self.conditioned_rasters[test_condition_B], ndata=ndata)
  846. if shuffled:
  847. rotation_A = np.arange(testing_array_A.shape[1]).astype(int)
  848. rotation_B = np.arange(testing_array_B.shape[1]).astype(int)
  849. np.random.shuffle(rotation_A)
  850. np.random.shuffle(rotation_B)
  851. testing_array_A = testing_array_A[:, rotation_A]
  852. testing_array_B = testing_array_B[:, rotation_A]
  853. if self._zscore:
  854. big_raster = np.vstack([testing_array_A, testing_array_B])
  855. big_mean = np.nanmean(big_raster, 0)
  856. big_std = np.nanstd(big_raster, 0)
  857. big_std[big_std == 0] = np.inf
  858. testing_array_A = (testing_array_A - big_mean) / big_std
  859. testing_array_B = (testing_array_B - big_mean) / big_std
  860. vA = np.nanmean(testing_array_A, 0)
  861. vB = np.nanmean(testing_array_B, 0)
  862. coding_directions.append(vB - vA)
  863. parallelism_scores = []
  864. for i in range(len(coding_directions)):
  865. for j in range(i + 1, len(coding_directions)):
  866. if method == 'pearson':
  867. parallelism_scores.append(scipy.stats.pearsonr(coding_directions[i], coding_directions[j])[0])
  868. elif method == 'cosine':
  869. parallelism_scores.append(cosine(coding_directions[i], coding_directions[j]))
  870. elif method == 'spearman':
  871. parallelism_scores.append(scipy.stats.spearmanr(coding_directions[i], coding_directions[j])[0])
  872. else:
  873. raise ValueError(
  874. "The specified method is not supported, please use one of: pearson, cosine, spearman")
  875. if return_combinations:
  876. return np.asarray(parallelism_scores)
  877. else:
  878. return np.nanmean(parallelism_scores)
  879. # Dichotomy analysis functions with null model
  880. def decode_multiclass(self,
  881. classes,
  882. training_fraction: float,
  883. cross_validations: int = 10,
  884. ndata: Optional[int] = None,
  885. subsample: Optional[int] = 0,
  886. shuffled: Optional[bool] = False):
  887. """
  888. Multiclass decoding of a single variable.
  889. Parameters
  890. ----------
  891. classes : str or list
  892. If str, interpreted as the name of a variable in self.conditions,
  893. and the class structure is obtained via self._balanced_classes(classes).
  894. If list, it should be a list of lists of condition keys, e.g.
  895. [['00', '01'], ['10', '11'], ['20', '21']].
  896. training_fraction : float
  897. Fraction of trials used for training in each cross-validation fold.
  898. cross_validations : int
  899. Number of cross-validation iterations.
  900. ndata : int, optional
  901. Number of data points sampled per condition for training and testing.
  902. If None, defaults are chosen as in decode_dichotomy.
  903. subsample : int, optional
  904. If >0, a random subset of neurons of size=subsample is used.
  905. shuffled : bool, optional
  906. If True, use the geometric null model implemented by
  907. ``self._shuffle_conditioned_arrays`` before decoding, and
  908. restore the original ordering afterwards.
  909. Returns
  910. -------
  911. performance : np.ndarray
  912. Array of decoding performance values for each cross-validation.
  913. """
  914. # define class groups
  915. if isinstance(classes, str):
  916. class_sets = self._balanced_classes(classes)
  917. else:
  918. class_sets = classes
  919. # choose default ndata as in decode_dichotomy
  920. if ndata is None and self.n_brains == 1:
  921. ndata = self._max_conditioned_data
  922. if ndata is None and self.n_brains > 1:
  923. ndata = max(self._max_conditioned_data, 2 * self.n_neurons)
  924. # null-model shuffling of conditioned arrays
  925. if shuffled:
  926. self._shuffle_conditioned_arrays([[],[]]) # empty dic returns 0 for dic_key so it shuffles as for XORs
  927. performance = []
  928. CM = []
  929. for k in range(cross_validations):
  930. if subsample:
  931. self._generate_random_subset(subsample)
  932. training_arrays = []
  933. testing_arrays = []
  934. training_labels = []
  935. testing_labels = []
  936. # loop over classes (values of the decoded variable)
  937. for class_id, cond_list in enumerate(class_sets):
  938. class_training = []
  939. class_testing = []
  940. # sample from each full condition in this class
  941. for d in cond_list:
  942. training, testing = sample_training_testing_from_rasters(
  943. rasters=self.conditioned_rasters[d],
  944. ndata=int(ndata),
  945. training_fraction=training_fraction,
  946. trials=self.conditioned_trial_index[d],
  947. debug=self._debug
  948. )
  949. class_training.append(training)
  950. class_testing.append(testing)
  951. if not len(class_training):
  952. continue
  953. class_training = np.vstack(class_training)
  954. class_testing = np.vstack(class_testing)
  955. training_arrays.append(class_training)
  956. testing_arrays.append(class_testing)
  957. training_labels.append(
  958. np.repeat(class_id, class_training.shape[0]).astype(object)
  959. )
  960. testing_labels.append(
  961. np.repeat(class_id, class_testing.shape[0]).astype(object)
  962. )
  963. # if nothing was sampled, skip
  964. if not len(training_arrays):
  965. performance.append(np.nan)
  966. continue
  967. training_raster = np.vstack(training_arrays)
  968. testing_raster = np.vstack(testing_arrays)
  969. training_labels = np.hstack(training_labels).astype(int)
  970. testing_labels = np.hstack(testing_labels).astype(int)
  971. # z-score on training data, apply to both
  972. if self._zscore:
  973. big_mean = np.nanmean(training_raster, 0)
  974. big_std = np.nanstd(training_raster, 0)
  975. big_std[big_std == 0] = np.inf
  976. training_raster = (training_raster - big_mean) / big_std
  977. testing_raster = (testing_raster - big_mean) / big_std
  978. # local clone of the classifier; keep self.classifier untouched
  979. classifier = sklearn.base.clone(self.classifier)
  980. training_raster_sub = training_raster[:, self.subset]
  981. testing_raster_sub = testing_raster[:, self.subset]
  982. if self._debug:
  983. print("decode_multiclass: fold %d" % k)
  984. print("Training shape:", training_raster_sub.shape)
  985. print("Testing shape:", testing_raster_sub.shape)
  986. classifier.fit(training_raster_sub, training_labels)
  987. perf = classifier.score(testing_raster_sub, testing_labels)
  988. from sklearn.metrics import confusion_matrix
  989. y_pred = classifier.predict(testing_raster_sub)
  990. cm = confusion_matrix(testing_labels, y_pred, labels=np.arange(len(class_sets)))
  991. performance.append(perf)
  992. CM.append(cm)
  993. if subsample:
  994. self._reset_random_subset()
  995. # restore original ordering of conditioned rasters after null shuffling
  996. if shuffled:
  997. self._order_conditioned_rasters()
  998. return np.asarray(performance), np.asarray(CM)
  999. def decode_multiclass_with_nullmodel(self, variable: Union[str, list],
  1000. training_fraction: float,
  1001. cross_validations: int = 10,
  1002. nshuffles: int = 10,
  1003. ndata: Optional[int] = None,
  1004. return_CV: bool = False,
  1005. plot: bool = False,
  1006. dic_key: Optional[str] = None,
  1007. subsample: Optional[int] = 0,
  1008. **kwargs):
  1009. if dic_key is None and type(variable) == str:
  1010. dic_key = variable
  1011. elif dic_key is None:
  1012. dic_key = 'Var'
  1013. if type(variable) == str:
  1014. classes = self._balanced_classes(variable)
  1015. else:
  1016. classes = variable
  1017. perf, cm = self.decode_multiclass(classes=classes,
  1018. training_fraction=training_fraction,
  1019. cross_validations=cross_validations,
  1020. ndata=ndata,
  1021. subsample=subsample,
  1022. shuffled=False)
  1023. perf_null = []
  1024. cm_null = []
  1025. for n in tqdm(range(nshuffles)):
  1026. perfs_n, cm_n = self.decode_multiclass(classes=classes,
  1027. training_fraction=training_fraction,
  1028. cross_validations=cross_validations,
  1029. ndata=ndata,
  1030. subsample=subsample,
  1031. shuffled=True)
  1032. perf_null.append(perfs_n)
  1033. cm_null.append(cm_n)
  1034. if plot:
  1035. f, axs = plt.subplots(1, 2, figsize=(8, 5), gridspec_kw={'width_ratios': [2, 1]})
  1036. plot_confusion_matrix(np.nanmean(cm, 0), ax=axs[0])
  1037. plot_perfs_null_model(perfs={dic_key: np.nanmean(perf)},
  1038. perfs_nullmodel={dic_key: np.nanmean(perf_null, 1)},
  1039. chance=0,
  1040. ylow=np.nanmin(perf_null) * 0.75,
  1041. yhigh=1.05,
  1042. ax=axs[1])
  1043. if not return_CV:
  1044. perf = np.nanmean(perf)
  1045. perf_null = np.nanmean(perf_null, 1)
  1046. cm = np.nanmean(cm, 0)
  1047. cm_null = np.nanmean(cm, 1)
  1048. return perf, perf_null, cm, cm_null
  1049. def decode_with_nullmodel(self, dichotomy: Union[str, list],
  1050. training_fraction: float,
  1051. cross_validations: int = 10,
  1052. nshuffles: int = 10,
  1053. ndata: Optional[int] = None,
  1054. parallel: bool = False,
  1055. return_CV: bool = False,
  1056. testing_trials: Optional[list] = None,
  1057. plot: bool = False,
  1058. dic_key: Optional[str] = None,
  1059. subsample: Optional[int] = 0,
  1060. **kwargs) -> Tuple[Union[list, ndarray], ndarray]:
  1061. """
  1062. Function that performs cross-validated decoding of a specific dichotomy and compares the resulting values with
  1063. a null model where the relationship between the neural data and the two sides of the dichotomy is
  1064. shuffled.
  1065. Decoding is performed by sampling a balanced amount of data points from each condition in each class of the
  1066. dichotomy, so to ensure that only the desired variable is analyzed by balancing confounds.
  1067. Before sampling, each condition is individually divided into training and testing bins
  1068. by using the ``self.trial`` array specified in the data structure when constructing the ``Decodanda`` object.
  1069. Parameters
  1070. ----------
  1071. dichotomy : str || list
  1072. The dichotomy to be decoded, expressed in a double-list binary format, e.g. [['10', '11'], ['01', '00']], or as a variable name.
  1073. training_fraction:
  1074. the fraction of trials used for training in each cross-validation fold.
  1075. cross_validations:
  1076. the number of cross-validations.
  1077. nshuffles:
  1078. the number of null-model iterations of the decoding procedure.
  1079. ndata:
  1080. the number of data points (population vectors) sampled for training and for testing for each condition.
  1081. parallel:
  1082. if True, each cross-validation is performed by a dedicated thread (experimental, use with caution).
  1083. return_CV:
  1084. if True, invidual cross-validation values are returned in a list. Otherwise, the average performance over the cross-validation folds is returned.
  1085. testing_trials:
  1086. if specified, data sampled from the specified trial numbers will be used for testing, and the remaining ones for training.
  1087. plot:
  1088. if True, a visualization of the decoding results is shown.
  1089. dic_key:
  1090. if specified, weights of the decoding analysis will be saved in self.decoding_weights using dic_key as the dictionary key.
  1091. subsample:
  1092. if >0, a random subsample of neurons of size=subsample will be used at each cross-validation
  1093. Returns
  1094. -------
  1095. performances, null_performances: list of decoding performance values for each cross-validation.
  1096. See Also
  1097. --------
  1098. Decodanda.decode_dichotomy : The method used for each decoding iteration.
  1099. Note
  1100. ----
  1101. ``dichotomy`` can be passed as a string or as a list.
  1102. If a string is passed, it has to be a name of one of the variables specified in the conditions dictionary.
  1103. If a list is passed, it needs to contain two lists in the shape [[...], [...]].
  1104. Each sub list contains the conditions used to define one of the two decoded classes
  1105. in binary notation.
  1106. For example, if the data set has two variables
  1107. ``stimulus`` :math:`\\in` {-1, 1} and ``action`` :math:`\\in` {-1, 1}, the condition
  1108. ``stimulus=-1`` & ``action=-1`` will correspond to the binary notation ``'00'``,
  1109. the condition ``stimulus=+1`` & ``action=-1`` will correspond to ``10`` and so on.
  1110. Therefore, the notation:
  1111. >>> dic = 'stimulus'
  1112. is equivalent to
  1113. >>> dic = [['00', '01'], ['10', '11']]
  1114. and
  1115. >>> dic = 'action'
  1116. is equivalent to
  1117. >>> dic = [['00', '10'], ['01', '11']]
  1118. However, not all dichotomies have names (are semantic). For example, the dichotomy
  1119. >>> [['01','10'], ['00', '11']]
  1120. can only be defined using the binary notation.
  1121. Note that this function gives you the flexibility to use sub-sets of conditions, for example
  1122. >>> dic = [['10'], ['01']]
  1123. will decode stimulus=1 & action=-1 vs. stimulus=-1 & action=1
  1124. Example
  1125. -------
  1126. >>> data = generate_synthetic_data(keyA='stimulus', keyB='action')
  1127. >>> dec = Decodanda(data=data, conditions={'stimulus': [-1, 1], 'action': [-1, 1]})
  1128. >>> perf, null = dec.decode_with_nullmodel('stimulus', training_fraction=0.75, cross_validations=10, nshuffles=20)
  1129. >>> perf
  1130. 0.88
  1131. >>> null
  1132. [0.51, 0.54, 0.48, ..., 0.46] # 25 values
  1133. """
  1134. if self._is_multiclass:
  1135. raise RuntimeError(f"{self.__class__.__name__}: this analysis is only defined for binary variables. ")
  1136. if type(dichotomy) == str:
  1137. dic = self._dichotomy_from_key(dichotomy)
  1138. else:
  1139. dic = dichotomy
  1140. d_performances = self.decode_dichotomy(dichotomy=dic,
  1141. training_fraction=training_fraction,
  1142. cross_validations=cross_validations,
  1143. ndata=ndata,
  1144. parallel=parallel,
  1145. testing_trials=testing_trials,
  1146. dic_key=dic_key,
  1147. subsample=subsample)
  1148. if return_CV:
  1149. data_performance = d_performances
  1150. else:
  1151. data_performance = np.nanmean(d_performances)
  1152. if self._verbose and nshuffles:
  1153. print(
  1154. "\n[decode_with_nullmodel]\t data <p> = %.2f" % np.nanmean(d_performances))
  1155. print('\n[decode_with_nullmodel]\tLooping over null model shuffles.')
  1156. count = tqdm(range(nshuffles))
  1157. else:
  1158. count = range(nshuffles)
  1159. null_model_performances = np.zeros(nshuffles)
  1160. for n in count:
  1161. performances = self.decode_dichotomy(dichotomy=dic,
  1162. training_fraction=training_fraction,
  1163. cross_validations=cross_validations,
  1164. ndata=ndata,
  1165. parallel=parallel,
  1166. testing_trials=testing_trials,
  1167. shuffled=True,
  1168. dic_key=dic_key,
  1169. subsample=subsample)
  1170. null_model_performances[n] = np.nanmean(performances)
  1171. if plot:
  1172. visualize_decoding(self, dic, d_performances, null_model_performances,
  1173. training_fraction=training_fraction, ndata=ndata, testing_trials=testing_trials)
  1174. return data_performance, null_model_performances
  1175. def CCGP_with_nullmodel(self, dichotomy: Union[str, list],
  1176. resamplings: int = 5,
  1177. nshuffles: int = 25,
  1178. ndata: Optional[int] = None,
  1179. max_semantic_dist: int = 1,
  1180. split_rule='OneOut',
  1181. return_combinations: bool = False,
  1182. **kwargs):
  1183. """
  1184. Function that performs the cross-condition generalization performance analysis (CCGP, Bernardi et al. 2020, Cell)
  1185. for a given variable, specified through its corresponding dichotomy.
  1186. This function tests how well a given coding strategy for the given variable generalizes
  1187. when the other variables are changed and compares the
  1188. resulting values with a geometrical null model that keeps variables decodable but randomly
  1189. displaces conditions in the neural activity space, hence breaking any coding parallelism and generizability.
  1190. See Bernardi et al 2020 & Boyle, Posani et al. 2023 for more details.
  1191. Parameters
  1192. ----------
  1193. dichotomy : str || list
  1194. The dichotomy corresponding to the variable to be tested, expressed in a double-list binary format, e.g. [['10', '11'], ['01', '00']], or as a variable name.
  1195. resamplings:
  1196. The number of iterations for each decoding analysis. The returned performance value is the average over these resamplings.
  1197. nshuffles:
  1198. The number of null-model iterations for the CCGP analysis.
  1199. ndata:
  1200. The number of data points (population vectors) sampled for training and for testing for each condition.
  1201. max_semantic_dist:
  1202. The maximum semantic distance (number of variables that change value) between conditions in the held-out pair used to test the classifier.
  1203. split_rule:
  1204. The way conditions are split in training and testing. OneOut (default), name of a variable, or dichotomy in the double-list binary format. If OneOut is used, one pair of conditions is held out and the rest is used to train the classifier; if a variable is specified, then CCGP is computed specifically across that variable, balancing any third (or further) variables during sampling.
  1205. return_combinations:
  1206. If True, returns all the individual performances for cross-conditions train-test splits, otherwise returns the average over combinations.
  1207. Returns
  1208. -------
  1209. ccgp: mean of performance values for each cross-condition training-testing split (or list, if ``return_combinations=True``).
  1210. null: a list of null values for the mean ccgp
  1211. Note
  1212. ----
  1213. This function trains the ``self._classifier`` to decode the given variable in a sub-set
  1214. of conditions, and tests it on the held-out set.
  1215. The split of training and testing conditions is decided by the ``max_semantic_dist`` parameter: if set to 1,
  1216. only pairs of conditions that have all variables in common except the specified one are held out to test the
  1217. classifier.
  1218. For example, if the data set has two variables
  1219. ``stimulus`` :math:`\\in` {-1, 1} and ``action`` :math:`\\in` {-1, 1}, to compute CCGP for ``stimulus``
  1220. with ``max_semantic_dist=1`` this function will train the classifier on
  1221. ``(stimulus = -1, action = -1)`` vs. ``(stimulus = 1, action = -1)``
  1222. And test it on
  1223. ``(stimulus = -1, action = 1)`` vs. ``(stimulus = 1, action = 1)``
  1224. note that action is kept fixed within the training and testing conditions.
  1225. If instead we use ``max_semantic_dist=2``, all possible combinations are used, including training on
  1226. ``(stimulus = -1, action = -1)`` vs. ``(stimulus = 1, action = 1)``
  1227. and testing on
  1228. ``(stimulus = -1, action = 1)`` vs. ``(stimulus = 1, action = -1)``
  1229. ``dichotomy`` can be passed as a string or as a list.
  1230. If a string is passed, it has to be a name of one of the variables specified in the conditions dictionary.
  1231. If a list is passed, it needs to contain two lists in the shape [[...], [...]].
  1232. Each sub list contains the conditions used to define one of the two decoded classes
  1233. in binary notation.
  1234. For example, if the data set has two variables
  1235. ``stimulus`` :math:`\\in` {-1, 1} and ``action`` :math:`\\in` {-1, 1}, the condition
  1236. ``stimulus=-1`` & ``action=-1`` will correspond to the binary notation ``'00'``,
  1237. the condition ``stimulus=+1`` & ``action=-1`` will correspond to ``10`` and so on.
  1238. Therefore, if ``stimulus`` is the first variable in the conditions dictionary, its corresponding dichotomy is
  1239. >>> stimulus = [['00', '01'], ['10', '11']]
  1240. Example
  1241. -------
  1242. >>> data = generate_synthetic_data(keyA='stimulus', keyB='action')
  1243. >>> dec = Decodanda(data=data, conditions={'stimulus': [-1, 1], 'action': [-1, 1]})
  1244. >>> perf, null = dec.CCGP_with_nullmodel('stimulus', nshuffles=10)
  1245. >>> perf
  1246. 0.85
  1247. >>> null
  1248. [0.44, 0.48, ..., 0.54] # 10 values
  1249. """
  1250. if self._is_multiclass:
  1251. raise RuntimeError(f"{self.__class__.__name__}: this analysis is only defined for binary variables. ")
  1252. performances = self.CCGP_dichotomy(dichotomy=dichotomy, resamplings=resamplings, ndata=ndata,
  1253. max_semantic_dist=max_semantic_dist, split_rule=split_rule)
  1254. if return_combinations:
  1255. ccgp = performances
  1256. else:
  1257. ccgp = np.nanmean(performances)
  1258. if self._verbose and nshuffles:
  1259. print("\t\t[CCGP_with_nullmodel]\t\t----- Data: <p> = %.2f -----\n" % np.nanmean(performances))
  1260. count = tqdm(range(nshuffles))
  1261. else:
  1262. count = range(nshuffles)
  1263. shuffled_ccgp = []
  1264. for n in count:
  1265. performances = self.CCGP_dichotomy(dichotomy=dichotomy,
  1266. resamplings=resamplings,
  1267. ndata=ndata,
  1268. max_semantic_dist=max_semantic_dist,
  1269. split_rule=split_rule,
  1270. shuffled=True)
  1271. if return_combinations:
  1272. shuffled_ccgp.append(performances)
  1273. else:
  1274. shuffled_ccgp.append(np.nanmean(performances))
  1275. return ccgp, shuffled_ccgp
  1276. def PS_with_nullmodel(self, dichotomy: Union[str, list],
  1277. nshuffles: int = 25,
  1278. max_semantic_dist: int = 1,
  1279. method: str = 'pearson',
  1280. return_combinations: bool = False,
  1281. **kwargs):
  1282. scores = self.parallelism_score_dichotomy(dichotomy=dichotomy,
  1283. method=method,
  1284. max_semantic_dist=max_semantic_dist,
  1285. return_combinations=return_combinations)
  1286. if self._verbose and nshuffles:
  1287. print("\t\t[PS_with_nullmodel]\t\t----- Data: <p> = %.2f -----\n" % np.nanmean(scores))
  1288. count = tqdm(range(nshuffles))
  1289. else:
  1290. count = range(nshuffles)
  1291. shuffled_scores = []
  1292. for n in count:
  1293. scores_null = self.parallelism_score_dichotomy(dichotomy=dichotomy,
  1294. method=method,
  1295. max_semantic_dist=max_semantic_dist,
  1296. return_combinations=return_combinations,
  1297. shuffled=True)
  1298. shuffled_scores.append(scores_null)
  1299. return scores, shuffled_scores
  1300. # Decoding analysis for semantic dichotomies
  1301. def decode(self, training_fraction: float,
  1302. cross_validations: int = 10,
  1303. nshuffles: int = 10,
  1304. ndata: Optional[int] = None,
  1305. subsample: Optional[int] = 0,
  1306. parallel: bool = False,
  1307. non_semantic: bool = False,
  1308. return_CV: bool = False,
  1309. testing_trials: Optional[list] = None,
  1310. plot: bool = False,
  1311. ax: Optional[plt.Axes] = None,
  1312. plot_all: bool = False,
  1313. **kwargs):
  1314. """
  1315. Main function to decode the variables specified in the ``conditions`` dictionary.
  1316. It returns a single decoding value per variable which represents the average over
  1317. the cross-validation folds.
  1318. It also returns an array of null-model values for each variable to test the significance of
  1319. the corresponding decoding result.
  1320. Notes
  1321. -----
  1322. Each decoding analysis is performed by first re-sampling an equal number of data points
  1323. from each condition (combination of variable values), so to ensure that possible confounds
  1324. due to correlated conditions are balanced out.
  1325. Before sampling, each condition is individually divided into training and testing bins
  1326. by using the ``self.trial`` array specified in the data structure when constructing the ``Decodanda`` object.
  1327. To generate the null model values, the relationship between the neural data and
  1328. the decoded variable is randomly shuffled. Eeach null model value corresponds to the
  1329. average across ``cross_validations``` iterations after a single data shuffle.
  1330. If ``non_semantic=True``, dichotomies that do not correspond to variables will also be decoded.
  1331. Note that, in the case of 2 variables, there is only one non-semantic dichotomy
  1332. (corresponding to grouping together conditions that have the same XOR value in the
  1333. binary notation: ``[['10', '01'], ['11', '00']]``). However, the number of non-semantic dichotomies
  1334. grows exponentially with the number of conditions, so use with caution if more than two variables
  1335. are specified in the conditions dictionary.
  1336. Parameters
  1337. ----------
  1338. training_fraction:
  1339. the fraction of trials used for training in each cross-validation fold.
  1340. cross_validations:
  1341. the number of cross-validations.
  1342. nshuffles:
  1343. the number of null-model iterations of the decoding procedure.
  1344. ndata:
  1345. the number of data points (population vectors) sampled for training and for testing for each condition.
  1346. subsample:
  1347. if >0, a random subsample of neurons of size=subsample will be used at each cross-validation
  1348. parallel:
  1349. if True, each cross-validation is performed by a dedicated thread (experimental, use with caution).
  1350. return_CV:
  1351. if True, invidual cross-validation values are returned in a list. Otherwise, the average performance over the cross-validation folds is returned.
  1352. testing_trials:
  1353. if specified, data sampled from the specified trial numbers will be used for testing, and the remaining ones for training.
  1354. non_semantic:
  1355. if True, non-semantic dichotomies (i.e., dichotomies that do not correspond to a variable) will also be decoded.
  1356. plot:
  1357. if True, a visualization of the decoding results is shown.
  1358. ax:
  1359. if specified and ``plot=True``, the results will be displayed in the specified axis instead of a new figure.
  1360. plot_all:
  1361. if True, a more in-depth visualization of the decoding results and of the decoded data is shown.
  1362. Returns
  1363. -------
  1364. perfs:
  1365. a dictionary containing the decoding performances for all variables in the form of ``{var_name_1: performance1, var_name_2: performance2, ...}``
  1366. null:
  1367. a dictionary containing an array of null model decoding performance for each variable in the form ``{var_name_1: [...], var_name_2: [...], ...}``.
  1368. See Also
  1369. --------
  1370. Decodanda.decode_with_nullmodel: The method used for each decoding analysis.
  1371. Example
  1372. -------
  1373. >>> from decodanda import Decodanda, generate_synthetic_data
  1374. >>> data = generate_synthetic_data(keyA='stimulus', keyB='action')
  1375. >>> dec = Decodanda(data=data, conditions={'stimulus': [-1, 1], 'action': [-1, 1]})
  1376. >>> perfs, null = dec.decode(training_fraction=0.75, cross_validations=10, nshuffles=20)
  1377. >>> perfs
  1378. {'stimulus': 0.88, 'action': 0.85} # mean over 10 cross-validation folds
  1379. >>> null
  1380. {'stimulus': [0.51, ..., 0.46], 'action': [0.48, ..., 0.55]} # null model means, 20 values each
  1381. """
  1382. perfs = {}
  1383. perfs_nullmodel = {}
  1384. if self._is_multiclass:
  1385. for key in self._semantic_keys:
  1386. print("\nTesting multiclass decoding performance for semantic variable: ", key)
  1387. performance, null_model_performances, cm_data, cm_null = self.decode_multiclass_with_nullmodel(
  1388. variable=key,
  1389. training_fraction=training_fraction,
  1390. cross_validations=cross_validations,
  1391. ndata=ndata,
  1392. nshuffles=nshuffles,
  1393. return_CV=return_CV,
  1394. plot=plot_all,
  1395. subsample=subsample)
  1396. perfs[key] = performance
  1397. perfs_nullmodel[key] = null_model_performances
  1398. if plot:
  1399. if not ax:
  1400. f, ax = plt.subplots(figsize=(0.5 + 1.8 * len(perfs.keys()), 3.5))
  1401. plot_perfs_null_model(perfs, perfs_nullmodel, ylabel='Decoding performance', ax=ax, marker='o', ylow=0,
  1402. chance=0, **kwargs)
  1403. return perfs, perfs_nullmodel
  1404. semantic_dics, semantic_keys = self._find_semantic_dichotomies()
  1405. for key, dic in zip(semantic_keys, semantic_dics):
  1406. if self._verbose:
  1407. print("\nTesting decoding performance for semantic dichotomy: ", key)
  1408. performance, null_model_performances = self.decode_with_nullmodel(
  1409. dic,
  1410. training_fraction,
  1411. cross_validations=cross_validations,
  1412. ndata=ndata,
  1413. nshuffles=nshuffles,
  1414. parallel=parallel,
  1415. return_CV=return_CV,
  1416. testing_trials=testing_trials,
  1417. plot=plot_all,
  1418. subsample=subsample)
  1419. perfs[key] = performance
  1420. perfs_nullmodel[key] = null_model_performances
  1421. if non_semantic and len(self.conditions) == 2:
  1422. xor_dic = [['01', '10'], ['00', '11']]
  1423. perfs_xor, perfs_null_xor = self.decode_with_nullmodel(dichotomy=xor_dic,
  1424. training_fraction=training_fraction,
  1425. cross_validations=cross_validations,
  1426. nshuffles=nshuffles,
  1427. parallel=parallel,
  1428. ndata=ndata,
  1429. return_CV=return_CV,
  1430. testing_trials=testing_trials,
  1431. plot=plot_all,
  1432. subsample=subsample)
  1433. perfs['XOR'] = perfs_xor
  1434. perfs_nullmodel['XOR'] = perfs_null_xor
  1435. if non_semantic and len(self.conditions) > 2:
  1436. dics = self._find_nonsemantic_dichotomies()
  1437. for dic in dics:
  1438. dic_key = '_'.join(dic[0]) + '__' + '_'.join(dic[1])
  1439. perfs_dic, null_dic = self.decode_with_nullmodel(dichotomy=dic,
  1440. training_fraction=training_fraction,
  1441. cross_validations=cross_validations,
  1442. nshuffles=nshuffles,
  1443. parallel=parallel,
  1444. ndata=ndata,
  1445. return_CV=return_CV,
  1446. testing_trials=testing_trials,
  1447. plot=plot_all,
  1448. subsample=subsample)
  1449. perfs[dic_key] = perfs_dic
  1450. perfs_nullmodel[dic_key] = null_dic
  1451. if plot:
  1452. if not ax:
  1453. f, ax = plt.subplots(figsize=(0.5 + 1.8 * len(perfs.keys()), 3.5))
  1454. plot_perfs_null_model(perfs, perfs_nullmodel, ylabel='Decoding performance', ax=ax, marker='o', **kwargs)
  1455. return perfs, perfs_nullmodel
  1456. # Geometrical analysis for semantic dichotomies
  1457. def CCGP(self, resamplings=5,
  1458. nshuffles: int = 25,
  1459. ndata: Optional[int] = None,
  1460. max_semantic_dist: int = 1,
  1461. plot: bool = False,
  1462. ax: Optional[plt.Axes] = None,
  1463. **kwargs):
  1464. """
  1465. Main function that performs the cross-condition generalization performance analysis (CCGP, Bernardi et al. 2020, Cell)
  1466. for the variables specified through the ``conditions`` dictionary.
  1467. It returns a single ccgp value per variable which represents the average over
  1468. all cross-condition train-test splits. This function uses split_rule='OneOut' as a default.
  1469. It also returns an array of null-model values for each variable to test the significance of
  1470. the corresponding ccgp result. The employed geometrical null model keeps variables decodable but randomly
  1471. displaces conditions in the neural activity space, hence breaking any coding parallelism and generizability.
  1472. See Bernardi et al 2020 & Boyle, Posani et al. 2023 for more details.
  1473. Parameters
  1474. ----
  1475. resamplings:
  1476. The number of iterations for each decoding analysis. The returned performance value is the average over these resamplings.
  1477. nshuffles:
  1478. The number of null-model iterations for the CCGP analysis.
  1479. ndata:
  1480. The number of data points (population vectors) sampled for training and for testing for each condition.
  1481. max_semantic_dist:
  1482. The maximum semantic distance (number of variables that change value) between conditions in the held-out pair used to test the classifier.
  1483. plot:
  1484. if True, a visualization of the decoding results is shown.
  1485. ax:
  1486. if specified and ``plot=True``, the results will be displayed in the specified axis instead of a new figure.
  1487. Returns
  1488. -------
  1489. performance: mean of performance values for each cross-condition training-testing split.
  1490. null: a list of null values for the generalization performance
  1491. See Also
  1492. --------
  1493. Decodanda.CCGP_with_nullmodel
  1494. Note
  1495. ----
  1496. For each variable, this function trains the ``self._classifier`` to decode the given variable in a sub-set
  1497. of conditions, and tests it on the held-out set.
  1498. The split of training and testing conditions is performed by keeping the semantic distance between held out
  1499. conditions to 1 (``max_semantic_dist=1`` in the CCGP_dichotomy function).
  1500. For example, if the data set has two variables:
  1501. ``stimulus`` :math:`\\in` {-1, 1} and ``action`` :math:`\\in` {-1, 1}, to compute CCGP for ``stimulus``
  1502. This function will train the classifier on
  1503. ``(stimulus = -1, action = -1)`` vs. ``(stimulus = 1, action = -1)``
  1504. And test it on
  1505. ``(stimulus = -1, action = 1)`` vs. ``(stimulus = 1, action = 1)``
  1506. And vice-versa. Note that action is kept fixed within the training and testing conditions.
  1507. Example
  1508. -------
  1509. >>> data = generate_synthetic_data(keyA='stimulus', keyB='action')
  1510. >>> dec = Decodanda(data=data, conditions={'stimulus': [-1, 1], 'action': [-1, 1]})
  1511. >>> perfs, null = dec.CCGP(nshuffles=10)
  1512. >>> perfs
  1513. {'stimulus': 0.81, 'action': 0.79} # each value is the mean over 2 cross-condition train-test splits
  1514. >>> null
  1515. {'stimulus': [0.51, ..., 0.46], 'action': [0.48, ..., 0.55]} # null model means, 10 values each
  1516. """
  1517. semantic_dics, semantic_keys = self._find_semantic_dichotomies()
  1518. ccgp = {}
  1519. ccgp_nullmodel = {}
  1520. for key, dic in zip(semantic_keys, semantic_dics):
  1521. if self._verbose:
  1522. print("\nTesting CCGP for semantic dichotomy: ", key)
  1523. data_ccgp, null_ccgps = self.CCGP_with_nullmodel(dichotomy=dic,
  1524. resamplings=resamplings,
  1525. nshuffles=nshuffles,
  1526. ndata=ndata,
  1527. max_semantic_dist=max_semantic_dist)
  1528. ccgp[key] = data_ccgp
  1529. ccgp_nullmodel[key] = null_ccgps
  1530. if plot:
  1531. if not ax:
  1532. f, ax = plt.subplots(figsize=(0.5 + 1.8 * len(semantic_dics), 3.5))
  1533. plot_perfs_null_model(ccgp, ccgp_nullmodel, ylabel='CCGP', ax=ax, marker='s', **kwargs)
  1534. return ccgp, ccgp_nullmodel
  1535. def PS(self, nshuffles: int = 25,
  1536. max_semantic_dist: int = 1,
  1537. method: str = 'pearson',
  1538. plot: bool = False,
  1539. ax: Optional[plt.Axes] = None,
  1540. **kwargs):
  1541. semantic_dics, semantic_keys = self._find_semantic_dichotomies()
  1542. ps = {}
  1543. ps_nullmodel = {}
  1544. for key, dic in zip(semantic_keys, semantic_dics):
  1545. if self._verbose:
  1546. print("\nTesting PS for semantic dichotomy: ", key)
  1547. data_ps, null_ps = self.PS_with_nullmodel(dichotomy=dic,
  1548. nshuffles=nshuffles,
  1549. method=method,
  1550. max_semantic_dist=max_semantic_dist)
  1551. ps[key] = data_ps
  1552. ps_nullmodel[key] = null_ps
  1553. if plot:
  1554. if not ax:
  1555. f, ax = plt.subplots(figsize=(0.5 + 1.8 * len(semantic_dics), 3.5))
  1556. plot_perfs_null_model(ps, ps_nullmodel, ylabel='Parallelism Score', ax=ax, ylow=-1.05, yhigh=1.05, chance=0,
  1557. **kwargs)
  1558. return ps, ps_nullmodel
  1559. def semantic_score_geometry(self,
  1560. training_fraction: float = 0.75,
  1561. cross_validations: int = 10,
  1562. nshuffles: int = 10,
  1563. ndata: Optional[int] = None,
  1564. visualize=True):
  1565. """
  1566. This function performs a balanced decoding analysis for each possible dichotomy, and
  1567. plots the result sorted by a semantic score that tells how close each dichotomy is to
  1568. any of the specified variables. A semantic dichotomy has ``semantic_score = 1``, the
  1569. XOR dichotomy has ``semantic_score = 0``.
  1570. Parameters
  1571. ----------
  1572. training_fraction:
  1573. the fraction of trials used for training in each cross-validation fold.
  1574. cross_validations:
  1575. the number of cross-validations.
  1576. nshuffles:
  1577. the number of null-model iterations of the decoding procedure.
  1578. ndata:
  1579. the number of data points (population vectors) sampled for training and for testing for each condition.
  1580. visualize:
  1581. if ``True``, the decoding results are shown in a figure.
  1582. Returns
  1583. -------
  1584. dichotomies_data:
  1585. Two lists, one containing all the dichotomies in binary notation
  1586. and one containing the corresponding semantic score.
  1587. decoding_data:
  1588. Two dictionaries, one containing the decoding performances for all dichotomies
  1589. and one containing all the corresponding lists of null model performances.
  1590. CCGP_data:
  1591. Two dictionaries, one containing the CCGP values for all dichotomies
  1592. and one containing all the corresponding lists of null model values.
  1593. """
  1594. all_dics = generate_dichotomies(self.n_conditions)[1]
  1595. semantic_overlap = []
  1596. dic_name = []
  1597. for i, dic in enumerate(all_dics):
  1598. semantic_overlap.append(semantic_score(dic))
  1599. dic_name.append(str(self._dic_key(dic)))
  1600. semantic_overlap = np.asarray(semantic_overlap)
  1601. # sorting dichotomies wrt semantic overlap
  1602. dic_name = np.asarray(dic_name)[np.argsort(semantic_overlap)[::-1]]
  1603. all_dics = list(np.asarray(all_dics)[np.argsort(semantic_overlap)[::-1]])
  1604. semantic_overlap = semantic_overlap[np.argsort(semantic_overlap)[::-1]]
  1605. semantic_overlap = (semantic_overlap - np.min(semantic_overlap)) / (
  1606. np.max(semantic_overlap) - np.min(semantic_overlap))
  1607. # decoding all dichotomies
  1608. decoding_results = []
  1609. decoding_null = []
  1610. for i, dic in enumerate(all_dics):
  1611. res, null = self.decode_with_nullmodel(dic,
  1612. training_fraction=training_fraction,
  1613. cross_validations=cross_validations,
  1614. nshuffles=nshuffles,
  1615. ndata=ndata)
  1616. # print(i, res)
  1617. decoding_results.append(res)
  1618. decoding_null.append(null)
  1619. # CCGP all dichotomies
  1620. CCGP_results = []
  1621. CCGP_null = []
  1622. for i, dic in enumerate(all_dics):
  1623. # print(dic)
  1624. res, null = self.CCGP_with_nullmodel(dic,
  1625. nshuffles=nshuffles,
  1626. ndata=ndata,
  1627. max_semantic_dist=self.n_conditions)
  1628. # print(i, res)
  1629. CCGP_results.append(res)
  1630. CCGP_null.append(null)
  1631. # plotting
  1632. if visualize:
  1633. if self.n_conditions > 2:
  1634. f, axs = plt.subplots(2, 1, figsize=(6, 6))
  1635. axs[0].set_xlabel('Dichotomy (ordered by semantic score)')
  1636. axs[1].set_xlabel('Dichotomy (ordered by semantic score)')
  1637. else:
  1638. f, axs = plt.subplots(1, 2, figsize=(6, 3.5))
  1639. axs[0].set_xlabel('Dichotomy')
  1640. axs[1].set_xlabel('Dichotomy')
  1641. axs[0].set_xlim([-0.5, 2.5])
  1642. axs[1].set_xlim([-0.5, 2.5])
  1643. axs[0].set_ylabel('Decoding Performance')
  1644. axs[1].set_ylabel('CCGP')
  1645. axs[0].axhline([0.5], color='k', linestyle='--', alpha=0.5)
  1646. axs[1].axhline([0.5], color='k', linestyle='--', alpha=0.5)
  1647. axs[0].set_xticks([])
  1648. axs[1].set_xticks([])
  1649. axs[0].set_ylim([0, 1.05])
  1650. axs[1].set_ylim([0, 1.05])
  1651. sns.despine(f)
  1652. # visualize Decoding
  1653. for i in range(len(all_dics)):
  1654. if z_pval(decoding_results[i], decoding_null[i])[1] < 0.01:
  1655. axs[0].scatter(i, decoding_results[i], marker='o',
  1656. color=cm.cool(int(semantic_overlap[i] * 255)), edgecolor='k',
  1657. s=110, linewidth=2, linestyle='-')
  1658. elif z_pval(decoding_results[i], decoding_null[i])[1] < 0.05:
  1659. axs[0].scatter(i, decoding_results[i], marker='o',
  1660. color=cm.cool(int(semantic_overlap[i] * 255)), edgecolor='k',
  1661. s=110, linewidth=2, linestyle='--')
  1662. elif z_pval(decoding_results[i], decoding_null[i])[1] > 0.05:
  1663. axs[0].scatter(i, decoding_results[i], marker='o',
  1664. color=cm.cool(int(semantic_overlap[i] * 255)), edgecolor='k',
  1665. s=110, linewidth=2, linestyle='dotted')
  1666. axs[0].errorbar(i, np.nanmean(decoding_null[i]), np.nanstd(decoding_null[i]), color='k', alpha=0.3)
  1667. if dic_name[i] != '0':
  1668. axs[0].text(i, decoding_results[i] + 0.08, dic_name[i], rotation=90, fontsize=6, color='k',
  1669. ha='center', fontweight='bold')
  1670. # visualize CCGP
  1671. for i in range(len(all_dics)):
  1672. if z_pval(CCGP_results[i], CCGP_null[i])[1] < 0.01:
  1673. axs[1].scatter(i, CCGP_results[i], marker='s',
  1674. color=cm.cool(int(semantic_overlap[i] * 255)), edgecolor='k',
  1675. s=110, linewidth=2, linestyle='-')
  1676. elif z_pval(CCGP_results[i], CCGP_null[i])[1] < 0.05:
  1677. axs[1].scatter(i, CCGP_results[i], marker='s',
  1678. color=cm.cool(int(semantic_overlap[i] * 255)), edgecolor='k',
  1679. s=110, linewidth=2, linestyle='--')
  1680. elif z_pval(CCGP_results[i], CCGP_null[i])[1] > 0.05:
  1681. axs[1].scatter(i, CCGP_results[i], marker='s',
  1682. color=cm.cool(int(semantic_overlap[i] * 255)), edgecolor='k',
  1683. s=110, linewidth=2, linestyle='dotted')
  1684. axs[1].errorbar(i, np.nanmean(CCGP_null[i]), np.nanstd(CCGP_null[i]), color='k', alpha=0.3)
  1685. if dic_name[i] != '0':
  1686. axs[1].text(i, CCGP_results[i] + 0.08, dic_name[i], rotation=90, fontsize=6, color='k',
  1687. ha='center', fontweight='bold')
  1688. dichotomies_data = [all_dics, semantic_overlap]
  1689. decoding_data = [decoding_results, decoding_null]
  1690. CCGP_data = [CCGP_results, CCGP_null]
  1691. return dichotomies_data, decoding_data, CCGP_data
  1692. def shattering_dimensionality(self,
  1693. training_fraction: float = 0.75,
  1694. cross_validations: int = 10,
  1695. nshuffles: int = 10,
  1696. ndata: Optional[int] = None,
  1697. subsample: Optional[int] = 0,
  1698. p_threshold: float = 0.01,
  1699. visualize: bool = True,
  1700. semantic_names: Optional[dict] = None,
  1701. **kwargs):
  1702. """
  1703. This function computes shattering dimensionality as defined in Bernardi et al. 2020, i.e., as
  1704. the number of balanced dichotomies that a linear decoder can classify above chance levels.
  1705. Parameters
  1706. ----------
  1707. training_fraction:
  1708. the fraction of trials used for training in each cross-validation fold.
  1709. cross_validations:
  1710. the number of cross-validations.
  1711. nshuffles:
  1712. the number of null-model iterations of the decoding procedure.
  1713. ndata:
  1714. the number of data points (population vectors) sampled for training and for testing for each condition.
  1715. subsample:
  1716. if >0, a random subsample of neurons of size=subsample will be used at each cross-validation.
  1717. p_threshold:
  1718. p-value threshold (z-score from the null model) to consider a performance as statistically significant.
  1719. visualize:
  1720. if ``True``, the decoding results are shown in a figure.
  1721. Returns
  1722. -------
  1723. shattering_dim:
  1724. shattering dimensionality
  1725. perfs:
  1726. dictionary of decoding performance per dichotomy
  1727. null:
  1728. dictionary of lists of null model values per dichotomy
  1729. """
  1730. all_dics_names, all_dics = generate_dichotomies(self.n_conditions)
  1731. semantic_overlap = []
  1732. dic_name = []
  1733. is_semantic = []
  1734. perfs = {}
  1735. nulls = {}
  1736. for i, dic in enumerate(all_dics):
  1737. semantic_overlap.append(semantic_score(dic))
  1738. is_semantic.append(str(self._dic_key(dic)))
  1739. dic_name.append(all_dics_names[i])
  1740. semantic_overlap = np.asarray(semantic_overlap)
  1741. # sorting dichotomies wrt semantic overlap
  1742. dic_name = np.asarray(dic_name)[np.argsort(semantic_overlap)[::-1]]
  1743. all_dics = list(np.asarray(all_dics)[np.argsort(semantic_overlap)[::-1]])
  1744. is_semantic = list(np.asarray(is_semantic)[np.argsort(semantic_overlap)[::-1]])
  1745. semantic_overlap = semantic_overlap[np.argsort(semantic_overlap)[::-1]]
  1746. semantic_overlap = (semantic_overlap - np.min(semantic_overlap)) / (
  1747. np.max(semantic_overlap) - np.min(semantic_overlap))
  1748. # decoding all dichotomies
  1749. for i, dic in tqdm(enumerate(all_dics)):
  1750. res, null = self.decode_with_nullmodel(dic,
  1751. training_fraction=training_fraction,
  1752. cross_validations=cross_validations,
  1753. nshuffles=nshuffles,
  1754. ndata=ndata,
  1755. subsample=subsample)
  1756. perfs[dic_name[i]] = res
  1757. nulls[dic_name[i]] = null
  1758. ps = np.asarray([z_pval(perfs[dic_name[i]], nulls[dic_name[i]])[1] for i in range(len(dic_name))])
  1759. shattering_dim = np.nanmean(ps < p_threshold)
  1760. # plotting
  1761. if visualize:
  1762. f, ax = plt.subplots(figsize=(6, 3))
  1763. if self.n_conditions > 2:
  1764. ax.set_xlabel('Dichotomy (ordered by semantic score)')
  1765. else:
  1766. ax.set_xlabel('Dichotomy')
  1767. ax.set_ylabel('Decoding Performance')
  1768. ax.axhline([0.5], color='k', linestyle='--', alpha=0.5)
  1769. ax.set_xticks([])
  1770. ax.set_ylim([0, 1.05])
  1771. sns.despine(f)
  1772. # visualize Decoding
  1773. for i in range(len(all_dics)):
  1774. if z_pval(perfs[dic_name[i]], nulls[dic_name[i]])[1] < 0.01:
  1775. ax.scatter(i, perfs[dic_name[i]], marker='o',
  1776. color=cm.cool(int(semantic_overlap[i] * 255)), edgecolor='k',
  1777. s=110, linewidth=2, linestyle='-')
  1778. else:
  1779. ax.scatter(i, perfs[dic_name[i]], marker='o',
  1780. color=cm.cool(int(semantic_overlap[i] * 255)), edgecolor='none',
  1781. s=110, linewidth=0, linestyle='')
  1782. ax.errorbar(i, np.nanmean(nulls[dic_name[i]]), np.nanstd(nulls[dic_name[i]]), color='k', alpha=0.3)
  1783. if semantic_names is not None:
  1784. if dic_name[i] in list(semantic_names.keys()):
  1785. ax.text(i, max(0.61, perfs[dic_name[i]] + 0.06), semantic_names[dic_name[i]], rotation=90,
  1786. fontsize=6, color='k',
  1787. ha='center', fontweight='bold')
  1788. return shattering_dim, perfs, nulls, f
  1789. else:
  1790. return shattering_dim, perfs, nulls
  1791. def shattering_generalization(self,
  1792. nshuffles: int = 10,
  1793. ndata: Optional[int] = None,
  1794. p_threshold: float = 0.01,
  1795. visualize: bool = True,
  1796. semantic_names: Optional[dict] = None,
  1797. max_semantic_dist = 99,
  1798. **kwargs):
  1799. """
  1800. This function computes shattering generalization defined as the number of balanced dichotomies that
  1801. have a above-chance CCGP.
  1802. Parameters
  1803. ----------
  1804. nshuffles:
  1805. the number of null-model iterations of the decoding procedure.
  1806. ndata:
  1807. the number of data points (population vectors) sampled for training and for testing for each condition.
  1808. p_threshold:
  1809. p-value threshold (z-score from the null model) to consider a performance as statistically significant.
  1810. visualize:
  1811. if ``True``, the decoding results are shown in a figure.
  1812. Returns
  1813. -------
  1814. shattering_gen:
  1815. shattering dimensionality
  1816. perfs:
  1817. dictionary of decoding performance per dichotomy
  1818. null:
  1819. dictionary of lists of null model values per dichotomy
  1820. """
  1821. all_dics_names, all_dics = generate_dichotomies(self.n_conditions)
  1822. semantic_overlap = []
  1823. dic_name = []
  1824. is_semantic = []
  1825. perfs = {}
  1826. nulls = {}
  1827. for i, dic in enumerate(all_dics):
  1828. semantic_overlap.append(semantic_score(dic))
  1829. is_semantic.append(str(self._dic_key(dic)))
  1830. dic_name.append(all_dics_names[i])
  1831. semantic_overlap = np.asarray(semantic_overlap)
  1832. # sorting dichotomies wrt semantic overlap
  1833. dic_name = np.asarray(dic_name)[np.argsort(semantic_overlap)[::-1]]
  1834. all_dics = list(np.asarray(all_dics)[np.argsort(semantic_overlap)[::-1]])
  1835. is_semantic = list(np.asarray(is_semantic)[np.argsort(semantic_overlap)[::-1]])
  1836. semantic_overlap = semantic_overlap[np.argsort(semantic_overlap)[::-1]]
  1837. semantic_overlap = (semantic_overlap - np.min(semantic_overlap)) / (
  1838. np.max(semantic_overlap) - np.min(semantic_overlap))
  1839. # CCGP for all dichotomies
  1840. for i, dic in tqdm(enumerate(all_dics)):
  1841. res, null = self.CCGP_with_nullmodel(dic, resamplings=2,
  1842. nshuffles=nshuffles, ndata=ndata,
  1843. max_semantic_dist=max_semantic_dist)
  1844. perfs[dic_name[i]] = res
  1845. nulls[dic_name[i]] = null
  1846. ps = np.asarray([z_pval(perfs[dic_name[i]], nulls[dic_name[i]])[1] for i in range(len(dic_name))])
  1847. shattering_gen = np.nanmean(ps < p_threshold)
  1848. # plotting
  1849. if visualize:
  1850. f, ax = plt.subplots(figsize=(6, 3))
  1851. if self.n_conditions > 2:
  1852. ax.set_xlabel('Dichotomy (ordered by semantic score)')
  1853. else:
  1854. ax.set_xlabel('Dichotomy')
  1855. ax.set_ylabel('CCGP')
  1856. ax.axhline([0.5], color='k', linestyle='--', alpha=0.5)
  1857. ax.set_xticks([])
  1858. ax.set_ylim([0, 1.05])
  1859. sns.despine(f)
  1860. # visualize Decoding
  1861. for i in range(len(all_dics)):
  1862. if z_pval(perfs[dic_name[i]], nulls[dic_name[i]])[1] < 0.01:
  1863. ax.scatter(i, perfs[dic_name[i]], marker='o',
  1864. color=cm.cool(int(semantic_overlap[i] * 255)), edgecolor='k',
  1865. s=110, linewidth=2, linestyle='-')
  1866. else:
  1867. ax.scatter(i, perfs[dic_name[i]], marker='o',
  1868. color=cm.cool(int(semantic_overlap[i] * 255)), edgecolor='none',
  1869. s=110, linewidth=0, linestyle='')
  1870. ax.errorbar(i, np.nanmean(nulls[dic_name[i]]), np.nanstd(nulls[dic_name[i]]), color='k', alpha=0.3)
  1871. if semantic_names is not None:
  1872. if dic_name[i] in list(semantic_names.keys()):
  1873. ax.text(i, max(0.61, perfs[dic_name[i]] + 0.06), semantic_names[dic_name[i]], rotation=90,
  1874. fontsize=6, color='k',
  1875. ha='center', fontweight='bold')
  1876. return shattering_gen, perfs, nulls, f
  1877. else:
  1878. return shattering_gen, perfs, nulls
  1879. def CVI(self, training_fraction: float = 0.75,
  1880. cross_validations: int = 10,
  1881. nshuffles: int = 10,
  1882. ndata: Optional[int] = None,
  1883. return_splits: bool = False,
  1884. signed=False
  1885. ):
  1886. if ndata is None:
  1887. ndata = 2 * self._max_conditioned_data
  1888. dics, vars = self._find_semantic_dichotomies()
  1889. # data
  1890. results = {}
  1891. for v1 in range(len(vars)):
  1892. var1 = vars[v1]
  1893. dic1 = dics[v1]
  1894. for v2 in range(len(vars)):
  1895. if v2 != v1:
  1896. var2 = vars[v2]
  1897. dic2 = dics[v2]
  1898. results[f'{var1}-{var2}'] = []
  1899. for k in range(cross_validations):
  1900. perf = self._one_X_cv_step(dic1, dic2, training_fraction, ndata)
  1901. results[f'{var1}-{var2}'].append(perf)
  1902. if signed:
  1903. results[f'{var1}-{var2}'] = np.nanmean(np.asarray(results[f'{var1}-{var2}']))
  1904. else:
  1905. results[f'{var1}-{var2}'] = 0.5 + np.abs(
  1906. np.nanmean(np.asarray(results[f'{var1}-{var2}']) - 0.5))
  1907. # null
  1908. null = {key: [] for key in results}
  1909. for n in range(nshuffles):
  1910. self._shuffle_conditioned_arrays(dic='XOR')
  1911. for v1 in range(len(vars)):
  1912. var1 = vars[v1]
  1913. dic1 = dics[v1]
  1914. for v2 in range(len(vars)):
  1915. if v2 != v1:
  1916. var2 = vars[v2]
  1917. dic2 = dics[v2]
  1918. null_n = []
  1919. for k in range(cross_validations):
  1920. perf = self._one_X_cv_step(dic1, dic2, training_fraction, ndata)
  1921. null_n.append(perf)
  1922. null[f'{var1}-{var2}'].append(np.nanmean(null_n))
  1923. self._order_conditioned_rasters()
  1924. if not return_splits:
  1925. megakey = '-'.join(vars)
  1926. results_combined = {megakey: np.nanmean([results[key] for key in results])}
  1927. null_combined = {megakey: [np.nanmean([null[key][i] for key in null]) for i in range(nshuffles)]}
  1928. return results_combined, null_combined
  1929. return results, null
  1930. # Utilities
  1931. def visualize_PCA(self, **kwargs):
  1932. fig = visualize_PCA(self, **kwargs)
  1933. return fig
  1934. # __init__ utilities
  1935. def _divide_data_into_conditions(self, sessions):
  1936. # TODO: rename sessions into datasets?
  1937. for si, session in enumerate(sessions):
  1938. if self._verbose:
  1939. if hasattr(session, 'name'):
  1940. print("\t\t[Decodanda]\tbuilding conditioned rasters for session %s" % session.name)
  1941. else:
  1942. print("\t\t[Decodanda]\tbuilding conditioned rasters for session %u" % si)
  1943. session_conditioned_rasters = {}
  1944. session_conditioned_trial_index = {}
  1945. # exclude inactive neurons across the specified conditions
  1946. array = getattr(session, self._neural_attr)
  1947. T = array.shape[0]
  1948. total_mask = np.zeros(T, dtype=bool)
  1949. cond_of_bin = np.full(T, -1.0, dtype=float)
  1950. local_trial_of_bin = np.full(T, -1.0, dtype=float)
  1951. for ci, condition_vec in enumerate(self._condition_vectors):
  1952. mask = np.ones(T, dtype=bool)
  1953. for i, sk in enumerate(self._semantic_keys):
  1954. semantic_values = list(self.conditions[sk])
  1955. mask_i = self.conditions[sk][semantic_values[condition_vec[i]]](session)
  1956. mask &= mask_i
  1957. total_mask |= mask
  1958. if not np.any(mask):
  1959. continue
  1960. # local trials
  1961. if self._trial_attr is not None:
  1962. local_trials = np.asarray(getattr(session, self._trial_attr))[mask]
  1963. else:
  1964. if self._trial_chunk is None:
  1965. local_trials = contiguous_chunking(mask)[mask].astype(float)
  1966. else:
  1967. local_trials = contiguous_chunking(mask, self._trial_chunk)[mask].astype(float)
  1968. local_trials[np.isnan(local_trials)] = -1.0
  1969. # overlap check (assigned bins are those with cond_of_bin != -1)
  1970. if np.any(cond_of_bin[mask] != -1.0):
  1971. raise ValueError("Conditions overlap in time; cannot build a unique session-wide trial vector.")
  1972. cond_of_bin[mask] = float(ci)
  1973. local_trial_of_bin[mask] = local_trials
  1974. # factorize (ci, local_trial) -> global id
  1975. valid = (cond_of_bin != -1.0) & (local_trial_of_bin != -1.0)
  1976. pairs = np.stack([cond_of_bin[valid], local_trial_of_bin[valid]], axis=1)
  1977. _, inv = np.unique(pairs, axis=0, return_inverse=True)
  1978. session_trial_vector = np.full(T, -1.0, dtype=float)
  1979. session_trial_vector[valid] = inv.astype(float)
  1980. if self._time_attr is not None:
  1981. time_vector = getattr(session, self._time_attr)
  1982. time_separation_mask = enforce_min_time_separation(session_trial_vector, self._min_time_separation, time_vector)
  1983. if self._debug:
  1984. print('[Enforcing time separation] selected %u time bin over %u total time bins' % (np.sum(time_separation_mask), np.sum(total_mask)))
  1985. else:
  1986. time_separation_mask = np.ones(T, dtype=bool)
  1987. self._time_separation_masks.append(time_separation_mask)
  1988. self._session_trial_vectors.append(session_trial_vector)
  1989. min_activity_mask = np.sum(array[total_mask & time_separation_mask] != 0, 0) >= self._min_activations_per_cell
  1990. for condition_vec in self._condition_vectors:
  1991. # get the array from the session object
  1992. array = getattr(session, self._neural_attr)
  1993. array = array[:, min_activity_mask]
  1994. # create a mask that becomes more and more restrictive by iterating on semanting conditions
  1995. mask = np.ones(T, dtype=bool)
  1996. for i, sk in enumerate(self._semantic_keys):
  1997. semantic_values = list(self.conditions[sk])
  1998. mask_i = self.conditions[sk][semantic_values[condition_vec[i]]](session)
  1999. mask = mask & mask_i
  2000. mask = mask & time_separation_mask
  2001. # select bins conditioned on the semantic behavioural vector
  2002. conditioned_raster = array[mask, :]
  2003. # select trial numbers
  2004. conditioned_trial = session_trial_vector[mask]
  2005. # exclude empty time bins (only for binary discrete decoding)
  2006. if self._exclude_silent:
  2007. active_mask = np.sum(conditioned_raster, 1) > 0
  2008. conditioned_raster = conditioned_raster[active_mask, :]
  2009. conditioned_trial = conditioned_trial[active_mask]
  2010. # squeeze into trials
  2011. if self._trial_average:
  2012. unique_trials = np.unique(conditioned_trial[conditioned_trial != -1.0])
  2013. squeezed_raster = []
  2014. squeezed_trial_index = []
  2015. for t in unique_trials:
  2016. trial_raster = conditioned_raster[conditioned_trial == t]
  2017. squeezed_raster.append(np.nanmean(trial_raster, 0))
  2018. squeezed_trial_index.append(t)
  2019. # set the new arrays
  2020. conditioned_raster = np.asarray(squeezed_raster)
  2021. conditioned_trial = np.asarray(squeezed_trial_index)
  2022. # set the conditioned neural data in the conditioned_rasters dictionary
  2023. session_conditioned_rasters[string_digits(condition_vec)] = conditioned_raster
  2024. session_conditioned_trial_index[string_digits(condition_vec)] = conditioned_trial
  2025. if self._verbose:
  2026. semantic_vector_string = []
  2027. for i, sk in enumerate(self._semantic_keys):
  2028. semantic_values = list(self.conditions[sk])
  2029. semantic_vector_string.append("%s = %s" % (sk, semantic_values[condition_vec[i]]))
  2030. semantic_vector_string = ', '.join(semantic_vector_string)
  2031. if len(conditioned_raster):
  2032. print("\t\t\t(%s):\tSelected %u time bin out of %u, divided into %u trials - %u neurons"
  2033. % (semantic_vector_string, conditioned_raster.shape[0], len(array),
  2034. len(np.unique(conditioned_trial)), conditioned_raster.shape[1]))
  2035. else:
  2036. print("\t\t\t(%s):\tNo data found" % semantic_vector_string)
  2037. session_conditioned_data = [r.shape[0] for r in list(session_conditioned_rasters.values())]
  2038. session_conditioned_trials = [len(np.unique(c)) for c in list(session_conditioned_trial_index.values())]
  2039. self._max_conditioned_data = max([self._max_conditioned_data, np.max(session_conditioned_data)])
  2040. self._min_conditioned_data = min([self._min_conditioned_data, np.min(session_conditioned_data)])
  2041. # if the session has enough data for each condition, append it to the main data dictionary
  2042. if np.min(session_conditioned_data) >= self._min_data_per_condition and \
  2043. np.min(session_conditioned_trials) >= self._min_trials_per_condition:
  2044. for cv in self._condition_vectors:
  2045. self.conditioned_rasters[string_digits(cv)].append(session_conditioned_rasters[string_digits(cv)])
  2046. self.conditioned_trial_index[string_digits(cv)].append(
  2047. session_conditioned_trial_index[string_digits(cv)])
  2048. if self._verbose:
  2049. print('\n')
  2050. self.n_brains += 1
  2051. self.n_neurons += list(session_conditioned_rasters.values())[0].shape[1]
  2052. self.which_brain.append(np.ones(list(session_conditioned_rasters.values())[0].shape[1]) * self.n_brains)
  2053. else:
  2054. if self._verbose:
  2055. print('\t\t\t===> Session discarded for insufficient data.\n')
  2056. if len(self.which_brain):
  2057. self.which_brain = np.hstack(self.which_brain)
  2058. def _find_semantic_dichotomies(self):
  2059. d_keys, dics = generate_dichotomies(self.n_conditions)
  2060. semantic_dics = []
  2061. semantic_keys = []
  2062. for i, dic in enumerate(dics):
  2063. d = [string_digits(x) for x in dic[0]]
  2064. col_sum = np.sum(d, 0)
  2065. if (0 in col_sum) or (len(dic[0]) in col_sum):
  2066. semantic_dics.append(dic)
  2067. semantic_keys.append(self._semantic_keys[np.where(col_sum == len(dic[0]))[0][0]])
  2068. return semantic_dics, semantic_keys
  2069. def _find_nonsemantic_dichotomies(self):
  2070. d_keys, dics = generate_dichotomies(self.n_conditions)
  2071. nonsemantic_dics = []
  2072. for i, dic in enumerate(dics):
  2073. d = [string_digits(x) for x in dic[0]]
  2074. col_sum = np.sum(d, 0)
  2075. if not ((0 in col_sum) or (len(dic[0]) in col_sum)):
  2076. nonsemantic_dics.append(dic)
  2077. return nonsemantic_dics
  2078. def all_dichotomies(self, balanced=True, semantic_names=False):
  2079. if balanced:
  2080. dichotomies = {}
  2081. sem, keys = self._find_semantic_dichotomies()
  2082. nsem = self._find_nonsemantic_dichotomies()
  2083. if (self.n_conditions == 2) and semantic_names:
  2084. dichotomies[keys[0]] = sem[0]
  2085. dichotomies[keys[1]] = sem[1]
  2086. dichotomies['XOR'] = nsem[0]
  2087. else:
  2088. for i in range(len(sem)):
  2089. dichotomies[keys[i]] = sem[i]
  2090. for dic in nsem:
  2091. dichotomies[_powerchotomy_to_key(dic)] = dic
  2092. else:
  2093. powerchotomies = self._powerchotomies()
  2094. dichotomies = {}
  2095. for dk in powerchotomies:
  2096. k = self._dic_key(powerchotomies[dk])
  2097. if k and semantic_names:
  2098. dichotomies[k] = powerchotomies[dk]
  2099. else:
  2100. dichotomies[dk] = powerchotomies[dk]
  2101. if self.n_conditions == 2:
  2102. dichotomies['XOR'] = dichotomies['00_11_v_01_10']
  2103. del dichotomies['00_11_v_01_10']
  2104. return dichotomies
  2105. def _powerchotomies(self):
  2106. conditions = list(self._semantic_vectors.keys())
  2107. powerset = list(chain.from_iterable(combinations(conditions, r) for r in range(1, len(conditions))))
  2108. dichotomies = {}
  2109. for i in range(len(powerset)):
  2110. for j in range(i + 1, len(powerset)):
  2111. if len(np.unique(powerset[i] + powerset[j])) == len(conditions):
  2112. if len(powerset[i] + powerset[j]) == len(conditions):
  2113. dic = [list(powerset[i]), list(powerset[j])]
  2114. dichotomies[_powerchotomy_to_key(dic)] = dic
  2115. return dichotomies
  2116. def _dic_key(self, dic):
  2117. if len(dic[0]) == 2 ** (self.n_conditions - 1) and len(dic[1]) == 2 ** (self.n_conditions - 1):
  2118. for i in range(len(dic)):
  2119. d = [string_digits(x) for x in dic[i]]
  2120. col_sum = np.sum(d, 0)
  2121. if len(dic[0]) in col_sum:
  2122. return self._semantic_keys[np.where(col_sum == len(dic[0]))[0][0]]
  2123. return 0
  2124. def _dichotomy_from_key(self, key):
  2125. dics, keys = self._find_semantic_dichotomies()
  2126. if key in keys:
  2127. dic = dics[np.where(np.asarray(keys) == key)[0][0]]
  2128. else:
  2129. raise RuntimeError(
  2130. "\n[dichotomy_from_key] The specified key does not correspond to a semantic dichotomy. Check the key value.")
  2131. return dic
  2132. def _generate_semantic_vectors(self):
  2133. self._semantic_vectors = {}
  2134. for condition_vec in self._condition_vectors:
  2135. semantic_vector = '('
  2136. for i, sk in enumerate(self._semantic_keys):
  2137. semantic_values = list(self.conditions[sk]) # keys of the dict
  2138. semantic_vector += semantic_values[condition_vec[i]] + ' '
  2139. semantic_vector = semantic_vector[:-1] + ')'
  2140. self._semantic_vectors[string_digits(condition_vec)] = semantic_vector
  2141. def _balanced_classes(self, key):
  2142. """
  2143. For a given semantic variable name ``key``, return a list of lists of
  2144. condition keys that define the classes for multiclass decoding.
  2145. Each inner list contains all condition keys that share the same value
  2146. of variable ``key``, while all other variables are free to vary.
  2147. Example:
  2148. conditions:
  2149. var1: 3 values
  2150. var2: 2 values
  2151. condition_vectors (var1,var2) ->
  2152. [0,0],[0,1],[1,0],[1,1],[2,0],[2,1]
  2153. _balanced_classes('var1') ->
  2154. [['00','01'], ['10','11'], ['20','21']]
  2155. """
  2156. if key not in self._semantic_keys:
  2157. raise KeyError(f"[Decodanda] Variable {key} not found in semantic keys.")
  2158. key_index = self._semantic_keys.index(key)
  2159. # group condition keys by the value of variable `key`
  2160. class_groups = {}
  2161. for condition_vec in self._condition_vectors:
  2162. class_id = condition_vec[key_index]
  2163. cond_key = string_digits(condition_vec)
  2164. # in principle all condition vectors should be present, but be robust
  2165. if cond_key not in self.conditioned_rasters:
  2166. continue
  2167. if class_id not in class_groups:
  2168. class_groups[class_id] = []
  2169. class_groups[class_id].append(cond_key)
  2170. # return groups ordered by class index (0,1,2,...) so they match
  2171. # the implicit ordering of values in self.conditions[key]
  2172. return [class_groups[cid] for cid in sorted(class_groups.keys())]
  2173. def _compute_centroids(self):
  2174. self.centroids = {w: np.hstack([np.nanmean(r, 0) for r in self.conditioned_rasters[w]])
  2175. for w in self.conditioned_rasters.keys()}
  2176. def _zscore_activity(self):
  2177. keys = [string_digits(w) for w in self._condition_vectors]
  2178. for n in range(self.n_brains):
  2179. n_neurons = self.conditioned_rasters[keys[0]][n].shape[1]
  2180. for i in range(n_neurons):
  2181. r = np.hstack([self.conditioned_rasters[key][n][:, i] for key in keys])
  2182. m = np.nanmean(r)
  2183. std = np.nanstd(r)
  2184. if std:
  2185. for key in keys:
  2186. self.conditioned_rasters[key][n][:, i] = (self.conditioned_rasters[key][n][:, i] - m) / std
  2187. def _print(self, string):
  2188. if self._verbose:
  2189. print(string)
  2190. # null model utilities
  2191. def _generate_random_subset(self, n):
  2192. if n < self.n_neurons:
  2193. self.subset = np.random.choice(self.n_neurons, n, replace=False)
  2194. else:
  2195. self.subset = np.arange(self.n_neurons)
  2196. def _reset_random_subset(self):
  2197. self.subset = np.arange(self.n_neurons)
  2198. def _shuffle_conditioned_arrays(self, dic):
  2199. """
  2200. the null model is built by interchanging trials between conditioned arrays that are in different
  2201. dichotomies but have only hamming distance = 1. This ensures that even in the null model the other
  2202. conditions (i.e., the one that do not define the dichotomy), are balanced during sampling.
  2203. So if my dichotomy is [1A, 1B] vs [0A, 0B], I will change trials between 1A and 0A, so that,
  2204. with oversampling, I will then ensure balance between A and B.
  2205. If the dichotomy is not semantic, then I'll probably have to interchange between conditions regardless
  2206. (to be implemented).
  2207. :param dic: The dichotomy to be decoded
  2208. """
  2209. # if the dichotomy is semantic, shuffle between rasters at semantic distance=1
  2210. if self._dic_key(dic):
  2211. set_A = dic[0]
  2212. set_B = dic[1]
  2213. for i in range(len(set_A)):
  2214. for j in range(len(set_B)):
  2215. test_condition_A = set_A[i]
  2216. test_condition_B = set_B[j]
  2217. if hamming(string_digits(test_condition_A), string_digits(test_condition_B)) == 1:
  2218. for n in range(self.n_brains):
  2219. # select conditioned rasters
  2220. arrayA = np.copy(self.conditioned_rasters[test_condition_A][n])
  2221. arrayB = np.copy(self.conditioned_rasters[test_condition_B][n])
  2222. # select conditioned trial index
  2223. trialA = np.copy(self.conditioned_trial_index[test_condition_A][n])
  2224. trialB = np.copy(self.conditioned_trial_index[test_condition_B][n])
  2225. n_trials_A = len(np.unique(trialA))
  2226. n_trials_B = len(np.unique(trialB))
  2227. # assign randomly trials between the two conditioned rasters, keeping the same
  2228. # number of trials between the two conditions
  2229. all_rasters = []
  2230. all_trials = []
  2231. for index in np.unique(trialA):
  2232. all_rasters.append(arrayA[trialA == index, :])
  2233. all_trials.append(trialA[trialA == index])
  2234. for index in np.unique(trialB):
  2235. all_rasters.append(arrayB[trialB == index, :])
  2236. all_trials.append(trialB[trialB == index])
  2237. all_trial_index = np.arange(n_trials_A + n_trials_B).astype(int)
  2238. np.random.shuffle(all_trial_index)
  2239. new_rasters_A = [all_rasters[iA] for iA in all_trial_index[:n_trials_A]]
  2240. new_rasters_B = [all_rasters[iB] for iB in all_trial_index[n_trials_A:]]
  2241. new_trials_A = [all_trials[iA] for iA in all_trial_index[:n_trials_A]]
  2242. new_trials_B = [all_trials[iB] for iB in all_trial_index[n_trials_A:]]
  2243. self.conditioned_rasters[test_condition_A][n] = np.vstack(new_rasters_A)
  2244. self.conditioned_rasters[test_condition_B][n] = np.vstack(new_rasters_B)
  2245. self.conditioned_trial_index[test_condition_A][n] = np.hstack(new_trials_A)
  2246. self.conditioned_trial_index[test_condition_B][n] = np.hstack(new_trials_B)
  2247. else:
  2248. for n in range(self.n_brains):
  2249. # select conditioned rasters
  2250. for iteration in range(10):
  2251. all_conditions = list(self._semantic_vectors.keys())
  2252. all_data = np.vstack([self.conditioned_rasters[cond][n] for cond in all_conditions])
  2253. all_trials = np.hstack([self.conditioned_trial_index[cond][n] for cond in all_conditions])
  2254. all_n_trials = {cond: len(np.unique(self.conditioned_trial_index[cond][n])) for cond in
  2255. all_conditions}
  2256. unique_trials = np.unique(all_trials)
  2257. np.random.shuffle(unique_trials)
  2258. i = 0
  2259. for cond in all_conditions:
  2260. cond_trials = unique_trials[i:i + all_n_trials[cond]]
  2261. new_cond_array = []
  2262. new_cond_trial = []
  2263. for trial in cond_trials:
  2264. new_cond_array.append(all_data[all_trials == trial])
  2265. new_cond_trial.append(all_trials[all_trials == trial])
  2266. self.conditioned_rasters[cond][n] = np.vstack(new_cond_array)
  2267. self.conditioned_trial_index[cond][n] = np.hstack(new_cond_trial)
  2268. i += all_n_trials[cond]
  2269. if not self._check_trial_availability(): # if the trial distribution is not cross validatable, redo the shuffling
  2270. print("Note: re-shuffling arrays")
  2271. self._order_conditioned_rasters()
  2272. self._shuffle_conditioned_arrays(dic)
  2273. def _rototraslate_conditioned_rasters(self):
  2274. # DEPCRECATED
  2275. for i in range(self.n_brains):
  2276. # brain_means = np.vstack([np.nanmean(self.conditioned_rasters[key][i], 0) for key in self.conditioned_rasters.keys()])
  2277. # mean_centroid = np.nanmean(brain_means, axis=0)
  2278. for w in self.conditioned_rasters.keys():
  2279. raster = self.conditioned_rasters[w][i]
  2280. rotation = np.arange(raster.shape[1]).astype(int)
  2281. np.random.shuffle(rotation)
  2282. raster = raster[:, rotation]
  2283. # mean = np.nanmean(raster, 0)
  2284. # randomdir = np.random.rand()-0.5
  2285. # randomdir = randomdir/np.sqrt(np.dot(randomdir, randomdir))
  2286. # vector_from_mean_centroid = mean - mean_centroid
  2287. # distance_from_mean_centroid = np.sqrt(np.dot(vector_from_mean_centroid, vector_from_mean_centroid))
  2288. # raster = raster - vector_from_mean_centroid + randomdir*distance_from_mean_centroid
  2289. self.conditioned_rasters[w][i] = raster
  2290. def _order_conditioned_rasters(self):
  2291. for w in self.conditioned_rasters.keys():
  2292. self.conditioned_rasters[w] = self.ordered_conditioned_rasters[w].copy()
  2293. self.conditioned_trial_index[w] = self.ordered_conditioned_trial_index[w].copy()
  2294. def _check_trial_availability(self):
  2295. if self._debug:
  2296. print('\nCheck trial availability')
  2297. for k in self.conditioned_trial_index:
  2298. for i, ti in enumerate(self.conditioned_trial_index[k]):
  2299. if self._debug:
  2300. print(k, 'raster %u:' % i, np.unique(ti).shape[0])
  2301. print(ti)
  2302. if np.unique(ti).shape[0] < 2:
  2303. return False
  2304. return True
  2305. def _reset_weight_arrays(self):
  2306. self.decoding_weights = {}
  2307. self.decoding_weights_null = {}
  2308. # Wrapper for decoding
  2309. def decoding_analysis(data, conditions, decodanda_params, analysis_params, parallel=False, plot=False, ax=None):
  2310. """
  2311. Function that performs a balanced decoding analyses of the
  2312. data set passed in the ``data`` argument, using variables and values
  2313. specified in the ``conditions`` dictionary.
  2314. This functions is a shortcut for building a ``Decodanda`` object
  2315. with ``decodanda_params`` as arguments and calling the ``Decodanda.decode`` function
  2316. with ``analysis_params`` as arguments.
  2317. See Also
  2318. --------
  2319. Decodanda
  2320. Decodanda.decode
  2321. Notes
  2322. -----
  2323. This function is equivalent to
  2324. >>> Decodanda(data, conditions, **decodanda_params).decode(**analysis_params)
  2325. Parameters
  2326. ----------
  2327. data
  2328. The data set used by the ``Decodanda`` object.
  2329. conditions
  2330. The conditions dictionary for the ``Decodanda`` object.
  2331. decodanda_params
  2332. A dictionary specifying the values for the ``Decodanda`` constructor parameters.
  2333. analysis_params
  2334. A dictionary specifying the values for the ``Decodanda.decode`` function parameters.
  2335. parallel
  2336. [Experimental] if ``True``, null model iterations are performed on separated threads.
  2337. plot
  2338. If ``True``, the decoding results are shown in a figure.
  2339. ax
  2340. If specified, and ``plot=True`` the results are shown in the specified axis.
  2341. Returns
  2342. -------
  2343. performances, null
  2344. """
  2345. an_params = copy.deepcopy(analysis_params)
  2346. if parallel:
  2347. # Data
  2348. null_iterations = an_params['nshuffles']
  2349. an_params['nshuffles'] = 0
  2350. performances, _ = Decodanda(data, conditions, **decodanda_params).decode(**an_params)
  2351. # Null
  2352. del an_params['nshuffles']
  2353. pool = Pool()
  2354. null_performances = pool.map(_NullmodelIterator(data, conditions, decodanda_params, an_params),
  2355. range(null_iterations))
  2356. null = {key: np.stack([p[key] for p in null_performances]) for key in null_performances[0].keys()}
  2357. else:
  2358. performances, null = Decodanda(data, conditions, **decodanda_params).decode(**an_params)
  2359. if plot:
  2360. plot_perfs_null_model(performances, null, ax=ax, ptype='zscore')
  2361. return performances, null
  2362. # Utilities
  2363. def check_session_requirements(session, conditions, **decodanda_params):
  2364. d = Decodanda(session, conditions, fault_tolerance=True, **decodanda_params)
  2365. if d.n_brains:
  2366. return True
  2367. else:
  2368. return False
  2369. def check_requirements_two_conditions(sessions, conditions_1, conditions_2, **decodanda_params):
  2370. good_sessions = []
  2371. for s in sessions:
  2372. if check_session_requirements(s, conditions_1, **decodanda_params) and check_session_requirements(s,
  2373. conditions_2,
  2374. **decodanda_params):
  2375. good_sessions.append(s)
  2376. return good_sessions
  2377. def balance_decodandas(ds):
  2378. for i in range(len(ds)):
  2379. for j in range(i + 1, len(ds)):
  2380. _balance_two_decodandas(ds[i], ds[j])
  2381. def _balance_two_decodandas(d1, d2, sampling_strategy='random'):
  2382. assert d1.n_brains == d2.n_brains, "The two decodandas do not have the same number of brains."
  2383. assert d1.n_conditions == d2.n_conditions, "The two decodanda do not have the same number of semantic conditions."
  2384. n_brains = d1.n_brains
  2385. n_conditioned_rasters = len(list(d1.conditioned_rasters.values()))
  2386. for n in range(n_brains):
  2387. for i in range(n_conditioned_rasters):
  2388. t1 = list(d1.conditioned_rasters.values())[i][n].shape[0]
  2389. t2 = list(d2.conditioned_rasters.values())[i][n].shape[0]
  2390. t = min(t1, t2)
  2391. if t1 > t2:
  2392. if sampling_strategy == 'random':
  2393. sampling = np.random.choice(t1, t2, replace=False)
  2394. if sampling_strategy == 'ordered':
  2395. sampling = np.arange(t2, dtype=int)
  2396. list(d1.conditioned_rasters.values())[i][n] = list(d1.conditioned_rasters.values())[i][n][sampling, :]
  2397. list(d1.conditioned_trial_index.values())[i][n] = list(d1.conditioned_trial_index.values())[i][n][
  2398. sampling]
  2399. if t2 > t1:
  2400. if sampling_strategy == 'random':
  2401. sampling = np.random.choice(t2, t1, replace=False)
  2402. if sampling_strategy == 'ordered':
  2403. sampling = np.arange(t1, dtype=int)
  2404. list(d2.conditioned_rasters.values())[i][n] = list(d2.conditioned_rasters.values())[i][n][sampling, :]
  2405. list(d2.conditioned_trial_index.values())[i][n] = list(d2.conditioned_trial_index.values())[i][n][
  2406. sampling]
  2407. if d1._verbose:
  2408. print("Balancing data for d1: %u, d2: %u - now d1: %u, d2: %u" % (
  2409. t1, t2, list(d1.conditioned_rasters.values())[i][n].shape[0],
  2410. list(d2.conditioned_rasters.values())[i][n].shape[0]))
  2411. for w in d1.conditioned_rasters.keys():
  2412. d1.ordered_conditioned_rasters[w] = d1.conditioned_rasters[w].copy()
  2413. d1.ordered_conditioned_trial_index[w] = d1.conditioned_trial_index[w].copy()
  2414. d2.ordered_conditioned_rasters[w] = d2.conditioned_rasters[w].copy()
  2415. d2.ordered_conditioned_trial_index[w] = d2.conditioned_trial_index[w].copy()
  2416. print("\n")
  2417. def _generate_binary_condition(var_key, value1, value2, key1=None, key2=None, var_key_plot=None):
  2418. if key1 is None:
  2419. key1 = '%s' % value1
  2420. if key2 is None:
  2421. key2 = '%s' % value2
  2422. if var_key_plot is None:
  2423. var_key_plot = var_key
  2424. conditions = {
  2425. var_key_plot: {
  2426. key1: lambda d, x=value1: d[var_key] == x,
  2427. key2: lambda d, x=value2: d[var_key] == x,
  2428. }
  2429. }
  2430. return conditions
  2431. def _generate_conditions_from_dic(discrete_dict):
  2432. conditions = {}
  2433. for key in discrete_dict.keys():
  2434. conditions[key] = {}
  2435. for v in discrete_dict[key]:
  2436. conditions[key]['%s' % v] = (lambda d, k=key, vv=v: getattr(d, k) == vv)
  2437. return conditions
  2438. def _powerchotomy_to_key(dic):
  2439. return '_'.join(dic[0]) + '_v_' + '_'.join(dic[1])
  2440. class _NullmodelIterator(object): # necessary for parallelization of null model iterations
  2441. def __init__(self, data, conditions, decodanda_params, analysis_params):
  2442. self.data = data
  2443. self.conditions = conditions
  2444. self.decodanda_params = decodanda_params
  2445. self.analysis_params = analysis_params
  2446. def __call__(self, i):
  2447. self.i = i
  2448. self.randomstate = RandomState(i)
  2449. dec = Decodanda(data=self.data, conditions=self.conditions, **self.decodanda_params)
  2450. semantic_dics, semantic_keys = dec._find_semantic_dichotomies()
  2451. if 'non_semantic' in self.analysis_params.keys():
  2452. if self.analysis_params['non_semantic'] and len(self.conditions) == 2:
  2453. semantic_dics.append([['01', '10'], ['00', '11']])
  2454. semantic_keys.append('XOR')
  2455. if self.analysis_params['non_semantic'] and len(self.conditions) > 2:
  2456. dics = dec.all_dichotomies(balanced=True)
  2457. semantic_dics = list(dics.values())
  2458. semantic_keys = list(dics.keys())
  2459. perfs = {}
  2460. for key, dic in zip(semantic_keys, semantic_dics):
  2461. if dec._verbose:
  2462. print("\nTesting null decoding performance for semantic dichotomy: ", key)
  2463. dec._shuffle_conditioned_arrays(dic)
  2464. performance = dec.decode_dichotomy(dic, **self.analysis_params)
  2465. perfs[key] = np.nanmean(performance)
  2466. dec._order_conditioned_rasters()
  2467. return perfs

classes.py at commit 4f4ad07, under GPL-3.0 · at the source

Overview

Authors: Pia-Kelsey O’Neill1,2,3, Lorenzo Posani1,2,4,5, Jozsef Meszaros6, Phebe Warren1, Carl E Schoonover1,2,7,8, Andrew J P Fink1,2,7,9, Stefano Fusi1,2,4,10, C Daniel Salzman1,2,6,10,11
  1. The Mortimer B. Zuckerman Mind Brain Behavior Institute, Columbia University, New York, NY USA
  2. Department of Neuroscience, Columbia University, New York, NY USA
  3. Present Address: Department of Psychological and Brain Sciences, Dartmouth College, Hanover, NH USA
  4. Center for Theoretical Neuroscience, Columbia University, New York, NY USA
  5. Present Address: ICM Paris Brain Institute, Hôpital de la Pitié Salpêtrière, Paris, France
  6. Department of Psychiatry, Columbia University, New York, NY USA
  7. Howard Hughes Medical Institute, Columbia University, New York, NY USA
  8. Present Address: Allen Institute for Neural Dynamics, Seattle, WA USA
  9. Present Address: Department of Neurobiology, Northwestern University, Evanston, IL USA
  10. Kavli Institute for Brain Science, Columbia University, New York, NY USA
  11. New York State Psychiatric Institute, New York, NY USA
Journal: Nature neuroscience, volume 29, issue 7, pages 1654-1666
Dates: received 11 September 2025; accepted 22 April 2026; published online 3 June 2026; in print 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1038/s41593-026-02315-y · PMID 42237032 · PMCID PMC13337481 · OpenAlex W4386998197
Open access: hybrid, a free copy (OpenAlex)
Status: code verified
Categories: mouse (organism), cognitive (subfield)
Methods: Connectivity, Statistics, Smoothing, state filtering, decompositions, Machine learning, Preprocessing, Spectral & time-frequency, fMRI & imaging, Single-unit activity, calcium imaging
Keywords: Amygdala, Neural decoding
MeSH: Amygdala*, Basolateral Nuclear Complex*, Emotions*, Neurons*, Action Potentials, Animals, Conditioning, Classical, Fear, Male, Mice, Mice, Inbred C57BL, Muscimol (* major topic)
Topic: Memory and Neural Mechanisms (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: Simons Foundation (Global Brain Initiative)
Citations: cited by 6 papers (Europe PMC); 69 references in the paper

Abstract

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

Repositories

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

lposani/decodanda

License: GPL-3.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 4f4ad07c98549a58a7b814c0a98f5caed8c2539c, 27 April 2026
Languages: Python (15), Jupyter (5), Shell (1)
Size: 49 files, 21 scripts
Software Heritage: not archived
Found in: the text, “Neural decoding analysis”
Holds: README, license file, environment (setup.py, docs/requirements.txt), tests, documentation, 5 notebooks
Not found: CITATION.cff, continuous integration
Tools: NumPy (15 files), Matplotlib (12 files), seaborn (6 files), scikit-learn (5 files), SciPy (4 files), h5py (2 files), pandas (2 files)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
23 files

pkoneill/Burrow-Amygdala-Code

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 867d34fe2798a937404a926d0ee02b9b14980486, 20 April 2026
Size: 1 file, 0 scripts
Software Heritage: not archived
Found in: “Code availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers

Code availability statement

The paper has a code availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41593-026-02315-y.

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:

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

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

Data

Datasets cited

Data availability statement

The paper has a data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41593-026-02315-y.

Versions

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

Version 2, 28 September 2026

  • Publisher: n/a → Nature Portfolio

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 8 authors, 2 keywords, 12 MeSH terms, 1 funder, 67 references.

Cite

This paper

O’Neill, P.-K., Posani, L., Meszaros, J., Warren, P., Schoonover, C. E., Fink, A. J. P., Fusi, S., & Salzman, C. D. (2026). The representational geometry of emotional states in basolateral amygdala. Nature neuroscience, 29(7), 1654-1666. https://doi.org/10.1038/s41593-026-02315-y

BibTeX

@article{oneill2026representational,
author = {O’Neill, Pia-Kelsey and Posani, Lorenzo and Meszaros, Jozsef and Warren, Phebe and Schoonover, Carl E and Fink, Andrew J P and Fusi, Stefano and Salzman, C Daniel},
title = {{The representational geometry of emotional states in basolateral amygdala}},
journal = {Nature neuroscience},
year = {2026},
month = jun,
volume = {29},
number = {7},
pages = {1654--1666},
publisher = {Nature Portfolio},
issn = {1097-6256},
doi = {10.1038/s41593-026-02315-y},
url = {https://doi.org/10.1038/s41593-026-02315-y},
pmid = {42237032},
pmcid = {PMC13337481}
}

RIS

TY - JOUR
AU - O’Neill, Pia-Kelsey
AU - Posani, Lorenzo
AU - Meszaros, Jozsef
AU - Warren, Phebe
AU - Schoonover, Carl E
AU - Fink, Andrew J P
AU - Fusi, Stefano
AU - Salzman, C Daniel
TI - The representational geometry of emotional states in basolateral amygdala
T2 - Nature neuroscience
J2 - Nat Neurosci
PY - 2026
DA - 2026/06/03
VL - 29
IS - 7
SP - 1654
EP - 1666
SN - 1097-6256
PB - Nature Portfolio
DO - 10.1038/s41593-026-02315-y
UR - https://doi.org/10.1038/s41593-026-02315-y
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41593-026-02315-y",
"type": "article-journal",
"title": "The representational geometry of emotional states in basolateral amygdala",
"container-title": "Nature neuroscience",
"author": [
{
"family": "O’Neill",
"given": "Pia-Kelsey"
},
{
"family": "Posani",
"given": "Lorenzo"
},
{
"family": "Meszaros",
"given": "Jozsef"
},
{
"family": "Warren",
"given": "Phebe"
},
{
"family": "Schoonover",
"given": "Carl E"
},
{
"family": "Fink",
"given": "Andrew J P"
},
{
"family": "Fusi",
"given": "Stefano"
},
{
"family": "Salzman",
"given": "C Daniel"
}
],
"container-title-short": "Nat Neurosci",
"volume": "29",
"issue": "7",
"page": "1654-1666",
"DOI": "10.1038/s41593-026-02315-y",
"PMID": "42237032",
"PMCID": "PMC13337481",
"ISSN": "1097-6256",
"publisher": "Nature Portfolio",
"URL": "https://doi.org/10.1038/s41593-026-02315-y",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
3
]
]
}
}

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

Similar papers

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

[1] doi:10.1038/s41593-026-02333-w [code]
Learning shapes neural geometry in the primate prefrontal cortex.
Journal: Nature neuroscience
In common: seaborn, scikit-learn, pandas, 3 other tools, 6 references
[2] doi:10.7554/elife.105528 [code]
Functional specialization of mPFC-BLA and mPFC-NAc pathways in affective state representation.
Journal: eLife
In common: h5py, pandas, SciPy, 2 other tools, mouse, 4 references
[3] doi:10.1371/journal.pcbi.1014162 [code]
Exploring neural manifolds across a wide range of intrinsic dimensions.
Journal: PLoS computational biology
In common: h5py, scikit-learn, pandas, 3 other tools, 3 references
[4] doi:10.1038/s41467-026-74347-8 [code]
Compositionality of social gaze in the prefrontal-amygdala circuits.
Journal: Nature communications
In common: seaborn, scikit-learn, pandas, 3 other tools, 3 references
[5] doi:10.1038/s41467-026-74818-y [code]
Stable readout of visual representations mediates flexible generalization.
Journal: Nature communications
In common: SciPy, Matplotlib, NumPy, cognitive, 4 references
[6] doi:10.1038/s41467-026-74566-z [code]
Low-dimensional and optimised representations of high-level information in the expert brain.
Journal: Nature communications
In common: seaborn, scikit-learn, pandas, 3 other tools, cognitive, 2 references
[7] doi:10.1162/imag.a.1266 [code]
Multimodal subspace independent vector analysis effectively captures latent relationships between brain structure and function.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: seaborn, scikit-learn, pandas, 3 other tools, 2 references
[8] doi:10.3389/fnsys.2026.1822122 [code]
Convergence-divergence circuits for multimodal integration of innate and learned opponent valences.
Journal: Frontiers in systems neuroscience
In common: h5py, seaborn, scikit-learn, 4 other tools, 1 reference
[9] doi:10.1038/s41467-026-72057-9 [code]
Sex-specific behavioral feedback modulates sensorimotor processing and drives flexible social behavior.
Journal: Nature communications
In common: h5py, seaborn, scikit-learn, 4 other tools, 1 reference
[10] doi:10.1038/s41467-026-76104-3 [code]
Sensorimotor remapping drives task specialization in prefrontal cortex.
Journal: Nature communications
In common: seaborn, scikit-learn, pandas, 3 other tools, 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.