OSCR

Learning shapes neural geometry in the primate prefrontal cortex.

Code ↔ Paper

6 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 6 matches
  1. [1] § Methods › Models › Generative models ↔ fun_lib.py, lines 3274–3380 · score 0.67 · generated firing rates, linear regression, covariance matrix, multivariate, population, noise
  2. [2] § Methods › Data and task › Adaptive trial termination and switch costs ↔ supp_fig_3.py, lines 14–69 · score 0.65 · Switch costs, fixation breaking, rewarded trial, errors
  3. [3] § Methods › Analysis methods › Decoding ↔ fun_lib.py, lines 206–274 · score 0.62 · cross validated, random splits, classifiers, binary, temporally, matrix
  4. [4] § Methods › Models › Multiple linear regression ↔ fun_lib.py, lines 3274–3380 · score 0.58 · linear model, design matrix, firing rate, populated, coefficients, regression
  5. [5] § Methods › Data and task › Adaptive trial termination and switch costs ↔ fun_lib.py, lines 4197–4295 · score 0.55 · Switch costs, fixation breaking, adaptive, rewarded
  6. [6] § Methods › Analysis methods › Measuring similarity between selectivity distributions ↔ fun_lib.py, lines 2867–2941 · score 0.55 · Euclidean distance, random selectivity, divergence, KL, metrics, models

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 · 4,746 lines · 195 KB · no license · 5 matches

  1. import numpy as np
  2. import pandas as pd
  3. import matplotlib.pyplot as plt
  4. import seaborn as sns
  5. import scipy as sp
  6. from scipy.stats import spearmanr, stats
  7. from tqdm import tqdm
  8. from itertools import groupby
  9. from sklearn.metrics import r2_score
  10. from sklearn import linear_model
  11. import random
  12. import itertools
  13. from matplotlib import colors
  14. from mne.decoding import SlidingEstimator, GeneralizingEstimator
  15. from sklearn.discriminant_analysis import LinearDiscriminantAnalysis as LDA
  16. from sklearn.pipeline import make_pipeline
  17. from sklearn.preprocessing import StandardScaler
  18. from sklearn.svm import LinearSVC as SVM
  19. from sklearn.svm import SVC
  20. from scipy.ndimage import label
  21. import pickle
  22. from matplotlib import patches
  23. from numpy.linalg import norm
  24. import matplotlib.gridspec as gridspec
  25. import matplotlib
  26. matplotlib.use('TkAgg')
  27. plt.rcParams['svg.fonttype'] = 'none'
  28. from itertools import chain
  29. import scipy.io as io
  30. import yaml
  31. from scipy import signal
  32. from sklearn.base import BaseEstimator, ClassifierMixin
  33. from sklearn.utils.validation import check_X_y, check_array, check_is_fitted
  34. from sklearn.utils.multiclass import unique_labels
  35. from scipy.stats import zscore
  36. from joblib import Parallel, delayed
  37. from numba import jit, prange
  38. import multiprocessing as mp
  39. import os
  40. @jit(nopython=True, parallel=True)
  41. def fast_pearsonr_matrix(X, Y):
  42. """
  43. Fast computation of Pearson correlation matrix using Numba.
  44. X: (n_features, n_timepoints1)
  45. Y: (n_features, n_timepoints2)
  46. Returns: (n_timepoints1, n_timepoints2)
  47. """
  48. n_features, n_timepoints1 = X.shape
  49. _, n_timepoints2 = Y.shape
  50. corr_matrix = np.zeros((n_timepoints1, n_timepoints2))
  51. for i in prange(n_timepoints1):
  52. for j in prange(n_timepoints2):
  53. x = X[:, i]
  54. y = Y[:, j]
  55. # Compute correlation manually for speed
  56. x_mean = np.mean(x)
  57. y_mean = np.mean(y)
  58. num = np.sum((x - x_mean) * (y - y_mean))
  59. den_x = np.sqrt(np.sum((x - x_mean) ** 2))
  60. den_y = np.sqrt(np.sum((y - y_mean) ** 2))
  61. if den_x == 0 or den_y == 0:
  62. corr_matrix[i, j] = 0
  63. else:
  64. corr_matrix[i, j] = num / (den_x * den_y)
  65. return corr_matrix
  66. @jit(nopython=True, parallel=True)
  67. def fast_pearsonr_vector(X, Y):
  68. """
  69. Fast computation of Pearson correlation vector using Numba.
  70. X: (n_features, n_timepoints)
  71. Y: (n_features, n_timepoints)
  72. Returns: (n_timepoints,)
  73. """
  74. n_features, n_timepoints = X.shape
  75. corr_vector = np.zeros(n_timepoints)
  76. for t in prange(n_timepoints):
  77. x = X[:, t]
  78. y = Y[:, t]
  79. x_mean = np.mean(x)
  80. y_mean = np.mean(y)
  81. num = np.sum((x - x_mean) * (y - y_mean))
  82. den_x = np.sqrt(np.sum((x - x_mean) ** 2))
  83. den_y = np.sqrt(np.sum((y - y_mean) ** 2))
  84. if den_x == 0 or den_y == 0:
  85. corr_vector[t] = 0
  86. else:
  87. corr_vector[t] = num / (den_x * den_y)
  88. return corr_vector
  89. def temp_dec_stages_permutation(data_eq, labels_eq, variable_mapping, tranc_window=[40, 160], n_permutations=100,
  90. random_state=42, method="Pearson"):
  91. """
  92. Compute temporal generalization matrices with original and scrambled labels.
  93. Uses the same permuted labels across all windows for each permutation.
  94. Parameters:
  95. -----------
  96. data_eq : list of ndarrays
  97. List of data arrays for each stage
  98. labels_eq : list of ndarrays
  99. List of label arrays for each stage
  100. variable_mapping : list or array
  101. Mapping of labels to factors
  102. tranc_window : list
  103. Time window to analyze [start, end]
  104. n_permutations : int
  105. Number of permutation iterations to perform
  106. random_state : int
  107. Random seed for reproducibility
  108. Returns:
  109. --------
  110. decoding : ndarray
  111. Original temporal generalization matrices
  112. decoding_perm : ndarray
  113. Permuted temporal generalization matrices
  114. """
  115. n_windows = data_eq[0].shape[0]
  116. n_stages = len(data_eq)
  117. window_size = tranc_window[1] - tranc_window[0] + 1
  118. rng = np.random.RandomState(random_state)
  119. # Initialize arrays for original and permuted decoding
  120. decoding = np.zeros((n_stages, n_windows, window_size, window_size))
  121. decoding_perm = np.zeros((n_permutations, n_stages, n_windows, window_size, window_size))
  122. # Convert mapping to numpy array if it's not already
  123. colour_fac = np.array(variable_mapping)
  124. # Compute original decoding matrices first
  125. for i_stage in range(n_stages):
  126. y_stage = labels_eq[i_stage][0, :]
  127. for i_window in range(n_windows):
  128. # Extract data for this stage and window
  129. X_stage = data_eq[i_stage][i_window, :, :, tranc_window[0]:tranc_window[1] + 1]
  130. # Compute original labels
  131. y_colour = np.array(assign_lables(y_stage, colour_fac))
  132. if method == "SVM":
  133. decoding[i_stage, i_window, :, :] = decode_time(X_stage, y_colour, n_inter=1)
  134. elif method == "Pearson":
  135. decoder = NeuralCorrelationDecoder(across_time=True)
  136. decoder.fit(X_stage, y_colour)
  137. decoding[i_stage, i_window, :, :] = decoder.get_correlation()
  138. print(f'Constructing null')
  139. # Now compute permuted decoding matrices
  140. for perm in tqdm(range(n_permutations)):
  141. # For each stage, generate ONE permutation to use across all windows
  142. for i_stage in range(n_stages):
  143. # Get the first window's labels to determine permutation structure
  144. y_stage_first = labels_eq[i_stage][0, :]
  145. y_colour_first = np.array(assign_lables(y_stage_first, colour_fac))
  146. # Create a single permutation for this stage
  147. y_perm_indices = rng.permutation(len(y_colour_first))
  148. # Use this permutation for all windows in this stage
  149. for i_window in range(n_windows):
  150. # Extract data for this stage and window
  151. X_stage = data_eq[i_stage][i_window, :, :, tranc_window[0]:tranc_window[1] + 1]
  152. y_stage = labels_eq[i_stage][i_window, :]
  153. # Get original color labels
  154. y_colour = np.array(assign_lables(y_stage, colour_fac))
  155. # Apply the same permutation pattern to these labels
  156. y_perm = y_colour[y_perm_indices]
  157. # Compute decoding matrix with permuted labels
  158. if method == "SVM":
  159. decoding_perm[perm, i_stage, i_window, :, :] = decode_time(X_stage, y_perm, n_inter=1)
  160. elif method == "Pearson":
  161. decoder = NeuralCorrelationDecoder(across_time=True)
  162. decoder.fit(X_stage, y_perm)
  163. decoding_perm[perm, i_stage, i_window, :, :] = decoder.get_correlation()
  164. return decoding, decoding_perm
  165. def _single_iteration(X, y, iteration_seed, across_time):
  166. """
  167. Single iteration of cross-validation for parallelization.
  168. Parameters
  169. ----------
  170. X : ndarray
  171. Input data
  172. y : ndarray
  173. Labels
  174. iteration_seed : int
  175. Random seed for this iteration
  176. across_time : bool
  177. Whether to compute across-time correlations
  178. Returns
  179. -------
  180. correlation_result : ndarray
  181. Correlation matrix or vector for this iteration
  182. """
  183. def condi_avg(data, labels):
  184. """Optimized for binary classification."""
  185. mask = labels.astype(bool)
  186. condition_0 = data[~mask].mean(axis=0)
  187. condition_1 = data[mask].mean(axis=0)
  188. return np.array([condition_0, condition_1])
  189. # Set seed for this iteration (convert numpy int to Python int)
  190. seed_value = int(iteration_seed) # Convert to native Python int
  191. np.random.seed(seed_value)
  192. random.seed(seed_value)
  193. n_trials = X.shape[0]
  194. # Create random split
  195. idx_rnd = np.concatenate([np.zeros(n_trials // 2), np.ones(n_trials // 2)])
  196. if (n_trials % 2) > 0:
  197. idx_rnd = np.concatenate([idx_rnd, [1.0]])
  198. np.random.shuffle(idx_rnd)
  199. # Split data
  200. X1 = X[idx_rnd < 1, :, :]
  201. y1 = y[idx_rnd < 1]
  202. X2 = X[idx_rnd > 0, :, :]
  203. y2 = y[idx_rnd > 0]
  204. # Z-score
  205. X1 = zscore(X1, axis=1)
  206. X2 = zscore(X2, axis=1)
  207. # Compute condition averages
  208. condition_averages1 = condi_avg(X1, y1)
  209. condition_averages2 = condi_avg(X2, y2)
  210. # Compute differences between conditions
  211. diffTrain = condition_averages1[0, :, :] - condition_averages1[1, :, :]
  212. diffTest = condition_averages2[0, :, :] - condition_averages2[1, :, :]
  213. # Compute correlations using fast functions
  214. if across_time:
  215. # Temporal generalization: correlate across all time point pairs
  216. corr1 = fast_pearsonr_matrix(diffTrain, diffTest)
  217. corr2 = fast_pearsonr_matrix(diffTest, diffTrain)
  218. return np.tanh(np.mean(np.arctanh(np.array([corr1, corr2])), axis=0))
  219. else:
  220. # Within-time correlations only
  221. return fast_pearsonr_vector(diffTrain, diffTest)
  222. class NeuralCorrelationDecoder(BaseEstimator, ClassifierMixin):
  223. """
  224. Optimized neural correlation decoder with parallel processing and JIT compilation.
  225. Parameters
  226. ----------
  227. n_iterations : int, default=10
  228. Number of cross-validation iterations to perform
  229. across_time : bool, default=True
  230. If True, compute correlations across all time point pairs (temporal generalization)
  231. If False, compute correlations only within the same time points
  232. random_state : int, RandomState instance or None, default=None
  233. Controls the randomness of the cross-validation splits
  234. n_jobs : int, default=-1
  235. Number of parallel jobs. -1 means use all available cores
  236. batch_size : int, default=None
  237. Batch size for parallel processing. If None, uses n_iterations // n_cores
  238. """
  239. def __init__(self, n_iterations=10, across_time=True, random_state=None,
  240. n_jobs=-1, batch_size=None):
  241. self.n_iterations = n_iterations
  242. self.across_time = across_time
  243. self.random_state = random_state
  244. self.n_jobs = n_jobs
  245. self.batch_size = batch_size
  246. def _validate_input(self, X, y=None):
  247. """Validate input data format."""
  248. if X.ndim != 3:
  249. raise ValueError(f"Expected 3D array (n_trials, n_features, n_timepoints), "
  250. f"got {X.ndim}D array")
  251. if y is not None:
  252. if len(np.unique(y)) != 2:
  253. raise ValueError("This decoder only supports binary classification "
  254. f"(2 classes), got {len(np.unique(y))} classes")
  255. return X, y
  256. def fit(self, X, y):
  257. """
  258. Fit the neural correlation decoder using parallel processing.
  259. Parameters
  260. ----------
  261. X : array-like of shape (n_trials, n_features, n_timepoints)
  262. Training data
  263. y : array-like of shape (n_trials,)
  264. Target values (binary classification)
  265. Returns
  266. -------
  267. self : object
  268. Returns the instance itself
  269. """
  270. # Validate inputs
  271. X, y = check_X_y(X, y, allow_nd=True)
  272. X, y = self._validate_input(X, y)
  273. # Store classes
  274. self.classes_ = unique_labels(y)
  275. # Generate seeds for each iteration to ensure reproducibility
  276. if self.random_state is not None:
  277. np.random.seed(self.random_state)
  278. seeds = np.random.randint(0, 2 ** 31, self.n_iterations)
  279. else:
  280. seeds = np.random.randint(0, 2 ** 31, self.n_iterations)
  281. # Determine number of jobs
  282. n_jobs = self.n_jobs
  283. if n_jobs == -1:
  284. n_jobs = mp.cpu_count()
  285. elif n_jobs <= 0:
  286. n_jobs = max(1, mp.cpu_count() + n_jobs)
  287. # Parallel computation of iterations
  288. # Handle batch_size properly for joblib
  289. parallel_kwargs = {'n_jobs': n_jobs}
  290. if self.batch_size is not None:
  291. parallel_kwargs['batch_size'] = self.batch_size
  292. results = Parallel(**parallel_kwargs)(
  293. delayed(_single_iteration)(X, y, seed, self.across_time)
  294. for seed in seeds
  295. )
  296. # Convert results to numpy array and average using Fisher z-transform
  297. corrs_iters = np.array(results)
  298. # Handle potential NaN/inf values before arctanh
  299. corrs_iters = np.clip(corrs_iters, -0.99999, 0.99999)
  300. self.correlation_matrix_ = np.tanh(np.mean(np.arctanh(corrs_iters), axis=0))
  301. self.is_fitted_ = True
  302. return self
  303. def get_correlation(self):
  304. """Get the fitted correlation matrix."""
  305. check_is_fitted(self, 'is_fitted_')
  306. return self.correlation_matrix_
  307. def set_params(self, **params):
  308. """Set the parameters of this estimator."""
  309. for key, value in params.items():
  310. if hasattr(self, key):
  311. setattr(self, key, value)
  312. else:
  313. raise ValueError(f"Invalid parameter {key}")
  314. return self
  315. def get_params(self, deep=True):
  316. """Get parameters for this estimator."""
  317. return {
  318. 'n_iterations': self.n_iterations,
  319. 'across_time': self.across_time,
  320. 'random_state': self.random_state,
  321. 'n_jobs': self.n_jobs,
  322. 'batch_size': self.batch_size
  323. }
  324. def decode_time(X, y, n_inter=1, return_inter=False):
  325. y = np.array(y)
  326. scores_all = []
  327. for n in range(n_inter):
  328. clf = make_pipeline(
  329. StandardScaler(),
  330. SVM(C=5e-4))
  331. time_gen = GeneralizingEstimator(clf, scoring="roc_auc", n_jobs=-1, verbose=True)
  332. n_trls = X.shape[0]
  333. idx_rnd = np.concatenate([np.zeros(n_trls // 2), np.ones(n_trls // 2)])
  334. if (n_trls % 2) > 0:
  335. idx_rnd = np.concatenate([idx_rnd, [1.0]])
  336. random.shuffle(idx_rnd)
  337. X1 = X[idx_rnd < 1, :, :]
  338. y1 = y[idx_rnd < 1]
  339. X2 = X[idx_rnd > 0, :, :]
  340. y2 = y[idx_rnd > 0]
  341. time_gen.fit(X1, y1)
  342. score1 = time_gen.score(X2, y2)
  343. time_gen.fit(X2, y2)
  344. score2 = time_gen.score(X1, y1)
  345. scores_all.append(np.array([score1, score2]).mean(0))
  346. if return_inter:
  347. return scores_all
  348. else:
  349. return np.mean(np.array(scores_all), axis=0)
  350. def dy_mask(obs, rnd, clusteralpha=0.05):
  351. '''
  352. looks for significant off diagonal reduction
  353. Params:
  354. obs: 2d observations
  355. rnd: 3d randomizations (randomizations are in trailing dimension)
  356. Outputs:
  357. dynaMask: 1d dynamism index (time-resolved)
  358. '''
  359. # Computing the mask
  360. # get the diagonal twice, reshape to allow easy comparison with full matrix
  361. dynaObsA = np.diag(obs)[:, np.newaxis] - obs
  362. dynaObsB = np.diag(obs)[np.newaxis, :] - obs
  363. # do the same for the randomizations
  364. dynaRndA = np.diagonal(rnd)[:, :, np.newaxis] - rnd.T
  365. dynaRndB = np.diagonal(rnd)[:, np.newaxis, :] - rnd.T
  366. dynaRndA = dynaRndA.T
  367. dynaRndB = dynaRndB.T
  368. pDynaA, labDynaA = permutation_test(obsdat=dynaObsA, rnddat=dynaRndA, tail=1, clusteralpha=clusteralpha)
  369. pDynaB, labDynaB = permutation_test(dynaObsB, dynaRndB, tail=1, clusteralpha=clusteralpha)
  370. dynaMask = (pDynaA < 0.05) & (pDynaB < 0.05)
  371. di = (np.mean(dynaMask, axis=0) + np.mean(dynaMask, axis=1)) / 2
  372. return di, dynaMask
  373. def plot_time_generalisation(gs, fig, dec, rnd=None, times_sec=None, vmin=-1, vmax=1.0,
  374. matrix_limits=None, tick_positions=None,
  375. reference_lines=None, cmap="RdBu_r", title= "Coding stability", stat_test ='off-diagonal', clusteralpha=0.05, tail=0):
  376. """
  377. Plot time generalisation matrix with both dynamism and stability indices
  378. This function creates a comprehensive visualisation of temporal generalisation
  379. within a specified grid position.
  380. Parameters:
  381. -----------
  382. gs : matplotlib GridSpec slice
  383. GridSpec slice (e.g., gs[0,0] for single cell or gs[0,0:2] for spanning cells)
  384. fig : matplotlib Figure
  385. The figure object to add the subplot to
  386. dec : ndarray
  387. The decoding matrix to display (observation)
  388. rnd : ndarray, optional
  389. Randomization matrix with shape (time, time, n_randomizations)
  390. times_sec : ndarray, optional
  391. Time points in seconds. If None, creates a default range
  392. vmin, vmax : float
  393. Min and max values for color mapping
  394. matrix_limits : tuple, optional
  395. (min_idx, max_idx) limits for both x and y axes
  396. tick_positions : list or ndarray, optional
  397. Indices where to place ticks
  398. reference_lines : list, optional
  399. Time points where to place reference lines
  400. cmap : str
  401. Colormap to use
  402. """
  403. import numpy as np
  404. import matplotlib.pyplot as plt
  405. from scipy.ndimage import label
  406. # Get current figure (use the provided fig parameter)
  407. # fig is already provided as parameter
  408. # Create default time range if not provided
  409. if times_sec is None:
  410. times_sec = np.linspace(-0.1, 1.1, dec.shape[-1])
  411. # Create default tick positions if not provided
  412. if tick_positions is None:
  413. tick_positions = np.array([10, 60, 110])
  414. # Create default reference lines if not provided
  415. if reference_lines is None:
  416. reference_lines = [0, 0.5, 1.0]
  417. # Create subplot using the provided GridSpec
  418. ax = fig.add_subplot(gs)
  419. # Plot the main time generalization matrix
  420. im = ax.matshow(
  421. dec,
  422. vmin=vmin,
  423. vmax=vmax,
  424. cmap=cmap,
  425. origin="lower",
  426. )
  427. # Compute dynamism and stability if randomization data is provided
  428. if rnd is not None:
  429. if stat_test == 'off-diagonal':
  430. # Get dynamism index and significance mask
  431. _, sig_mask = dy_mask(dec, rnd, clusteralpha=clusteralpha)
  432. elif stat_test == 'on-diagonal':
  433. p_vals, cluster_labels = permutation_test(
  434. obsdat=dec,
  435. rnddat=rnd,
  436. clustercorrect=True,
  437. clusteralpha=clusteralpha,
  438. tail = tail,
  439. )
  440. # Create mask for significant areas (p <= 0.05)
  441. sig_mask = p_vals <= 0.05
  442. # Plot significance masks on the main plot
  443. if np.any(sig_mask):
  444. # Use contour to outline significant dynamism areas
  445. ax.contour(sig_mask, levels=[0.5], colors='black',
  446. linestyles='-', linewidths=0.8)
  447. # Set tick labels
  448. tick_labels = [f"{times_sec[i]:.1f}" for i in tick_positions]
  449. # Set the ticks and labels
  450. ax.set_xticks(tick_positions)
  451. ax.set_xticklabels(tick_labels)
  452. ax.set_yticks(tick_positions)
  453. ax.set_yticklabels(tick_labels)
  454. # Set matrix display limits if provided
  455. if matrix_limits is not None:
  456. min_idx, max_idx = matrix_limits
  457. ax.set_xlim(min_idx, max_idx)
  458. ax.set_ylim(min_idx, max_idx)
  459. # Make sure ticks are at the bottom
  460. ax.xaxis.set_ticks_position("bottom")
  461. # Add reference lines to main plot
  462. for time_val in reference_lines:
  463. time_idx = np.argmin(np.abs(times_sec - time_val))
  464. ax.axhline(time_idx, color="k", linestyle="--", linewidth=0.8)
  465. ax.axvline(time_idx, color="k", linestyle="--", linewidth=0.8)
  466. # Add axis labels
  467. ax.set_xlabel('time (s)')
  468. ax.set_ylabel('time (s)')
  469. # Add title
  470. ax.set_title(title)
  471. # Add colorbar
  472. cbar = fig.colorbar(im, ax=ax)
  473. cbar.set_label("Pearson r")
  474. def plot_magnitude(gs, fig, mag_obs, mag_null_within, mag_null_ler,
  475. title, y_lim=[-1,1], alpha=0.05):
  476. """
  477. Plot cross-generalization magnitude across learning stages as bars.
  478. Parameters
  479. ----------
  480. mag_obs : array, shape (n_stages,)
  481. Observed magnitude values
  482. mag_null_within : array, shape (n_permutations, n_stages)
  483. Null distribution where only cross-gen is shuffled
  484. mag_null_ler : array, shape (n_permutations, n_stages)
  485. Full permutation null (both within and cross-gen shuffled)
  486. """
  487. x_data = np.arange(len(mag_obs)) + 1
  488. ax = fig.add_subplot(gs)
  489. # Bar plot
  490. ax.bar(x_data, mag_obs, color='black', width=0.6, zorder=5)
  491. ax.set_xticks(x_data)
  492. ax.set_xticklabels(x_data)
  493. sns.despine(right=True, top=True)
  494. ax.set_ylabel('% of ceiling')
  495. ax.set_xlabel('learning stage')
  496. ax.set_title(title)
  497. # --- Null band from mag_null_within ---
  498. null_flat = mag_null_within.flatten()
  499. null_mean = null_flat.mean()
  500. # Center around mean
  501. null_centered = null_flat - null_mean
  502. crit = np.quantile(np.abs(null_centered), 1 - alpha)
  503. lower = null_mean - crit
  504. upper = null_mean + crit
  505. # Mean line
  506. ax.axhline(null_mean, linestyle='--', linewidth=1,
  507. color='black', zorder=-6, alpha=0.7)
  508. ax.axhline(1, linestyle='--', linewidth=1,
  509. color='black', zorder=-6, alpha=0.7)
  510. # Null band
  511. rect = patches.Rectangle(
  512. (x_data[0] - 0.5, lower),
  513. len(x_data),
  514. upper - lower,
  515. edgecolor=None,
  516. facecolor='lightgrey',
  517. zorder=-7,
  518. )
  519. ax.add_patch(rect)
  520. # --- Horizontal bracket: Learning effect (Stage 1 vs Stage 4) ---
  521. # Bracket displayed from BELOW with arms pointing UP
  522. y_range = y_lim[1] - y_lim[0]
  523. bracket_base = y_lim[0] + 0.2 * y_range # Changed from y_lim[1] - 0.11
  524. bracket_height = bracket_base - 0.06 * y_range # Changed sign to go down
  525. bracket_center = (x_data[0] + x_data[-1]) / 2
  526. gap_size = 1.2
  527. print('Learning effect (magnitude increase):')
  528. p_value_learning = compute_p_value(mag_obs[-1], mag_obs[0],
  529. mag_null_ler[:, -1], mag_null_ler[:, 0],
  530. tail='greater')
  531. # Draw bracket - arms now point UP from below
  532. ax.plot([x_data[0], x_data[0]], [bracket_height, bracket_base],
  533. color='black', linewidth=1)
  534. ax.plot([x_data[-1], x_data[-1]], [bracket_height, bracket_base],
  535. color='black', linewidth=1)
  536. ax.plot([x_data[0], bracket_center - gap_size / 2],
  537. [bracket_height, bracket_height], color='black', linewidth=1)
  538. ax.plot([bracket_center + gap_size / 2, x_data[-1]],
  539. [bracket_height, bracket_height], color='black', linewidth=1)
  540. ax.text(bracket_center, bracket_height * 1.2, p_into_stars(p_value_learning),
  541. fontsize=10, color='black', ha='center', va='bottom') # Changed va to 'bottom'
  542. ax.set_ylim(y_lim)
  543. return ax, p_value_learning
  544. def plot_significance_stars(ax, scores_real, scores_null, time_window, ylim, tail='two', y_position=None, dec_type=0):
  545. """
  546. Compute p-value and plot significance stars on a time-resolved decoding plot.
  547. Parameters:
  548. -----------
  549. ax : matplotlib axis object
  550. The axis to plot on
  551. scores_real : array
  552. Real decoding scores of shape [decoding_type, time_points]
  553. scores_null : array
  554. Null distribution scores of shape [n_reps, decoding_type, time_points]
  555. time_window : tuple or list
  556. (start, end) in milliseconds for the time window
  557. ylim : tuple or list
  558. (ymin, ymax) of the plot
  559. tail : str
  560. 'two', 'greater', or 'less' for the statistical test
  561. y_position : float, optional
  562. Y-position for stars. If None, uses 90% of ylim range
  563. """
  564. # Compute p-value comparing first and last timepoint
  565. p_val = compute_p_value(scores_real[dec_type, 0], scores_real[dec_type, -1],
  566. scores_null[:, dec_type, 0], scores_null[:, dec_type, -1],
  567. tail=tail)
  568. # Convert p-value to stars
  569. if p_val <= 0.001:
  570. stars = '***'
  571. font_size = 15
  572. scaler_position = 0.8
  573. elif p_val <= 0.01:
  574. stars = '**'
  575. font_size = 15
  576. scaler_position = 0.8
  577. elif p_val <= 0.05:
  578. stars = '*'
  579. font_size = 15
  580. scaler_position = 0.8
  581. elif p_val <= 0.1:
  582. stars = '†'
  583. font_size = 10
  584. scaler_position = 0.8
  585. else:
  586. stars = 'ns'
  587. font_size = 10
  588. scaler_position = 0.9
  589. # Calculate center of time window in seconds
  590. tw_center = (((time_window[0] + time_window[1]) / 2) / 100) - 0.5
  591. # Calculate y-position (default to 90% of y-range for better spacing)
  592. if y_position is None:
  593. y_position = ylim[0] + scaler_position * (ylim[1] - ylim[0])
  594. # Plot stars
  595. ax.text(tw_center, y_position, stars,
  596. ha='center', va='center',
  597. fontsize=font_size,
  598. color='black')
  599. return ax, p_val
  600. def KLdivergence(x, y):
  601. """Compute the Kullback-Leibler divergence between two multivariate samples.
  602. Parameters
  603. ----------
  604. x : 2D array (n,d)
  605. Samples from distribution P, which typically represents the true
  606. distribution.
  607. y : 2D array (m,d)
  608. Samples from distribution Q, which typically represents the approximate
  609. distribution.
  610. Returns
  611. -------
  612. out : float
  613. The estimated Kullback-Leibler divergence D(P||Q).
  614. References
  615. ----------
  616. Pérez-Cruz, F. Kullback-Leibler divergence estimation of
  617. continuous distributions IEEE International Symposium on Information
  618. Theory, 2008.
  619. """
  620. from scipy.spatial import cKDTree as KDTree
  621. # Check the dimensions are consistent
  622. x = np.atleast_2d(x)
  623. y = np.atleast_2d(y)
  624. n, d = x.shape
  625. m, dy = y.shape
  626. assert (d == dy)
  627. # Build a KD tree representation of the samples and find the nearest neighbour
  628. # of each point in x.
  629. xtree = KDTree(x)
  630. ytree = KDTree(y)
  631. # Get the first two nearest neighbours for x, since the closest one is the
  632. # sample itself.
  633. r = xtree.query(x, k=2, eps=.01, p=2)[0][:, 1]
  634. s = ytree.query(x, k=1, eps=.01, p=2)[0]
  635. return -np.log(r / s).sum() * d / n + np.log(m / (n - 1.))
  636. def permutation_test(obsdat, rnddat, clustercorrect=True, clusteralpha=0.05, tail=1):
  637. """
  638. Performs an (optionally cluster-corrected) permutation test of the observed
  639. data, given the pre-computed randomizations. rnddat must have one extra
  640. trailing dimension compared to obsdat.
  641. """
  642. if tail == 0:
  643. alpha_2tail = clusteralpha / 2
  644. clusterThreshold_right = np.percentile(rnddat, 100 * (1 - alpha_2tail), axis=obsdat.ndim)
  645. clusterThreshold_left = np.percentile(rnddat, 100 * alpha_2tail, axis=obsdat.ndim)
  646. elif tail == -1:
  647. clusterThreshold = np.percentile(rnddat, 100 * clusteralpha, axis=obsdat.ndim)
  648. elif tail == 1:
  649. clusterThreshold = np.percentile(rnddat, 100 * (1 - clusteralpha), axis=obsdat.ndim)
  650. p = np.ones_like(obsdat, dtype='float64')
  651. # uncorrected 'test'
  652. if not clustercorrect:
  653. for inds, value in np.ndenumerate(obsdat):
  654. rnd = np.sort(rnddat[inds])
  655. if tail == 0:
  656. pval = 2 * min(
  657. np.searchsorted(rnd, value, side='left'),
  658. rnd.shape[0] - np.searchsorted(rnd, value, side='right')
  659. ) / rnd.shape[0]
  660. elif tail == -1:
  661. pval = np.searchsorted(rnd, value, side='right') / rnd.shape[0]
  662. elif tail == 1:
  663. pval = 1 - np.searchsorted(rnd, value, side='left') / rnd.shape[0]
  664. p[inds] = pval
  665. return p
  666. # subfunction to compute clusterstats in one dataset (observed/randomized)
  667. def compute_clusterstats(dat, getinds=True):
  668. if tail == 0:
  669. clusterCandidates_right = dat > clusterThreshold_right
  670. clusterCandidates_left = dat < clusterThreshold_left
  671. clusterCandidates = np.logical_or(clusterCandidates_right, clusterCandidates_left)
  672. elif tail == -1:
  673. clusterCandidates = dat < clusterThreshold
  674. elif tail == 1:
  675. clusterCandidates = dat > clusterThreshold
  676. # label connected tiles
  677. labelled, numfeat = label(clusterCandidates, output='uint32')
  678. # compute aggregate cluster statistic for each cluster
  679. # this is a quick way to use cluster size as the clusterstat.
  680. # for other clusterstats (e.g. summed stat) a few more lines are needed
  681. clusterNums, clusterStats = np.unique(labelled.ravel(), return_counts=True)
  682. # remove 0, which corresponds to a non-cluster
  683. if clusterNums[0] == 0:
  684. # note that the check is necessary because it can happen that the
  685. # entire observe data matrix exceeds the threshold, in that case we
  686. # don't want to remove the first element (which will be 1 instead
  687. # of zero)
  688. clusterNums = clusterNums[1:]
  689. clusterStats = clusterStats[1:]
  690. # use cluster sum instead of size
  691. clusterStats = [np.sum(dat[labelled == x]) for x in clusterNums]
  692. if getinds:
  693. return clusterStats, clusterNums, labelled
  694. else:
  695. return clusterStats
  696. # get observed clusters and maximum randomized clusterstats
  697. clusObs, clusNums, labelled = compute_clusterstats(obsdat)
  698. clusRnd = [compute_clusterstats(rnddat[..., x], False)
  699. for x in range(rnddat.shape[-1])]
  700. # treat randomizations with 0 cluster candidates as if their max was 0
  701. mymax = lambda x: 0 if len(x) == 0 else np.max(x)
  702. mymin = lambda x: 0 if len(x) == 0 else np.min(x)
  703. if tail == 0:
  704. clusRnd = [mymax(np.abs(x)) for x in clusRnd]
  705. clusObs = [np.abs(s) for s in clusObs]
  706. elif tail == -1:
  707. clusRnd = [mymin(x) for x in clusRnd]
  708. elif tail == 1:
  709. clusRnd = [mymax(x) for x in clusRnd]
  710. clusRnd.sort()
  711. for stat, num in zip(clusObs, clusNums):
  712. if tail == 0:
  713. pval = 1 - np.searchsorted(clusRnd, stat, side='left') / rnddat.shape[-1]
  714. elif tail == -1:
  715. pval = np.searchsorted(clusRnd, stat, side='right') / rnddat.shape[-1]
  716. elif tail == 1:
  717. pval = 1 - np.searchsorted(clusRnd, stat, side='left') / rnddat.shape[-1]
  718. p[labelled == num] = pval
  719. return p, labelled
  720. def assign_lables(labels, factor):
  721. conditions = np.unique(labels)
  722. res = dict(zip(conditions, factor))
  723. return list(map(res.get, labels))
  724. def bias_corr(sel):
  725. conds = np.zeros((4, 3)) # conditions x selectivity (pure color, pure shape, interaction)
  726. conds[0, 0] = 1
  727. conds[1, 1] = 1 # Not rewarded
  728. conds[2, :] = 0
  729. conds[3, :] = 1 # rewarded
  730. const = np.array([1, 1, 1, 1])
  731. X = np.vstack([const, conds[:, 0], conds[:, 1], conds[:, 2]]).T
  732. rates = sel @ conds.T
  733. coeffs = sp.linalg.lstsq(X, rates.T)[0][1:, :]
  734. cov_new = np.cov(coeffs)
  735. new_sel = np.random.multivariate_normal(np.zeros(sel.shape[1]),
  736. cov_new,
  737. sel.shape[0])
  738. return new_sel
  739. def downsample_data(*arg, fact=10):
  740. retval = []
  741. # downsample trailing dimension (i.e. time axis) by factor 10
  742. for dat in arg:
  743. shape = list(dat.shape)
  744. shape[-1] = shape[-1] // fact
  745. dat = dat.reshape(shape + [fact])
  746. retval.append(np.mean(dat, axis=len(shape)))
  747. return np.array(retval)[0, :, :, :]
  748. def get_data(session_list, path_spikes, path_meta, window, cut_off=False):
  749. parts = session_list
  750. data_all_parts = []
  751. labels_all_parts = []
  752. for p in range(len(parts)):
  753. sessions = parts[p]
  754. print('Loading and combining the spiking data: part ', p + 1, '/', len(parts))
  755. units_dat = []
  756. units_labels = []
  757. for s in tqdm(range(len(sessions))):
  758. spks = np.load(path_spikes.format(sessions[s]))
  759. spks = spks[:, :, 500:3000]
  760. spks = downsample_data(spks)
  761. trials = np.load(path_meta.format(sessions[s]))
  762. labels = trials[trials != 0]
  763. spks_data = spks[trials != 0, :, window[0]:window[1]]
  764. spks_data = spks_data[:, :, :]
  765. if cut_off:
  766. spks_data = spks_data[:cut_off, :, :]
  767. labels = labels[:cut_off]
  768. units_dat.append(spks_data)
  769. units_labels.append(labels)
  770. data_all_parts.append(units_dat)
  771. labels_all_parts.append(units_labels)
  772. return data_all_parts, labels_all_parts
  773. def exclude_neurons(data, session_list, path_locations, path_sel_exclude, loc=None, non_sig=False, threshold=0.0,
  774. plot=False, times=np.linspace(-0.5, 2.0, 250), ylim=50):
  775. n_parts = len(data)
  776. data_new_ses_parts = []
  777. n_exc_parts = []
  778. n_neurons_parts = []
  779. exc_parts = []
  780. parts = session_list
  781. for p in range(n_parts):
  782. data_new_ses = []
  783. n_exc_ses = []
  784. n_neurons_ses = []
  785. for s in range(len(data[p])):
  786. data_ses = data[p][s]
  787. data_avg = data_ses.mean(0)
  788. mean_firing = data_avg.mean(1)
  789. if loc:
  790. cell_vals = pd.read_csv(path_locations.format(parts[p][s]))['Area'].values
  791. exc_idc = (cell_vals >= loc[0]) & (cell_vals <= loc[1])
  792. if threshold:
  793. exc_thr = mean_firing > threshold
  794. exc_idc = np.logical_and(exc_idc, exc_thr)
  795. if non_sig:
  796. sel_exc_idc = ~np.load(path_sel_exclude.format(parts[p][s]))
  797. exc_idc = np.logical_and(exc_idc, sel_exc_idc)
  798. else:
  799. exc_idc = mean_firing > -1
  800. data_new = data_ses[:, exc_idc, :]
  801. data_new_ses.append(data_new)
  802. n_exc_ses.append(np.sum(exc_idc))
  803. n_neurons_ses.append(len(exc_idc))
  804. exc_proc = round(1 - np.sum(n_exc_ses) / np.sum(n_neurons_ses), 2)
  805. exc_parts.append(exc_proc)
  806. print('Excluded', exc_parts[p] * 100, '% of neurons from part', p + 1)
  807. if plot:
  808. colors = sns.xkcd_palette(['pale red'])
  809. data_avg_trl = [data_new_ses[_].mean(0) for _ in range(len(data_new_ses))]
  810. data_plot = np.concatenate(data_avg_trl, axis=0)
  811. df_plot = pd.DataFrame(data_plot.T, columns=list(range(1, data_plot.shape[0] + 1)))
  812. df_plot['Times'] = times
  813. df_melted = pd.melt(df_plot, id_vars=['Times'])
  814. df_melted['Rate'] = df_melted['value']
  815. df_melted['Neurons'] = df_melted['variable']
  816. plt.figure()
  817. sns.lineplot(data=df_melted, x="Times", y='Rate', hue="Neurons")
  818. plt.ylim([None, ylim])
  819. plt.axvline(0, linestyle="--", linedecode_epoch=0.8, color='black')
  820. sns.despine(right=True, top=True)
  821. plt.axvline(0.5, linestyle="--", linewidth=0.8, color='black')
  822. plt.axvline(1, linestyle="--", linewidth=0.8, color='black')
  823. mean_rate = df_melted.groupby('Neurons').mean()['Rate'].values
  824. plt.figure()
  825. sns.distplot(mean_rate, bins=20, kde=False, color='grey')
  826. sns.despine(right=True, top=True)
  827. median_firing_rate = round(np.median(data_plot.mean(1)), 2)
  828. plt.title("Median firing rate = " + str(median_firing_rate))
  829. plt.xlabel('Rate')
  830. plt.ylabel('Count')
  831. plt.axvline(median_firing_rate, 0, 1, linewidth=1, linestyle='--', color='black')
  832. exc_proc = round(1 - np.sum(n_exc_ses) / np.sum(n_neurons_ses), 2)
  833. plt.annotate('Exc % = ' + str(exc_proc), xy=(1, 40), color=colors[0])
  834. plt.xlim([None, 25])
  835. data_new_ses_parts.append(data_new_ses)
  836. n_neurons_parts.append(np.sum(n_neurons_ses))
  837. n_exc_parts.append(np.sum(n_neurons_ses) - np.sum(n_exc_ses))
  838. return data_new_ses_parts, exc_parts, n_neurons_parts
  839. def condi_avg(data, labels):
  840. conditions = np.unique(labels)
  841. condition_averages = []
  842. for _ in range(len(conditions)):
  843. condition_dat = data[labels == conditions[_], :, :]
  844. condition_averages.append(condition_dat.mean(0))
  845. return np.array(condition_averages)
  846. def plot_covs(data, model_names, if_diff=False):
  847. n_models = len(data)
  848. n_epochs = len(data[0])
  849. if if_diff:
  850. covs = np.zeros((n_models, n_epochs + 1, 3, 3))
  851. else:
  852. covs = np.zeros((n_models, n_epochs, 3, 3))
  853. for m in range(n_models):
  854. for p in range(n_epochs):
  855. covs[m, p, :, :] = np.cov(data[m][p][:, :, 0].T)
  856. if if_diff:
  857. covs[:, n_epochs, :, :] = covs[:, 0, :, :] - covs[:, -1, :, :]
  858. n_epochs = n_epochs + 1
  859. max_cov = np.max(covs)
  860. min_cov = -max_cov
  861. print(min_cov)
  862. print(max_cov)
  863. norm = colors.TwoSlopeNorm(vmin=min_cov, vcenter=0, vmax=max_cov)
  864. covs_flat = np.resize(covs, (covs.shape[0] * covs.shape[1], covs.shape[2], covs.shape[3]))
  865. epoch_label = list(range(1, n_epochs + 1)) * 2
  866. fig, big_axes = plt.subplots(figsize=(8.0, 4.0), nrows=2, ncols=1, sharey=True)
  867. for row, big_ax in enumerate(big_axes, start=1):
  868. big_ax.set_title(model_names[row - 1] + " \n", fontsize=14)
  869. # Turn off axis lines and ticks of the big subplot
  870. # obs alpha is 0 in RGBA string!
  871. big_ax.tick_params(labelcolor=(1., 1., 1., 0.0), top='off', bottom='off', left='off', right='off')
  872. big_ax.axis('off')
  873. # removes the white frame
  874. big_ax._frameon = False
  875. for i in range(1, n_epochs * 2 + 1):
  876. ax = fig.add_subplot(2, n_epochs, i)
  877. ax.imshow(covs_flat[i - 1, :, :], cmap=sns.diverging_palette(230, 20, as_cmap=True), norm=norm)
  878. if if_diff:
  879. if i == n_epochs or i == (n_epochs * 2):
  880. plt.title('difference (last vs first)')
  881. else:
  882. plt.title('epoch ' + str(epoch_label[i - 1]))
  883. else:
  884. plt.title('epoch ' + str(epoch_label[i - 1]))
  885. plt.axis('off')
  886. fig.set_facecolor('w')
  887. plt.tight_layout()
  888. plt.show()
  889. return covs
  890. def get_betas_cross_val(data, labels, condition_labels, normalisation='mean_centred', time_window=[130, 150],
  891. add_constant=True,
  892. design_model='0/1', cross_val=False, if_xgen=False, task='task_1', full_model=False):
  893. n_parts = len(data)
  894. betas_part = []
  895. for p in range(n_parts):
  896. betas = []
  897. for s in range(len(data[p])):
  898. firing_rates_ses = data[p][s]
  899. firing_rates_ses = firing_rates_ses[:, :, time_window[0]:time_window[1]]
  900. labels_ses = labels[p][s]
  901. if task == 'task_1':
  902. firing_rates_ses = firing_rates_ses[labels_ses < 9, :, :]
  903. labels_ses = labels_ses[labels_ses < 9]
  904. elif task == 'task_2':
  905. firing_rates_ses = firing_rates_ses[labels_ses > 8, :, :]
  906. labels_ses = labels_ses[labels_ses > 8]
  907. if cross_val:
  908. firing_rates_ses1 = firing_rates_ses[::2, :, :]
  909. firing_rates_ses2 = firing_rates_ses[1::2, :, :]
  910. labels_ses1 = labels_ses[::2]
  911. labels_ses2 = labels_ses[1::2]
  912. else:
  913. firing_rates_ses1 = firing_rates_ses
  914. firing_rates_ses2 = firing_rates_ses
  915. labels_ses1 = labels_ses
  916. labels_ses2 = labels_ses
  917. if if_xgen:
  918. firing_rates_ses1 = firing_rates_ses[labels_ses1 < 9, :, :]
  919. firing_rates_ses2 = firing_rates_ses[labels_ses2 > 8, :, :]
  920. labels_ses1 = labels_ses1[labels_ses1 < 9]
  921. labels_ses2 = labels_ses2[labels_ses2 > 8]
  922. firing_rates_ses = [firing_rates_ses1, firing_rates_ses2]
  923. labels_ses = [labels_ses1, labels_ses2]
  924. betas_split = []
  925. n_splits = len(labels_ses)
  926. for split in range(n_splits):
  927. firing_mean = condi_avg(firing_rates_ses[split], labels_ses[split])
  928. if normalisation == 'zscore':
  929. firing_mean = (firing_mean - np.mean(firing_mean, axis=0, keepdims=True)) / np.std(firing_mean,
  930. axis=0,
  931. keepdims=True)
  932. elif normalisation == 'mean_centred':
  933. firing_mean = firing_mean - np.mean(firing_mean, axis=0, keepdims=True)
  934. elif normalisation == 'none':
  935. firing_mean = firing_mean
  936. elif normalisation == 'soft':
  937. firing_mean = firing_mean / (
  938. (np.max(firing_mean, axis=0, keepdims=True) - np.min(firing_mean, axis=0,
  939. keepdims=True)) + 5)
  940. if design_model == '0/1':
  941. labels_cue = np.array(condition_labels[0])
  942. labels_target = np.array(condition_labels[1])
  943. labels_int1 = np.array(condition_labels[2])
  944. if full_model:
  945. labels_target2 = np.array(condition_labels[2])
  946. labels_int1 = np.array(condition_labels[3])
  947. labels_int2 = np.array(condition_labels[4])
  948. elif design_model == '+1/-1':
  949. labels_cue = np.array(condition_labels[0])
  950. labels_cue = np.where(labels_cue == 0, -1, labels_cue)
  951. labels_target = np.array(condition_labels[1])
  952. labels_target = np.where(labels_target == 0, -1, labels_target)
  953. labels_int1 = np.array(condition_labels[2])
  954. labels_int1 = np.where(labels_int1 == 0, -1, labels_int1)
  955. if full_model:
  956. labels_target2 = np.array(condition_labels[2])
  957. labels_target2 = np.where(labels_target2 == 0, -1, labels_target2)
  958. labels_int1 = np.array(condition_labels[3])
  959. labels_int1 = np.where(labels_int1 == 0, -1, labels_int1)
  960. labels_int2 = np.array(condition_labels[4])
  961. labels_int2 = np.where(labels_int2 == 0, -1, labels_int2)
  962. if full_model:
  963. if add_constant:
  964. constant = np.ones_like(labels_cue).astype(float)
  965. design_matrix = np.vstack(
  966. [constant, labels_cue, labels_target, labels_int1, labels_target2, labels_int2]).T
  967. else:
  968. design_matrix = np.vstack(
  969. [labels_cue, labels_target, labels_int1, labels_target2, labels_int2]).T
  970. else:
  971. if add_constant:
  972. constant = np.ones_like(labels_cue).astype(float)
  973. design_matrix = np.vstack(
  974. [constant, labels_cue, labels_target, labels_int1]).T
  975. else:
  976. design_matrix = np.vstack(
  977. [labels_cue, labels_target, labels_int1]).T
  978. design_matrix = design_matrix.astype(float)
  979. n_neurons = firing_mean.shape[1]
  980. n_times = firing_mean.shape[-1]
  981. n_models = design_matrix.shape[1]
  982. design_matrix_stacked = np.concatenate([design_matrix] * n_times, axis=0)
  983. firing_rates_ses_stacked = np.concatenate(np.array_split(firing_mean, n_times, axis=-1), axis=0)
  984. betas_ses = np.zeros((n_neurons, n_models, 1))
  985. for cell in range(n_neurons):
  986. betas_ses[cell, :, :] = \
  987. sp.linalg.lstsq(design_matrix_stacked, firing_rates_ses_stacked[:, cell, :])[0][
  988. :, :]
  989. betas_split.append(betas_ses)
  990. betas.append(np.array(betas_split)[:, :, :, 0])
  991. betas_part.append(betas)
  992. if full_model:
  993. epochs_rel = []
  994. for p in range(n_parts):
  995. if add_constant:
  996. epoch = np.concatenate(betas_part[p], axis=1)[:, :, 1:4].transpose(
  997. (0, 2, 1)) # splits x coeffs x neurons
  998. else:
  999. epoch = np.concatenate(betas_part[p], axis=0)[:, :, :3].transpose((0, 2, 1))[0, :,
  1000. :] # coeffs x neurons
  1001. epoch -= np.mean(epoch, axis=2, keepdims=True)
  1002. epochs_rel.append(epoch.T)
  1003. epochs_irrel = []
  1004. for p in range(n_parts):
  1005. if add_constant:
  1006. epoch_1 = np.concatenate(betas_part[p], axis=1)[:, :, 1][:, :, None].transpose(
  1007. (0, 2, 1)) # splits x coeffs x neurons
  1008. epoch_2 = np.concatenate(betas_part[p], axis=1)[:, :, 4:].transpose(
  1009. (0, 2, 1)) # splits x coeffs x neurons
  1010. epoch = np.concatenate([epoch_1, epoch_2], axis=1)
  1011. else:
  1012. epoch_1 = np.concatenate(betas_part[p], axis=1)[:, :, 0][:, :, None].transpose(
  1013. (0, 2, 1)) # splits x coeffs x neurons
  1014. epoch_2 = np.concatenate(betas_part[p], axis=1)[:, :, 3:].transpose(
  1015. (0, 2, 1)) # splits x coeffs x neurons
  1016. epoch = np.concatenate([epoch_1, epoch_2], axis=1)
  1017. epoch -= np.mean(epoch, axis=2, keepdims=True)
  1018. epochs_irrel.append(epoch.T)
  1019. return epochs_rel, epochs_irrel
  1020. else:
  1021. epochs = []
  1022. for p in range(n_parts):
  1023. if add_constant:
  1024. epoch = np.concatenate(betas_part[p], axis=1)[:, :, 1:4].transpose(
  1025. (0, 2, 1)) # splits x coeffs x neurons
  1026. else:
  1027. epoch = np.concatenate(betas_part[p], axis=0)[:, :, :3].transpose((0, 2, 1))[0, :,
  1028. :] # coeffs x neurons
  1029. epoch -= np.mean(epoch, axis=2, keepdims=True)
  1030. epochs.append(epoch.T)
  1031. return epochs
  1032. def get_freqs(x):
  1033. return {value: len(list(freq)) for value, freq in groupby(sorted(list(x)))}
  1034. def dist_random(epochs, n_bootstraps=1000, rnd_model='gaussian (spherical)', model_names=['cue + shape', 'cue + width'],
  1035. design_model='0/1', bon_correction=False, metric='KL divergance estimate', relative_dist=False):
  1036. n_epochs = len(epochs[0])
  1037. dfs = []
  1038. p_vals = []
  1039. for model_n in range(len(model_names)):
  1040. KL = np.zeros((n_epochs, n_bootstraps))
  1041. KL_r = np.zeros((n_epochs, n_bootstraps))
  1042. KL_opt = np.zeros((n_epochs, n_bootstraps))
  1043. for e in range(n_epochs):
  1044. for bootstrap in tqdm(range(n_bootstraps)):
  1045. data_train = epochs[model_n][e][:, :, 0]
  1046. data_test = epochs[model_n][e][:, :, 1]
  1047. if design_model == '0/1':
  1048. opt_cov = np.zeros((3, 3))
  1049. m = np.mean(np.diag(np.cov(data_train.T)))
  1050. opt_cov[:2, :2] = 0.5 * m
  1051. opt_cov[2, 2] = 2 * m
  1052. opt_cov[2, :2] = -m;
  1053. opt_cov[:2, 2] = -m
  1054. elif design_model == '+1/-1':
  1055. # m = np.mean(np.diag(np.cov(data_train.T)))
  1056. # opt_cov = np.diag([0, 0, m * 3])
  1057. opt_cov = np.diag([0, 0, np.cov(data_train.T)[2, 2]])
  1058. s_opt = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
  1059. opt_cov,
  1060. data_train.shape[0])
  1061. if rnd_model == 'gaussian (spherical)':
  1062. m = np.mean(np.diag(np.cov(data_train.T)))
  1063. shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
  1064. np.diag([m, m, m]),
  1065. data_train.shape[0])
  1066. m = np.mean(np.diag(np.cov(data_test.T)))
  1067. shuffled_2 = np.random.multivariate_normal(np.zeros(data_test.shape[1]),
  1068. np.diag([m, m, m]),
  1069. data_test.shape[0])
  1070. if metric == 'KL divergance estimate':
  1071. kl_itr = 0.5 * (KLdivergence(data_test, shuffled_1) + KLdivergence(shuffled_1, data_test))
  1072. kl_r_itr = 0.5 * (KLdivergence(shuffled_2, shuffled_1) + KLdivergence(shuffled_1, shuffled_2))
  1073. elif metric == 'euclidean distance':
  1074. kl_itr = euclidean_distance(data_test, shuffled_1)
  1075. kl_r_itr = euclidean_distance(shuffled_2, shuffled_1)
  1076. elif metric == 'epairs':
  1077. kl_itr = epairs_metric(data_test, shuffled_1)
  1078. kl_r_itr = epairs_metric(shuffled_2, shuffled_1)
  1079. elif rnd_model == 'gaussian (tied)':
  1080. shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
  1081. np.diag(np.diag(np.cov(data_train.T))),
  1082. data_train.shape[0])
  1083. shuffled_2 = np.random.multivariate_normal(np.zeros(data_test.shape[1]),
  1084. np.diag(np.diag(np.cov(data_test.T))),
  1085. data_test.shape[0])
  1086. if metric == 'KL divergance estimate':
  1087. kl_itr = 0.5 * (KLdivergence(data_test, shuffled_1) + KLdivergence(shuffled_1, data_test))
  1088. kl_r_itr = 0.5 * (KLdivergence(shuffled_2, shuffled_1) + KLdivergence(shuffled_1, shuffled_2))
  1089. elif metric == 'euclidean distance':
  1090. kl_itr = euclidean_distance(data_test, shuffled_1)
  1091. kl_r_itr = euclidean_distance(shuffled_2, shuffled_1)
  1092. elif metric == 'epairs':
  1093. kl_itr = epairs_metric(data_test, shuffled_1)
  1094. kl_r_itr = epairs_metric(shuffled_2, shuffled_1)
  1095. KL[e, bootstrap] = kl_itr
  1096. KL_r[e, bootstrap] = kl_r_itr
  1097. if metric == 'KL divergance estimate':
  1098. KL_opt[e, bootstrap] = 0.5 * (KLdivergence(s_opt, shuffled_1) + KLdivergence(shuffled_1, s_opt))
  1099. elif metric == 'euclidean distance':
  1100. KL_opt[e, bootstrap] = euclidean_distance(s_opt, shuffled_1)
  1101. elif metric == 'epairs':
  1102. KL_opt[e, bootstrap] = epairs_metric(s_opt, shuffled_1)
  1103. if relative_dist:
  1104. KL_r_avg = np.mean(KL_r, keepdims=True, axis=-1)
  1105. KL -= KL_r_avg
  1106. KL_r -= KL_r_avg
  1107. KL_opt -= KL_r_avg
  1108. KL_opt_avg = np.mean(KL_opt, keepdims=True, axis=-1)
  1109. KL /= KL_opt_avg
  1110. KL_r /= KL_opt_avg
  1111. KL_opt /= KL_opt_avg
  1112. p = 1 * (np.sum(KL_r >= np.mean(KL, axis=-1, keepdims=True), axis=-1) / n_bootstraps)
  1113. if bon_correction:
  1114. p = p * n_epochs
  1115. p_vals.append(p)
  1116. epoch = np.sort(list(range(1, n_epochs + 1)) * n_bootstraps)
  1117. epoch_labels = np.concatenate([epoch, epoch, epoch])
  1118. dist_label = np.concatenate([['observed'] * (n_bootstraps * n_epochs), [rnd_model] * (n_bootstraps * n_epochs),
  1119. ['structured'] * (n_bootstraps * n_epochs)])
  1120. KL_df = np.reshape(KL, KL.shape[0] * KL.shape[1])
  1121. KL_r_df = np.reshape(KL_r, KL_r.shape[0] * KL_r.shape[1])
  1122. KL_opt_df = np.reshape(KL_opt, KL_opt.shape[0] * KL_opt.shape[1])
  1123. KL_all_df = np.concatenate([KL_df, KL_r_df, KL_opt_df])
  1124. df = pd.DataFrame(np.array([KL_all_df, epoch_labels, dist_label]).T,
  1125. columns=[metric, 'learning epoch', 'distribution'])
  1126. df[metric] = df[metric].astype(float)
  1127. df['model'] = model_names[model_n]
  1128. dfs.append(df)
  1129. df_all = pd.concat(dfs)
  1130. df_all['divergence from'] = 'random selectivity'
  1131. print(p_vals)
  1132. return df_all, p_vals, [KL, KL_r, KL_opt]
  1133. def dist_structured(epochs, n_bootstraps=1000, rnd_model='gaussian (spherical)',
  1134. model_names=['cue + shape', 'cue + width'],
  1135. design_model='0/1', bon_correction=False, metric='KL divergance estimate', relative_dist=False):
  1136. n_epochs = len(epochs[0])
  1137. dfs = []
  1138. p_vals = []
  1139. for model_n in range(len(model_names)):
  1140. KL = np.zeros((n_epochs, n_bootstraps))
  1141. KL_r = np.zeros((n_epochs, n_bootstraps))
  1142. KL_opt = np.zeros((n_epochs, n_bootstraps))
  1143. for e in range(n_epochs):
  1144. for bootstrap in tqdm(range(n_bootstraps)):
  1145. data_train = epochs[model_n][e][:, :, 0]
  1146. data_test = epochs[model_n][e][:, :, 1]
  1147. if design_model == '0/1':
  1148. opt_cov1 = np.zeros((3, 3))
  1149. m = np.mean(np.diag(np.cov(data_train.T)))
  1150. opt_cov1[:2, :2] = 0.5 * m
  1151. opt_cov1[2, 2] = 2 * m
  1152. opt_cov1[2, :2] = -m;
  1153. opt_cov1[:2, 2] = -m
  1154. opt_cov2 = np.zeros((3, 3))
  1155. m = np.mean(np.diag(np.cov(data_test.T)))
  1156. opt_cov2[:2, :2] = 0.5 * m
  1157. opt_cov2[2, 2] = 2 * m
  1158. opt_cov2[2, :2] = -m;
  1159. opt_cov2[:2, 2] = -m
  1160. elif design_model == '+1/-1':
  1161. opt_cov1 = np.diag([0, 0, np.cov(data_train.T)[2, 2]])
  1162. opt_cov2 = np.diag([0, 0, np.cov(data_test.T)[2, 2]])
  1163. s_opt_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
  1164. opt_cov1,
  1165. data_train.shape[0])
  1166. s_opt_2 = np.random.multivariate_normal(np.zeros(data_test.shape[1]),
  1167. opt_cov2,
  1168. data_test.shape[0])
  1169. if rnd_model == 'gaussian (tied)':
  1170. shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
  1171. np.diag(np.diag(np.cov(data_train.T))),
  1172. data_train.shape[0])
  1173. if metric == 'KL divergance estimate':
  1174. KL[e, bootstrap] = 0.5 * (KLdivergence(data_test, s_opt_1) + KLdivergence(s_opt_1, data_test))
  1175. KL_r[e, bootstrap] = 0.5 * (
  1176. KLdivergence(shuffled_1, s_opt_1) + KLdivergence(s_opt_1, shuffled_1))
  1177. KL_opt[e, bootstrap] = 0.5 * (KLdivergence(s_opt_2, s_opt_1) + KLdivergence(s_opt_1, s_opt_2))
  1178. elif metric == 'euclidean distance':
  1179. KL[e, bootstrap] = euclidean_distance(data_test, s_opt_1)
  1180. KL_r[e, bootstrap] = euclidean_distance(shuffled_1, s_opt_1)
  1181. KL_opt[e, bootstrap] = euclidean_distance(s_opt_2, s_opt_1)
  1182. elif metric == 'epairs':
  1183. KL[e, bootstrap] = epairs_metric(data_test, s_opt_1)
  1184. KL_r[e, bootstrap] = epairs_metric(shuffled_1, s_opt_1)
  1185. KL_opt[e, bootstrap] = epairs_metric(s_opt_2, s_opt_1)
  1186. if rnd_model == 'gaussian (spherical)':
  1187. m = np.mean(np.diag(np.cov(data_train.T)))
  1188. shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
  1189. np.diag([m, m, m]),
  1190. data_train.shape[0])
  1191. if metric == 'KL divergance estimate':
  1192. KL[e, bootstrap] = 0.5 * (KLdivergence(data_test, s_opt_1) + KLdivergence(s_opt_1, data_test))
  1193. KL_r[e, bootstrap] = 0.5 * (
  1194. KLdivergence(shuffled_1, s_opt_1) + KLdivergence(s_opt_1, shuffled_1))
  1195. KL_opt[e, bootstrap] = 0.5 * (KLdivergence(s_opt_2, s_opt_1) + KLdivergence(s_opt_1, s_opt_2))
  1196. elif metric == 'euclidean distance':
  1197. KL[e, bootstrap] = euclidean_distance(data_test, s_opt_1)
  1198. KL_r[e, bootstrap] = euclidean_distance(shuffled_1, s_opt_1)
  1199. KL_opt[e, bootstrap] = euclidean_distance(s_opt_2, s_opt_1)
  1200. elif metric == 'epairs':
  1201. KL[e, bootstrap] = epairs_metric(data_test, s_opt_1)
  1202. KL_r[e, bootstrap] = epairs_metric(shuffled_1, s_opt_1)
  1203. KL_opt[e, bootstrap] = epairs_metric(s_opt_2, s_opt_1)
  1204. if relative_dist:
  1205. KL_opt_avg = np.mean(KL_opt, keepdims=True, axis=-1)
  1206. KL -= KL_opt_avg
  1207. KL_r -= KL_opt_avg
  1208. KL_opt -= KL_opt_avg
  1209. KL_r_avg = np.mean(KL_r, keepdims=True, axis=-1)
  1210. KL /= KL_r_avg
  1211. KL_r /= KL_r_avg
  1212. KL_opt /= KL_r_avg
  1213. p = 1 * (np.sum(KL_r <= np.mean(KL, axis=-1, keepdims=True), axis=-1) / n_bootstraps)
  1214. if bon_correction:
  1215. p = p * n_epochs
  1216. p_vals.append(p)
  1217. epoch = np.sort(list(range(1, n_epochs + 1)) * n_bootstraps)
  1218. epoch_labels = np.concatenate([epoch, epoch, epoch])
  1219. dist_label = np.concatenate([['observed'] * (n_bootstraps * n_epochs),
  1220. [rnd_model] * (n_bootstraps * n_epochs),
  1221. ['structured'] * (n_bootstraps * n_epochs)])
  1222. KL_df = np.reshape(KL, KL.shape[0] * KL.shape[1])
  1223. KL_r_df = np.reshape(KL_r, KL_r.shape[0] * KL_r.shape[1])
  1224. KL_opt_df = np.reshape(KL_opt, KL_opt.shape[0] * KL_opt.shape[1])
  1225. KL_all_df = np.concatenate([KL_df, KL_r_df, KL_opt_df])
  1226. df = pd.DataFrame(np.array([KL_all_df, epoch_labels, dist_label]).T,
  1227. columns=[metric, 'learning epoch', 'distribution'])
  1228. df[metric] = df[metric].astype(float)
  1229. df['model'] = model_names[model_n]
  1230. dfs.append(df)
  1231. df_all = pd.concat(dfs)
  1232. df_all['divergence from'] = 'structured selectivity'
  1233. print(p_vals)
  1234. return df_all, p_vals, [KL, KL_r, KL_opt]
  1235. def fit_lines(data1, data2, if_xgen=False):
  1236. r_best_n, r_opt_n = [], []
  1237. for side in range(2):
  1238. if side == 0:
  1239. data_train = data1
  1240. data_test = data2
  1241. elif side == 1:
  1242. data_train = data2
  1243. data_test = data1
  1244. N = data_train.shape[0]
  1245. # colour as a regressor
  1246. A1 = np.ones((N, 2)) # With intercept/constant
  1247. A1[:, 1] = data_train[:, 0] # Use color selectivities as the regressor
  1248. x_shape_c, _, _, _ = np.linalg.lstsq(A1, data_train[:, 1], rcond=-1) # Fit to the shape selectivies
  1249. x_interaction_c, _, _, _ = np.linalg.lstsq(A1, data_train[:, 2], rcond=-1) # Fit to the interaction selectivies
  1250. A1[:, 1] = data_test[:, 0] # Use color selectivities as the regressor
  1251. pred_best1 = np.array([A1[:, 1] * x_shape_c[1], A1[:, 1] * x_interaction_c[1]]).T
  1252. r_best1 = r2_score(data_test[:, 1:], pred_best1)
  1253. # shape as a regressor
  1254. A2 = np.ones((N, 2)) # With intercept/constant
  1255. A2[:, 1] = data_train[:, 1] # Use shape selectivities as the regressor
  1256. x_colour_s, _, _, _ = np.linalg.lstsq(A2, data_train[:, 0], rcond=-1) # Fit to the colour selectivies
  1257. x_interaction_s, _, _, _ = np.linalg.lstsq(A2, data_train[:, 2], rcond=-1) # Fit to the interaction selectivies
  1258. A2[:, 1] = data_test[:, 1] # Use color selectivities as the regressor
  1259. pred_best2 = np.array([A2[:, 1] * x_colour_s[1], A2[:, 1] * x_interaction_s[1]]).T
  1260. data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, -1][:, None]], axis=1)
  1261. r_best2 = r2_score(data_compare, pred_best2)
  1262. # interaction as a regressor
  1263. A3 = np.ones((N, 2)) # With intercept/constant
  1264. A3[:, 1] = data_train[:, 2] # Use interaction selectivities as the regressor
  1265. x_colour_int, _, _, _ = np.linalg.lstsq(A3, data_train[:, 0], rcond=-1) # Fit to the colour selectivies
  1266. x_shape_int, _, _, _ = np.linalg.lstsq(A3, data_train[:, 1], rcond=-1) # Fit to the shape selectivies
  1267. A3[:, 1] = data_test[:, 2] # Use interaction selectivities as the regressor
  1268. pred_best3 = np.array([A3[:, 1] * x_colour_int[1], A3[:, 1] * x_shape_int[1]]).T
  1269. data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, 1][:, None]], axis=1)
  1270. r_best3 = r2_score(data_compare, pred_best3)
  1271. if if_xgen:
  1272. pred_opt1 = np.zeros((data_train.shape[0], 2))
  1273. pred_opt1[:, 0] = data_train[:, 0]
  1274. pred_opt1[:, 1] = -2 * data_train[:, 0]
  1275. r_opt1 = r2_score(data_test[:, 1:], pred_opt1)
  1276. else:
  1277. pred_opt1 = np.zeros((data_test.shape[0], 2))
  1278. pred_opt1[:, 0] = data_test[:, 0]
  1279. pred_opt1[:, 1] = -2 * data_test[:, 0]
  1280. r_opt1 = r2_score(data_test[:, 1:], pred_opt1)
  1281. if if_xgen:
  1282. pred_opt2 = np.zeros((data_train.shape[0], 2))
  1283. pred_opt2[:, 0] = data_train[:, 1]
  1284. pred_opt2[:, 1] = -2 * data_train[:, 1]
  1285. data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, -1][:, None]], axis=1)
  1286. r_opt2 = r2_score(data_compare, pred_opt2)
  1287. else:
  1288. pred_opt2 = np.zeros((data_test.shape[0], 2))
  1289. pred_opt2[:, 0] = data_test[:, 1]
  1290. pred_opt2[:, 1] = -2 * data_test[:, 1]
  1291. data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, -1][:, None]], axis=1)
  1292. r_opt2 = r2_score(data_compare, pred_opt2)
  1293. if if_xgen:
  1294. pred_opt3 = np.zeros((data_train.shape[0], 2))
  1295. pred_opt3[:, 0] = data_train[:, -1] / -2
  1296. pred_opt3[:, 1] = data_train[:, -1] / -2
  1297. data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, 1][:, None]], axis=1)
  1298. r_opt3 = r2_score(data_compare, pred_opt3)
  1299. else:
  1300. pred_opt3 = np.zeros((data_test.shape[0], 2))
  1301. pred_opt3[:, 0] = data_test[:, -1] / -2
  1302. pred_opt3[:, 1] = data_test[:, -1] / -2
  1303. data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, 1][:, None]], axis=1)
  1304. r_opt3 = r2_score(data_compare, pred_opt3)
  1305. r_best_n.append(np.mean([r_best1, r_best2, r_best3]))
  1306. r_opt_n.append(np.mean([r_opt1, r_opt2, r_opt3]))
  1307. r_best = np.mean(r_best_n)
  1308. r_opt = np.mean(r_opt_n)
  1309. return r_best, r_opt
  1310. def r_squared_pop(epochs, rnd_model='gaussian (spherical)', n_bootstraps=1000,
  1311. model_names=['cue + shape', 'cue + width'], break_axis=1.5, if_xgen=False):
  1312. dfs = []
  1313. n_epochs = len(epochs[0])
  1314. for model_n in range(len(model_names)):
  1315. fits_data = np.zeros((n_epochs, 2))
  1316. fits_opt = np.zeros((n_bootstraps, n_epochs, 2))
  1317. fits_rnd = np.zeros((n_bootstraps, n_epochs, 2))
  1318. for e in range(n_epochs):
  1319. s_train = epochs[model_n][e][:, :, 0]
  1320. s_test = epochs[model_n][e][:, :, 1]
  1321. # remove mean across neurons so we don't have to fit an intercept
  1322. s_train -= np.mean(s_train, axis=0, keepdims=True) # length 3 vector
  1323. s_test -= np.mean(s_test, axis=0, keepdims=True) # length 3 vector
  1324. fits_data[e, :] = fit_lines(s_train, s_test, if_xgen=if_xgen)
  1325. for bootstrap in tqdm(range(n_bootstraps)):
  1326. opt_cov = np.eye(3)
  1327. m = np.mean(np.diag(np.cov(s_train.T)))
  1328. opt_cov[:2, :2] = 0.5 * m
  1329. opt_cov[2, 2] = 2 * m
  1330. opt_cov[2, :2] = -m
  1331. opt_cov[:2, 2] = -m
  1332. opt_cov = np.eye(3)
  1333. m = np.mean(np.diag(np.cov(s_test.T)))
  1334. opt_cov[:2, :2] = 0.5 * m
  1335. opt_cov[2, 2] = 2 * m
  1336. opt_cov[2, :2] = -m
  1337. opt_cov[:2, 2] = -m
  1338. s_train_opt = np.random.multivariate_normal(np.zeros(s_train.shape[1]),
  1339. opt_cov,
  1340. s_train.shape[0])
  1341. s_test_opt = np.random.multivariate_normal(np.zeros(s_test.shape[1]),
  1342. opt_cov,
  1343. s_test.shape[0])
  1344. fits_opt[bootstrap, e, :] = fit_lines(s_train_opt, s_test_opt)
  1345. fits_opt[bootstrap, e, 0] += 0.05
  1346. fits_opt[bootstrap, e, 1] -= 0.05
  1347. if rnd_model == 'gaussian (spherical)':
  1348. m = np.mean(np.diag(np.cov(s_train.T)))
  1349. s_train_rnd = np.random.multivariate_normal(np.zeros(s_train.shape[1]),
  1350. np.diag([m, m, m]),
  1351. s_train.shape[0])
  1352. m = np.mean(np.diag(np.cov(s_test.T)))
  1353. s_test_rnd = np.random.multivariate_normal(np.zeros(s_test.shape[1]),
  1354. np.diag([m, m, m]),
  1355. s_test.shape[0])
  1356. elif rnd_model == 'gaussian (tied)':
  1357. s_train_rnd = np.random.multivariate_normal(np.zeros(s_train.shape[1]),
  1358. np.diag(np.diag(np.cov(s_train.T))),
  1359. s_train.shape[0])
  1360. s_test_rnd = np.random.multivariate_normal(np.zeros(s_test.shape[1]),
  1361. np.diag(np.diag(np.cov(s_test.T))),
  1362. s_test.shape[0])
  1363. fits_rnd[bootstrap, e, :] = fit_lines(s_train_rnd, s_test_rnd)
  1364. if break_axis:
  1365. fits_rnd[bootstrap, e, 1] += break_axis
  1366. epoch_data = list(range(1, n_epochs + 1)) * 2
  1367. epoch = np.concatenate([np.sort(list(range(1, n_epochs + 1)) * n_bootstraps)] * 2)
  1368. epoch_labels = np.concatenate([epoch_data, epoch, epoch])
  1369. gen_model = np.concatenate([['data'] * (n_epochs * 2),
  1370. ['gaussian'] * (n_epochs * 2 * n_bootstraps),
  1371. ['optimal'] * (n_epochs * 2 * n_bootstraps),
  1372. ])
  1373. dist_label = np.concatenate([['best fit line'] * n_epochs,
  1374. ['optimal XOR line'] * n_epochs,
  1375. ['best fit line'] * (n_epochs * n_bootstraps),
  1376. ['optimal XOR line'] * (n_epochs * n_bootstraps),
  1377. ['best fit line'] * (n_epochs * n_bootstraps),
  1378. ['optimal XOR line'] * (n_epochs * n_bootstraps),
  1379. ])
  1380. data_df = np.reshape(fits_data, fits_data.shape[0] * fits_data.shape[1], order='F')
  1381. rnd_df = np.reshape(fits_rnd, fits_rnd.shape[0] * fits_rnd.shape[1] * fits_rnd.shape[2], order='F')
  1382. opt_df = np.reshape(fits_opt, fits_opt.shape[0] * fits_opt.shape[1] * fits_opt.shape[2], order='F')
  1383. all_df = np.concatenate([data_df, rnd_df, opt_df])
  1384. df = pd.DataFrame(np.array([all_df, epoch_labels, dist_label, gen_model]).T,
  1385. columns=['r squared', 'learning epoch', 'fitted model', 'generative model'])
  1386. df['r squared'] = df['r squared'].astype(float)
  1387. df['model'] = model_names[model_n]
  1388. dfs.append(df)
  1389. return pd.concat(dfs)
  1390. def xor_sim_noise(sig_min, sig_max, sig_n, n_neurons=400):
  1391. s_list = np.linspace(sig_min, sig_max, sig_n)
  1392. s_tot = s_list.shape[0]
  1393. N = n_neurons
  1394. n_coeffs = 3
  1395. conds = np.ones((4, n_coeffs))
  1396. conds[0, :2] = -1
  1397. conds[1, [0, 2]] = -1
  1398. conds[2, 1:] = -1
  1399. targets = [1, -1, -1, 1] # XOR
  1400. clf1 = linear_model.LogisticRegression() # 1 is no reg
  1401. n = 100
  1402. perf = np.zeros((s_tot, n, 2))
  1403. print('XOr decoding performance as a function of noice (sigma) - running simulation ')
  1404. for s_c, sig in tqdm(enumerate(s_list)):
  1405. for i in range(n):
  1406. # Random selectivity
  1407. cov = np.eye(3)
  1408. s = np.random.multivariate_normal(np.zeros(3), cov, N)
  1409. r = s @ conds.T
  1410. r_train = r.T + sig * np.random.normal(0, 1, (4, N))
  1411. r_train = r_train - np.mean(r_train, axis=1, keepdims=True)
  1412. clf1.fit(r_train, targets)
  1413. r_test = r.T + sig * np.random.normal(0, 1, (4, N))
  1414. r_test = r_test - np.mean(r_test, axis=1, keepdims=True)
  1415. perf[s_c, i, 0] = clf1.score(r_test, targets)
  1416. # Structured selectivity
  1417. opt_cov = np.diag([0, 0, 3])
  1418. s = np.random.multivariate_normal(np.zeros(3), opt_cov, N)
  1419. r = s @ conds.T
  1420. r_train = r.T + sig * np.random.normal(0, 1, (4, N))
  1421. r_train = r_train - np.mean(r_train, axis=1, keepdims=True)
  1422. clf1.fit(r_train, targets)
  1423. r_test = r.T + sig * np.random.normal(0, 1, (4, N))
  1424. r_test = r_test - np.mean(r_test, axis=1, keepdims=True)
  1425. perf[s_c, i, 1] = clf1.score(r_test, targets)
  1426. return perf, s_list
  1427. def xor_sim_units(n_neurons, step_neurons, noise_sig):
  1428. N_list = np.arange(1, n_neurons, step_neurons)
  1429. N_tot = N_list.shape[0]
  1430. conds = np.ones((4, 3))
  1431. conds[0, :2] = -1
  1432. conds[1, [0, 2]] = -1
  1433. conds[2, 1:] = -1
  1434. targets = [1, -1, -1, 1] # XOR
  1435. clf1 = linear_model.LogisticRegression() # 1 is no reg
  1436. sig = noise_sig
  1437. n = 100
  1438. perf_units = np.zeros((N_tot, n, 2))
  1439. for N_c, N in enumerate(N_list):
  1440. print(N, end=',')
  1441. for i in range(n):
  1442. s = np.random.normal(0, 1, (N, 3))
  1443. r = s @ conds.T
  1444. clf1.fit(r.T + sig * np.random.normal(0, 1, (4, N)), targets)
  1445. perf_units[N_c, i, 0] = clf1.score(r.T + sig * np.random.normal(0, 1, (4, N)), targets)
  1446. # Structured selectivity
  1447. opt_cov = np.diag([0, 0, 3])
  1448. s = np.random.multivariate_normal(np.zeros(3), opt_cov, N)
  1449. r = s @ conds.T
  1450. clf1.fit(r.T + sig * np.random.normal(0, 1, (4, N)), targets)
  1451. perf_units[N_c, i, 1] = clf1.score(r.T + sig * np.random.normal(0, 1, (4, N)), targets)
  1452. return perf_units, N_list
  1453. def generate_boundary(sel, sig=0.9, n=100):
  1454. conds = np.ones((4, 3))
  1455. conds[0, :2] = -1
  1456. conds[1, [0, 2]] = -1
  1457. conds[2, 1:] = -1
  1458. labels = [1, -1, -1, 1] # XOR
  1459. neurons = sel[:2, :] # take the first two neurons
  1460. rates = neurons @ conds.T
  1461. rates = np.concatenate([rates] * n, axis=-1)
  1462. rates = rates + np.random.normal(0, sig, (2, 4 * n))
  1463. rates -= np.mean(rates, axis=1, keepdims=True)
  1464. labels = np.concatenate([labels] * n, axis=0)
  1465. clf = LDA()
  1466. clf.fit(rates.T, labels)
  1467. w = clf.coef_[0]
  1468. a = -w[0] / w[1]
  1469. xx = np.linspace(-10, 10)
  1470. yy = a * xx - (clf.intercept_[0]) / w[1]
  1471. return rates[:, :4], xx, yy
  1472. def KL_optimal_2(epochs, n_bootstraps=1000, rnd_model='gaussian (tied)', model_names=['cue + shape', 'cue + width'],
  1473. bon_correction=False):
  1474. n_epochs = len(epochs[0])
  1475. dfs = []
  1476. p_vals = []
  1477. for model_n in range(len(model_names)):
  1478. KL = np.zeros((n_epochs, n_bootstraps))
  1479. KL_r = np.zeros((n_epochs, n_bootstraps))
  1480. for e in range(n_epochs):
  1481. for bootstrap in tqdm(range(n_bootstraps)):
  1482. data_train = epochs[model_n][e][:, :, 0]
  1483. data_test = epochs[model_n][e][:, :, 1]
  1484. s_opt = np.zeros(data_test.shape)
  1485. s_opt[:, 0] = data_test[:, 0]
  1486. s_opt[:, 1] = data_test[:, 0]
  1487. s_opt[:, 2] = data_test[:, 0] * -2
  1488. if rnd_model == 'gaussian (tied)':
  1489. shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
  1490. np.diag(np.diag(np.cov(data_train.T))),
  1491. data_train.shape[0])
  1492. KL[e, bootstrap] = 0.5 * (KLdivergence(data_test, s_opt) + KLdivergence(s_opt, data_test))
  1493. KL_r[e, bootstrap] = 0.5 * (KLdivergence(shuffled_1, s_opt) + KLdivergence(s_opt, shuffled_1))
  1494. if rnd_model == 'gaussian (spherical)':
  1495. m = np.mean(np.diag(np.cov(data_train.T)))
  1496. shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
  1497. np.diag([m, m, m]),
  1498. data_train.shape[0])
  1499. KL[e, bootstrap] = 0.5 * (KLdivergence(data_train, s_opt) + KLdivergence(s_opt, data_train))
  1500. KL_r[e, bootstrap] = 0.5 * (KLdivergence(shuffled_1, s_opt) + KLdivergence(s_opt, shuffled_1))
  1501. p = 2 * (np.sum(KL_r <= np.mean(KL, axis=-1, keepdims=True), axis=-1) / n_bootstraps)
  1502. if bon_correction:
  1503. p = p * n_epochs
  1504. p_vals.append(p)
  1505. epoch = np.sort(list(range(1, n_epochs + 1)) * n_bootstraps)
  1506. epoch_labels = np.concatenate([epoch, epoch])
  1507. dist_label = np.concatenate([['observed'] * (n_bootstraps * n_epochs), [rnd_model] * (n_bootstraps * n_epochs)])
  1508. KL_df = np.reshape(KL, KL.shape[0] * KL.shape[1])
  1509. KL_r_df = np.reshape(KL_r, KL_r.shape[0] * KL_r.shape[1])
  1510. KL_all_df = np.concatenate([KL_df, KL_r_df])
  1511. df = pd.DataFrame(np.array([KL_all_df, epoch_labels, dist_label]).T,
  1512. columns=['KL divergance estimate', 'learning epoch', 'distribution'])
  1513. df['KL divergance estimate'] = df['KL divergance estimate'].astype(float)
  1514. df['model'] = model_names[model_n]
  1515. dfs.append(df)
  1516. df_all = pd.concat(dfs)
  1517. df_all['divergence from'] = 'optimal selectivity'
  1518. print(p_vals)
  1519. return df_all, p_vals
  1520. def plot_pvals(df_data, p, ax, offset=0.0001, tail=1, n_bootstraps=1000, metric='KL divergance estimate'):
  1521. y_vals = np.resize(df_data[df_data["distribution"] == 'observed'][metric].values,
  1522. (len(p), n_bootstraps))
  1523. y_avgs = y_vals.mean(-1)
  1524. y_stds = y_vals.std(-1)
  1525. if tail == 1:
  1526. y_pos = y_avgs + y_stds + offset
  1527. elif tail == -1:
  1528. y_pos = y_avgs - y_stds - (offset * 12)
  1529. for e in range(len(p)):
  1530. if p[e] > 0.05:
  1531. star = 'ns'
  1532. size = 10
  1533. elif (p[e] <= 0.05) & (p[e] > 0.01):
  1534. star = '*'
  1535. size = 14
  1536. elif (p[e] <= 0.01) & (p[e] > 0.001):
  1537. star = '**'
  1538. size = 14
  1539. elif p[e] <= 0.001:
  1540. star = '***'
  1541. size = 14
  1542. ax.text(e, y_pos[e], star, ha='center', size=size)
  1543. return
  1544. def get_betas_cross_val_2(data, labels, condition_labels, normalisation='zscore', time_window=[140, 150],
  1545. cross_val=False, if_xgen=False, task='task_1', shuffle=False):
  1546. n_parts = len(data)
  1547. betas_part = []
  1548. for p in range(n_parts):
  1549. betas = []
  1550. for s in range(len(data[p])):
  1551. firing_rates_ses = data[p][s]
  1552. firing_rates_ses = firing_rates_ses[:, :, time_window[0]:time_window[1]]
  1553. labels_ses = labels[p][s]
  1554. if task == 'task_1':
  1555. firing_rates_ses = firing_rates_ses[labels_ses < 9, :, :]
  1556. labels_ses = labels_ses[labels_ses < 9]
  1557. elif task == 'task_2':
  1558. firing_rates_ses = firing_rates_ses[labels_ses > 8, :, :]
  1559. labels_ses = labels_ses[labels_ses > 8]
  1560. elif task == 'context_1':
  1561. fac_cxt1 = np.array([1, 2, 3, 4, 9, 10, 11, 12])
  1562. idc_cxt1 = np.isin(labels_ses, fac_cxt1)
  1563. firing_rates_ses = firing_rates_ses[idc_cxt1, :, :]
  1564. labels_ses = labels_ses[idc_cxt1]
  1565. elif task == 'context_2':
  1566. fac_cxt2 = np.array([5, 6, 7, 8, 13, 14, 15, 16])
  1567. idc_cxt2 = np.isin(labels_ses, fac_cxt2)
  1568. firing_rates_ses = firing_rates_ses[idc_cxt2, :, :]
  1569. labels_ses = labels_ses[idc_cxt2]
  1570. if cross_val:
  1571. n_trls = firing_rates_ses.shape[0]
  1572. idc_0 = np.zeros(n_trls // 2)
  1573. idc_1 = np.ones(n_trls // 2)
  1574. if (n_trls % 2) > 0:
  1575. idc_1 = np.concatenate([idc_1, [1.0]])
  1576. idc_rnd = np.concatenate([idc_0, idc_1])
  1577. random.shuffle(idc_rnd)
  1578. firing_rates_ses1 = firing_rates_ses[idc_rnd < 1, :, :]
  1579. firing_rates_ses2 = firing_rates_ses[idc_rnd > 0, :, :]
  1580. labels_ses1 = labels_ses[idc_rnd < 1]
  1581. labels_ses2 = labels_ses[idc_rnd > 0]
  1582. else:
  1583. firing_rates_ses1 = firing_rates_ses
  1584. firing_rates_ses2 = firing_rates_ses
  1585. labels_ses1 = labels_ses
  1586. labels_ses2 = labels_ses
  1587. if if_xgen:
  1588. firing_rates_ses1 = firing_rates_ses[labels_ses1 < 9, :, :]
  1589. firing_rates_ses2 = firing_rates_ses[labels_ses2 > 8, :, :]
  1590. labels_ses1 = labels_ses1[labels_ses1 < 9]
  1591. labels_ses2 = labels_ses2[labels_ses2 > 8]
  1592. firing_rates_ses = [firing_rates_ses1, firing_rates_ses2]
  1593. labels_ses = [labels_ses1, labels_ses2]
  1594. betas_split = []
  1595. n_splits = len(labels_ses)
  1596. for split in range(n_splits):
  1597. firing_mean = np.mean(firing_rates_ses[split], axis=-1, keepdims=True)
  1598. labels_ses_split = labels_ses[split]
  1599. if shuffle:
  1600. random.shuffle(labels_ses_split)
  1601. if normalisation == 'zscore':
  1602. firing_mean = (firing_mean - np.mean(firing_mean, axis=0, keepdims=True)) / (
  1603. np.std(firing_mean, axis=0,
  1604. keepdims=True))
  1605. elif normalisation == 'mean_centred':
  1606. firing_mean = firing_mean - np.mean(firing_mean, axis=0, keepdims=True)
  1607. elif normalisation == 'none':
  1608. firing_mean = firing_mean
  1609. elif normalisation == 'soft':
  1610. firing_mean = firing_mean / (
  1611. (np.max(firing_mean, axis=0, keepdims=True) - np.min(firing_mean, axis=0,
  1612. keepdims=True)) + 5)
  1613. labels_cue = np.array(assign_lables(labels_ses_split, condition_labels[0]))
  1614. labels_target = np.array(assign_lables(labels_ses_split, condition_labels[1]))
  1615. labels_int1 = np.array(assign_lables(labels_ses_split, condition_labels[2]))
  1616. labels_target2 = np.array(assign_lables(labels_ses_split, condition_labels[3]))
  1617. labels_int2 = np.array(assign_lables(labels_ses_split, condition_labels[4]))
  1618. constant = np.ones_like(labels_cue).astype(float)
  1619. design_matrix = np.vstack(
  1620. [constant, labels_cue, labels_target, labels_int1, labels_target2, labels_int2]).T
  1621. if task == 'all':
  1622. labels_context = np.array(assign_lables(labels_ses_split, condition_labels[0]))
  1623. labels_shape = np.array(assign_lables(labels_ses_split, condition_labels[1]))
  1624. labels_rew = np.array(assign_lables(labels_ses_split, condition_labels[2]))
  1625. labels_task = np.array(assign_lables(labels_ses_split, condition_labels[3]))
  1626. labels_width = np.array(assign_lables(labels_ses_split, condition_labels[4]))
  1627. labels_int_irr = np.array(assign_lables(labels_ses_split, condition_labels[5]))
  1628. constant = np.ones_like(labels_cue).astype(float)
  1629. design_matrix = np.vstack(
  1630. [constant, labels_context, labels_shape, labels_rew, labels_task, labels_width,
  1631. labels_int_irr]).T
  1632. design_matrix = design_matrix.astype(float)
  1633. n_neurons = firing_mean.shape[1]
  1634. n_times = firing_mean.shape[-1]
  1635. n_models = design_matrix.shape[1]
  1636. design_matrix_stacked = np.concatenate([design_matrix] * n_times, axis=0)
  1637. firing_rates_ses_stacked = np.concatenate(np.array_split(firing_mean, n_times, axis=-1), axis=0)
  1638. betas_ses = np.zeros((n_neurons, n_models, 1))
  1639. for cell in range(n_neurons):
  1640. betas_ses[cell, :, :] = \
  1641. sp.linalg.lstsq(design_matrix_stacked, firing_rates_ses_stacked[:, cell, :])[0][:, :]
  1642. betas_split.append(betas_ses)
  1643. betas.append(np.array(betas_split)[:, :, :, 0])
  1644. betas_part.append(betas)
  1645. if task == 'all':
  1646. epochs_rel = []
  1647. for p in range(n_parts):
  1648. epoch = np.concatenate(betas_part[p], axis=1)[:, :, 1:4].transpose((0, 2, 1)) # splits x coeffs x neurons
  1649. epoch -= np.mean(epoch, axis=2, keepdims=True)
  1650. epochs_rel.append(epoch.T)
  1651. epochs_irrel = []
  1652. for p in range(n_parts):
  1653. epoch = np.concatenate(betas_part[p], axis=1)[:, :, 4:].transpose((0, 2, 1)) # splits x coeffs x neurons
  1654. epoch -= np.mean(epoch, axis=2, keepdims=True)
  1655. epochs_irrel.append(epoch.T)
  1656. else:
  1657. epochs_rel = []
  1658. for p in range(n_parts):
  1659. epoch = np.concatenate(betas_part[p], axis=1)[:, :, 1:4].transpose((0, 2, 1)) # splits x coeffs x neurons
  1660. epoch -= np.mean(epoch, axis=2, keepdims=True)
  1661. epochs_rel.append(epoch.T)
  1662. epochs_irrel = []
  1663. for p in range(n_parts):
  1664. epoch_1 = np.concatenate(betas_part[p], axis=1)[:, :, 1][:, :, None].transpose(
  1665. (0, 2, 1)) # splits x coeffs x neurons
  1666. epoch_2 = np.concatenate(betas_part[p], axis=1)[:, :, 4:].transpose((0, 2, 1)) # splits x coeffs x neurons
  1667. epoch = np.concatenate([epoch_1, epoch_2], axis=1)
  1668. epoch -= np.mean(epoch, axis=2, keepdims=True)
  1669. epochs_irrel.append(epoch.T)
  1670. return epochs_rel, epochs_irrel
  1671. def get_betas_cross_val_3(data, labels, condition_labels, normalisation='zscore', time_window=[140, 150],
  1672. add_constant=True, n_splits=10):
  1673. n_parts = len(data)
  1674. betas_part = []
  1675. for p in range(n_parts):
  1676. betas_ses_all = []
  1677. for s in tqdm(range(len(data[p]))):
  1678. firing_rates_ses = data[p][s]
  1679. firing_rates_ses = firing_rates_ses[:, :, time_window[0]:time_window[1]]
  1680. labels_ses = labels[p][s]
  1681. n_trl = firing_rates_ses.shape[0]
  1682. firing_rates_ses_all = []
  1683. labels_ses_all = []
  1684. for split in range(n_splits):
  1685. idc_1 = np.zeros(n_trl // 2)
  1686. idc_2 = np.ones(n_trl // 2)
  1687. if (n_trl % 2) > 0:
  1688. idc_2 = np.append(idc_2, [1.0])
  1689. idc_rnd = np.concatenate([idc_1, idc_2])
  1690. random.shuffle(idc_rnd)
  1691. firing_rates_ses1 = firing_rates_ses[idc_rnd < 1.0, :, :]
  1692. firing_rates_ses2 = firing_rates_ses[idc_rnd > 0.0, :, :]
  1693. labels_ses1 = labels_ses[idc_rnd < 1]
  1694. labels_ses2 = labels_ses[idc_rnd > 0]
  1695. firing_rates_ses_all.append([firing_rates_ses1, firing_rates_ses2])
  1696. labels_ses_all.append([labels_ses1, labels_ses2])
  1697. betas_split = []
  1698. for split in range(n_splits):
  1699. betas_run = []
  1700. for run in range(2):
  1701. firing_mean = np.mean(firing_rates_ses_all[split][run], axis=-1, keepdims=True)
  1702. labels_ses_split = labels_ses_all[split][run]
  1703. if normalisation == 'zscore':
  1704. firing_mean = (firing_mean - np.mean(firing_mean, axis=0, keepdims=True)) / (
  1705. np.std(firing_mean, axis=0,
  1706. keepdims=True) + 1)
  1707. elif normalisation == 'mean_centred':
  1708. firing_mean = firing_mean - np.mean(firing_mean, axis=0, keepdims=True)
  1709. elif normalisation == 'none':
  1710. firing_mean = firing_mean
  1711. elif normalisation == 'soft':
  1712. firing_mean = firing_mean / (
  1713. (np.max(firing_mean, axis=0, keepdims=True) - np.min(firing_mean, axis=0,
  1714. keepdims=True)) + 5)
  1715. labels_cue = np.array(assign_lables(labels_ses_split, condition_labels[0]))
  1716. labels_target = np.array(assign_lables(labels_ses_split, condition_labels[1]))
  1717. labels_int1 = np.array(assign_lables(labels_ses_split, condition_labels[2]))
  1718. labels_cue2 = np.array(assign_lables(labels_ses_split, condition_labels[3]))
  1719. labels_target2 = np.array(assign_lables(labels_ses_split, condition_labels[4]))
  1720. labels_int2 = np.array(assign_lables(labels_ses_split, condition_labels[5]))
  1721. if add_constant:
  1722. constant = np.ones_like(labels_cue).astype(float)
  1723. design_matrix = np.vstack(
  1724. [constant, labels_cue, labels_target, labels_int1, labels_cue2, labels_target2,
  1725. labels_int2]).T
  1726. else:
  1727. design_matrix = np.vstack(
  1728. [labels_cue, labels_cue, labels_target, labels_int1, labels_cue2, labels_target2,
  1729. labels_int2]).T
  1730. design_matrix = design_matrix.astype(float)
  1731. n_neurons = firing_mean.shape[1]
  1732. n_times = firing_mean.shape[-1]
  1733. n_models = design_matrix.shape[1]
  1734. design_matrix_stacked = np.concatenate([design_matrix] * n_times, axis=0)
  1735. firing_rates_ses_stacked = np.concatenate(np.array_split(firing_mean, n_times, axis=-1), axis=0)
  1736. betas_ses = np.zeros((n_neurons, n_models, 1))
  1737. for cell in range(n_neurons):
  1738. betas_ses[cell, :, :] = \
  1739. sp.linalg.lstsq(design_matrix_stacked, firing_rates_ses_stacked[:, cell, :])[0][:, :]
  1740. betas_run.append(betas_ses)
  1741. betas_split.append(np.mean(np.array(betas_run), axis=0)[:, :, 0])
  1742. betas_ses_all.append(np.mean(np.array(betas_split), axis=0))
  1743. betas_part.append(np.concatenate(betas_ses_all, axis=0))
  1744. epochs_rel = []
  1745. for p in range(n_parts):
  1746. if add_constant:
  1747. epoch = betas_part[p][:, 1:4]
  1748. else:
  1749. epoch = betas_part[p][:, :3]
  1750. epoch = np.concatenate([epoch[:, :, None], epoch[:, :, None]], axis=-1)
  1751. epoch -= np.mean(epoch, axis=0, keepdims=True)
  1752. epochs_rel.append(epoch)
  1753. epochs_irrel = []
  1754. for p in range(n_parts):
  1755. if add_constant:
  1756. epoch = betas_part[p][:, 4:]
  1757. else:
  1758. epoch = betas_part[p][:, 3:]
  1759. epoch = np.concatenate([epoch[:, :, None], epoch[:, :, None]], axis=-1)
  1760. epoch -= np.mean(epoch, axis=0, keepdims=True)
  1761. epochs_irrel.append(epoch)
  1762. return epochs_rel, epochs_irrel
  1763. def plot_pvals_2(data, ax, n_bootstraps, tail=1, offset=0.05, bon_correction=False, side=1,
  1764. metric='KL divergance estimate'):
  1765. df_data = data[data['distribution'] == 'observed']
  1766. n_epochs = len(np.unique(df_data['learning epoch'].values))
  1767. rel = df_data[df_data['model'] == 'cue + shape'][metric].values
  1768. rel = np.resize(rel, (n_epochs, n_bootstraps))
  1769. rel_avg = np.mean(rel, axis=-1, keepdims=True)
  1770. irrel = df_data[df_data['model'] == 'cue + width'][metric].values
  1771. irrel = np.resize(irrel, (n_epochs, n_bootstraps))
  1772. if side == 1:
  1773. side_fac = 1
  1774. elif side == 2:
  1775. side_fac = 2
  1776. y_avgs = rel_avg[:, 0]
  1777. y_stds = rel.std(-1)
  1778. if tail == 1:
  1779. p = side_fac * (np.sum(irrel >= rel_avg, axis=-1) / n_bootstraps)
  1780. y_pos = y_avgs + y_stds + offset
  1781. elif tail == -1:
  1782. p = side_fac * (np.sum(irrel <= rel_avg, axis=-1) / n_bootstraps)
  1783. y_pos = y_avgs - y_stds - (offset * 12)
  1784. if bon_correction:
  1785. p = p * n_epochs
  1786. for e in range(n_epochs):
  1787. if p[e] > 0.05:
  1788. star = 'ns'
  1789. size = 12
  1790. elif (p[e] <= 0.05) & (p[e] > 0.01):
  1791. star = '*'
  1792. size = 20
  1793. elif (p[e] <= 0.01) & (p[e] > 0.001):
  1794. star = '**'
  1795. size = 20
  1796. elif p[e] <= 0.001:
  1797. star = '***'
  1798. size = 20
  1799. ax.text(e, y_pos[e], star, ha='center', size=size)
  1800. return
  1801. def plot_pvals_3(y, x, p, ax, col, offset=0.05):
  1802. if p > 0.05:
  1803. star = 'ns'
  1804. size = 10
  1805. elif (p <= 0.05) & (p > 0.01):
  1806. star = '*'
  1807. size = 15
  1808. elif (p <= 0.01) & (p > 0.001):
  1809. star = '**'
  1810. size = 15
  1811. elif p <= 0.001:
  1812. star = '***'
  1813. size = 15
  1814. ax.text(x, y + offset, star, ha='center', size=size, color=col)
  1815. return
  1816. def compare_r2s(data1, data2, n_bootstraps, tail=1, mode='diff'):
  1817. n_units_1 = data1.shape[0]
  1818. n_units_2 = data2.shape[0]
  1819. data_all = np.concatenate([data1, data2], axis=0)
  1820. fit_diff_rnd = np.zeros((n_bootstraps))
  1821. fit_1_obs = fit_lines(data1[:, :, 0], data1[:, :, 1], if_xgen=True)
  1822. fit_2_obs = fit_lines(data2[:, :, 0], data2[:, :, 1], if_xgen=True)
  1823. if mode == 'diff':
  1824. fit_diff1 = fit_1_obs[0] - fit_1_obs[1]
  1825. fit_diff2 = fit_2_obs[0] - fit_2_obs[1]
  1826. elif mode == 'best':
  1827. fit_diff1 = fit_1_obs[0]
  1828. fit_diff2 = fit_2_obs[0]
  1829. elif mode == 'optimal':
  1830. fit_diff1 = fit_1_obs[1]
  1831. fit_diff2 = fit_2_obs[1]
  1832. fit_diff = fit_diff1 - fit_diff2
  1833. for n_bootstrap in range(n_bootstraps):
  1834. rnd_idc1 = np.zeros(n_units_1)
  1835. rnd_idc2 = np.ones(n_units_2)
  1836. rnd_idx = np.concatenate([rnd_idc1, rnd_idc2])
  1837. random.shuffle(rnd_idx)
  1838. data1_rnd = data_all[rnd_idx < 1, :, :]
  1839. data2_rnd = data_all[rnd_idx > 0, :, :]
  1840. fit_1_rnd = fit_lines(data1_rnd[:, :, 0], data1_rnd[:, :, 1], if_xgen=True)
  1841. fit_2_rnd = fit_lines(data2_rnd[:, :, 0], data2_rnd[:, :, 1], if_xgen=True)
  1842. if mode == 'diff':
  1843. fit_diff1_rnd = fit_1_rnd[0] - fit_1_rnd[1]
  1844. fit_diff2_rnd = fit_2_rnd[0] - fit_2_rnd[1]
  1845. elif mode == 'best':
  1846. fit_diff1_rnd = fit_1_rnd[0]
  1847. fit_diff2_rnd = fit_2_rnd[0]
  1848. elif mode == 'optimal':
  1849. fit_diff1_rnd = fit_1_rnd[1]
  1850. fit_diff2_rnd = fit_2_rnd[1]
  1851. fit_diff_rnd[n_bootstrap] = fit_diff1_rnd - fit_diff2_rnd
  1852. if tail == -1:
  1853. p = (np.sum(fit_diff_rnd >= fit_diff) / n_bootstraps)
  1854. elif tail == 1:
  1855. p = (np.sum(fit_diff_rnd <= fit_diff) / n_bootstraps)
  1856. return p
  1857. def data_handle(data, labels, min_n, do_rand=False):
  1858. conditions = np.unique(labels)
  1859. conditions_ = []
  1860. for _ in range(len(conditions)):
  1861. condition_dat = data[labels == conditions[_], :, :]
  1862. if do_rand:
  1863. idx_re = np.random.choice(np.array(list(range(0, condition_dat.shape[0]))), min_n, replace=False)
  1864. conditions_.append(condition_dat[idx_re, :, :])
  1865. else:
  1866. conditions_.append(condition_dat[:min_n, :, :])
  1867. conditions_ = np.array(conditions_)
  1868. return np.concatenate(conditions_, axis=0)
  1869. def euclidean_distance(data1, data2):
  1870. return np.sqrt(np.sum((np.cov(data2.T) - np.cov(data1.T)) ** 2))
  1871. def prepare_data(data, labels, which_trl='beginning', set_min_trl=None):
  1872. n_trls_ses = []
  1873. for sess in range(len(data)):
  1874. n_trls_ses.append(data[sess].shape[0])
  1875. n_trls = np.min(n_trls_ses)
  1876. n_condi_min = []
  1877. for sess in range(len(data)):
  1878. if which_trl == 'beginning':
  1879. freqs = get_freqs(labels[sess][:n_trls])
  1880. elif which_trl == 'end':
  1881. freqs = get_freqs(labels[sess][-n_trls:])
  1882. elif which_trl == 'rnd':
  1883. freqs = get_freqs(labels[sess])
  1884. elif which_trl == 'middle':
  1885. labels_sess = labels[sess]
  1886. idc_half = int(len(labels_sess) / 2)
  1887. # check whether even or odd number of trials
  1888. if (len(labels_sess) % 2) > 0:
  1889. labels_sess = labels_sess[idc_half - int(n_trls / 2):idc_half + int(n_trls / 2) + 1]
  1890. else:
  1891. labels_sess = labels_sess[idc_half - int(n_trls / 2):idc_half + int(n_trls / 2)]
  1892. freqs = get_freqs(labels_sess)
  1893. n_condi_min.append(freqs[min(freqs, key=freqs.get)])
  1894. n_trl_condi = np.min(n_condi_min)
  1895. if set_min_trl:
  1896. n_trl_condi = set_min_trl
  1897. dat_combined = []
  1898. for sess in range(len(data)):
  1899. if which_trl == 'beginning':
  1900. data_sess = data[sess][:n_trls, :, :]
  1901. labels_sess = labels[sess][:n_trls]
  1902. dat_combined.append(data_handle(data_sess, labels_sess, n_trl_condi))
  1903. elif which_trl == 'end':
  1904. data_sess = data[sess][-n_trls:, :, :]
  1905. labels_sess = labels[sess][-n_trls:]
  1906. dat_combined.append(data_handle(data_sess, labels_sess, n_trl_condi))
  1907. elif which_trl == 'rnd':
  1908. data_sess = data_handle(data[sess], labels[sess], n_trl_condi, do_rand=False)
  1909. dat_combined.append(data_sess)
  1910. elif which_trl == 'middle':
  1911. labels_sess = labels[sess]
  1912. idc_half = int(len(labels_sess) / 2)
  1913. # check whether even or odd number of trials
  1914. if (len(labels_sess) % 2) > 0:
  1915. labels_sess = labels_sess[idc_half - int(n_trls / 2):idc_half + int(n_trls / 2) + 1]
  1916. data_sess = data[sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2) + 1, :, :]
  1917. elif (len(labels_sess) % 2) == 0:
  1918. labels_sess = labels_sess[idc_half - int(n_trls / 2):idc_half + int(n_trls / 2)]
  1919. data_sess = data[sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2), :, :]
  1920. dat_combined.append(data_handle(data_sess, labels_sess, n_trl_condi))
  1921. dat_combined = np.concatenate(dat_combined, axis=1)
  1922. labels_combined = np.sort(list(range(0, len(np.unique(labels[0])))) * n_trl_condi)
  1923. return dat_combined, labels_combined
  1924. def get_shattering_ids(cue, target, width, reward):
  1925. n_condi = len(cue)
  1926. all_combos = list(itertools.combinations(list(range(n_condi)), int(n_condi / 2)))
  1927. n_combos = int(len(all_combos) / 2)
  1928. decoding_targets = np.zeros((n_combos, n_condi))
  1929. cue_tgt_width_reward_ids = np.zeros(int(n_condi / 2))
  1930. for i in range(n_combos):
  1931. decoding_targets[i, all_combos[i]] = 1
  1932. # Find cue ids
  1933. if np.sum(np.abs(decoding_targets[i, :] - cue)) == 0 or np.sum(np.abs(decoding_targets[i, :] - (1 - cue))) == 0:
  1934. cue_tgt_width_reward_ids[0] = i
  1935. # Find target ids
  1936. if np.sum(np.abs(decoding_targets[i, :] - target)) == 0 or np.sum(
  1937. np.abs(decoding_targets[i, :] - (1 - target))) == 0:
  1938. cue_tgt_width_reward_ids[1] = i
  1939. # Find width ids
  1940. if np.sum(np.abs(decoding_targets[i, :] - width)) == 0 or np.sum(
  1941. np.abs(decoding_targets[i, :] - (1 - width))) == 0:
  1942. cue_tgt_width_reward_ids[2] = i
  1943. # Find reward ids
  1944. if np.sum(np.abs(decoding_targets[i, :] - reward)) == 0 or np.sum(
  1945. np.abs(decoding_targets[i, :] - (1 - reward))) == 0:
  1946. cue_tgt_width_reward_ids[3] = i
  1947. cross_gen_decoding_train_ids = []
  1948. cross_gen_decoding_test_ids = []
  1949. for i in range(n_combos):
  1950. current_zeros = np.squeeze(np.where(decoding_targets[i, :] == 0))
  1951. current_combs1 = list(itertools.combinations(current_zeros.tolist(), 2))
  1952. current_ones = np.squeeze(np.where(decoding_targets[i, :] == 1))
  1953. current_combs2 = list(itertools.combinations(current_ones.tolist(), 2))
  1954. n_combs = len(current_combs1)
  1955. current_train = []
  1956. current_test = []
  1957. for j in range(n_combs):
  1958. for k in range(n_combs):
  1959. current_train.append([current_combs1[j], current_combs2[k]])
  1960. current_test_zeros = list(np.setdiff1d(current_zeros, np.array(current_combs1[j])))
  1961. current_test_ones = list(np.setdiff1d(current_ones, np.array(current_combs2[k])))
  1962. current_test.append([current_test_zeros, current_test_ones])
  1963. cross_gen_decoding_train_ids.append(current_train)
  1964. cross_gen_decoding_test_ids.append(current_test)
  1965. return decoding_targets, cue_tgt_width_reward_ids, cross_gen_decoding_train_ids, cross_gen_decoding_test_ids
  1966. def decode(X, y, method='svm', n_inter=5, return_inter=False, n_jobs=None):
  1967. y = np.array(y)
  1968. scores_all = []
  1969. for n in range(n_inter):
  1970. if method == 'svm':
  1971. clf1 = make_pipeline(StandardScaler(), SVM(C=5e-4))
  1972. clf2 = make_pipeline(StandardScaler(), SVM(C=5e-4))
  1973. if n_jobs != None:
  1974. clf1 = SlidingEstimator(clf1, verbose='warning', n_jobs=n_jobs)
  1975. clf2 = SlidingEstimator(clf2, verbose='warning', n_jobs=n_jobs)
  1976. elif method == 'lda':
  1977. clf1 = make_pipeline(StandardScaler(), LDA())
  1978. if n_jobs != None:
  1979. clf1 = SlidingEstimator(clf1, verbose='warning', n_jobs=n_jobs)
  1980. clf2 = make_pipeline(StandardScaler(), LDA())
  1981. if n_jobs != None:
  1982. clf2 = SlidingEstimator(clf2, verbose='warning', n_jobs=n_jobs)
  1983. if method == 'nonlin_svm':
  1984. clf1 = make_pipeline(StandardScaler(), SVC(kernel='poly', C=1.0))
  1985. clf2 = make_pipeline(StandardScaler(), SVC(kernel='poly', C=1.0))
  1986. if n_jobs != None:
  1987. clf1 = SlidingEstimator(clf1, verbose='warning', n_jobs=n_jobs)
  1988. clf2 = SlidingEstimator(clf2, verbose='warning', n_jobs=n_jobs)
  1989. n_trls = X.shape[0]
  1990. idx_rnd = np.concatenate([np.zeros(n_trls // 2), np.ones(n_trls // 2)])
  1991. if (n_trls % 2) > 0:
  1992. idx_rnd = np.concatenate([idx_rnd, [1.0]])
  1993. random.shuffle(idx_rnd)
  1994. if n_jobs != None:
  1995. X1 = X[idx_rnd < 1, :, :]
  1996. y1 = y[idx_rnd < 1]
  1997. X2 = X[idx_rnd > 0, :, :]
  1998. y2 = y[idx_rnd > 0]
  1999. else:
  2000. X1 = X[idx_rnd < 1, :]
  2001. y1 = y[idx_rnd < 1]
  2002. X2 = X[idx_rnd > 0, :]
  2003. y2 = y[idx_rnd > 0]
  2004. clf1.fit(X1, y1)
  2005. score1 = clf1.score(X2, y2)
  2006. clf2.fit(X2, y2)
  2007. score2 = clf2.score(X1, y1)
  2008. scores_all.append(np.array([score1, score2]).mean(0))
  2009. if return_inter:
  2010. return scores_all
  2011. else:
  2012. return np.mean(scores_all)
  2013. def get_decoding(data, labels, variables, time_window, method, n_jobs=None):
  2014. decoding_targets, cue_tgt_width_reward_ids, _, _ = get_shattering_ids(variables[0], variables[1], variables[2],
  2015. variables[3])
  2016. n_combos = decoding_targets.shape[0]
  2017. scores = []
  2018. print('Cutting along all possible axes (shattering dimensionality)')
  2019. for combo in tqdm(range(n_combos)):
  2020. if n_jobs != None:
  2021. X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
  2022. else:
  2023. X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=False)
  2024. y = assign_lables(labels, decoding_targets[combo, :])
  2025. scores.append(decode(X, y, method=method, n_jobs=n_jobs))
  2026. score_cue = scores[int(cue_tgt_width_reward_ids[0])]
  2027. score_target = scores[int(cue_tgt_width_reward_ids[1])]
  2028. score_width = scores[int(cue_tgt_width_reward_ids[2])]
  2029. score_rew = scores[int(cue_tgt_width_reward_ids[3])]
  2030. score_shattering = scores
  2031. var_ids = [int(cue_tgt_width_reward_ids) for cue_tgt_width_reward_ids in cue_tgt_width_reward_ids]
  2032. scores_irrel = np.delete(scores, obj=var_ids)
  2033. return [score_cue, score_target, score_width, score_rew, score_shattering], scores_irrel
  2034. def get_decoding_null(data, labels, time_window, method, n_reps=100, n_jobs=None):
  2035. if len(np.unique(labels)) < 9:
  2036. fac = [0, 0, 0, 0, 1, 1, 1, 1]
  2037. else:
  2038. fac = [0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1]
  2039. print('Generating the null decoding distribution...')
  2040. scores = []
  2041. for rep in tqdm(range(n_reps)):
  2042. if n_jobs != None:
  2043. X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
  2044. else:
  2045. X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=False)
  2046. y = assign_lables(labels, factor=fac).copy()
  2047. if len(np.unique(labels)) < 3:
  2048. y = labels.copy()
  2049. random.shuffle(y)
  2050. scores.append(decode(X, y, method=method, n_jobs=n_jobs))
  2051. return np.array(scores)
  2052. def decode_xgen(X1, X2, y1, y2, method='svm', n_jobs=-1):
  2053. # prepare a series of classifier applied at each time sample
  2054. if method == 'svm':
  2055. clf1 = make_pipeline(StandardScaler(), SVM(C=5e-4))
  2056. if n_jobs != None:
  2057. clf1 = SlidingEstimator(clf1, verbose=False, n_jobs=-1)
  2058. clf2 = make_pipeline(StandardScaler(), SVM(C=5e-4))
  2059. if n_jobs != None:
  2060. clf2 = SlidingEstimator(clf2, verbose=False, n_jobs=-1)
  2061. elif method == 'lda':
  2062. clf1 = make_pipeline(StandardScaler(), LDA())
  2063. if n_jobs != None:
  2064. clf1 = SlidingEstimator(clf1, verbose=False, n_jobs=-1)
  2065. clf2 = make_pipeline(StandardScaler(), LDA())
  2066. if n_jobs != None:
  2067. clf2 = SlidingEstimator(clf2, verbose=False, n_jobs=-1)
  2068. clf1.fit(X1, y1)
  2069. score1 = clf1.score(X2, y2)
  2070. clf2.fit(X2, y2)
  2071. score2 = clf2.score(X1, y1)
  2072. return np.array([score1, score2]).mean(0)
  2073. def get_x_gen_null(data, labels, method, n_reps=100, n_jobs=None):
  2074. if len(np.unique(labels)) < 9:
  2075. fac = [0, 0, 0, 0, 1, 1, 1, 1]
  2076. else:
  2077. fac = [0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1]
  2078. print('Generating the null x-gen distribution...')
  2079. scores = []
  2080. for rep in tqdm(range(n_reps)):
  2081. X = data
  2082. y = assign_lables(labels, factor=fac).copy()
  2083. if n_jobs != None:
  2084. X1 = X[::2, :, :]
  2085. X2 = X[1::2, :, :]
  2086. else:
  2087. X1 = X[::2, :]
  2088. X2 = X[1::2, :]
  2089. y1 = y[::2]
  2090. y2 = y[1::2]
  2091. random.shuffle(y1)
  2092. random.shuffle(y2)
  2093. scores.append(decode_xgen(X1, X2, y1, y2, method=method, n_jobs=n_jobs))
  2094. return np.array(scores)
  2095. def get_xgen(data, labels, variables, time_window, method='svm', mode='full', n_jobs=None, verbose=False):
  2096. decoding_targets, cue_tgt_width_reward_ids, cross_gen_decoding_train_ids, cross_gen_decoding_test_ids = get_shattering_ids(
  2097. variables[0], variables[1], variables[2], variables[3])
  2098. var_ids = [int(cue_tgt_width_reward_ids) for cue_tgt_width_reward_ids in cue_tgt_width_reward_ids]
  2099. n_combos = len(cross_gen_decoding_train_ids)
  2100. if mode == 'only_rel':
  2101. combo_list = var_ids
  2102. elif mode == 'full':
  2103. combo_list = list(range(n_combos))
  2104. scores = []
  2105. if verbose:
  2106. print('Cutting along all possible axes (xgen)')
  2107. for combo in tqdm(combo_list, disable=not verbose):
  2108. xgen_train = cross_gen_decoding_train_ids[combo]
  2109. xgen_test = cross_gen_decoding_test_ids[combo]
  2110. if n_jobs != None:
  2111. X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
  2112. else:
  2113. X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=False)
  2114. n_axes = len(xgen_train)
  2115. scores_axes = []
  2116. for axis in range(n_axes):
  2117. y1 = np.array(xgen_train[axis]).flatten()
  2118. X1 = []
  2119. for _ in range(len(y1)):
  2120. if n_jobs != None:
  2121. X1.append(X[labels == y1[_], :, :])
  2122. else:
  2123. X1.append(X[labels == y1[_], :])
  2124. X1 = np.concatenate(X1, axis=0)
  2125. y2 = np.array(xgen_test[axis]).flatten()
  2126. X2 = []
  2127. for _ in range(len(y1)):
  2128. if n_jobs != None:
  2129. X2.append(X[labels == y2[_], :, :])
  2130. else:
  2131. X2.append(X[labels == y2[_], :])
  2132. X2 = np.concatenate(X2, axis=0)
  2133. y = np.concatenate([np.zeros(int(X1.shape[0] / 2)), np.ones(int(X1.shape[0] / 2))])
  2134. scores_axes.append(decode_xgen(X1, X2, y, y, method=method, n_jobs=n_jobs))
  2135. scores.append(np.mean(scores_axes))
  2136. if mode == 'full':
  2137. scores_rel = [scores[var_ids[0]], scores[var_ids[1]], scores[var_ids[2]], scores[var_ids[3]]]
  2138. scores_irrel = np.delete(scores, obj=var_ids)
  2139. return scores_rel, scores_irrel
  2140. elif mode == 'only_rel':
  2141. return np.array(scores)
  2142. def get_decoding_models(n_neurons=400, noise_std=1, n_trials=20, seed=42):
  2143. np.random.seed(seed)
  2144. conds = np.ones((4, 3))
  2145. conds[0, :2] = -1
  2146. conds[1, [0, 2]] = -1
  2147. conds[2, 1:] = -1
  2148. conds = np.array([conds] * n_trials)
  2149. clfc = SVM(C=1e-5)
  2150. clfs = SVM(C=1e-5)
  2151. clfr = SVM(C=1e-5)
  2152. N_std = [[n_neurons, noise_std]] # number of neurons and std of noise
  2153. targets_r = [1, -1, -1, 1]
  2154. targets_s = [-1, 1, -1, 1]
  2155. targets_c = [-1, -1, 1, 1]
  2156. decoding_tot = np.zeros((len(N_std), 100, 3, 2))
  2157. cross_decoding_tot = np.zeros((len(N_std), 100, 3, 2))
  2158. for counter, n_std in enumerate(N_std):
  2159. N = n_std[0]
  2160. noise_std = n_std[1]
  2161. for model in range(100):
  2162. for rand_opt in range(2):
  2163. if rand_opt == 0:
  2164. betas = np.random.normal(0, 1, (3, N))
  2165. else:
  2166. cov = np.diag([0, 0, 3])
  2167. betas = np.random.multivariate_normal(np.zeros(3), cov, N).T
  2168. x_train0 = conds @ betas + noise_std * np.random.normal(0, 1, (n_trials, 4, N))
  2169. x_test0 = conds @ betas + noise_std * np.random.normal(0, 1, (n_trials, 4, N))
  2170. x_train0_long = np.concatenate(np.array_split(x_train0, n_trials, axis=0), axis=1)[0, :, :]
  2171. x_test0_long = np.concatenate(np.array_split(x_test0, n_trials, axis=0), axis=1)[0, :, :]
  2172. clfc.fit(x_train0_long, np.concatenate([targets_c] * n_trials))
  2173. clfs.fit(x_train0_long, np.concatenate([targets_s] * n_trials))
  2174. clfr.fit(x_train0_long, np.concatenate([targets_r] * n_trials))
  2175. decoding_tot[counter, model, 0, rand_opt] = clfc.score(x_test0_long,
  2176. np.concatenate([targets_c] * n_trials))
  2177. decoding_tot[counter, model, 1, rand_opt] = clfs.score(x_test0_long,
  2178. np.concatenate([targets_s] * n_trials))
  2179. decoding_tot[counter, model, 2, rand_opt] = clfr.score(x_test0_long,
  2180. np.concatenate([targets_r] * n_trials))
  2181. # Cross decoding
  2182. r_train_id = [[0, 1], [0, 2], [3, 1], [3, 2]]
  2183. r_test_id = [[3, 2], [3, 1], [0, 2], [0, 1]]
  2184. s_train_id = [[0, 1], [0, 3], [2, 1], [2, 3]]
  2185. s_test_id = [[2, 3], [2, 1], [0, 3], [0, 1]]
  2186. c_train_id = [[0, 2], [0, 3], [1, 2], [1, 3]]
  2187. c_test_id = [[1, 3], [1, 2], [0, 3], [0, 2]]
  2188. for i in range(4):
  2189. # color
  2190. x_train1 = x_train0[:, c_train_id[i], :]
  2191. x_test1 = x_test0[:, c_test_id[i], :]
  2192. x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
  2193. x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
  2194. clfc.fit(x_train1, [0, 1] * n_trials)
  2195. cross_decoding_tot[counter, model, 0, rand_opt] += 0.25 * clfc.score(x_test1, [0, 1] * n_trials)
  2196. # shape
  2197. x_train1 = x_train0[:, s_train_id[i], :]
  2198. x_test1 = x_test0[:, s_test_id[i], :]
  2199. x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
  2200. x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
  2201. clfs.fit(x_train1, [0, 1] * n_trials)
  2202. cross_decoding_tot[counter, model, 1, rand_opt] += 0.25 * clfs.score(x_test1, [0, 1] * n_trials)
  2203. # reward
  2204. x_train1 = x_train0[:, r_train_id[i], :]
  2205. x_test1 = x_test0[:, r_test_id[i], :]
  2206. x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
  2207. x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
  2208. clfr.fit(x_train1, [0, 1] * n_trials)
  2209. cross_decoding_tot[counter, model, 2, rand_opt] += 0.25 * clfr.score(x_test1, [0, 1] * n_trials)
  2210. m_decode = np.mean(decoding_tot[0, :], axis=0)
  2211. m_cross_gen = np.mean(cross_decoding_tot[0, :], axis=0)
  2212. return m_decode, m_cross_gen
  2213. def get_decoding_models(n_neurons=400, noise_std=1, n_trials=20, seed=42):
  2214. np.random.seed(seed)
  2215. conds = np.ones((4, 3))
  2216. conds[0, :2] = -1
  2217. conds[1, [0, 2]] = -1
  2218. conds[2, 1:] = -1
  2219. conds = np.array([conds] * n_trials)
  2220. clfc = SVM(C=1e-5)
  2221. clfs = SVM(C=1e-5)
  2222. clfr = SVM(C=1e-5)
  2223. N_std = [[n_neurons, noise_std]] # number of neurons and std of noise
  2224. targets_r = [1, -1, -1, 1]
  2225. targets_s = [-1, 1, -1, 1]
  2226. targets_c = [-1, -1, 1, 1]
  2227. decoding_tot = np.zeros((len(N_std), 100, 3, 2))
  2228. cross_decoding_tot = np.zeros((len(N_std), 100, 3, 2))
  2229. for counter, n_std in enumerate(N_std):
  2230. N = n_std[0]
  2231. noise_std = n_std[1]
  2232. for model in range(100):
  2233. for rand_opt in range(2):
  2234. if rand_opt == 0:
  2235. betas = np.random.normal(0, 1, (3, N))
  2236. else:
  2237. cov = np.diag([0, 0, 3])
  2238. betas = np.random.multivariate_normal(np.zeros(3), cov, N).T
  2239. x_train0 = conds @ betas + noise_std * np.random.normal(0, 1, (n_trials, 4, N))
  2240. x_test0 = conds @ betas + noise_std * np.random.normal(0, 1, (n_trials, 4, N))
  2241. x_train0_long = np.concatenate(np.array_split(x_train0, n_trials, axis=0), axis=1)[0, :, :]
  2242. x_test0_long = np.concatenate(np.array_split(x_test0, n_trials, axis=0), axis=1)[0, :, :]
  2243. clfc.fit(x_train0_long, np.concatenate([targets_c] * n_trials))
  2244. clfs.fit(x_train0_long, np.concatenate([targets_s] * n_trials))
  2245. clfr.fit(x_train0_long, np.concatenate([targets_r] * n_trials))
  2246. decoding_tot[counter, model, 0, rand_opt] = clfc.score(x_test0_long,
  2247. np.concatenate([targets_c] * n_trials))
  2248. decoding_tot[counter, model, 1, rand_opt] = clfs.score(x_test0_long,
  2249. np.concatenate([targets_s] * n_trials))
  2250. decoding_tot[counter, model, 2, rand_opt] = clfr.score(x_test0_long,
  2251. np.concatenate([targets_r] * n_trials))
  2252. # Cross decoding
  2253. r_train_id = [[0, 1], [0, 2], [3, 1], [3, 2]]
  2254. r_test_id = [[3, 2], [3, 1], [0, 2], [0, 1]]
  2255. s_train_id = [[0, 1], [0, 3], [2, 1], [2, 3]]
  2256. s_test_id = [[2, 3], [2, 1], [0, 3], [0, 1]]
  2257. c_train_id = [[0, 2], [0, 3], [1, 2], [1, 3]]
  2258. c_test_id = [[1, 3], [1, 2], [0, 3], [0, 2]]
  2259. for i in range(4):
  2260. # color
  2261. x_train1 = x_train0[:, c_train_id[i], :]
  2262. x_test1 = x_test0[:, c_test_id[i], :]
  2263. x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
  2264. x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
  2265. clfc.fit(x_train1, [0, 1] * n_trials)
  2266. cross_decoding_tot[counter, model, 0, rand_opt] += 0.25 * clfc.score(x_test1, [0, 1] * n_trials)
  2267. # shape
  2268. x_train1 = x_train0[:, s_train_id[i], :]
  2269. x_test1 = x_test0[:, s_test_id[i], :]
  2270. x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
  2271. x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
  2272. clfs.fit(x_train1, [0, 1] * n_trials)
  2273. cross_decoding_tot[counter, model, 1, rand_opt] += 0.25 * clfs.score(x_test1, [0, 1] * n_trials)
  2274. # reward
  2275. x_train1 = x_train0[:, r_train_id[i], :]
  2276. x_test1 = x_test0[:, r_test_id[i], :]
  2277. x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
  2278. x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
  2279. clfr.fit(x_train1, [0, 1] * n_trials)
  2280. cross_decoding_tot[counter, model, 2, rand_opt] += 0.25 * clfr.score(x_test1, [0, 1] * n_trials)
  2281. m_decode = np.mean(decoding_tot[0, :], axis=0)
  2282. m_cross_gen = np.mean(cross_decoding_tot[0, :], axis=0)
  2283. return m_decode, m_cross_gen
  2284. def pearsonr(X, Y, axis=None, keepdims=False, strip_nans=False):
  2285. """
  2286. Pearson correlation across a specific axis.
  2287. """
  2288. if strip_nans:
  2289. mymean = np.nanmean
  2290. mystd = np.nanstd
  2291. mysum = np.nansum
  2292. else:
  2293. mymean = np.mean
  2294. mystd = np.std
  2295. mysum = np.sum
  2296. should_squeeze = axis is not None and not keepdims
  2297. if axis is None:
  2298. X = X.ravel()
  2299. Y = Y.ravel()
  2300. axis = 0
  2301. xbar = mymean(X, axis=axis, keepdims=True)
  2302. ybar = mymean(Y, axis=axis, keepdims=True)
  2303. ssx = mysum((X - xbar) ** 2, axis=axis, keepdims=True)
  2304. ssy = mysum((Y - ybar) ** 2, axis=axis, keepdims=True)
  2305. # the following else-block is equivalent to:
  2306. # num = np.sum( (X-xbar)*(Y-ybar) , axis=axis, keepdims=True)
  2307. # but use MUCH less memory as they accumulate the sum in a loop.
  2308. # the two approaches use the same amount of cputime
  2309. if strip_nans:
  2310. # use the memory-inefficient way in case we have to deal with nans
  2311. num = mysum((X - xbar) * (Y - ybar), axis=axis, keepdims=True)
  2312. else:
  2313. tmpX = X.take(0, axis=axis).reshape(xbar.shape)
  2314. tmpY = Y.take(0, axis=axis).reshape(ybar.shape)
  2315. num = (tmpX - xbar) * (tmpY - ybar)
  2316. for k in range(1, X.shape[axis]):
  2317. tmpX = X.take(k, axis=axis).reshape(xbar.shape)
  2318. tmpY = Y.take(k, axis=axis).reshape(ybar.shape)
  2319. num += (tmpX - xbar) * (tmpY - ybar)
  2320. denom = np.sqrt(ssx) * np.sqrt(ssy)
  2321. r = num / denom
  2322. if should_squeeze:
  2323. s = list(r.shape)
  2324. assert (s[axis] == 1) # the axis dimension should now be singleton
  2325. s.pop(axis) # so remove it
  2326. r = r.reshape(s)
  2327. return r
  2328. def dist_random_data2(const_coeffs, metric='euclidean distance', rnd_model='gaussian (spherical)', n_bootstraps=1000,
  2329. relative_dist=True, bon_correction=False):
  2330. n_epochs = len(const_coeffs)
  2331. n_pairs = const_coeffs[0].shape[0]
  2332. dist_data = np.zeros((n_epochs, n_pairs, n_bootstraps))
  2333. dist_rnd = np.zeros((n_epochs, n_pairs, n_bootstraps))
  2334. dist_str = np.zeros((n_epochs, n_pairs, n_bootstraps))
  2335. for part in range(n_epochs):
  2336. for n_bootstrap in tqdm(range(n_bootstraps)):
  2337. for pair in range(n_pairs):
  2338. data_coeffs = const_coeffs[part][pair, :, :]
  2339. opt_cov = np.diag([0, 0, np.cov(data_coeffs.T)[2, 2]])
  2340. s_opt = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
  2341. opt_cov,
  2342. data_coeffs.shape[0])
  2343. m = np.mean(np.diag(np.cov(data_coeffs.T)))
  2344. s_rnd_1 = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
  2345. np.diag([m, m, m]),
  2346. data_coeffs.shape[0])
  2347. s_rnd_2 = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
  2348. np.diag([m, m, m]),
  2349. data_coeffs.shape[0])
  2350. if metric == 'euclidean distance':
  2351. dist_data[part, pair, n_bootstrap] = euclidean_distance(s_rnd_1, data_coeffs)
  2352. dist_rnd[part, pair, n_bootstrap] = euclidean_distance(s_rnd_1, s_rnd_2)
  2353. dist_str[part, pair, n_bootstrap] = euclidean_distance(s_rnd_1, s_opt)
  2354. elif metric == 'KL divergance estimate':
  2355. dist_data[part, pair, n_bootstrap] = 0.5 * (
  2356. KLdivergence(s_rnd_1, data_coeffs) + KLdivergence(data_coeffs, s_rnd_1))
  2357. dist_rnd[part, pair, n_bootstrap] = 0.5 * (
  2358. KLdivergence(s_rnd_1, s_rnd_2) + KLdivergence(s_rnd_2, s_rnd_1))
  2359. dist_str[part, pair, n_bootstrap] = 0.5 * (
  2360. KLdivergence(s_rnd_1, s_opt) + KLdivergence(s_opt, s_rnd_1))
  2361. dist_data = np.reshape(dist_data, (dist_data.shape[0], dist_data.shape[1] * dist_data.shape[2]))
  2362. dist_rnd = np.reshape(dist_rnd, (dist_rnd.shape[0], dist_rnd.shape[1] * dist_rnd.shape[2]))
  2363. dist_str = np.reshape(dist_str, (dist_str.shape[0], dist_str.shape[1] * dist_str.shape[2]))
  2364. if relative_dist:
  2365. dist_rnd_avg = np.mean(dist_rnd, keepdims=True, axis=-1)
  2366. dist_str_avg = np.mean(dist_str, keepdims=True, axis=-1)
  2367. dist_data -= dist_rnd_avg
  2368. dist_rnd -= dist_rnd_avg
  2369. dist_str -= dist_rnd_avg
  2370. dist_data /= dist_str_avg
  2371. dist_rnd /= dist_str_avg
  2372. dist_str /= dist_str_avg
  2373. p = 2 * (np.sum(dist_rnd >= np.mean(dist_data, axis=-1, keepdims=True), axis=-1) / n_bootstraps)
  2374. if bon_correction:
  2375. p = p * n_epochs
  2376. print('p-values:')
  2377. print(p)
  2378. epoch = np.sort(list(range(1, n_epochs + 1)) * n_bootstraps * n_pairs)
  2379. epoch_labels = np.concatenate([epoch, epoch, epoch])
  2380. dist_label = np.concatenate([['observed'] * (n_bootstraps * n_epochs * n_pairs),
  2381. [rnd_model] * (n_bootstraps * n_epochs * n_pairs),
  2382. ['structured'] * (n_bootstraps * n_epochs * n_pairs)])
  2383. data_df = np.reshape(dist_data, dist_data.shape[0] * dist_data.shape[1])
  2384. rnd_df = np.reshape(dist_rnd, dist_rnd.shape[0] * dist_rnd.shape[1])
  2385. str_df = np.reshape(dist_str, dist_str.shape[0] * dist_str.shape[1])
  2386. all_df = np.concatenate([data_df, rnd_df, str_df])
  2387. df = pd.DataFrame(np.array([all_df, epoch_labels, dist_label]).T,
  2388. columns=[metric, 'learning epoch', 'distribution'])
  2389. df[metric] = df[metric].astype(float)
  2390. df['divergence from'] = 'random selectivity'
  2391. return df, p, np.array([dist_data, dist_rnd, dist_str])
  2392. def dist_structured_data2(const_coeffs, metric='euclidean distance', rnd_model='gaussian (spherical)',
  2393. n_bootstraps=1000, relative_dist=True, bon_correction=False):
  2394. n_epochs = len(const_coeffs)
  2395. n_pairs = const_coeffs[0].shape[0]
  2396. dist_data = np.zeros((n_epochs, n_pairs, n_bootstraps))
  2397. dist_rnd = np.zeros((n_epochs, n_pairs, n_bootstraps))
  2398. dist_str = np.zeros((n_epochs, n_pairs, n_bootstraps))
  2399. for part in range(n_epochs):
  2400. for n_bootstrap in tqdm(range(n_bootstraps)):
  2401. for pair in range(n_pairs):
  2402. data_coeffs = const_coeffs[part][pair, :, :]
  2403. opt_cov = np.diag([0, 0, np.cov(data_coeffs.T)[2, 2]])
  2404. s_opt_1 = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
  2405. opt_cov,
  2406. data_coeffs.shape[0])
  2407. s_opt_2 = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
  2408. opt_cov,
  2409. data_coeffs.shape[0])
  2410. m = np.mean(np.diag(np.cov(data_coeffs.T)))
  2411. s_rnd = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
  2412. np.diag([m, m, m]),
  2413. data_coeffs.shape[0])
  2414. if metric == 'euclidean distance':
  2415. dist_data[part, pair, n_bootstrap] = euclidean_distance(s_opt_1, data_coeffs)
  2416. dist_rnd[part, pair, n_bootstrap] = euclidean_distance(s_opt_1, s_rnd)
  2417. dist_str[part, pair, n_bootstrap] = euclidean_distance(s_opt_1, s_opt_2)
  2418. elif metric == 'KL divergance estimate':
  2419. dist_data[part, pair, n_bootstrap] = 0.5 * (
  2420. KLdivergence(s_opt_1, data_coeffs) + KLdivergence(data_coeffs, s_opt_1))
  2421. dist_rnd[part, pair, n_bootstrap] = 0.5 * (
  2422. KLdivergence(s_opt_1, s_rnd) + KLdivergence(s_rnd, s_opt_1))
  2423. dist_str[part, pair, n_bootstrap] = 0.5 * (
  2424. KLdivergence(s_opt_1, s_opt_2) + KLdivergence(s_opt_2, s_opt_1))
  2425. dist_data = np.reshape(dist_data, (dist_data.shape[0], dist_data.shape[1] * dist_data.shape[2]))
  2426. dist_rnd = np.reshape(dist_rnd, (dist_rnd.shape[0], dist_rnd.shape[1] * dist_rnd.shape[2]))
  2427. dist_str = np.reshape(dist_str, (dist_str.shape[0], dist_str.shape[1] * dist_str.shape[2]))
  2428. if relative_dist:
  2429. dist_str_avg = np.mean(dist_str, keepdims=True, axis=-1)
  2430. dist_rnd_avg = np.mean(dist_rnd, keepdims=True, axis=-1)
  2431. dist_data -= dist_str_avg
  2432. dist_rnd -= dist_str_avg
  2433. dist_str -= dist_str_avg
  2434. dist_data /= dist_rnd_avg
  2435. dist_rnd /= dist_rnd_avg
  2436. dist_str /= dist_rnd_avg
  2437. p = 2 * (np.sum(dist_rnd <= np.mean(dist_data, axis=-1, keepdims=True), axis=-1) / n_bootstraps)
  2438. if bon_correction:
  2439. p = p * n_epochs
  2440. print('p-values:')
  2441. print(p)
  2442. epoch = np.sort(list(range(1, n_epochs + 1)) * n_bootstraps * n_pairs)
  2443. epoch_labels = np.concatenate([epoch, epoch, epoch])
  2444. dist_label = np.concatenate([['observed'] * (n_bootstraps * n_epochs * n_pairs),
  2445. [rnd_model] * (n_bootstraps * n_epochs * n_pairs),
  2446. ['structured'] * (n_bootstraps * n_epochs * n_pairs)])
  2447. data_df = np.reshape(dist_data, dist_data.shape[0] * dist_data.shape[1])
  2448. rnd_df = np.reshape(dist_rnd, dist_rnd.shape[0] * dist_rnd.shape[1])
  2449. str_df = np.reshape(dist_str, dist_str.shape[0] * dist_str.shape[1])
  2450. all_df = np.concatenate([data_df, rnd_df, str_df])
  2451. df = pd.DataFrame(np.array([all_df, epoch_labels, dist_label]).T,
  2452. columns=[metric, 'learning epoch', 'distribution'])
  2453. df[metric] = df[metric].astype(float)
  2454. df['divergence from'] = 'structured selectivity'
  2455. return df, p, np.array([dist_data, dist_rnd, dist_str])
  2456. def get_decoding_exp2(data, labels, variables, time_window, method):
  2457. n_combos = len(variables)
  2458. scores = []
  2459. print('Decoding task variables')
  2460. for combo in tqdm(range(n_combos)):
  2461. X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
  2462. y = assign_lables(labels, variables[combo])
  2463. scores.append(decode(X, y, method=method, return_inter=True))
  2464. return np.array(scores)[:, :, 0]
  2465. def shattering_dim_rel(data, labels, fac_new, time_window, method='svm'):
  2466. labels_new = assign_lables(labels, factor=fac_new)
  2467. n_condi = len(np.unique(fac_new))
  2468. all_combos = list(itertools.combinations(list(range(n_condi)), int(n_condi / 2)))
  2469. n_combos = int(len(all_combos) / 2)
  2470. combo_targets = np.ones((n_combos, n_condi))
  2471. for combo in range(n_combos):
  2472. combo_targets[combo, all_combos[combo][0]] = 0
  2473. combo_targets[combo, all_combos[combo][1]] = 0
  2474. scores = []
  2475. print('Cutting along all possible axes (shattering dimensionality)')
  2476. for combo in tqdm(range(n_combos)):
  2477. X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
  2478. y = assign_lables(labels_new, combo_targets[combo, :])
  2479. scores.append(decode(X, y, method=method))
  2480. return np.array(scores).mean()
  2481. def cross_decoding_1axis(data, labels, target_ax, splitting_ax, method='svm', n_jobs=-1, if_rnd=False):
  2482. labels_splitting_ax = np.array(assign_lables(labels, splitting_ax))
  2483. labels_target_ax = np.array(assign_lables(labels, target_ax))
  2484. if if_rnd:
  2485. random.shuffle(labels_target_ax)
  2486. X1 = data[labels_splitting_ax < 1, :, :]
  2487. y1 = labels_target_ax[labels_splitting_ax < 1]
  2488. X2 = data[labels_splitting_ax > 0, :, :]
  2489. y2 = labels_target_ax[labels_splitting_ax > 0]
  2490. score = decode_xgen(X1, X2, y1, y2, method=method, n_jobs=n_jobs)
  2491. return score
  2492. def get_cell_frate(dat, lab):
  2493. cells_lis_cue = []
  2494. cells_lis_shape = []
  2495. cells_lis_xor = []
  2496. cue = np.array([0, 0, 0, 0, 1, 1, 1, 1])
  2497. shape = np.array([0, 0, 1, 1, 0, 0, 1, 1])
  2498. width = np.array([0, 1, 0, 1, 0, 1, 0, 1])
  2499. xor = np.array([1, 1, 0, 0, 0, 0, 1, 1])
  2500. for _ in range(len(dat)):
  2501. lables_fac = assign_lables(lab[_], factor=cue)
  2502. cells_lis_cue.append(condi_avg(dat[_], lables_fac))
  2503. lables_fac = assign_lables(lab[_], factor=shape)
  2504. cells_lis_shape.append(condi_avg(dat[_], lables_fac))
  2505. lables_fac = assign_lables(lab[_], factor=xor)
  2506. cells_lis_xor.append(condi_avg(dat[_], lables_fac))
  2507. cells_arr_cue = np.concatenate(cells_lis_cue, axis=1)
  2508. cells_lis_shape = np.concatenate(cells_lis_shape, axis=1)
  2509. cells_lis_xor = np.concatenate(cells_lis_xor, axis=1)
  2510. return cells_arr_cue, cells_lis_shape, cells_lis_xor
  2511. def decode_epoch(data, labels, n_reps=100, method='svm', n_inter=1, n_jobs=-1):
  2512. obs = np.array(decode(data, labels, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
  2513. rnd = np.zeros((n_reps, data.shape[-1]))
  2514. for _ in range(n_reps):
  2515. n_trls = data.shape[0]
  2516. idx_rnd = np.concatenate([np.zeros(n_trls // 2), np.ones(n_trls // 2)])
  2517. if (n_trls % 2) > 0:
  2518. idx_rnd = np.concatenate([idx_rnd, [1.0]])
  2519. random.shuffle(idx_rnd)
  2520. rnd[_, :] = np.array(
  2521. decode(data, idx_rnd, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
  2522. return obs, rnd
  2523. def decode_epoch_diff(data1, data2, labels, n_reps=100, method='svm', n_inter=10, tail=1, n_jobs=-1):
  2524. n_cells_1 = data1.shape[1]
  2525. n_cells_2 = data2.shape[1]
  2526. data_all = np.concatenate([data1, data2], axis=1)
  2527. obs1 = np.array(decode(data1, labels, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
  2528. obs2 = np.array(decode(data2, labels, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
  2529. if tail == 1:
  2530. obs = obs1 - obs2
  2531. elif tail == -1:
  2532. obs = obs2 - obs1
  2533. rnd = np.zeros((n_reps, data1.shape[-1]))
  2534. for _ in range(n_reps):
  2535. cell_idx = np.concatenate([np.zeros(n_cells_1), np.ones(n_cells_2)])
  2536. random.shuffle(cell_idx)
  2537. data1_rnd = data_all[:, cell_idx == 0, :]
  2538. data2_rnd = data_all[:, cell_idx == 1, :]
  2539. rnd_1 = np.array(
  2540. decode(data1_rnd, labels, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
  2541. rnd_2 = np.array(
  2542. decode(data2_rnd, labels, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
  2543. if tail == 1:
  2544. rnd[_, :] = rnd_1 - rnd_2
  2545. elif tail == -1:
  2546. rnd[_, :] = rnd_2 - rnd_1
  2547. return obs, rnd
  2548. def smooth(x, window_len=5, window='hanning'):
  2549. s = np.r_[x[window_len - 1:0:-1], x, x[-2:-window_len - 1:-1]]
  2550. if window == 'flat': # moving average
  2551. w = np.ones(window_len, 'd')
  2552. else:
  2553. w = eval('np.' + window + '(window_len)')
  2554. y = np.convolve(w / w.sum(), s, mode='valid')
  2555. return y[int(window_len / 2):-int((window_len / 2))]
  2556. def compute_perm_stats(obs_parts, rnd_parts, tails, if_smooth=False):
  2557. clu_times, clu_lables = [], []
  2558. for _ in range(obs_parts.shape[0]):
  2559. if if_smooth:
  2560. obs = smooth(obs_parts[_, :])
  2561. rnd = np.array([smooth(rnd_parts[_, rep, :]) for rep in range(rnd_parts.shape[1])])
  2562. else:
  2563. obs = obs_parts[_, :]
  2564. rnd = rnd_parts[_, :, :]
  2565. clu_t, clu_lable = permutation_test(obs, rnd.T, tail=tails[_])
  2566. clu_times.append(clu_t)
  2567. clu_lables.append(clu_lable)
  2568. return clu_times, clu_lables
  2569. def melt_data(data, models, times):
  2570. dat_models = []
  2571. for a, m in enumerate(models):
  2572. dat = pd.DataFrame(data[a, :].T)
  2573. # melt data into a long format
  2574. dat['time (s)'] = times
  2575. dat['Model'] = m # add the model label before melting
  2576. dat_models.append(pd.melt(dat, id_vars=['time (s)', 'Model']))
  2577. # combine every model into long format
  2578. data_melt = pd.concat(dat_models)
  2579. return data_melt
  2580. def get_slices(clu_t, clu_lable, times):
  2581. clusters = np.unique(clu_lable)
  2582. clusters = clusters[clusters > 0]
  2583. result = np.zeros((len(clusters), 3))
  2584. for _ in range(len(clusters)):
  2585. idc = clu_lable == _ + 1
  2586. p_value = np.unique(clu_t[idc])[0]
  2587. slices = times[idc]
  2588. result[_, :] = p_value, slices[0], slices[-1]
  2589. return result
  2590. def plot_perm_results(obs, clt_times, clt_labels, model_names, part_names, times=np.linspace(-0.5, 2, 250),
  2591. threshold=0.050001):
  2592. data_all_parts = []
  2593. for m in range(len(model_names)):
  2594. data_melted = melt_data(obs[m], models=part_names, times=times)
  2595. data_melted['variable'] = model_names[m]
  2596. data_all_parts.append(data_melted)
  2597. df = pd.concat(data_all_parts)
  2598. df['decoding accuracy'] = df['value']
  2599. sns.set_style("ticks")
  2600. sns.set_context("notebook", rc={"lines.linewidth": 3})
  2601. cols = sns.cubehelix_palette(len(clt_times[0]), rot=-.25, light=.7)
  2602. g = sns.FacetGrid(df, col="variable", hue="Model", palette=cols, legend_out=True)
  2603. g.map(sns.lineplot, "time (s)", "decoding accuracy")
  2604. for m in range(len(model_names)):
  2605. g.axes[0][m].axhline(0, linestyle='--', linewidth=0.8, color='black')
  2606. g.axes[0][m].axhline(0.5, linestyle='--', linewidth=0.8, color='black')
  2607. g.axes[0][m].axvline(0, linestyle='--', linewidth=0.8, color='black')
  2608. g.axes[0][m].axvline(0.5, linestyle='--', linewidth=0.8, color='black')
  2609. g.axes[0][m].axvline(1., linestyle='--', linewidth=0.8, color='black')
  2610. for p in range(len(clt_times[0])):
  2611. slices = get_slices(clt_times[m][p], clt_labels[m][p], times)
  2612. for s_i in range(slices.shape[0]):
  2613. if slices[s_i, 0] <= threshold:
  2614. g.axes[0][m].hlines(xmin=slices[s_i, 1], xmax=slices[s_i, 2], colors=cols[p],
  2615. y=0.45 - (p / 30),
  2616. linewidth=4)
  2617. # g.set_titles(row_template='', col_template='')
  2618. # plt.ylim([-10, 10]
  2619. g.set_titles("{col_name}")
  2620. # g.axes[0][-1].legend(loc='lower left')
  2621. return
  2622. def plot_clusters(ax, clt_times, clt_labels, times, epoch_names, colour_lis, p_threshold=0.05, plot_chance_lvl=0.45):
  2623. for p in range(len(epoch_names)):
  2624. slices = get_slices(clt_times[p], clt_labels[p], times)
  2625. for s_i in range(slices.shape[0]):
  2626. if slices[s_i, 0] <= p_threshold:
  2627. if p == 2:
  2628. ax.hlines(xmin=slices[s_i, 1], xmax=slices[s_i, 2], colors=colour_lis[p],
  2629. y=plot_chance_lvl - (p / 80), linestyles=(0, (1, 0.5)),
  2630. linewidth=2)
  2631. else:
  2632. ax.hlines(xmin=slices[s_i, 1], xmax=slices[s_i, 2], colors=colour_lis[p],
  2633. y=plot_chance_lvl - (p / 80),
  2634. linewidth=2)
  2635. return
  2636. def p_into_stars(p_val):
  2637. if p_val <= 0.001:
  2638. stars = "***"
  2639. elif p_val <= 0.01:
  2640. stars = "**"
  2641. elif p_val <= 0.05:
  2642. stars = "*"
  2643. elif (p_val > 0.05) & (p_val <= 0.1):
  2644. stars = '†'
  2645. elif p_val > 0.1:
  2646. stars = 'ns'
  2647. return stars
  2648. def reg_par_recovery(n_neurons=100, max_noise=5, n_trials=101, fit_to_noise=False, n_bootstraps=100):
  2649. noise_levels = np.linspace(0.0001, max_noise, 9)
  2650. corrs_min = np.zeros((2, len(noise_levels), n_bootstraps))
  2651. r2_min = np.zeros((2, len(noise_levels), n_bootstraps))
  2652. corrs_rnd = np.zeros((2, len(noise_levels), n_bootstraps))
  2653. r2_rnd = np.zeros((2, len(noise_levels), n_bootstraps))
  2654. coefs_min_lis = []
  2655. coefs_rnd_lis = []
  2656. print('Recovering the underlying covariance matrix')
  2657. for sig in tqdm(range(len(noise_levels))):
  2658. for n_bootstrap in range(n_bootstraps):
  2659. # Structured selectivity
  2660. cov = np.zeros((3, 3))
  2661. cov[2, 2] = 1
  2662. # simulate neuronal selectivity profiles of minimal
  2663. s_opt = np.random.multivariate_normal(np.zeros(3), cov, n_neurons)
  2664. # simulate neuronal selectivity profiles of random
  2665. s_rnd = np.random.multivariate_normal(np.zeros(3), np.diag([1 / 3, 1 / 3, 1 / 3]), n_neurons)
  2666. design = np.zeros((4, 3))
  2667. design[0, 0] = -0.5
  2668. design[0, 1] = 0.5
  2669. design[0, 2] = -0.5
  2670. design[1, 0] = 0.5
  2671. design[1, 1] = -0.5
  2672. design[1, 2] = -0.5
  2673. design[2, :] = -0.5
  2674. design[2, 2] = 0.5
  2675. design[3, :] = 0.5 # rewarded
  2676. # construct design matrix populated with orthogonalised coefficients
  2677. design_stack = np.tile(design.T, n_trials).T
  2678. # prepare the regression models
  2679. clf1 = linear_model.LinearRegression(fit_intercept=True)
  2680. clf2 = linear_model.LinearRegression(fit_intercept=True)
  2681. # generate firing rate for the optimal model using orthogonalised coefficients
  2682. r_min = s_opt @ design_stack.T
  2683. r_rnd = s_rnd @ design_stack.T
  2684. # add noise to the firing rates
  2685. r_train_min = r_min.T + np.random.normal(0, noise_levels[sig], (4 * n_trials, n_neurons))
  2686. r_train_rnd = r_rnd.T + np.random.normal(0, noise_levels[sig], (4 * n_trials, n_neurons))
  2687. # try to recover the model when only noise was supplied
  2688. if fit_to_noise:
  2689. r_train_min = np.random.normal(0, noise_levels[sig], (4 * n_trials, n_neurons))
  2690. r_train_rnd = np.random.normal(0, noise_levels[sig], (4 * n_trials, n_neurons))
  2691. design = np.zeros((4, 3))
  2692. design[0, 0] = -0.5;
  2693. design[0, 1] = 0.5;
  2694. design[0, 2] = -0.5
  2695. design[1, 0] = 0.5;
  2696. design[1, 1] = -0.5;
  2697. design[1, 2] = -0.5
  2698. design[2, :] = -0.5;
  2699. design[2, 2] = 0.5
  2700. design[3, :] = 0.5 # rewarded
  2701. design_stack = np.tile(design.T, n_trials).T
  2702. # fit regression to data (underlying structured selectivity) using orthogonalised coefficients
  2703. clf1.fit(design_stack, r_train_min)
  2704. clf2.fit(design_stack, r_train_rnd)
  2705. # get and mean-center orthogonalised coefficients
  2706. coefs_min = clf1.coef_
  2707. coefs_min = coefs_min - np.mean(coefs_min, axis=0, keepdims=True)
  2708. coefs_rnd = clf2.coef_
  2709. coefs_rnd = coefs_rnd - np.mean(coefs_rnd, axis=0, keepdims=True)
  2710. # get true covariance under min generative data
  2711. true_min_min = np.zeros((3, 3))
  2712. true_min_min[2, 2] = np.cov(coefs_min.T)[2, 2]
  2713. true_min_min = true_min_min.flatten()
  2714. m_rnd = np.mean(np.diag(np.cov(coefs_min.T)))
  2715. true_min_rnd = np.diag([m_rnd, m_rnd, m_rnd]).flatten()
  2716. r2_min[0, sig, n_bootstrap] = r2_score(true_min_min, np.cov(coefs_min.T).flatten())
  2717. corrs_min[0, sig, n_bootstrap] = pearsonr(true_min_min, np.cov(coefs_min.T).flatten())[0]
  2718. r2_min[1, sig, n_bootstrap] = r2_score(true_min_rnd, np.cov(coefs_min.T).flatten())
  2719. corrs_min[1, sig, n_bootstrap] = pearsonr(true_min_rnd, np.cov(coefs_min.T).flatten())[0]
  2720. # get true covariance under rnd generative data
  2721. true_rnd_min = np.zeros((3, 3))
  2722. true_rnd_min[2, 2] = np.cov(coefs_rnd.T)[2, 2]
  2723. true_rnd_min = true_rnd_min.flatten()
  2724. m_rnd = np.mean(np.diag(np.cov(coefs_rnd.T)))
  2725. true_rnd_rnd = np.diag([m_rnd, m_rnd, m_rnd]).flatten()
  2726. r2_rnd[0, sig, n_bootstrap] = r2_score(true_rnd_min, np.cov(coefs_rnd.T).flatten())
  2727. corrs_rnd[0, sig, n_bootstrap] = pearsonr(true_rnd_min, np.cov(coefs_rnd.T).flatten())[0]
  2728. r2_rnd[1, sig, n_bootstrap] = r2_score(true_rnd_rnd, np.cov(coefs_rnd.T).flatten())
  2729. corrs_rnd[1, sig, n_bootstrap] = pearsonr(true_rnd_rnd, np.cov(coefs_rnd.T).flatten())[0]
  2730. coefs_min_lis.append(coefs_min)
  2731. coefs_rnd_lis.append(coefs_rnd)
  2732. return corrs_min, r2_min, corrs_rnd, r2_rnd, coefs_min_lis, coefs_rnd_lis
  2733. def grab_variables(Name):
  2734. with open(Name + '.txt', "rb") as f:
  2735. data = pickle.load(f)
  2736. total_cost_over_time = data[1]
  2737. betas_final = data[0][0]
  2738. w_final = data[0][1]
  2739. w_b_final = data[0][2]
  2740. perf_cost_final = data[0][3]
  2741. reg_cost_final = data[0][4]
  2742. del data
  2743. return betas_final, w_final, w_b_final, perf_cost_final, reg_cost_final, total_cost_over_time
  2744. def cos_sim(v1, v2):
  2745. return np.dot(v1, v2) / (norm(v1) * norm(v2))
  2746. def epairs_metric(weights1, weights2, l=20):
  2747. def epairs(weights, l=5):
  2748. N = weights.shape[0]
  2749. angles = np.zeros((N))
  2750. for n in range(N):
  2751. cosine = weights[n, :] @ weights.T / (
  2752. np.linalg.norm(weights[n, :]) *
  2753. np.linalg.norm(weights, axis=1))
  2754. # cosine = np.abs(cosine)
  2755. cosine = np.delete(cosine, n)
  2756. cosine.sort()
  2757. angles[n] = np.median(np.arccos(cosine[-l:]))
  2758. return angles
  2759. return np.abs(np.mean(epairs(weights1, l=l)) - np.mean(epairs(weights2, l=l)))
  2760. def split_data(data, labels1, n_splits=10, min_trl=100, n_condi=8):
  2761. data1_splits = np.zeros((n_splits, int((min_trl / n_splits) * n_condi), data.shape[1], 250))
  2762. for c in range(n_condi):
  2763. data1 = data[labels1 == c, :, :]
  2764. trl_idc = np.repeat(list(range(n_splits)), int(min_trl / n_splits))
  2765. random.shuffle(trl_idc)
  2766. for split in range(n_splits):
  2767. data1_splits[split, c * 10:c * 10 + 10, :, :] = data1[trl_idc == split, :, :]
  2768. return data1_splits
  2769. def get_sd_dimensions(data, labels, time_window, method='svm', n_jobs=None, n_inter=1, n_condi=16):
  2770. all_combos = list(itertools.combinations(list(range(n_condi)), int(n_condi / 2)))
  2771. n_combos = int(len(all_combos) / 2)
  2772. decoding_targets = np.zeros((n_combos, n_condi))
  2773. for i in range(n_combos):
  2774. decoding_targets[i, all_combos[i]] = 1
  2775. n_combos = decoding_targets.shape[0]
  2776. scores = []
  2777. print('Cutting along all possible axes (shattering dimensionality)')
  2778. for combo in tqdm(range(n_combos)):
  2779. if n_jobs != None:
  2780. X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
  2781. else:
  2782. X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=False)
  2783. y = assign_lables(labels, decoding_targets[combo, :])
  2784. scores.append(decode(X, y, method=method, n_jobs=n_jobs, n_inter=n_inter))
  2785. return np.array(scores)
  2786. def get_sd_min_axis(data, mode='momentum'):
  2787. from scipy.optimize import curve_fit
  2788. def func_sigmoid(x, L, x0, k, b):
  2789. return L / (1 + np.exp(-k * (x - x0))) + b
  2790. axes_min = np.zeros(data.shape[0])
  2791. for rep in range(data.shape[0]):
  2792. if mode == 'momentum':
  2793. data_sorted = np.array(sorted(data[rep, :], reverse=True))
  2794. axes_min[rep] = next((i for i, j in enumerate(data_sorted < 0) if j), None)
  2795. elif mode == 'sigmoid':
  2796. data_sorted = np.array(sorted(data[rep, :], reverse=False))
  2797. x = data_sorted
  2798. y = np.array(list(range(1, data_sorted.shape[0] + 1)))
  2799. p0 = [max(y), np.median(x), 1, min(y)]
  2800. axes_min[rep] = curve_fit(func_sigmoid, x, y, p0, maxfev=5000)[0][1]
  2801. return axes_min
  2802. def get_xgen_exp2(data, labels, condi_fac, target_var_fac, method='svm'):
  2803. condi_lis1 = condi_fac[target_var_fac == 0]
  2804. condi_lis2 = condi_fac[target_var_fac == 1]
  2805. combos = []
  2806. for x in condi_lis1:
  2807. for y in condi_lis2:
  2808. combos.append([x, y])
  2809. train_test_combos_unfltr = list(itertools.combinations(combos, 2))
  2810. # filter for duplicates
  2811. train_test_combos = []
  2812. for _ in range(len(train_test_combos_unfltr)):
  2813. lis = list(np.array(train_test_combos_unfltr[_]).flatten())
  2814. has_dubs = len(lis) != len(np.unique(lis))
  2815. if not has_dubs:
  2816. train_test_combos.append(train_test_combos_unfltr[_])
  2817. n_combos = len(train_test_combos)
  2818. scores = []
  2819. print('Cutting along all possible axes (xgen)')
  2820. for combo in tqdm(range(n_combos)):
  2821. train_condi = train_test_combos[combo][0]
  2822. test_condi = train_test_combos[combo][1]
  2823. idc_train = np.in1d(labels, train_condi)
  2824. idc_test = np.in1d(labels, test_condi)
  2825. X1 = data[idc_train, :]
  2826. X2 = data[idc_test, :]
  2827. y = np.concatenate([np.zeros(int(X1.shape[0] / 2)), np.ones(int(X1.shape[0] / 2))])
  2828. scores.append(decode_xgen(X1, X2, y, y, method=method, n_jobs=None))
  2829. return np.mean(scores)
  2830. def cross_val_pca(data, labels, factor, n_comps, n_splits=10):
  2831. # zscore the data
  2832. data = (data - np.mean(data, axis=0, keepdims=True)) / (np.std(data, axis=0, keepdims=True) + 1)
  2833. labels = np.array(assign_lables(labels, factor))
  2834. n_trls = data.shape[0]
  2835. idx_rnd = np.concatenate([np.zeros(n_trls // 2), np.ones(n_trls // 2)])
  2836. if (n_trls % 2) > 0:
  2837. idx_rnd = np.concatenate([idx_rnd, [1.0]])
  2838. var_ratios_splits = []
  2839. for split in range(n_splits):
  2840. random.shuffle(idx_rnd)
  2841. data1 = data[idx_rnd == 0, :, :]
  2842. data2 = data[idx_rnd == 1, :, :]
  2843. labels1 = labels[idx_rnd == 0]
  2844. labels2 = labels[idx_rnd == 1]
  2845. data1 = condi_avg(data1, labels1)
  2846. data2 = condi_avg(data2, labels2)
  2847. dat1 = data1.mean(-1)
  2848. dat2 = data2.mean(-1)
  2849. pca1 = PCA(n_components=n_comps, random_state=42)
  2850. pca2 = PCA(n_components=n_comps, random_state=42)
  2851. pca1.fit(dat1)
  2852. comps = pca1.transform(dat2)
  2853. var_ratio1 = pca1.explained_variance_ratio_
  2854. pca2.fit(dat2)
  2855. comps = pca2.transform(dat1)
  2856. var_ratio2 = pca2.explained_variance_ratio_
  2857. var_ratios_splits.append(np.array([var_ratio1, var_ratio2]).mean(0))
  2858. var_ratio = np.array(var_ratios_splits).mean(0)
  2859. return var_ratio
  2860. def set_seed(seed=None):
  2861. """
  2862. Function that controls randomness. NumPy and random modules must be imported.
  2863. Args:
  2864. seed : Integer
  2865. A non-negative integer that defines the random state. Default is `None`.
  2866. Returns:
  2867. Nothing.
  2868. """
  2869. if seed is None:
  2870. seed = np.random.choice(2 ** 32)
  2871. random.seed(seed)
  2872. np.random.seed(seed)
  2873. print(f'Random seed {seed} has been set.')
  2874. def equalise_data_witihn_session(data, labels, which_trl='end', n_splits=4):
  2875. trl_min_ses = np.min([data[i].shape[0] for i in range(len(data))])
  2876. if which_trl == 'beginning':
  2877. data_re = [data[_][:trl_min_ses, :, :] for _ in range(len(data))]
  2878. labels_re = [labels[_][:trl_min_ses] for _ in range(len(data))]
  2879. elif which_trl == 'end':
  2880. data_re = [data[_][-trl_min_ses:, :, :] for _ in range(len(data))]
  2881. labels_re = [labels[_][-trl_min_ses:] for _ in range(len(data))]
  2882. elif which_trl == 'rnd':
  2883. idx_re = np.random.choice(np.array(list(range(0, data.shape[0]))), trl_min_ses, replace=False)
  2884. data_re = [data[_][idx_re, :, :] for _ in range(len(data))]
  2885. labels_re = [labels[_][idx_re] for _ in range(len(data))]
  2886. elif which_trl == 'middle':
  2887. labels_re = []
  2888. data_re = []
  2889. for i_sess in range(len(data)):
  2890. idc_half = int(len(labels[i_sess]) / 2)
  2891. n_trls = len(labels[i_sess])
  2892. if (n_trls % 2) > 0:
  2893. labels_re.append(labels[i_sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2) + 1])
  2894. data_re.append(data[i_sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2) + 1, :, :])
  2895. else:
  2896. labels_re.append(labels[i_sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2)])
  2897. data_re.append(data[i_sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2), :, :])
  2898. min_condi = int(
  2899. np.min([np.min(list(get_freqs(labels_re[i]).values())) for i in range(len(labels_re))]) / n_splits) - 1
  2900. n_trl_bck = int(trl_min_ses / n_splits)
  2901. trial_start = 0
  2902. trial_stop = n_trl_bck
  2903. labels_all_blcks = []
  2904. data_all_blcks = []
  2905. for split in range(n_splits):
  2906. data_block = [data_re[i][trial_start:trial_stop, :, :] for i in range(len(data_re))]
  2907. label_block = [labels_re[i][trial_start:trial_stop] for i in range(len(data_re))]
  2908. data_blck, label_blck = prepare_data(data_block, label_block, which_trl='beginning', set_min_trl=min_condi)
  2909. trial_start += n_trl_bck
  2910. trial_stop += n_trl_bck
  2911. data_all_blcks.append(data_blck)
  2912. labels_all_blcks.append(label_blck)
  2913. min_trl_blck_post = np.min([data_all_blcks[i].shape[0] for i in range(n_splits)])
  2914. data_eq = np.array([data_all_blcks[i][:min_trl_blck_post, :, :] for i in range(n_splits)])
  2915. labels_eq = np.array([labels_all_blcks[i][:min_trl_blck_post] for i in range(n_splits)])
  2916. return data_eq, labels_eq
  2917. def compute_95ci(data, n_bootstraps=10000):
  2918. # Bootstrap resampling
  2919. bootstrap_means = np.empty(n_bootstraps)
  2920. for i in range(n_bootstraps):
  2921. bootstrap_sample = np.random.choice(data, size=len(data), replace=True)
  2922. bootstrap_means[i] = np.mean(bootstrap_sample)
  2923. # Compute the 95% confidence interval
  2924. lower_bound = np.percentile(bootstrap_means, 2.5)
  2925. upper_bound = np.percentile(bootstrap_means, 97.5)
  2926. return lower_bound, upper_bound
  2927. def get_xgen_cross_set_null(X1, X2, labels1, labels2, n_reps=100, method='svm', n_jobs=None):
  2928. xgen_null = np.zeros(n_reps)
  2929. print('Computing null distribution (cross-variable generalisation)')
  2930. for _ in tqdm(range(n_reps)):
  2931. labels2_rnd = labels2.copy()
  2932. np.random.shuffle(labels2_rnd)
  2933. xgen_null[_] = decode_xgen_within_ses(X1, X2, labels1, labels2_rnd, method=method, n_jobs=n_jobs)
  2934. return xgen_null
  2935. def compute_p_value(obs1, obs2, rnd1, rnd2, tail='greater'):
  2936. diff = obs2 - obs1
  2937. diff_rnd = rnd2 - rnd1
  2938. if tail == 'smaller':
  2939. p_val = np.sum(diff_rnd > diff) / diff_rnd.shape[0]
  2940. elif tail == 'greater':
  2941. p_val = np.sum(diff_rnd < diff) / diff_rnd.shape[0]
  2942. elif tail == 'two':
  2943. p_val = np.sum(np.abs(diff_rnd) >= abs(diff)) / diff_rnd.shape[0]
  2944. else:
  2945. raise ValueError('tail must be greater, smaller or two')
  2946. print(' Stats: M1 = ', str(round(obs1, 3)), ', M2 = ', str(round(obs2, 3)), ' | p-value = ', str(round(p_val, 3)))
  2947. return p_val
  2948. def run_within_session_decoding(X, X_times, labels, variables, N_SPLITS, N_REPS, TIME_WINDOW, splitting_factor):
  2949. n_variables = len(variables)
  2950. splitting_labels = np.array(assign_lables(labels[0, :], factor=splitting_factor))
  2951. X_set1 = X[:, splitting_labels == 0, :]
  2952. labels_set1 = labels[:, splitting_labels == 0]
  2953. X_set2 = X[:, splitting_labels == 1, :]
  2954. labels_set2 = labels[:, splitting_labels == 1]
  2955. decoding = np.zeros((n_variables, 2, N_SPLITS))
  2956. decoding_nulls = np.zeros((n_variables, 2, N_SPLITS, N_REPS))
  2957. for i_var in range(n_variables):
  2958. for i_block in range(N_SPLITS):
  2959. decoding[i_var, 0, i_block] = decode(X[i_block, :, :],
  2960. assign_lables(labels[i_block], factor=variables[i_var]), method='svm',
  2961. n_inter=40)
  2962. decoding_nulls[i_var, 0, i_block, :] = get_decoding_null(X_times[i_block, :, :, :],
  2963. assign_lables(labels[i_block],
  2964. factor=variables[i_var]),
  2965. TIME_WINDOW,
  2966. n_jobs=None, method='svm', n_reps=N_REPS)
  2967. decoding[i_var, 1, i_block] = decode_xgen_within_ses(X_set1[i_block, :, :], X_set2[i_block, :, :],
  2968. assign_lables(labels_set1[i_block],
  2969. factor=variables[i_var][
  2970. splitting_factor == 0]),
  2971. assign_lables(labels_set2[i_block],
  2972. factor=variables[i_var][
  2973. splitting_factor == 0]),
  2974. method='svm', n_jobs=None)
  2975. decoding_nulls[i_var, 1, i_block, :] = get_xgen_cross_set_null(X_set1[i_block, :, :], X_set2[i_block, :, :],
  2976. assign_lables(labels_set1[i_block],
  2977. factor=variables[i_var][
  2978. splitting_factor == 0]),
  2979. assign_lables(labels_set2[i_block],
  2980. factor=variables[i_var][
  2981. splitting_factor == 0]),
  2982. method='svm', n_jobs=None, n_reps=N_REPS)
  2983. return decoding, decoding_nulls
  2984. def plot_within_session_decoding(decoding, decoding_nulls, variable_names, learning_stages, trails, ylim=[0.4, 0.8],
  2985. plot_p=True):
  2986. fig, ax = plt.subplots(1, len(variable_names), figsize=(2.5 * len(variable_names), 2.5))
  2987. for i_var, variables in enumerate(variable_names):
  2988. ax[i_var].plot(learning_stages, decoding[i_var, 0, :], label='decoding', color='black')
  2989. ax[i_var].plot(learning_stages, decoding[i_var, 1, :], label='cross-gen.\ndecoding', color='grey', zorder=-5)
  2990. ax[i_var].scatter(learning_stages, decoding[i_var, 0, :], color='black')
  2991. ax[i_var].scatter(learning_stages, decoding[i_var, 1, :], color='grey', zorder=-5)
  2992. ax[i_var].set_title(variables)
  2993. ax[i_var].set_ylim(ylim)
  2994. ax[i_var].set_xlabel('learning stage')
  2995. ax[i_var].set_ylabel('accuracy')
  2996. ax[i_var].axhline(0.5, color='black', linestyle='--')
  2997. sns.despine(top=True, right=True)
  2998. # plot p-values for decoding and cross-gen decoding
  2999. if plot_p:
  3000. p_val_dec = compute_p_value(decoding[i_var, 0, -1], decoding[i_var, 0, 0], decoding_nulls[i_var, 0, -1],
  3001. decoding_nulls[i_var, 0, 1], tail=trails[i_var])
  3002. p_val_cross = compute_p_value(decoding[i_var, 1, -1], decoding[i_var, 1, 0], decoding_nulls[i_var, 1, -1],
  3003. decoding_nulls[i_var, 1, 1], tail=trails[i_var])
  3004. ax[i_var].text(0.5, 0.75, 'p = ' + str(np.round(p_val_dec, 3)), fontsize=8, transform=ax[i_var].transAxes,
  3005. ha='center', color='black')
  3006. ax[i_var].text(0.5, 0.65, 'p = ' + str(np.round(p_val_cross, 3)), fontsize=8, transform=ax[i_var].transAxes,
  3007. ha='center', color='grey')
  3008. ax[0].legend()
  3009. plt.tight_layout()
  3010. plt.show()
  3011. def decode_xgen_within_ses(X1, X2, y1, y2, method='svm', n_jobs=-1, n_iter=10):
  3012. # prepare a series of classifier applied at each time sample
  3013. if method == 'svm':
  3014. clf1 = make_pipeline(StandardScaler(), SVM(C=5e-4))
  3015. if n_jobs != None:
  3016. clf1 = SlidingEstimator(clf1, verbose=False, n_jobs=-1)
  3017. clf2 = make_pipeline(StandardScaler(), SVM(C=5e-4))
  3018. if n_jobs != None:
  3019. clf2 = SlidingEstimator(clf2, verbose=False, n_jobs=-1)
  3020. elif method == 'lda':
  3021. clf1 = make_pipeline(StandardScaler(), LDA())
  3022. if n_jobs != None:
  3023. clf1 = SlidingEstimator(clf1, verbose=False, n_jobs=-1)
  3024. clf2 = make_pipeline(StandardScaler(), LDA())
  3025. if n_jobs != None:
  3026. clf2 = SlidingEstimator(clf2, verbose=False, n_jobs=-1)
  3027. scores = np.zeros((n_iter, 2))
  3028. for _ in range(n_iter):
  3029. idc1 = np.arange(X1.shape[0])
  3030. idc2 = np.arange(X2.shape[0])
  3031. idx_rnd1 = np.random.choice(idc1, X1.shape[0], replace=True)
  3032. idx_rnd2 = np.random.choice(idc2, X2.shape[0], replace=True)
  3033. X1_rnd = X1[idx_rnd1, :]
  3034. y1_rnd = np.array(y1)[idx_rnd1]
  3035. X2_rnd = X2[idx_rnd2, :]
  3036. y2_rnd = np.array(y2)[idx_rnd2]
  3037. clf1.fit(X1_rnd, y1_rnd)
  3038. scores[_, 0] = clf1.score(X2_rnd, y2_rnd)
  3039. clf2.fit(X2_rnd, y2_rnd)
  3040. scores[_, 1] = clf2.score(X1_rnd, y1_rnd)
  3041. return scores.mean((1, 0))
  3042. def split_into_blocks(dat, labels, n_splits):
  3043. n_trl_blck = int(dat.shape[0] / n_splits)
  3044. i_start = 0
  3045. i_stop = n_trl_blck
  3046. data_blocks = []
  3047. labels_blocks = []
  3048. for i_block in range(n_splits):
  3049. data_blocks.append(dat[i_start:i_stop, :, :])
  3050. labels_blocks.append(labels[i_start:i_stop])
  3051. i_start += n_trl_blck
  3052. i_stop += n_trl_blck
  3053. return data_blocks, labels_blocks
  3054. def split_vector(input_vector, k, d):
  3055. n = len(input_vector) # Total number of elements in the input vector
  3056. output_vectors = []
  3057. # Calculate the step between the starts of each output vector to distribute elements evenly
  3058. step = max(1, (n - d) // (k - 1)) if k > 1 else 0
  3059. # Generate the output vectors
  3060. for i in range(k):
  3061. start_index = i * step
  3062. end_index = start_index + d
  3063. # Adjust the end index if it goes beyond the input vector length
  3064. if end_index > n:
  3065. start_index = max(0, n - d) # Move back to fit the last vector
  3066. end_index = n
  3067. output_vector = input_vector[start_index:end_index]
  3068. output_vectors.append(output_vector)
  3069. if end_index == n: # Stop if the last vector reaches the end of the input vector
  3070. break
  3071. return output_vectors
  3072. def split_data_blocks_moveavg(data, labels, N_SPLITS, N_WINDOWS):
  3073. dat_split, labels_split = [], []
  3074. for i_sess in range(len(data)):
  3075. dat_split_ses, labels_split_ses = split_into_blocks(data[i_sess], labels[i_sess], N_SPLITS)
  3076. dat_split.append(dat_split_ses)
  3077. labels_split.append(labels_split_ses)
  3078. dat_blocks, labs_blocks = [], []
  3079. for i_block in range(N_SPLITS):
  3080. dat_blocks.append([dat_split[_][i_block] for _ in range(len(data))])
  3081. labs_blocks.append([labels_split[_][i_block] for _ in range(len(data))])
  3082. data_blocks, labels_blocks = [], []
  3083. for i_block in range(N_SPLITS):
  3084. dat = dat_blocks[i_block]
  3085. labs = labs_blocks[i_block]
  3086. trl_number_min = np.min([dat[_].shape[0] for _ in range(len(dat))])
  3087. dat_block_windows = []
  3088. labels_block_windows = []
  3089. for i_sess in range(len(dat)):
  3090. dat_ses = dat[i_sess]
  3091. labs_ses = labs[i_sess]
  3092. idc_trials = np.arange(dat_ses.shape[0])
  3093. idc_windows = split_vector(idc_trials, N_WINDOWS, trl_number_min - N_WINDOWS)
  3094. dat_block_windows.append([dat_ses[idc_windows[_], :, :] for _ in range(N_WINDOWS)])
  3095. labels_block_windows.append([labs_ses[idc_windows[_]] for _ in range(N_WINDOWS)])
  3096. data_pseudopop_wins, labs_pseudopop_wins = [], []
  3097. for i_window in range(N_WINDOWS):
  3098. dat_window = [dat_block_windows[_][i_window] for _ in range(len(dat))]
  3099. labels_window = [labels_block_windows[_][i_window] for _ in range(len(dat))]
  3100. data_pseudopop, labs_pseudopop = prepare_data(dat_window, labels_window, which_trl="end")
  3101. data_pseudopop_wins.append(data_pseudopop)
  3102. labs_pseudopop_wins.append(labs_pseudopop)
  3103. data_pseudopop_wins = np.array(data_pseudopop_wins)
  3104. labs_pseudopop_wins = np.array(labs_pseudopop_wins)
  3105. data_blocks.append(data_pseudopop_wins)
  3106. labels_blocks.append(labs_pseudopop_wins)
  3107. return data_blocks, labels_blocks
  3108. def split_data_blocks_moveavg_sel(data, labels, N_SPLITS):
  3109. dat_split, labels_split = [], []
  3110. for i_sess in range(len(data)):
  3111. dat_split_ses, labels_split_ses = split_into_blocks(data[i_sess], labels[i_sess], N_SPLITS)
  3112. dat_split.append(dat_split_ses)
  3113. labels_split.append(labels_split_ses)
  3114. dat_blocks, labs_blocks = [], []
  3115. for i_block in range(N_SPLITS):
  3116. dat_blocks.append([dat_split[_][i_block] for _ in range(len(data))])
  3117. labs_blocks.append([labels_split[_][i_block] for _ in range(len(data))])
  3118. return dat_blocks, labs_blocks
  3119. def compute_p_val_learning(data_first, data_last, n_perm=10000, tail='greater'):
  3120. mean_diff = np.mean(data_last) - np.mean(data_first)
  3121. data_all = np.concatenate((data_first, data_last))
  3122. mean_diff_rnd = []
  3123. for _ in range(n_perm):
  3124. idc_rnd = np.random.permutation(len(data_all))
  3125. data_first_rnd = data_all[idc_rnd[:len(data_first)]]
  3126. data_last_rnd = data_all[idc_rnd[len(data_first):]]
  3127. mean_diff_rnd.append(np.mean(data_last_rnd) - np.mean(data_first_rnd))
  3128. if tail == 'greater':
  3129. p_value = np.sum(mean_diff_rnd > mean_diff) / len(mean_diff_rnd)
  3130. elif tail == 'less':
  3131. p_value = np.sum(mean_diff_rnd < mean_diff) / len(mean_diff_rnd)
  3132. elif tail == 'two-sided':
  3133. p_value = np.sum(np.abs(mean_diff_rnd) > np.abs(mean_diff)) / len(mean_diff_rnd)
  3134. return p_value
  3135. def plot_fixation_breaks(reward_prop, animal):
  3136. plt.figure()
  3137. sessions = list(range(1, len(reward_prop) + 1))
  3138. plt.plot(sessions, reward_prop, label='No reward - reward')
  3139. plt.legend()
  3140. plt.xlabel('Session')
  3141. plt.ylabel('Proportion of fixation breaks')
  3142. plt.title(animal)
  3143. plt.show()
  3144. def sum_values(trial_types, keys):
  3145. return sum([trial_types.get(key, 0) for key in keys])
  3146. def compute_p_val_learning(data_first, data_last, n_perm=10000, tail='greater'):
  3147. mean_diff = np.mean(data_last) - np.mean(data_first)
  3148. data_all = np.concatenate((data_first, data_last))
  3149. mean_diff_rnd = []
  3150. for _ in range(n_perm):
  3151. idc_rnd = np.random.permutation(len(data_all))
  3152. data_first_rnd = data_all[idc_rnd[:len(data_first)]]
  3153. data_last_rnd = data_all[idc_rnd[len(data_first):]]
  3154. mean_diff_rnd.append(np.mean(data_last_rnd) - np.mean(data_first_rnd))
  3155. if tail == 'greater':
  3156. p_value = np.sum(mean_diff_rnd > mean_diff) / len(mean_diff_rnd)
  3157. elif tail == 'less':
  3158. p_value = np.sum(mean_diff_rnd < mean_diff) / len(mean_diff_rnd)
  3159. elif tail == 'two-sided':
  3160. p_value = np.sum(np.abs(mean_diff_rnd) > np.abs(mean_diff)) / len(mean_diff_rnd)
  3161. return p_value
  3162. def plot_mean_and_ci(reward, no_reward, n_perm=10000, tail='greater'):
  3163. plt.figure(figsize=(4, 3))
  3164. x = list(range(1, len(reward) + 1))
  3165. y1 = np.array([np.mean(reward[i], axis=0) for i in range(len(reward))])
  3166. y2 = np.array([np.mean(no_reward[i], axis=0) for i in range(len(no_reward))])
  3167. # plote 95% confidence interval
  3168. y1_95ci = np.array([1.96 * np.std(reward[i], axis=0) / np.sqrt(len(reward[i])) for i in range(len(reward))])
  3169. y2_95ci = np.array(
  3170. [1.96 * np.std(no_reward[i], axis=0) / np.sqrt(len(no_reward[i])) for i in range(len(no_reward))])
  3171. plt.plot(x, y1, label='reward', color='black')
  3172. plt.errorbar(x, y1, yerr=y1_95ci, fmt='o', color='black')
  3173. # plt.fill_between(x, y1 - y1_sem, y1 + y1_sem, alpha=0.5)
  3174. plt.plot(x, y2, label='no reward', color='grey')
  3175. plt.errorbar(x, y2, yerr=y2_95ci, fmt='o', color='grey')
  3176. # plt.fill_between(x, y2 - y2_sem, y2 + y2_sem, alpha=0.5)
  3177. p_value_norew = compute_p_val_learning(no_reward[0], no_reward[-1], n_perm=n_perm, tail=tail)
  3178. p_value_rew = compute_p_val_learning(reward[0], reward[-1], n_perm=n_perm, tail=tail)
  3179. # annotate the plot with p-value
  3180. plt.text(0.1, 0.9, f'p-value no reward = {round(p_value_norew, 3)}\np-value reward = {round(p_value_rew, 3)}',
  3181. horizontalalignment='left', verticalalignment='center', transform=plt.gca().transAxes)
  3182. plt.legend()
  3183. plt.xlabel('learning stage')
  3184. plt.xticks([1, 2, 3, 4])
  3185. plt.ylabel('fixation breaks (%)')
  3186. plt.legend(['reward', 'no reward'], loc='lower left')
  3187. sns.despine(top=True, right=True)
  3188. plt.tight_layout()
  3189. plt.show()
  3190. def plot_mean_and_ci_prop(gs, fig, reward_prop, n_perm=10000, tail='greater', stat='reg', plot_comp_dat=False, title='trial termination', y_label='no reward/reward\nproporiton', baseline_val=1.0, vmin=None, vmax=None, jitter_width=0.2):
  3191. ax = fig.add_subplot(gs)
  3192. x = list(range(1, len(reward_prop) + 1))
  3193. y1 = np.array([np.mean(reward_prop[i], axis=0) for i in range(len(reward_prop))])
  3194. y1_95ci = np.array(
  3195. [1.96 * np.std(reward_prop[i], axis=0) / np.sqrt(len(reward_prop[i])) for i in range(len(reward_prop))])
  3196. # plot individual datapoints as empty circles with jitter
  3197. for i, stage_data in enumerate(reward_prop):
  3198. jitter = np.random.uniform(-jitter_width, jitter_width, size=len(stage_data))
  3199. ax.scatter(np.full(len(stage_data), x[i]) + jitter, stage_data,
  3200. facecolors='none', edgecolors='black', s=20, zorder=1, alpha=0.5)
  3201. ax.plot(x, y1, color='black', zorder=2)
  3202. ax.errorbar(x, y1, yerr=y1_95ci, fmt='o', color='black', zorder=3)
  3203. if stat == 'reg':
  3204. y_flat = list(itertools.chain.from_iterable(reward_prop))
  3205. x_flat = [[i+1] * len(reward_prop[i]) for i in range(len(reward_prop))]
  3206. x_flat = list(itertools.chain.from_iterable(x_flat))
  3207. slope, intercept, p_value, r_value = compute_lin_reg(x_flat, y_flat, n_perm=n_perm, tail=tail)
  3208. ax.text(0.1, 0.9, f'p-value = {round(p_value, 3)}\nr = {round(r_value, 3)}',
  3209. horizontalalignment='left', verticalalignment='center', transform=plt.gca().transAxes)
  3210. elif stat == 'ttest':
  3211. p_value_rew_prop = compute_p_val_learning(reward_prop[0], reward_prop[-1], n_perm=n_perm, tail=tail)
  3212. ax.text(0.1, 0.9, f'p-value = {round(p_value_rew_prop, 3)}',
  3213. horizontalalignment='left', verticalalignment='center', transform=plt.gca().transAxes)
  3214. if plot_comp_dat != False:
  3215. y2 = np.array([np.mean(plot_comp_dat[i], axis=0) for i in range(len(plot_comp_dat))])
  3216. y2_95ci = np.array(
  3217. [1.96 * np.std(plot_comp_dat[i], axis=0) / np.sqrt(len(plot_comp_dat[i])) for i in
  3218. range(len(plot_comp_dat))])
  3219. # plot individual datapoints for comparison data
  3220. for i, stage_data in enumerate(plot_comp_dat):
  3221. jitter = np.random.uniform(-jitter_width, jitter_width, size=len(stage_data))
  3222. ax.scatter(np.full(len(stage_data), x[i]) + jitter, stage_data,
  3223. facecolors='none', edgecolors='grey', s=20, zorder=1, alpha=0.5)
  3224. ax.plot(x, y2, color='grey', zorder=2)
  3225. ax.errorbar(x, y2, yerr=y2_95ci, fmt='o', color='grey', zorder=3)
  3226. ax.set_xlabel('learning stage')
  3227. ax.set_xticks(x)
  3228. ax.set_ylabel(y_label)
  3229. ax.set_ylim([vmin, vmax])
  3230. sns.despine(top=True, right=True)
  3231. ax.axhline(y=baseline_val, color='black', linestyle='--', linewidth=0.8)
  3232. ax.set_title(title)
  3233. def get_fixation_breaks(sessions_animals, experiment_label='exp1'):
  3234. with open('config.yml', 'r') as f:
  3235. configs = yaml.safe_load(f)
  3236. code_labels = configs['TRIGGER_CODES']
  3237. fix_breaks_animals = []
  3238. for animal in range(len(sessions_animals)):
  3239. sessions = sessions_animals[animal]
  3240. fix_numers = np.zeros((len(sessions), 5))
  3241. for s in range(len(sessions)):
  3242. print('Session: ', s + 1)
  3243. beh = io.loadmat(configs['PATHS']['in_template_beh'].format(sessions[s]))
  3244. codes = beh['uecode']
  3245. codes = list(codes[0, :])
  3246. times = beh['timingms']
  3247. times = list(times[0, :])
  3248. size = len(codes)
  3249. idx_list = [idx for idx, val in
  3250. enumerate(codes) if val == 6116]
  3251. trials = [codes[i: j] for i, j in
  3252. zip([0] + idx_list, idx_list +
  3253. ([size] if idx_list[-1] != size else []))]
  3254. trials = trials[1:]
  3255. conditions = []
  3256. for _ in range(len(trials)):
  3257. if code_labels['CUE1_ON'] in trials[_] and code_labels['TARGET1_ON'] in trials[_] and (
  3258. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3259. conditions.append(1)
  3260. elif code_labels['CUE1_ON'] in trials[_] and code_labels['TARGET2_ON'] in trials[_] and (
  3261. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3262. conditions.append(2)
  3263. elif code_labels['CUE1_ON'] in trials[_] and code_labels['TARGET3_ON'] in trials[_] and (
  3264. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3265. conditions.append(3)
  3266. elif code_labels['CUE1_ON'] in trials[_] and code_labels['TARGET4_ON'] in trials[_] and (
  3267. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3268. conditions.append(4)
  3269. elif code_labels['CUE2_ON'] in trials[_] and code_labels['TARGET1_ON'] in trials[_] and (
  3270. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3271. conditions.append(5)
  3272. elif code_labels['CUE2_ON'] in trials[_] and code_labels['TARGET2_ON'] in trials[_] and (
  3273. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3274. conditions.append(6)
  3275. elif code_labels['CUE2_ON'] in trials[_] and code_labels['TARGET3_ON'] in trials[_] and (
  3276. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3277. conditions.append(7)
  3278. elif code_labels['CUE2_ON'] in trials[_] and code_labels['TARGET4_ON'] in trials[_] and (
  3279. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3280. conditions.append(8)
  3281. elif code_labels['CUE3_ON'] in trials[_] and code_labels['TARGET1_ON'] in trials[_] and (
  3282. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3283. conditions.append(9)
  3284. elif code_labels['CUE3_ON'] in trials[_] and code_labels['TARGET2_ON'] in trials[_] and (
  3285. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3286. conditions.append(10)
  3287. elif code_labels['CUE3_ON'] in trials[_] and code_labels['TARGET3_ON'] in trials[_] and (
  3288. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3289. conditions.append(11)
  3290. elif code_labels['CUE3_ON'] in trials[_] and code_labels['TARGET4_ON'] in trials[_] and (
  3291. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3292. conditions.append(12)
  3293. elif code_labels['CUE4_ON'] in trials[_] and code_labels['TARGET1_ON'] in trials[_] and (
  3294. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3295. conditions.append(13)
  3296. elif code_labels['CUE4_ON'] in trials[_] and code_labels['TARGET2_ON'] in trials[_] and (
  3297. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3298. conditions.append(14)
  3299. elif code_labels['CUE4_ON'] in trials[_] and code_labels['TARGET3_ON'] in trials[_] and (
  3300. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3301. conditions.append(15)
  3302. elif code_labels['CUE4_ON'] in trials[_] and code_labels['TARGET4_ON'] in trials[_] and (
  3303. code_labels['BREAK_TARGET_ERROR'] in trials[_]):
  3304. conditions.append(16)
  3305. if code_labels['BREAK_CUE_ERROR'] in trials[_] or code_labels['BREAK_TARGET_ERROR'] in trials[_] or (
  3306. code_labels['BREAK_ERROR'] in trials[_]) or (code_labels['FIXATION_ERROR'] in trials[_]):
  3307. conditions.append(99)
  3308. if code_labels['BREAK_CUE_ERROR'] in trials[_]:
  3309. conditions.append(999)
  3310. all_trls = []
  3311. for _ in range(len(trials)):
  3312. if code_labels['CUE1_ON'] in trials[_]:
  3313. all_trls.append(1)
  3314. elif code_labels['CUE2_ON'] in trials[_]:
  3315. all_trls.append(1)
  3316. elif code_labels['CUE3_ON'] in trials[_]:
  3317. all_trls.append(2)
  3318. elif code_labels['CUE4_ON'] in trials[_]:
  3319. all_trls.append(2)
  3320. trial_types = {value: len(list(freq)) for value, freq in groupby(sorted(conditions))}
  3321. trial_all = {value: len(list(freq)) for value, freq in groupby(sorted(all_trls))}
  3322. if experiment_label == 'exp1':
  3323. fix_numers[s, 0] = len(trials)
  3324. fix_numers[s, 1] = sum_values(trial_types,
  3325. [1, 2, 7, 8]) # number of fixation breaks in reward trials
  3326. fix_numers[s, 2] = sum_values(trial_types,
  3327. [3, 4, 5, 6]) # number of fixation breaks in no reward trials
  3328. fix_numers[s, 3] = trial_types[99] # number of trials with all fixation breaks
  3329. fix_numers[s, 4] = trial_types[999] # number of trials with cue fixation breaks
  3330. elif experiment_label == 'exp2':
  3331. fix_numers[s, 0] = len(trials)
  3332. fix_numers[s, 1] = sum_values(trial_types,
  3333. [9, 10, 15, 16]) # number of fixation breaks in reward trials
  3334. fix_numers[s, 2] = sum_values(trial_types,
  3335. [11, 12, 13, 14]) # number of fixation breaks in no reward trials
  3336. fix_numers[s, 3] = trial_types[99] # number of trials with all fixation breaks
  3337. fix_numers[s, 4] = trial_types[999] # number of trials with cue fixation breaks
  3338. fix_breaks_animals.append(fix_numers)
  3339. return fix_breaks_animals
  3340. def get_data_stages(observe_or_run='observe', file_name=None, return_data=False, session_list=None):
  3341. with open('config.yml', 'r') as file:
  3342. configs = yaml.safe_load(file)
  3343. if observe_or_run == 'run':
  3344. data_all_parts, labels = get_data(session_list=session_list,
  3345. path_spikes=configs['PATHS']['out_template_spks'],
  3346. path_meta=configs['PATHS']['out_template_meta'],
  3347. window=[0, 250],
  3348. cut_off=None
  3349. )
  3350. data, _, _ = exclude_neurons(data=data_all_parts,
  3351. session_list=session_list,
  3352. path_locations=configs['PATHS']['out_template_loc'],
  3353. path_sel_exclude=configs['PATHS']['out_template_sel_list'],
  3354. loc=configs['ANALYSIS_PARAMS']['SAMPLED_AREAS']
  3355. )
  3356. save_data([data, labels], ['data', 'labels'], configs['PATHS']['output_path'] + file_name + '.pickle')
  3357. elif observe_or_run == 'observe':
  3358. return_data = True
  3359. obj_loaded = load_data(configs['PATHS']['output_path'] + file_name + '.pickle')
  3360. data = obj_loaded['data']
  3361. labels = obj_loaded['labels']
  3362. if return_data:
  3363. return data, labels
  3364. def create_sliding_windows_adaptive(sessions, n_stages=5, min_window_ratio=0.25):
  3365. """
  3366. Create n overlapping windows with adaptive overlap, ensuring continuous coverage.
  3367. """
  3368. total = len(sessions)
  3369. # Window size based on minimum ratio
  3370. window_size = max(2, int(np.ceil(total * min_window_ratio)))
  3371. window_size = min(window_size, total)
  3372. if n_stages == 1:
  3373. return [sessions]
  3374. # Calculate step size to evenly distribute windows across all sessions
  3375. # This ensures the last window ends at total while maintaining overlap
  3376. step = (total - window_size) / (n_stages - 1)
  3377. windows = []
  3378. for i in range(n_stages):
  3379. start = int(i * step)
  3380. end = min(start + window_size, total)
  3381. # Ensure we don't go past the end
  3382. if end > total:
  3383. end = total
  3384. start = max(0, end - window_size)
  3385. windows.append(sessions[start:end])
  3386. return windows
  3387. def combine_session_lists(mode='time', which_exp='exp1', combine_all=True):
  3388. with open('config.yml', 'r') as file:
  3389. configs = yaml.safe_load(file)
  3390. if which_exp == 'exp1':
  3391. animal1_ses_labels = configs['SESSION_NAMES']['sessions_womble_1']
  3392. animal2_ses_labels = configs['SESSION_NAMES']['sessions_wilfred_1']
  3393. elif which_exp == 'exp2':
  3394. animal1_ses_labels = configs['SESSION_NAMES']['sessions_womble_2']
  3395. animal2_ses_labels = configs['SESSION_NAMES']['sessions_wilfred_2']
  3396. if mode == 'time':
  3397. min_num_sessions = min(len(animal1_ses_labels), len(animal2_ses_labels))
  3398. sesions_split_animal1 = np.array_split(animal1_ses_labels, min_num_sessions)
  3399. sesions_split_animal2 = np.array_split(animal2_ses_labels, min_num_sessions)
  3400. ses_com = []
  3401. for _ in range(min_num_sessions):
  3402. ses_com.append(list(sesions_split_animal1[_]))
  3403. ses_com.append(list(sesions_split_animal2[_]))
  3404. if combine_all:
  3405. ses_com = list(chain.from_iterable(ses_com))
  3406. elif mode == 'fix_bias':
  3407. session_labels = [animal1_ses_labels, animal2_ses_labels]
  3408. fix_breaks_animals = get_fixation_breaks(session_labels, experiment_label=which_exp)
  3409. sessions_sorted = []
  3410. for i_anim in range(2):
  3411. no_reward_fix = fix_breaks_animals[i_anim][:, 2] / fix_breaks_animals[i_anim][:, 3]
  3412. reward_fix = fix_breaks_animals[i_anim][:, 1] / fix_breaks_animals[i_anim][:, 3]
  3413. dat = no_reward_fix / reward_fix
  3414. # sort data_prop_s and get the indices of the sorting
  3415. sort_idx = np.argsort(dat)
  3416. sessions_sorted.append(np.array(session_labels[i_anim])[sort_idx])
  3417. ses_wom = np.array_split(sessions_sorted[0], configs['ANALYSIS_PARAMS']['N_STAGES'])
  3418. ses_wil = np.array_split(sessions_sorted[1], configs['ANALYSIS_PARAMS']['N_STAGES'])
  3419. ses_com = []
  3420. for _ in range(configs['ANALYSIS_PARAMS']['N_STAGES']):
  3421. ses_com.append(list(ses_wom[_]) + list(ses_wil[_]))
  3422. if combine_all:
  3423. ses_com = list(chain.from_iterable(ses_com))
  3424. elif mode == 'stages':
  3425. ses_wom = np.array_split(animal1_ses_labels, configs['ANALYSIS_PARAMS']['N_STAGES'])
  3426. ses_wil = np.array_split(animal2_ses_labels, configs['ANALYSIS_PARAMS']['N_STAGES'])
  3427. ses_com = []
  3428. for _ in range(configs['ANALYSIS_PARAMS']['N_STAGES']):
  3429. ses_com.append(list(ses_wom[_]) + list(ses_wil[_]))
  3430. if combine_all:
  3431. ses_com = list(chain.from_iterable(ses_com))
  3432. elif mode == 'cxt_cost':
  3433. session_labels = [animal1_ses_labels, animal2_ses_labels]
  3434. _, _, switch_costs_cxt = get_switch_costs(session_labels)
  3435. sessions_sorted = []
  3436. for i_anim in range(2):
  3437. dat = switch_costs_cxt[i_anim]
  3438. sort_idx = np.argsort(dat)
  3439. sessions_sorted.append(np.array(session_labels[i_anim])[sort_idx])
  3440. ses_wom = np.array_split(sessions_sorted[0], configs['ANALYSIS_PARAMS']['N_STAGES'])
  3441. ses_wil = np.array_split(sessions_sorted[1], configs['ANALYSIS_PARAMS']['N_STAGES'])
  3442. ses_com = []
  3443. for _ in range(configs['ANALYSIS_PARAMS']['N_STAGES']):
  3444. ses_com.append(list(ses_wom[_]) + list(ses_wil[_]))
  3445. if combine_all:
  3446. ses_com = list(chain.from_iterable(ses_com))
  3447. elif mode == 'colour_cost':
  3448. session_labels = [animal1_ses_labels, animal2_ses_labels]
  3449. switch_costs_col, _, _ = get_switch_costs(session_labels)
  3450. sessions_sorted = []
  3451. for i_anim in range(2):
  3452. dat = switch_costs_col[i_anim]
  3453. sort_idx = np.argsort(dat)
  3454. sessions_sorted.append(np.array(session_labels[i_anim])[sort_idx])
  3455. ses_wom = np.array_split(sessions_sorted[0], configs['ANALYSIS_PARAMS']['N_STAGES'])
  3456. ses_wil = np.array_split(sessions_sorted[1], configs['ANALYSIS_PARAMS']['N_STAGES'])
  3457. ses_com = []
  3458. for _ in range(configs['ANALYSIS_PARAMS']['N_STAGES']):
  3459. ses_com.append(list(ses_wom[_]) + list(ses_wil[_]))
  3460. if combine_all:
  3461. ses_com = list(chain.from_iterable(ses_com))
  3462. elif mode == 'sliding_window':
  3463. parts_animal1 = create_sliding_windows_adaptive(animal1_ses_labels, n_stages=configs['ANALYSIS_PARAMS']['N_STAGES'], min_window_ratio=0.25)
  3464. parts_animal2 = create_sliding_windows_adaptive(animal2_ses_labels, n_stages=configs['ANALYSIS_PARAMS']['N_STAGES'], min_window_ratio=0.25)
  3465. ses_com = [pw + pwi for pw, pwi in zip(parts_animal1, parts_animal2)]
  3466. if combine_all:
  3467. ses_com = list(chain.from_iterable(ses_com))
  3468. return ses_com
  3469. def plot_reg_prop(reward_prop, n_perm=10000, tail='greater'):
  3470. plt.figure(figsize=(3, 2))
  3471. # flatten reward_prop and construct list with stage labels
  3472. y_flat = list(itertools.chain.from_iterable(reward_prop))
  3473. x_flat = [1] * len(reward_prop[0]) + [2] * len(reward_prop[1]) + [3] * len(reward_prop[2]) + [4] * len(
  3474. reward_prop[3])
  3475. slope, intercept, p_value, r_value = compute_lin_reg(x_flat, y_flat, n_perm=n_perm, tail=tail)
  3476. # plot regression
  3477. sns.regplot(x=x_flat, y=y_flat, color='black', scatter_kws={'color': 'black'},
  3478. line_kws={'color': 'black'})
  3479. # annotate the plot with p-value
  3480. plt.text(0.1, 0.9, f'p-value = {round(p_value, 3)}\nr = {round(r_value, 3)}',
  3481. horizontalalignment='left', verticalalignment='center', transform=plt.gca().transAxes)
  3482. plt.xlabel('learning stage')
  3483. plt.xticks([1, 2, 3, 4])
  3484. plt.ylabel('no reward/reward\nproporiton')
  3485. sns.despine(top=True, right=True)
  3486. plt.axhline(y=1, color='black', linestyle='--')
  3487. plt.title('fixation breaks')
  3488. plt.tight_layout()
  3489. plt.show()
  3490. def compute_lin_reg(x, y, n_perm=1000, tail='greater'):
  3491. A = np.vstack([x, np.ones(len(x))]).T
  3492. results = np.linalg.lstsq(A, y, rcond=None)
  3493. slope = results[0][0]
  3494. intercept = results[0][1]
  3495. null_slopes = []
  3496. for _ in range(n_perm):
  3497. x_perm = np.random.permutation(x)
  3498. results = np.linalg.lstsq(np.vstack([x_perm, np.ones(len(x))]).T, y, rcond=None)
  3499. null_slopes.append(results[0][0])
  3500. if tail == 'greater':
  3501. p_value = np.sum(null_slopes > slope) / len(null_slopes)
  3502. elif tail == 'less':
  3503. p_value = np.sum(null_slopes < slope) / len(null_slopes)
  3504. elif tail == 'two-sided':
  3505. p_value = np.sum(np.abs(null_slopes) > np.abs(slope)) / len(null_slopes)
  3506. # compute the r value of the regression
  3507. r_value = spearmanr(x, y)[0]
  3508. return slope, intercept, p_value, r_value
  3509. def plot_regression(df, n_perm=100000):
  3510. # Unique animals
  3511. animals = df['animal'].unique()
  3512. # Create subplots
  3513. fig, axes = plt.subplots(1, len(animals), figsize=(5 * len(animals), 5), sharey=False)
  3514. # Check if there is only one subplot (axis) and make it iterable
  3515. if len(animals) == 1:
  3516. axes = [axes]
  3517. # Iterate over each animal and its corresponding axis
  3518. for animal, ax in zip(animals, axes.flatten()):
  3519. # Filter the DataFrame for the current animal
  3520. animal_df = df[df['animal'] == animal]
  3521. # Plot the regression for the current animal
  3522. sns.regplot(x='x', y='y', data=animal_df, ax=ax, color='black', scatter_kws={'color': 'black'},
  3523. line_kws={'color': 'black'})
  3524. # Compute the regression for the current animal
  3525. slope, intercept, p_value, r_value = compute_lin_reg(animal_df['x'], animal_df['y'], n_perm=n_perm,
  3526. tail='greater')
  3527. # Calculate buffer for x and y limits
  3528. x_buffer = (animal_df['x'].max() - animal_df['x'].min()) * 0.1
  3529. y_buffer = (animal_df['y'].max() - animal_df['y'].min()) * 0.1
  3530. # Set individual x and y limits with buffer
  3531. ax.set_xlim(animal_df['x'].min() - x_buffer, animal_df['x'].max() + x_buffer)
  3532. ax.set_ylim(animal_df['y'].min() - y_buffer, animal_df['y'].max() + y_buffer)
  3533. # Annotate the plot with p-value and r-value
  3534. ax.text(0.5, 0.9, f'p-value = {round(p_value, 3)}\nr = {round(r_value, 3)}',
  3535. horizontalalignment='center', verticalalignment='center', transform=ax.transAxes)
  3536. # Set title
  3537. ax.set_title(animal)
  3538. ax.set_xlabel('Session')
  3539. ax.set_ylabel('Proportion of fixation breaks')
  3540. sns.despine(ax=ax, top=True, right=True)
  3541. # Set common labels
  3542. plt.tight_layout()
  3543. plt.show()
  3544. def creat_plot_grid(n_rows, n_cols, size, width_ratios=None):
  3545. fig = plt.figure(figsize=(size * n_cols, (size * n_rows) * 0.8))
  3546. gs = gridspec.GridSpec(n_rows, n_cols, figure=fig, width_ratios=width_ratios)
  3547. return fig, gs
  3548. def split_data_stages_moveavg(data, labels, n_stages, n_windows, trl_min=None, if_rnd=False):
  3549. # collapse list of lists into a single list
  3550. data_lis = list(chain.from_iterable(data))
  3551. trl_number_min = np.min([data_lis[_].shape[0] for _ in range(len(data_lis))])
  3552. data_stages, labels_stages = [], []
  3553. for i_stage in range(n_stages):
  3554. dat = data[i_stage]
  3555. labs = labels[i_stage]
  3556. dat_stage_windows = []
  3557. labels_stage_windows = []
  3558. for i_sess in range(len(dat)):
  3559. dat_ses = dat[i_sess]
  3560. labs_ses = labs[i_sess]
  3561. idc_trials = np.arange(dat_ses.shape[0])
  3562. idc_windows = split_vector(idc_trials, n_windows, trl_number_min - n_windows)
  3563. dat_stage_windows.append([dat_ses[idc_windows[i_win], :, :] for i_win in range(n_windows)])
  3564. labels_stage_windows.append([labs_ses[idc_windows[i_win]] for i_win in range(n_windows)])
  3565. data_pseudopop_wins, labs_pseudopop_wins = [], []
  3566. for i_window in range(n_windows):
  3567. dat_window = [dat_stage_windows[_][i_window] for _ in range(len(dat))]
  3568. labels_window = [labels_stage_windows[_][i_window] for _ in range(len(dat))]
  3569. if trl_min is None:
  3570. data_pseudopop, labs_pseudopop = prepare_data(dat_window, labels_window, which_trl="end")
  3571. else:
  3572. if if_rnd:
  3573. data_pseudopop, labs_pseudopop = prepare_data(dat_window, labels_window, which_trl="rnd",
  3574. set_min_trl=trl_min)
  3575. else:
  3576. data_pseudopop, labs_pseudopop = prepare_data(dat_window, labels_window, which_trl="end",
  3577. set_min_trl=trl_min)
  3578. data_pseudopop_wins.append(data_pseudopop)
  3579. labs_pseudopop_wins.append(labs_pseudopop)
  3580. data_pseudopop_wins = np.array(data_pseudopop_wins)
  3581. labs_pseudopop_wins = np.array(labs_pseudopop_wins)
  3582. data_stages.append(data_pseudopop_wins)
  3583. labels_stages.append(labs_pseudopop_wins)
  3584. return data_stages, labels_stages
  3585. def run_moving_window_decoding(data_eq, labels_eq, variables, time_window, method='svm', n_jobs=None, if_xgen=True,
  3586. if_null=False, n_reps=None):
  3587. n_windows = data_eq[0].shape[0]
  3588. n_stages = len(data_eq)
  3589. shattering_dim = np.zeros((n_stages, 31, n_windows))
  3590. decoding = np.zeros((n_stages, 4, n_windows))
  3591. if if_xgen:
  3592. xgen_decoding = np.zeros((n_stages, 4, n_windows))
  3593. for i_stage in range(n_stages):
  3594. for i_window in range(n_windows):
  3595. X_stage = data_eq[i_stage][i_window, :, :, :]
  3596. y_stage = labels_eq[i_stage][i_window, :]
  3597. decoding_sess, decoding_dich_sess = get_decoding(data=X_stage,
  3598. labels=y_stage,
  3599. variables=variables,
  3600. time_window=time_window,
  3601. method=method,
  3602. n_jobs=n_jobs)
  3603. if if_xgen:
  3604. decoding_xgen_sess = get_xgen(data=X_stage,
  3605. labels=y_stage,
  3606. variables=variables,
  3607. time_window=time_window,
  3608. method=method,
  3609. mode='only_rel',
  3610. n_jobs=n_jobs,
  3611. )
  3612. xgen_decoding[i_stage, :, i_window] = decoding_xgen_sess
  3613. shattering_dim[i_stage, :, i_window] = decoding_dich_sess
  3614. decoding[i_stage, :, i_window] = decoding_sess[:4]
  3615. shattering_dim = shattering_dim.mean(-1)
  3616. decoding = decoding.mean(-1)
  3617. if if_xgen:
  3618. xgen_decoding = xgen_decoding.mean(-1)
  3619. if if_null:
  3620. decoding_null = np.zeros((n_stages, n_reps, n_windows))
  3621. if if_xgen:
  3622. xgen_null = np.zeros((n_stages, n_reps, 4, n_windows))
  3623. for i_rep in tqdm(range(n_reps)):
  3624. for i_stage in range(n_stages):
  3625. X_stage = data_eq[i_stage]
  3626. y_stage = labels_eq[i_stage]
  3627. idc_rnd = np.arange(X_stage.shape[1])
  3628. random.shuffle(idc_rnd)
  3629. y_stage_rnd = y_stage[:, idc_rnd]
  3630. for i_window in range(n_windows):
  3631. X_stage_win = X_stage[i_window, :, :, :]
  3632. y_stage_win = y_stage_rnd[i_window, :]
  3633. labels_random = assign_lables(y_stage_win, factor=[0, 0, 0, 0, 1, 1, 1, 1])
  3634. decoding_null[i_stage, i_rep, i_window] = decode(
  3635. X_stage_win[:, :, time_window[0]: time_window[1]].mean(-1), labels_random, method=method,
  3636. n_jobs=n_jobs)
  3637. if if_xgen:
  3638. xgen_null[i_stage, i_rep, :, i_window] = get_xgen(data=X_stage_win,
  3639. labels=y_stage_win,
  3640. variables=variables,
  3641. time_window=time_window,
  3642. method=method,
  3643. mode='only_rel',
  3644. n_jobs=n_jobs,
  3645. verbose=False
  3646. )
  3647. decoding_null = decoding_null.mean(-1)
  3648. if if_xgen:
  3649. xgen_null = xgen_null.mean(-1)
  3650. if if_xgen:
  3651. if if_null:
  3652. return shattering_dim, decoding, xgen_decoding, decoding_null, xgen_null
  3653. else:
  3654. return shattering_dim, decoding, xgen_decoding
  3655. else:
  3656. if if_null:
  3657. return shattering_dim, decoding, decoding_null
  3658. else:
  3659. return shattering_dim, decoding
  3660. def shuffle_stages(stage1, stage4, label1, seed=None):
  3661. if seed is None:
  3662. random.seed()
  3663. else:
  3664. random.seed(seed)
  3665. data_all = np.concatenate((stage1, stage4), axis=2)
  3666. n_cells_stage1 = stage1.shape[2]
  3667. idc_cells = np.arange(data_all.shape[2])
  3668. random.shuffle(idc_cells) # Python’s shuffle
  3669. dat_stage1_shuffled = data_all[:, :, idc_cells[:n_cells_stage1], :]
  3670. dat_stage4_shuffled = data_all[:, :, idc_cells[n_cells_stage1:], :]
  3671. return [dat_stage1_shuffled, dat_stage4_shuffled], [label1, label1]
  3672. def run_moving_window_decoding_ler_null(data1, data2, labels1, variables, time_window, n_reps=100, if_xgen=True, method='svm'):
  3673. if time_window is not None:
  3674. shattering_dim_rnd = np.zeros((n_reps, 2, 31))
  3675. decoding_rnd = np.zeros((n_reps, 2, 4))
  3676. xgen_decoding_rnd = np.zeros((n_reps, 2, 4))
  3677. for rep in range(n_reps):
  3678. print("Rep: ", rep + 1, " of ", n_reps)
  3679. data_epochs1and4, labels_epochs1and4 = shuffle_stages(data1, data2, labels1)
  3680. if if_xgen:
  3681. shattering_dim_rep, decoding_rep, xgen_decoding_rep = run_moving_window_decoding(data_epochs1and4,
  3682. labels_epochs1and4,
  3683. variables, time_window,
  3684. if_xgen=if_xgen, method=method)
  3685. xgen_decoding_rnd[rep, :, :] = xgen_decoding_rep
  3686. else:
  3687. shattering_dim_rep, decoding_rep = run_moving_window_decoding(data_epochs1and4, labels_epochs1and4,
  3688. variables, time_window, if_xgen=if_xgen, method=method)
  3689. shattering_dim_rnd[rep, :, :] = shattering_dim_rep
  3690. decoding_rnd[rep, :, :] = decoding_rep
  3691. if if_xgen:
  3692. return shattering_dim_rnd, decoding_rnd, xgen_decoding_rnd
  3693. else:
  3694. return shattering_dim_rnd, decoding_rnd
  3695. if time_window is None:
  3696. n_times = data1[0].shape[-1] - 80
  3697. shattering_dim_rnd = np.zeros((n_times, n_reps, 2, 31))
  3698. decoding_rnd = np.zeros((n_times, n_reps, 2, 4))
  3699. for i_rep in range(n_reps):
  3700. print("Rep: ", i_rep + 1, " of ", n_reps)
  3701. data_epochs1and4, labels_epochs1and4 = shuffle_stages(data1, data2, labels1)
  3702. for i_time in range(30, 200):
  3703. time_sliding = [i_time, i_time + 1]
  3704. print('Time point: ', i_time + 1, ' of ', data1[0].shape[-1] - 50)
  3705. shattering_dim_t, decoding_t = run_moving_window_decoding(data_epochs1and4, labels_epochs1and4,
  3706. variables, time_sliding, if_xgen=False, method=method)
  3707. shattering_dim_rnd[i_time-30, i_rep, :, :] = shattering_dim_t
  3708. decoding_rnd[i_time-30, i_rep, :, :] = decoding_t
  3709. return shattering_dim_rnd, decoding_rnd
  3710. def save_data(data_list, names_list, path):
  3711. data_dict = {}
  3712. for i, name in enumerate(names_list):
  3713. data_dict[name] = data_list[i]
  3714. with open(path, 'wb') as f:
  3715. pickle.dump(data_dict, f)
  3716. def load_data(path):
  3717. with open(path, 'rb') as f:
  3718. data_dict = pickle.load(f)
  3719. return data_dict
  3720. def line_plot_timevar(gs, fig, x, y, color, xlabel, ylabel, title, ylim, xlim, xticks=None, xticklabels=None,
  3721. baseline_line=None,
  3722. patch_pars=None, if_title=True, if_sem=False, xaxese_booo=False):
  3723. ax = fig.add_subplot(gs)
  3724. [ax.plot(x, y.mean(0)[_, :], color=color[_]) for _ in range(y.shape[1])]
  3725. sns.despine(top=True, right=True)
  3726. # plot the shaded area using standard error of the mean
  3727. if if_sem:
  3728. [ax.fill_between(x, y.mean(0)[_] - y.std(0)[_] / np.sqrt(y.shape[0]),
  3729. y.mean(0)[_] + y.std(0)[_] / np.sqrt(y.shape[0]), color=color[_], alpha=0.3) for _ in
  3730. range(y.shape[1])]
  3731. ax.set_ylabel(ylabel)
  3732. if xaxese_booo:
  3733. ax.set_xlabel(xlabel)
  3734. ax.set_ylim(ylim)
  3735. ax.set_xlim(xlim)
  3736. if if_title:
  3737. ax.set_title(title)
  3738. if xticks is not None:
  3739. ax.set_xticks(xticks)
  3740. [ax.axvline(xticks[_], linewidth=0.8, color='black', linestyle='--') for _ in range(len(xticks))]
  3741. if xticklabels is not None:
  3742. ax.set_xticklabels(xticklabels)
  3743. if baseline_line is not None:
  3744. ax.axhline(baseline_line, color='black', linestyle='--', linewidth=0.8)
  3745. if patch_pars is not None:
  3746. rect = patches.Rectangle(patch_pars['xy'], patch_pars['width'], patch_pars['height'], edgecolor=None,
  3747. facecolor='peachpuff',
  3748. zorder=-8)
  3749. ax.add_patch(rect)
  3750. return ax
  3751. def plot_sig_bars(ax, dat_obs, dat_rnd, times, tails_lis, colour_lis, p_threshold=0.05, plot_chance_lvl=0.5,
  3752. if_smooth=False, variable_name='[name here]', time_window=(-0.2, 1.4)):
  3753. """
  3754. Plot significance bars for cluster-corrected permutation tests.
  3755. Parameters:
  3756. -----------
  3757. ax : matplotlib axis
  3758. Axis to plot on
  3759. dat_obs : ndarray
  3760. Observed data
  3761. dat_rnd : ndarray
  3762. Randomization data
  3763. times : ndarray
  3764. Time points
  3765. tails_lis : list
  3766. Tail directions for tests
  3767. colour_lis : list
  3768. Colors for each condition
  3769. p_threshold : float
  3770. P-value threshold for significance
  3771. plot_chance_lvl : float
  3772. Y-position for significance bars
  3773. if_smooth : bool
  3774. Whether to smooth data
  3775. variable_name : str
  3776. Name for printing results
  3777. time_window : tuple or None
  3778. (start_time, end_time) to restrict analysis to specific time window.
  3779. If None, uses entire time range.
  3780. """
  3781. # Restrict to time window if specified
  3782. if time_window is not None:
  3783. start_time, end_time = time_window
  3784. # Find indices corresponding to time window
  3785. time_mask = (times >= start_time) & (times <= end_time)
  3786. time_indices = np.where(time_mask)[0]
  3787. # Subset data and times
  3788. dat_obs_subset = dat_obs[:, time_indices]
  3789. dat_rnd_subset = dat_rnd[:, :, time_indices]
  3790. times_subset = times[time_indices]
  3791. print(f'Restricting analysis to time window: {start_time:.3f}s to {end_time:.3f}s')
  3792. else:
  3793. dat_obs_subset = dat_obs
  3794. dat_rnd_subset = dat_rnd
  3795. times_subset = times
  3796. # Run permutation test on subset
  3797. clt_times, clt_labels = compute_perm_stats(dat_obs_subset,
  3798. dat_rnd_subset,
  3799. tails=tails_lis,
  3800. if_smooth=if_smooth
  3801. )
  3802. print('Cluster perm. test: ' + variable_name)
  3803. for p in range(dat_obs_subset.shape[0]):
  3804. slices = get_slices(clt_times[p], clt_labels[p], times_subset)
  3805. for s_i in range(slices.shape[0]):
  3806. if slices[s_i, 0] <= p_threshold:
  3807. start_time_sig = slices[s_i, 1]
  3808. end_time_sig = slices[s_i, 2]
  3809. p_value = slices[s_i, 0]
  3810. # Print time window for each significant cluster
  3811. print(
  3812. f" significant cluster found: {start_time_sig:.3f}s to {end_time_sig:.3f}s (duration: {end_time_sig - start_time_sig:.3f}s, p={p_value:.3f})")
  3813. if p == 2:
  3814. ax.hlines(xmin=start_time_sig, xmax=end_time_sig, colors=colour_lis[p],
  3815. y=plot_chance_lvl - (p / 80), linestyles=(0, (1, 0.5)),
  3816. linewidth=2)
  3817. else:
  3818. ax.hlines(xmin=start_time_sig, xmax=end_time_sig, colors=colour_lis[p],
  3819. y=plot_chance_lvl - (p / 80),
  3820. linewidth=2)
  3821. return
  3822. def plot_scatter(gs, fig, x, y, scale, yaxis_label, xaxis_label, offset_axis=20, overlay_model='contour', title=None,
  3823. plot_reg=False, dot_size=10, color='darkgrey', edgecolor='dimgrey', mass_levels=(0.68, 0.95, 0.997), out_factor=1.0, kde_smooth=1.4, reg_null=None):
  3824. from scipy.stats import gaussian_kde
  3825. ax = fig.add_subplot(gs)
  3826. ax.scatter(x, y, color=color, s=dot_size, zorder=-5,
  3827. edgecolor=edgecolor)
  3828. # move the left spine (y axis) to the right
  3829. ax.spines['left'].set_position(('axes', 0.5))
  3830. # move the bottom spine (x axis) up
  3831. ax.spines['bottom'].set_position(('axes', 0.5))
  3832. # turn off the right and top spines
  3833. ax.spines['right'].set_visible(False)
  3834. ax.spines['top'].set_visible(False)
  3835. ax.set_ylim([-scale, scale])
  3836. ax.set_xlim([-scale, scale])
  3837. ax.set_yticks([-scale, scale])
  3838. ax.set_xticks([-scale,

fun_lib.py at commit 48ada80, no license · at the source

Overview

Authors: Michał J. Wójcik1,2, Jake P. Stroud3,4, Dante Wasmuht2, Makoto Kusunoki2,5, Mikiko Kadohisa2,5, Mark J. Buckley2, Rui Ponte Costa1, Nicholas E. Myers2,6, Laurence T. Hunt2,7, John Duncan2,5, Mark G. Stokes2
  1. Centre for Neural Circuits and Behaviour, Department of Physiology, Anatomy and Genetics, University of Oxford,Oxford, UK
  2. Department of Experimental Psychology, University of Oxford,Oxford, UK
  3. Department of Engineering, University of Cambridge,Cambridge, UK
  4. Coherence Neuro Global,San Francisco, CA USA
  5. MRC Cognition and Brain Sciences Unit, University of Cambridge,Cambridge, UK
  6. Centre for Neurotechnology, Neuromodulation and Neurotherapeutics, University of Nottingham,Nottingham, UK
  7. Department of Psychiatry, University of Oxford,Oxford, UK
Institutions: University of Oxford (United Kingdom); University of Cambridge (United Kingdom); Montreal Neurological Institute and Hospital (Canada); University of Nottingham (United Kingdom)
Journal: Nature neuroscience, volume 29, issue 8, pages 1966-1975
Dates: received 16 June 2023; accepted 7 May 2026; published online 25 June 2026; in print 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41593-026-02333-w · PMID 42350815 · PMCID PMC13433252 · OpenAlex W7165849883
Open access: hybrid, a free copy (OpenAlex)
Status: code verified
Categories: non-human primate (organism)
Methods: Statistics, Smoothing, state filtering, decompositions, Machine learning, Preprocessing, Connectivity, Single-unit activity, calcium imaging
Keywords: Cognitive control, Neural decoding
MeSH: Learning*, Neurons*, Prefrontal Cortex*, Animals, Macaca mulatta, Male, Models, Neurological, Photic Stimulation (* major topic)
Topic: Memory and Neural Mechanisms (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: Wellcome Trust (Wellcome) (215909/Z/19/Z)
Citations: cited by 6 papers (Europe PMC); 47 references in the paper

Abstract

The relationship between the geometry of neural representations and the task being performed is a central question in neuroscience. The primate prefrontal cortex (PFC) is a primary focus of inquiry, as it can encode information with geometries that either rely on past experience or are experience agnostic. One hypothesis is that PFC representations should evolve with learning, from a format that supports exploration of all possible task rules to a format that minimizes the encoding of task-irrelevant features and supports generalization. Here we test this idea by recording neural activity from the macaque PFC when learning a new rule (‘XOR rule’) from scratch. We show that PFC representations progress from being high dimensional, nonlinear and randomly mixed to low dimensional and rule selective. Upon generalizing the rule to new stimuli, these representations further evolve into an abstract, stimulus-invariant geometry. These findings reconcile previously conflicting accounts of PFC function by demonstrating how neural representations adapt across distinct stages of learning.

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

Repository

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

m-j-wojcik/pfc_learning

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 48ada8054940f6a7ac26e8e83d150357a9f249d2, 27 July 2026
Languages: Python (12)
Size: 154 files, 12 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: Matplotlib (6 files), NumPy (2 files), MNE-Python (1 file), Numba (1 file), pandas (1 file), scikit-learn (1 file), SciPy (1 file), seaborn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
13 files

Code availability

All code for this study was custom written in Python using the NumPy, SciPy, Matplotlib, Scikit-learn and TensorFlow libraries. The code repository is publicly available on GitHub (https://github.com/m-j-wojcik/pfc_learning).

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

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 12 scripts, each with its path and the digest of its content;
  • 6 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

The dataset supporting the findings of this study is accessible through Dryad47. For any additional details, please contact the corresponding author. Source data are provided with this paper.

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

Versions

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

Version 2, 28 September 2026

  • Publisher: n/a → Nature Portfolio

Version 1, 27 September 2026: the first record

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

Cite

This paper

Wójcik, M. J., Stroud, J. P., Wasmuht, D., Kusunoki, M., Kadohisa, M., Buckley, M. J., Costa, R. P., Myers, N. E., Hunt, L. T., Duncan, J., & Stokes, M. G. (2026). Learning shapes neural geometry in the primate prefrontal cortex. Nature neuroscience, 29(8), 1966-1975. https://doi.org/10.1038/s41593-026-02333-w

BibTeX

@article{wojcik2026learning,
author = {Wójcik, Michał J. and Stroud, Jake P. and Wasmuht, Dante and Kusunoki, Makoto and Kadohisa, Mikiko and Buckley, Mark J. and Costa, Rui Ponte and Myers, Nicholas E. and Hunt, Laurence T. and Duncan, John and Stokes, Mark G.},
title = {{Learning shapes neural geometry in the primate prefrontal cortex}},
journal = {Nature neuroscience},
year = {2026},
month = jun,
volume = {29},
number = {8},
pages = {1966--1975},
publisher = {Nature Portfolio},
issn = {1097-6256},
doi = {10.1038/s41593-026-02333-w},
url = {https://doi.org/10.1038/s41593-026-02333-w},
pmid = {42350815},
pmcid = {PMC13433252}
}

RIS

TY - JOUR
AU - Wójcik, Michał J.
AU - Stroud, Jake P.
AU - Wasmuht, Dante
AU - Kusunoki, Makoto
AU - Kadohisa, Mikiko
AU - Buckley, Mark J.
AU - Costa, Rui Ponte
AU - Myers, Nicholas E.
AU - Hunt, Laurence T.
AU - Duncan, John
AU - Stokes, Mark G.
TI - Learning shapes neural geometry in the primate prefrontal cortex
T2 - Nature neuroscience
J2 - Nat Neurosci
PY - 2026
DA - 2026/06/25
VL - 29
IS - 8
SP - 1966
EP - 1975
SN - 1097-6256
PB - Nature Portfolio
DO - 10.1038/s41593-026-02333-w
UR - https://doi.org/10.1038/s41593-026-02333-w
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41593-026-02333-w",
"type": "article-journal",
"title": "Learning shapes neural geometry in the primate prefrontal cortex",
"container-title": "Nature neuroscience",
"author": [
{
"family": "Wójcik",
"given": "Michał J."
},
{
"family": "Stroud",
"given": "Jake P."
},
{
"family": "Wasmuht",
"given": "Dante"
},
{
"family": "Kusunoki",
"given": "Makoto"
},
{
"family": "Kadohisa",
"given": "Mikiko"
},
{
"family": "Buckley",
"given": "Mark J."
},
{
"family": "Costa",
"given": "Rui Ponte"
},
{
"family": "Myers",
"given": "Nicholas E."
},
{
"family": "Hunt",
"given": "Laurence T."
},
{
"family": "Duncan",
"given": "John"
},
{
"family": "Stokes",
"given": "Mark G."
}
],
"container-title-short": "Nat Neurosci",
"volume": "29",
"issue": "8",
"page": "1966-1975",
"DOI": "10.1038/s41593-026-02333-w",
"PMID": "42350815",
"PMCID": "PMC13433252",
"ISSN": "1097-6256",
"publisher": "Nature Portfolio",
"URL": "https://doi.org/10.1038/s41593-026-02333-w",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
25
]
]
}
}

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/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, non-human primate, 9 references
[2] 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, non-human primate, 7 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: scikit-learn, pandas, SciPy, 2 other tools, 8 references
[4] doi:10.1038/s41593-026-02315-y [code]
The representational geometry of emotional states in basolateral amygdala.
Journal: Nature neuroscience
In common: seaborn, scikit-learn, pandas, 3 other tools, 6 references
[5] doi:10.1016/j.isci.2026.117492 [code]
Neural subspace reorganization reflects value-based decision-making.
Journal: iScience
In common: scikit-learn, pandas, SciPy, 2 other tools, 5 references
[6] doi:10.7554/elife.108673 [code]
Adaptive behavior is guided by integrated representations of controlled and non-controlled information.
Journal: eLife
In common: MNE-Python, seaborn, scikit-learn, 4 other tools, 2 references
[7] 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, 3 references
[8] 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: Numba, seaborn, scikit-learn, 4 other tools, 2 references
[9] doi:10.1038/s41467-026-76109-y [code]
Assistive algorithms influence neural representations in motor brain-computer interfaces.
Journal: Nature communications
In common: Numba, seaborn, scikit-learn, 4 other tools, non-human primate, 1 reference
[10] doi:10.1371/journal.pone.0351053 [code]
Distinct roles of neuronal phenotypes during neurofeedback adaptation.
Journal: PloS one
In common: seaborn, scikit-learn, pandas, 3 other tools, non-human primate, 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.