Learning shapes neural geometry in the primate prefrontal cortex.
The 6 matches
- [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] § 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] § Methods › Analysis methods › Decoding ↔ fun_lib.py, lines 206–274 · score 0.62 · cross validated, random splits, classifiers, binary, temporally, matrix
- [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] § 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] § 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
- import numpy as np
- import pandas as pd
- import matplotlib.pyplot as plt
- import seaborn as sns
- import scipy as sp
- from scipy.stats import spearmanr, stats
- from tqdm import tqdm
- from itertools import groupby
- from sklearn.metrics import r2_score
- from sklearn import linear_model
- import random
- import itertools
- from matplotlib import colors
- from mne.decoding import SlidingEstimator, GeneralizingEstimator
- from sklearn.discriminant_analysis import LinearDiscriminantAnalysis as LDA
- from sklearn.pipeline import make_pipeline
- from sklearn.preprocessing import StandardScaler
- from sklearn.svm import LinearSVC as SVM
- from sklearn.svm import SVC
- from scipy.ndimage import label
- import pickle
- from matplotlib import patches
- from numpy.linalg import norm
- import matplotlib.gridspec as gridspec
- import matplotlib
- matplotlib.use('TkAgg')
- plt.rcParams['svg.fonttype'] = 'none'
- from itertools import chain
- import scipy.io as io
- import yaml
- from scipy import signal
- from sklearn.base import BaseEstimator, ClassifierMixin
- from sklearn.utils.validation import check_X_y, check_array, check_is_fitted
- from sklearn.utils.multiclass import unique_labels
- from scipy.stats import zscore
- from joblib import Parallel, delayed
- from numba import jit, prange
- import multiprocessing as mp
- import os
- @jit(nopython=True, parallel=True)
- def fast_pearsonr_matrix(X, Y):
- """
- Fast computation of Pearson correlation matrix using Numba.
- X: (n_features, n_timepoints1)
- Y: (n_features, n_timepoints2)
- Returns: (n_timepoints1, n_timepoints2)
- """
- n_features, n_timepoints1 = X.shape
- _, n_timepoints2 = Y.shape
- corr_matrix = np.zeros((n_timepoints1, n_timepoints2))
- for i in prange(n_timepoints1):
- for j in prange(n_timepoints2):
- x = X[:, i]
- y = Y[:, j]
- # Compute correlation manually for speed
- x_mean = np.mean(x)
- y_mean = np.mean(y)
- num = np.sum((x - x_mean) * (y - y_mean))
- den_x = np.sqrt(np.sum((x - x_mean) ** 2))
- den_y = np.sqrt(np.sum((y - y_mean) ** 2))
- if den_x == 0 or den_y == 0:
- corr_matrix[i, j] = 0
- else:
- corr_matrix[i, j] = num / (den_x * den_y)
- return corr_matrix
- @jit(nopython=True, parallel=True)
- def fast_pearsonr_vector(X, Y):
- """
- Fast computation of Pearson correlation vector using Numba.
- X: (n_features, n_timepoints)
- Y: (n_features, n_timepoints)
- Returns: (n_timepoints,)
- """
- n_features, n_timepoints = X.shape
- corr_vector = np.zeros(n_timepoints)
- for t in prange(n_timepoints):
- x = X[:, t]
- y = Y[:, t]
- x_mean = np.mean(x)
- y_mean = np.mean(y)
- num = np.sum((x - x_mean) * (y - y_mean))
- den_x = np.sqrt(np.sum((x - x_mean) ** 2))
- den_y = np.sqrt(np.sum((y - y_mean) ** 2))
- if den_x == 0 or den_y == 0:
- corr_vector[t] = 0
- else:
- corr_vector[t] = num / (den_x * den_y)
- return corr_vector
- def temp_dec_stages_permutation(data_eq, labels_eq, variable_mapping, tranc_window=[40, 160], n_permutations=100,
- random_state=42, method="Pearson"):
- """
- Compute temporal generalization matrices with original and scrambled labels.
- Uses the same permuted labels across all windows for each permutation.
- Parameters:
- -----------
- data_eq : list of ndarrays
- List of data arrays for each stage
- labels_eq : list of ndarrays
- List of label arrays for each stage
- variable_mapping : list or array
- Mapping of labels to factors
- tranc_window : list
- Time window to analyze [start, end]
- n_permutations : int
- Number of permutation iterations to perform
- random_state : int
- Random seed for reproducibility
- Returns:
- --------
- decoding : ndarray
- Original temporal generalization matrices
- decoding_perm : ndarray
- Permuted temporal generalization matrices
- """
- n_windows = data_eq[0].shape[0]
- n_stages = len(data_eq)
- window_size = tranc_window[1] - tranc_window[0] + 1
- rng = np.random.RandomState(random_state)
- # Initialize arrays for original and permuted decoding
- decoding = np.zeros((n_stages, n_windows, window_size, window_size))
- decoding_perm = np.zeros((n_permutations, n_stages, n_windows, window_size, window_size))
- # Convert mapping to numpy array if it's not already
- colour_fac = np.array(variable_mapping)
- # Compute original decoding matrices first
- for i_stage in range(n_stages):
- y_stage = labels_eq[i_stage][0, :]
- for i_window in range(n_windows):
- # Extract data for this stage and window
- X_stage = data_eq[i_stage][i_window, :, :, tranc_window[0]:tranc_window[1] + 1]
- # Compute original labels
- y_colour = np.array(assign_lables(y_stage, colour_fac))
- if method == "SVM":
- decoding[i_stage, i_window, :, :] = decode_time(X_stage, y_colour, n_inter=1)
- elif method == "Pearson":
- decoder = NeuralCorrelationDecoder(across_time=True)
- decoder.fit(X_stage, y_colour)
- decoding[i_stage, i_window, :, :] = decoder.get_correlation()
- print(f'Constructing null')
- # Now compute permuted decoding matrices
- for perm in tqdm(range(n_permutations)):
- # For each stage, generate ONE permutation to use across all windows
- for i_stage in range(n_stages):
- # Get the first window's labels to determine permutation structure
- y_stage_first = labels_eq[i_stage][0, :]
- y_colour_first = np.array(assign_lables(y_stage_first, colour_fac))
- # Create a single permutation for this stage
- y_perm_indices = rng.permutation(len(y_colour_first))
- # Use this permutation for all windows in this stage
- for i_window in range(n_windows):
- # Extract data for this stage and window
- X_stage = data_eq[i_stage][i_window, :, :, tranc_window[0]:tranc_window[1] + 1]
- y_stage = labels_eq[i_stage][i_window, :]
- # Get original color labels
- y_colour = np.array(assign_lables(y_stage, colour_fac))
- # Apply the same permutation pattern to these labels
- y_perm = y_colour[y_perm_indices]
- # Compute decoding matrix with permuted labels
- if method == "SVM":
- decoding_perm[perm, i_stage, i_window, :, :] = decode_time(X_stage, y_perm, n_inter=1)
- elif method == "Pearson":
- decoder = NeuralCorrelationDecoder(across_time=True)
- decoder.fit(X_stage, y_perm)
- decoding_perm[perm, i_stage, i_window, :, :] = decoder.get_correlation()
- return decoding, decoding_perm
- def _single_iteration(X, y, iteration_seed, across_time):
- """
- Single iteration of cross-validation for parallelization.
- Parameters
- ----------
- X : ndarray
- Input data
- y : ndarray
- Labels
- iteration_seed : int
- Random seed for this iteration
- across_time : bool
- Whether to compute across-time correlations
- Returns
- -------
- correlation_result : ndarray
- Correlation matrix or vector for this iteration
- """
- def condi_avg(data, labels):
- """Optimized for binary classification."""
- mask = labels.astype(bool)
- condition_0 = data[~mask].mean(axis=0)
- condition_1 = data[mask].mean(axis=0)
- return np.array([condition_0, condition_1])
- # Set seed for this iteration (convert numpy int to Python int)
- seed_value = int(iteration_seed) # Convert to native Python int
- np.random.seed(seed_value)
- random.seed(seed_value)
- n_trials = X.shape[0]
- # Create random split
- idx_rnd = np.concatenate([np.zeros(n_trials // 2), np.ones(n_trials // 2)])
- if (n_trials % 2) > 0:
- idx_rnd = np.concatenate([idx_rnd, [1.0]])
- np.random.shuffle(idx_rnd)
- # Split data
- X1 = X[idx_rnd < 1, :, :]
- y1 = y[idx_rnd < 1]
- X2 = X[idx_rnd > 0, :, :]
- y2 = y[idx_rnd > 0]
- # Z-score
- X1 = zscore(X1, axis=1)
- X2 = zscore(X2, axis=1)
- # Compute condition averages
- condition_averages1 = condi_avg(X1, y1)
- condition_averages2 = condi_avg(X2, y2)
- # Compute differences between conditions
- diffTrain = condition_averages1[0, :, :] - condition_averages1[1, :, :]
- diffTest = condition_averages2[0, :, :] - condition_averages2[1, :, :]
- # Compute correlations using fast functions
- if across_time:
- # Temporal generalization: correlate across all time point pairs
- corr1 = fast_pearsonr_matrix(diffTrain, diffTest)
- corr2 = fast_pearsonr_matrix(diffTest, diffTrain)
- return np.tanh(np.mean(np.arctanh(np.array([corr1, corr2])), axis=0))
- else:
- # Within-time correlations only
- return fast_pearsonr_vector(diffTrain, diffTest)
- class NeuralCorrelationDecoder(BaseEstimator, ClassifierMixin):
- """
- Optimized neural correlation decoder with parallel processing and JIT compilation.
- Parameters
- ----------
- n_iterations : int, default=10
- Number of cross-validation iterations to perform
- across_time : bool, default=True
- If True, compute correlations across all time point pairs (temporal generalization)
- If False, compute correlations only within the same time points
- random_state : int, RandomState instance or None, default=None
- Controls the randomness of the cross-validation splits
- n_jobs : int, default=-1
- Number of parallel jobs. -1 means use all available cores
- batch_size : int, default=None
- Batch size for parallel processing. If None, uses n_iterations // n_cores
- """
- def __init__(self, n_iterations=10, across_time=True, random_state=None,
- n_jobs=-1, batch_size=None):
- self.n_iterations = n_iterations
- self.across_time = across_time
- self.random_state = random_state
- self.n_jobs = n_jobs
- self.batch_size = batch_size
- def _validate_input(self, X, y=None):
- """Validate input data format."""
- if X.ndim != 3:
- raise ValueError(f"Expected 3D array (n_trials, n_features, n_timepoints), "
- f"got {X.ndim}D array")
- if y is not None:
- if len(np.unique(y)) != 2:
- raise ValueError("This decoder only supports binary classification "
- f"(2 classes), got {len(np.unique(y))} classes")
- return X, y
- def fit(self, X, y):
- """
- Fit the neural correlation decoder using parallel processing.
- Parameters
- ----------
- X : array-like of shape (n_trials, n_features, n_timepoints)
- Training data
- y : array-like of shape (n_trials,)
- Target values (binary classification)
- Returns
- -------
- self : object
- Returns the instance itself
- """
- # Validate inputs
- X, y = check_X_y(X, y, allow_nd=True)
- X, y = self._validate_input(X, y)
- # Store classes
- self.classes_ = unique_labels(y)
- # Generate seeds for each iteration to ensure reproducibility
- if self.random_state is not None:
- np.random.seed(self.random_state)
- seeds = np.random.randint(0, 2 ** 31, self.n_iterations)
- else:
- seeds = np.random.randint(0, 2 ** 31, self.n_iterations)
- # Determine number of jobs
- n_jobs = self.n_jobs
- if n_jobs == -1:
- n_jobs = mp.cpu_count()
- elif n_jobs <= 0:
- n_jobs = max(1, mp.cpu_count() + n_jobs)
- # Parallel computation of iterations
- # Handle batch_size properly for joblib
- parallel_kwargs = {'n_jobs': n_jobs}
- if self.batch_size is not None:
- parallel_kwargs['batch_size'] = self.batch_size
- results = Parallel(**parallel_kwargs)(
- delayed(_single_iteration)(X, y, seed, self.across_time)
- for seed in seeds
- )
- # Convert results to numpy array and average using Fisher z-transform
- corrs_iters = np.array(results)
- # Handle potential NaN/inf values before arctanh
- corrs_iters = np.clip(corrs_iters, -0.99999, 0.99999)
- self.correlation_matrix_ = np.tanh(np.mean(np.arctanh(corrs_iters), axis=0))
- self.is_fitted_ = True
- return self
- def get_correlation(self):
- """Get the fitted correlation matrix."""
- check_is_fitted(self, 'is_fitted_')
- return self.correlation_matrix_
- def set_params(self, **params):
- """Set the parameters of this estimator."""
- for key, value in params.items():
- if hasattr(self, key):
- setattr(self, key, value)
- else:
- raise ValueError(f"Invalid parameter {key}")
- return self
- def get_params(self, deep=True):
- """Get parameters for this estimator."""
- return {
- 'n_iterations': self.n_iterations,
- 'across_time': self.across_time,
- 'random_state': self.random_state,
- 'n_jobs': self.n_jobs,
- 'batch_size': self.batch_size
- }
- def decode_time(X, y, n_inter=1, return_inter=False):
- y = np.array(y)
- scores_all = []
- for n in range(n_inter):
- clf = make_pipeline(
- StandardScaler(),
- SVM(C=5e-4))
- time_gen = GeneralizingEstimator(clf, scoring="roc_auc", n_jobs=-1, verbose=True)
- n_trls = X.shape[0]
- idx_rnd = np.concatenate([np.zeros(n_trls // 2), np.ones(n_trls // 2)])
- if (n_trls % 2) > 0:
- idx_rnd = np.concatenate([idx_rnd, [1.0]])
- random.shuffle(idx_rnd)
- X1 = X[idx_rnd < 1, :, :]
- y1 = y[idx_rnd < 1]
- X2 = X[idx_rnd > 0, :, :]
- y2 = y[idx_rnd > 0]
- time_gen.fit(X1, y1)
- score1 = time_gen.score(X2, y2)
- time_gen.fit(X2, y2)
- score2 = time_gen.score(X1, y1)
- scores_all.append(np.array([score1, score2]).mean(0))
- if return_inter:
- return scores_all
- else:
- return np.mean(np.array(scores_all), axis=0)
- def dy_mask(obs, rnd, clusteralpha=0.05):
- '''
- looks for significant off diagonal reduction
- Params:
- obs: 2d observations
- rnd: 3d randomizations (randomizations are in trailing dimension)
- Outputs:
- dynaMask: 1d dynamism index (time-resolved)
- '''
- # Computing the mask
- # get the diagonal twice, reshape to allow easy comparison with full matrix
- dynaObsA = np.diag(obs)[:, np.newaxis] - obs
- dynaObsB = np.diag(obs)[np.newaxis, :] - obs
- # do the same for the randomizations
- dynaRndA = np.diagonal(rnd)[:, :, np.newaxis] - rnd.T
- dynaRndB = np.diagonal(rnd)[:, np.newaxis, :] - rnd.T
- dynaRndA = dynaRndA.T
- dynaRndB = dynaRndB.T
- pDynaA, labDynaA = permutation_test(obsdat=dynaObsA, rnddat=dynaRndA, tail=1, clusteralpha=clusteralpha)
- pDynaB, labDynaB = permutation_test(dynaObsB, dynaRndB, tail=1, clusteralpha=clusteralpha)
- dynaMask = (pDynaA < 0.05) & (pDynaB < 0.05)
- di = (np.mean(dynaMask, axis=0) + np.mean(dynaMask, axis=1)) / 2
- return di, dynaMask
- def plot_time_generalisation(gs, fig, dec, rnd=None, times_sec=None, vmin=-1, vmax=1.0,
- matrix_limits=None, tick_positions=None,
- reference_lines=None, cmap="RdBu_r", title= "Coding stability", stat_test ='off-diagonal', clusteralpha=0.05, tail=0):
- """
- Plot time generalisation matrix with both dynamism and stability indices
- This function creates a comprehensive visualisation of temporal generalisation
- within a specified grid position.
- Parameters:
- -----------
- gs : matplotlib GridSpec slice
- GridSpec slice (e.g., gs[0,0] for single cell or gs[0,0:2] for spanning cells)
- fig : matplotlib Figure
- The figure object to add the subplot to
- dec : ndarray
- The decoding matrix to display (observation)
- rnd : ndarray, optional
- Randomization matrix with shape (time, time, n_randomizations)
- times_sec : ndarray, optional
- Time points in seconds. If None, creates a default range
- vmin, vmax : float
- Min and max values for color mapping
- matrix_limits : tuple, optional
- (min_idx, max_idx) limits for both x and y axes
- tick_positions : list or ndarray, optional
- Indices where to place ticks
- reference_lines : list, optional
- Time points where to place reference lines
- cmap : str
- Colormap to use
- """
- import numpy as np
- import matplotlib.pyplot as plt
- from scipy.ndimage import label
- # Get current figure (use the provided fig parameter)
- # fig is already provided as parameter
- # Create default time range if not provided
- if times_sec is None:
- times_sec = np.linspace(-0.1, 1.1, dec.shape[-1])
- # Create default tick positions if not provided
- if tick_positions is None:
- tick_positions = np.array([10, 60, 110])
- # Create default reference lines if not provided
- if reference_lines is None:
- reference_lines = [0, 0.5, 1.0]
- # Create subplot using the provided GridSpec
- ax = fig.add_subplot(gs)
- # Plot the main time generalization matrix
- im = ax.matshow(
- dec,
- vmin=vmin,
- vmax=vmax,
- cmap=cmap,
- origin="lower",
- )
- # Compute dynamism and stability if randomization data is provided
- if rnd is not None:
- if stat_test == 'off-diagonal':
- # Get dynamism index and significance mask
- _, sig_mask = dy_mask(dec, rnd, clusteralpha=clusteralpha)
- elif stat_test == 'on-diagonal':
- p_vals, cluster_labels = permutation_test(
- obsdat=dec,
- rnddat=rnd,
- clustercorrect=True,
- clusteralpha=clusteralpha,
- tail = tail,
- )
- # Create mask for significant areas (p <= 0.05)
- sig_mask = p_vals <= 0.05
- # Plot significance masks on the main plot
- if np.any(sig_mask):
- # Use contour to outline significant dynamism areas
- ax.contour(sig_mask, levels=[0.5], colors='black',
- linestyles='-', linewidths=0.8)
- # Set tick labels
- tick_labels = [f"{times_sec[i]:.1f}" for i in tick_positions]
- # Set the ticks and labels
- ax.set_xticks(tick_positions)
- ax.set_xticklabels(tick_labels)
- ax.set_yticks(tick_positions)
- ax.set_yticklabels(tick_labels)
- # Set matrix display limits if provided
- if matrix_limits is not None:
- min_idx, max_idx = matrix_limits
- ax.set_xlim(min_idx, max_idx)
- ax.set_ylim(min_idx, max_idx)
- # Make sure ticks are at the bottom
- ax.xaxis.set_ticks_position("bottom")
- # Add reference lines to main plot
- for time_val in reference_lines:
- time_idx = np.argmin(np.abs(times_sec - time_val))
- ax.axhline(time_idx, color="k", linestyle="--", linewidth=0.8)
- ax.axvline(time_idx, color="k", linestyle="--", linewidth=0.8)
- # Add axis labels
- ax.set_xlabel('time (s)')
- ax.set_ylabel('time (s)')
- # Add title
- ax.set_title(title)
- # Add colorbar
- cbar = fig.colorbar(im, ax=ax)
- cbar.set_label("Pearson r")
- def plot_magnitude(gs, fig, mag_obs, mag_null_within, mag_null_ler,
- title, y_lim=[-1,1], alpha=0.05):
- """
- Plot cross-generalization magnitude across learning stages as bars.
- Parameters
- ----------
- mag_obs : array, shape (n_stages,)
- Observed magnitude values
- mag_null_within : array, shape (n_permutations, n_stages)
- Null distribution where only cross-gen is shuffled
- mag_null_ler : array, shape (n_permutations, n_stages)
- Full permutation null (both within and cross-gen shuffled)
- """
- x_data = np.arange(len(mag_obs)) + 1
- ax = fig.add_subplot(gs)
- # Bar plot
- ax.bar(x_data, mag_obs, color='black', width=0.6, zorder=5)
- ax.set_xticks(x_data)
- ax.set_xticklabels(x_data)
- sns.despine(right=True, top=True)
- ax.set_ylabel('% of ceiling')
- ax.set_xlabel('learning stage')
- ax.set_title(title)
- # --- Null band from mag_null_within ---
- null_flat = mag_null_within.flatten()
- null_mean = null_flat.mean()
- # Center around mean
- null_centered = null_flat - null_mean
- crit = np.quantile(np.abs(null_centered), 1 - alpha)
- lower = null_mean - crit
- upper = null_mean + crit
- # Mean line
- ax.axhline(null_mean, linestyle='--', linewidth=1,
- color='black', zorder=-6, alpha=0.7)
- ax.axhline(1, linestyle='--', linewidth=1,
- color='black', zorder=-6, alpha=0.7)
- # Null band
- rect = patches.Rectangle(
- (x_data[0] - 0.5, lower),
- len(x_data),
- upper - lower,
- edgecolor=None,
- facecolor='lightgrey',
- zorder=-7,
- )
- ax.add_patch(rect)
- # --- Horizontal bracket: Learning effect (Stage 1 vs Stage 4) ---
- # Bracket displayed from BELOW with arms pointing UP
- y_range = y_lim[1] - y_lim[0]
- bracket_base = y_lim[0] + 0.2 * y_range # Changed from y_lim[1] - 0.11
- bracket_height = bracket_base - 0.06 * y_range # Changed sign to go down
- bracket_center = (x_data[0] + x_data[-1]) / 2
- gap_size = 1.2
- print('Learning effect (magnitude increase):')
- p_value_learning = compute_p_value(mag_obs[-1], mag_obs[0],
- mag_null_ler[:, -1], mag_null_ler[:, 0],
- tail='greater')
- # Draw bracket - arms now point UP from below
- ax.plot([x_data[0], x_data[0]], [bracket_height, bracket_base],
- color='black', linewidth=1)
- ax.plot([x_data[-1], x_data[-1]], [bracket_height, bracket_base],
- color='black', linewidth=1)
- ax.plot([x_data[0], bracket_center - gap_size / 2],
- [bracket_height, bracket_height], color='black', linewidth=1)
- ax.plot([bracket_center + gap_size / 2, x_data[-1]],
- [bracket_height, bracket_height], color='black', linewidth=1)
- ax.text(bracket_center, bracket_height * 1.2, p_into_stars(p_value_learning),
- fontsize=10, color='black', ha='center', va='bottom') # Changed va to 'bottom'
- ax.set_ylim(y_lim)
- return ax, p_value_learning
- def plot_significance_stars(ax, scores_real, scores_null, time_window, ylim, tail='two', y_position=None, dec_type=0):
- """
- Compute p-value and plot significance stars on a time-resolved decoding plot.
- Parameters:
- -----------
- ax : matplotlib axis object
- The axis to plot on
- scores_real : array
- Real decoding scores of shape [decoding_type, time_points]
- scores_null : array
- Null distribution scores of shape [n_reps, decoding_type, time_points]
- time_window : tuple or list
- (start, end) in milliseconds for the time window
- ylim : tuple or list
- (ymin, ymax) of the plot
- tail : str
- 'two', 'greater', or 'less' for the statistical test
- y_position : float, optional
- Y-position for stars. If None, uses 90% of ylim range
- """
- # Compute p-value comparing first and last timepoint
- p_val = compute_p_value(scores_real[dec_type, 0], scores_real[dec_type, -1],
- scores_null[:, dec_type, 0], scores_null[:, dec_type, -1],
- tail=tail)
- # Convert p-value to stars
- if p_val <= 0.001:
- stars = '***'
- font_size = 15
- scaler_position = 0.8
- elif p_val <= 0.01:
- stars = '**'
- font_size = 15
- scaler_position = 0.8
- elif p_val <= 0.05:
- stars = '*'
- font_size = 15
- scaler_position = 0.8
- elif p_val <= 0.1:
- stars = '†'
- font_size = 10
- scaler_position = 0.8
- else:
- stars = 'ns'
- font_size = 10
- scaler_position = 0.9
- # Calculate center of time window in seconds
- tw_center = (((time_window[0] + time_window[1]) / 2) / 100) - 0.5
- # Calculate y-position (default to 90% of y-range for better spacing)
- if y_position is None:
- y_position = ylim[0] + scaler_position * (ylim[1] - ylim[0])
- # Plot stars
- ax.text(tw_center, y_position, stars,
- ha='center', va='center',
- fontsize=font_size,
- color='black')
- return ax, p_val
- def KLdivergence(x, y):
- """Compute the Kullback-Leibler divergence between two multivariate samples.
- Parameters
- ----------
- x : 2D array (n,d)
- Samples from distribution P, which typically represents the true
- distribution.
- y : 2D array (m,d)
- Samples from distribution Q, which typically represents the approximate
- distribution.
- Returns
- -------
- out : float
- The estimated Kullback-Leibler divergence D(P||Q).
- References
- ----------
- Pérez-Cruz, F. Kullback-Leibler divergence estimation of
- continuous distributions IEEE International Symposium on Information
- Theory, 2008.
- """
- from scipy.spatial import cKDTree as KDTree
- # Check the dimensions are consistent
- x = np.atleast_2d(x)
- y = np.atleast_2d(y)
- n, d = x.shape
- m, dy = y.shape
- assert (d == dy)
- # Build a KD tree representation of the samples and find the nearest neighbour
- # of each point in x.
- xtree = KDTree(x)
- ytree = KDTree(y)
- # Get the first two nearest neighbours for x, since the closest one is the
- # sample itself.
- r = xtree.query(x, k=2, eps=.01, p=2)[0][:, 1]
- s = ytree.query(x, k=1, eps=.01, p=2)[0]
- return -np.log(r / s).sum() * d / n + np.log(m / (n - 1.))
- def permutation_test(obsdat, rnddat, clustercorrect=True, clusteralpha=0.05, tail=1):
- """
- Performs an (optionally cluster-corrected) permutation test of the observed
- data, given the pre-computed randomizations. rnddat must have one extra
- trailing dimension compared to obsdat.
- """
- if tail == 0:
- alpha_2tail = clusteralpha / 2
- clusterThreshold_right = np.percentile(rnddat, 100 * (1 - alpha_2tail), axis=obsdat.ndim)
- clusterThreshold_left = np.percentile(rnddat, 100 * alpha_2tail, axis=obsdat.ndim)
- elif tail == -1:
- clusterThreshold = np.percentile(rnddat, 100 * clusteralpha, axis=obsdat.ndim)
- elif tail == 1:
- clusterThreshold = np.percentile(rnddat, 100 * (1 - clusteralpha), axis=obsdat.ndim)
- p = np.ones_like(obsdat, dtype='float64')
- # uncorrected 'test'
- if not clustercorrect:
- for inds, value in np.ndenumerate(obsdat):
- rnd = np.sort(rnddat[inds])
- if tail == 0:
- pval = 2 * min(
- np.searchsorted(rnd, value, side='left'),
- rnd.shape[0] - np.searchsorted(rnd, value, side='right')
- ) / rnd.shape[0]
- elif tail == -1:
- pval = np.searchsorted(rnd, value, side='right') / rnd.shape[0]
- elif tail == 1:
- pval = 1 - np.searchsorted(rnd, value, side='left') / rnd.shape[0]
- p[inds] = pval
- return p
- # subfunction to compute clusterstats in one dataset (observed/randomized)
- def compute_clusterstats(dat, getinds=True):
- if tail == 0:
- clusterCandidates_right = dat > clusterThreshold_right
- clusterCandidates_left = dat < clusterThreshold_left
- clusterCandidates = np.logical_or(clusterCandidates_right, clusterCandidates_left)
- elif tail == -1:
- clusterCandidates = dat < clusterThreshold
- elif tail == 1:
- clusterCandidates = dat > clusterThreshold
- # label connected tiles
- labelled, numfeat = label(clusterCandidates, output='uint32')
- # compute aggregate cluster statistic for each cluster
- # this is a quick way to use cluster size as the clusterstat.
- # for other clusterstats (e.g. summed stat) a few more lines are needed
- clusterNums, clusterStats = np.unique(labelled.ravel(), return_counts=True)
- # remove 0, which corresponds to a non-cluster
- if clusterNums[0] == 0:
- # note that the check is necessary because it can happen that the
- # entire observe data matrix exceeds the threshold, in that case we
- # don't want to remove the first element (which will be 1 instead
- # of zero)
- clusterNums = clusterNums[1:]
- clusterStats = clusterStats[1:]
- # use cluster sum instead of size
- clusterStats = [np.sum(dat[labelled == x]) for x in clusterNums]
- if getinds:
- return clusterStats, clusterNums, labelled
- else:
- return clusterStats
- # get observed clusters and maximum randomized clusterstats
- clusObs, clusNums, labelled = compute_clusterstats(obsdat)
- clusRnd = [compute_clusterstats(rnddat[..., x], False)
- for x in range(rnddat.shape[-1])]
- # treat randomizations with 0 cluster candidates as if their max was 0
- mymax = lambda x: 0 if len(x) == 0 else np.max(x)
- mymin = lambda x: 0 if len(x) == 0 else np.min(x)
- if tail == 0:
- clusRnd = [mymax(np.abs(x)) for x in clusRnd]
- clusObs = [np.abs(s) for s in clusObs]
- elif tail == -1:
- clusRnd = [mymin(x) for x in clusRnd]
- elif tail == 1:
- clusRnd = [mymax(x) for x in clusRnd]
- clusRnd.sort()
- for stat, num in zip(clusObs, clusNums):
- if tail == 0:
- pval = 1 - np.searchsorted(clusRnd, stat, side='left') / rnddat.shape[-1]
- elif tail == -1:
- pval = np.searchsorted(clusRnd, stat, side='right') / rnddat.shape[-1]
- elif tail == 1:
- pval = 1 - np.searchsorted(clusRnd, stat, side='left') / rnddat.shape[-1]
- p[labelled == num] = pval
- return p, labelled
- def assign_lables(labels, factor):
- conditions = np.unique(labels)
- res = dict(zip(conditions, factor))
- return list(map(res.get, labels))
- def bias_corr(sel):
- conds = np.zeros((4, 3)) # conditions x selectivity (pure color, pure shape, interaction)
- conds[0, 0] = 1
- conds[1, 1] = 1 # Not rewarded
- conds[2, :] = 0
- conds[3, :] = 1 # rewarded
- const = np.array([1, 1, 1, 1])
- X = np.vstack([const, conds[:, 0], conds[:, 1], conds[:, 2]]).T
- rates = sel @ conds.T
- coeffs = sp.linalg.lstsq(X, rates.T)[0][1:, :]
- cov_new = np.cov(coeffs)
- new_sel = np.random.multivariate_normal(np.zeros(sel.shape[1]),
- cov_new,
- sel.shape[0])
- return new_sel
- def downsample_data(*arg, fact=10):
- retval = []
- # downsample trailing dimension (i.e. time axis) by factor 10
- for dat in arg:
- shape = list(dat.shape)
- shape[-1] = shape[-1] // fact
- dat = dat.reshape(shape + [fact])
- retval.append(np.mean(dat, axis=len(shape)))
- return np.array(retval)[0, :, :, :]
- def get_data(session_list, path_spikes, path_meta, window, cut_off=False):
- parts = session_list
- data_all_parts = []
- labels_all_parts = []
- for p in range(len(parts)):
- sessions = parts[p]
- print('Loading and combining the spiking data: part ', p + 1, '/', len(parts))
- units_dat = []
- units_labels = []
- for s in tqdm(range(len(sessions))):
- spks = np.load(path_spikes.format(sessions[s]))
- spks = spks[:, :, 500:3000]
- spks = downsample_data(spks)
- trials = np.load(path_meta.format(sessions[s]))
- labels = trials[trials != 0]
- spks_data = spks[trials != 0, :, window[0]:window[1]]
- spks_data = spks_data[:, :, :]
- if cut_off:
- spks_data = spks_data[:cut_off, :, :]
- labels = labels[:cut_off]
- units_dat.append(spks_data)
- units_labels.append(labels)
- data_all_parts.append(units_dat)
- labels_all_parts.append(units_labels)
- return data_all_parts, labels_all_parts
- def exclude_neurons(data, session_list, path_locations, path_sel_exclude, loc=None, non_sig=False, threshold=0.0,
- plot=False, times=np.linspace(-0.5, 2.0, 250), ylim=50):
- n_parts = len(data)
- data_new_ses_parts = []
- n_exc_parts = []
- n_neurons_parts = []
- exc_parts = []
- parts = session_list
- for p in range(n_parts):
- data_new_ses = []
- n_exc_ses = []
- n_neurons_ses = []
- for s in range(len(data[p])):
- data_ses = data[p][s]
- data_avg = data_ses.mean(0)
- mean_firing = data_avg.mean(1)
- if loc:
- cell_vals = pd.read_csv(path_locations.format(parts[p][s]))['Area'].values
- exc_idc = (cell_vals >= loc[0]) & (cell_vals <= loc[1])
- if threshold:
- exc_thr = mean_firing > threshold
- exc_idc = np.logical_and(exc_idc, exc_thr)
- if non_sig:
- sel_exc_idc = ~np.load(path_sel_exclude.format(parts[p][s]))
- exc_idc = np.logical_and(exc_idc, sel_exc_idc)
- else:
- exc_idc = mean_firing > -1
- data_new = data_ses[:, exc_idc, :]
- data_new_ses.append(data_new)
- n_exc_ses.append(np.sum(exc_idc))
- n_neurons_ses.append(len(exc_idc))
- exc_proc = round(1 - np.sum(n_exc_ses) / np.sum(n_neurons_ses), 2)
- exc_parts.append(exc_proc)
- print('Excluded', exc_parts[p] * 100, '% of neurons from part', p + 1)
- if plot:
- colors = sns.xkcd_palette(['pale red'])
- data_avg_trl = [data_new_ses[_].mean(0) for _ in range(len(data_new_ses))]
- data_plot = np.concatenate(data_avg_trl, axis=0)
- df_plot = pd.DataFrame(data_plot.T, columns=list(range(1, data_plot.shape[0] + 1)))
- df_plot['Times'] = times
- df_melted = pd.melt(df_plot, id_vars=['Times'])
- df_melted['Rate'] = df_melted['value']
- df_melted['Neurons'] = df_melted['variable']
- plt.figure()
- sns.lineplot(data=df_melted, x="Times", y='Rate', hue="Neurons")
- plt.ylim([None, ylim])
- plt.axvline(0, linestyle="--", linedecode_epoch=0.8, color='black')
- sns.despine(right=True, top=True)
- plt.axvline(0.5, linestyle="--", linewidth=0.8, color='black')
- plt.axvline(1, linestyle="--", linewidth=0.8, color='black')
- mean_rate = df_melted.groupby('Neurons').mean()['Rate'].values
- plt.figure()
- sns.distplot(mean_rate, bins=20, kde=False, color='grey')
- sns.despine(right=True, top=True)
- median_firing_rate = round(np.median(data_plot.mean(1)), 2)
- plt.title("Median firing rate = " + str(median_firing_rate))
- plt.xlabel('Rate')
- plt.ylabel('Count')
- plt.axvline(median_firing_rate, 0, 1, linewidth=1, linestyle='--', color='black')
- exc_proc = round(1 - np.sum(n_exc_ses) / np.sum(n_neurons_ses), 2)
- plt.annotate('Exc % = ' + str(exc_proc), xy=(1, 40), color=colors[0])
- plt.xlim([None, 25])
- data_new_ses_parts.append(data_new_ses)
- n_neurons_parts.append(np.sum(n_neurons_ses))
- n_exc_parts.append(np.sum(n_neurons_ses) - np.sum(n_exc_ses))
- return data_new_ses_parts, exc_parts, n_neurons_parts
- def condi_avg(data, labels):
- conditions = np.unique(labels)
- condition_averages = []
- for _ in range(len(conditions)):
- condition_dat = data[labels == conditions[_], :, :]
- condition_averages.append(condition_dat.mean(0))
- return np.array(condition_averages)
- def plot_covs(data, model_names, if_diff=False):
- n_models = len(data)
- n_epochs = len(data[0])
- if if_diff:
- covs = np.zeros((n_models, n_epochs + 1, 3, 3))
- else:
- covs = np.zeros((n_models, n_epochs, 3, 3))
- for m in range(n_models):
- for p in range(n_epochs):
- covs[m, p, :, :] = np.cov(data[m][p][:, :, 0].T)
- if if_diff:
- covs[:, n_epochs, :, :] = covs[:, 0, :, :] - covs[:, -1, :, :]
- n_epochs = n_epochs + 1
- max_cov = np.max(covs)
- min_cov = -max_cov
- print(min_cov)
- print(max_cov)
- norm = colors.TwoSlopeNorm(vmin=min_cov, vcenter=0, vmax=max_cov)
- covs_flat = np.resize(covs, (covs.shape[0] * covs.shape[1], covs.shape[2], covs.shape[3]))
- epoch_label = list(range(1, n_epochs + 1)) * 2
- fig, big_axes = plt.subplots(figsize=(8.0, 4.0), nrows=2, ncols=1, sharey=True)
- for row, big_ax in enumerate(big_axes, start=1):
- big_ax.set_title(model_names[row - 1] + " \n", fontsize=14)
- # Turn off axis lines and ticks of the big subplot
- # obs alpha is 0 in RGBA string!
- big_ax.tick_params(labelcolor=(1., 1., 1., 0.0), top='off', bottom='off', left='off', right='off')
- big_ax.axis('off')
- # removes the white frame
- big_ax._frameon = False
- for i in range(1, n_epochs * 2 + 1):
- ax = fig.add_subplot(2, n_epochs, i)
- ax.imshow(covs_flat[i - 1, :, :], cmap=sns.diverging_palette(230, 20, as_cmap=True), norm=norm)
- if if_diff:
- if i == n_epochs or i == (n_epochs * 2):
- plt.title('difference (last vs first)')
- else:
- plt.title('epoch ' + str(epoch_label[i - 1]))
- else:
- plt.title('epoch ' + str(epoch_label[i - 1]))
- plt.axis('off')
- fig.set_facecolor('w')
- plt.tight_layout()
- plt.show()
- return covs
- def get_betas_cross_val(data, labels, condition_labels, normalisation='mean_centred', time_window=[130, 150],
- add_constant=True,
- design_model='0/1', cross_val=False, if_xgen=False, task='task_1', full_model=False):
- n_parts = len(data)
- betas_part = []
- for p in range(n_parts):
- betas = []
- for s in range(len(data[p])):
- firing_rates_ses = data[p][s]
- firing_rates_ses = firing_rates_ses[:, :, time_window[0]:time_window[1]]
- labels_ses = labels[p][s]
- if task == 'task_1':
- firing_rates_ses = firing_rates_ses[labels_ses < 9, :, :]
- labels_ses = labels_ses[labels_ses < 9]
- elif task == 'task_2':
- firing_rates_ses = firing_rates_ses[labels_ses > 8, :, :]
- labels_ses = labels_ses[labels_ses > 8]
- if cross_val:
- firing_rates_ses1 = firing_rates_ses[::2, :, :]
- firing_rates_ses2 = firing_rates_ses[1::2, :, :]
- labels_ses1 = labels_ses[::2]
- labels_ses2 = labels_ses[1::2]
- else:
- firing_rates_ses1 = firing_rates_ses
- firing_rates_ses2 = firing_rates_ses
- labels_ses1 = labels_ses
- labels_ses2 = labels_ses
- if if_xgen:
- firing_rates_ses1 = firing_rates_ses[labels_ses1 < 9, :, :]
- firing_rates_ses2 = firing_rates_ses[labels_ses2 > 8, :, :]
- labels_ses1 = labels_ses1[labels_ses1 < 9]
- labels_ses2 = labels_ses2[labels_ses2 > 8]
- firing_rates_ses = [firing_rates_ses1, firing_rates_ses2]
- labels_ses = [labels_ses1, labels_ses2]
- betas_split = []
- n_splits = len(labels_ses)
- for split in range(n_splits):
- firing_mean = condi_avg(firing_rates_ses[split], labels_ses[split])
- if normalisation == 'zscore':
- firing_mean = (firing_mean - np.mean(firing_mean, axis=0, keepdims=True)) / np.std(firing_mean,
- axis=0,
- keepdims=True)
- elif normalisation == 'mean_centred':
- firing_mean = firing_mean - np.mean(firing_mean, axis=0, keepdims=True)
- elif normalisation == 'none':
- firing_mean = firing_mean
- elif normalisation == 'soft':
- firing_mean = firing_mean / (
- (np.max(firing_mean, axis=0, keepdims=True) - np.min(firing_mean, axis=0,
- keepdims=True)) + 5)
- if design_model == '0/1':
- labels_cue = np.array(condition_labels[0])
- labels_target = np.array(condition_labels[1])
- labels_int1 = np.array(condition_labels[2])
- if full_model:
- labels_target2 = np.array(condition_labels[2])
- labels_int1 = np.array(condition_labels[3])
- labels_int2 = np.array(condition_labels[4])
- elif design_model == '+1/-1':
- labels_cue = np.array(condition_labels[0])
- labels_cue = np.where(labels_cue == 0, -1, labels_cue)
- labels_target = np.array(condition_labels[1])
- labels_target = np.where(labels_target == 0, -1, labels_target)
- labels_int1 = np.array(condition_labels[2])
- labels_int1 = np.where(labels_int1 == 0, -1, labels_int1)
- if full_model:
- labels_target2 = np.array(condition_labels[2])
- labels_target2 = np.where(labels_target2 == 0, -1, labels_target2)
- labels_int1 = np.array(condition_labels[3])
- labels_int1 = np.where(labels_int1 == 0, -1, labels_int1)
- labels_int2 = np.array(condition_labels[4])
- labels_int2 = np.where(labels_int2 == 0, -1, labels_int2)
- if full_model:
- if add_constant:
- constant = np.ones_like(labels_cue).astype(float)
- design_matrix = np.vstack(
- [constant, labels_cue, labels_target, labels_int1, labels_target2, labels_int2]).T
- else:
- design_matrix = np.vstack(
- [labels_cue, labels_target, labels_int1, labels_target2, labels_int2]).T
- else:
- if add_constant:
- constant = np.ones_like(labels_cue).astype(float)
- design_matrix = np.vstack(
- [constant, labels_cue, labels_target, labels_int1]).T
- else:
- design_matrix = np.vstack(
- [labels_cue, labels_target, labels_int1]).T
- design_matrix = design_matrix.astype(float)
- n_neurons = firing_mean.shape[1]
- n_times = firing_mean.shape[-1]
- n_models = design_matrix.shape[1]
- design_matrix_stacked = np.concatenate([design_matrix] * n_times, axis=0)
- firing_rates_ses_stacked = np.concatenate(np.array_split(firing_mean, n_times, axis=-1), axis=0)
- betas_ses = np.zeros((n_neurons, n_models, 1))
- for cell in range(n_neurons):
- betas_ses[cell, :, :] = \
- sp.linalg.lstsq(design_matrix_stacked, firing_rates_ses_stacked[:, cell, :])[0][
- :, :]
- betas_split.append(betas_ses)
- betas.append(np.array(betas_split)[:, :, :, 0])
- betas_part.append(betas)
- if full_model:
- epochs_rel = []
- for p in range(n_parts):
- if add_constant:
- epoch = np.concatenate(betas_part[p], axis=1)[:, :, 1:4].transpose(
- (0, 2, 1)) # splits x coeffs x neurons
- else:
- epoch = np.concatenate(betas_part[p], axis=0)[:, :, :3].transpose((0, 2, 1))[0, :,
- :] # coeffs x neurons
- epoch -= np.mean(epoch, axis=2, keepdims=True)
- epochs_rel.append(epoch.T)
- epochs_irrel = []
- for p in range(n_parts):
- if add_constant:
- epoch_1 = np.concatenate(betas_part[p], axis=1)[:, :, 1][:, :, None].transpose(
- (0, 2, 1)) # splits x coeffs x neurons
- epoch_2 = np.concatenate(betas_part[p], axis=1)[:, :, 4:].transpose(
- (0, 2, 1)) # splits x coeffs x neurons
- epoch = np.concatenate([epoch_1, epoch_2], axis=1)
- else:
- epoch_1 = np.concatenate(betas_part[p], axis=1)[:, :, 0][:, :, None].transpose(
- (0, 2, 1)) # splits x coeffs x neurons
- epoch_2 = np.concatenate(betas_part[p], axis=1)[:, :, 3:].transpose(
- (0, 2, 1)) # splits x coeffs x neurons
- epoch = np.concatenate([epoch_1, epoch_2], axis=1)
- epoch -= np.mean(epoch, axis=2, keepdims=True)
- epochs_irrel.append(epoch.T)
- return epochs_rel, epochs_irrel
- else:
- epochs = []
- for p in range(n_parts):
- if add_constant:
- epoch = np.concatenate(betas_part[p], axis=1)[:, :, 1:4].transpose(
- (0, 2, 1)) # splits x coeffs x neurons
- else:
- epoch = np.concatenate(betas_part[p], axis=0)[:, :, :3].transpose((0, 2, 1))[0, :,
- :] # coeffs x neurons
- epoch -= np.mean(epoch, axis=2, keepdims=True)
- epochs.append(epoch.T)
- return epochs
- def get_freqs(x):
- return {value: len(list(freq)) for value, freq in groupby(sorted(list(x)))}
- def dist_random(epochs, n_bootstraps=1000, rnd_model='gaussian (spherical)', model_names=['cue + shape', 'cue + width'],
- design_model='0/1', bon_correction=False, metric='KL divergance estimate', relative_dist=False):
- n_epochs = len(epochs[0])
- dfs = []
- p_vals = []
- for model_n in range(len(model_names)):
- KL = np.zeros((n_epochs, n_bootstraps))
- KL_r = np.zeros((n_epochs, n_bootstraps))
- KL_opt = np.zeros((n_epochs, n_bootstraps))
- for e in range(n_epochs):
- for bootstrap in tqdm(range(n_bootstraps)):
- data_train = epochs[model_n][e][:, :, 0]
- data_test = epochs[model_n][e][:, :, 1]
- if design_model == '0/1':
- opt_cov = np.zeros((3, 3))
- m = np.mean(np.diag(np.cov(data_train.T)))
- opt_cov[:2, :2] = 0.5 * m
- opt_cov[2, 2] = 2 * m
- opt_cov[2, :2] = -m;
- opt_cov[:2, 2] = -m
- elif design_model == '+1/-1':
- # m = np.mean(np.diag(np.cov(data_train.T)))
- # opt_cov = np.diag([0, 0, m * 3])
- opt_cov = np.diag([0, 0, np.cov(data_train.T)[2, 2]])
- s_opt = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
- opt_cov,
- data_train.shape[0])
- if rnd_model == 'gaussian (spherical)':
- m = np.mean(np.diag(np.cov(data_train.T)))
- shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
- np.diag([m, m, m]),
- data_train.shape[0])
- m = np.mean(np.diag(np.cov(data_test.T)))
- shuffled_2 = np.random.multivariate_normal(np.zeros(data_test.shape[1]),
- np.diag([m, m, m]),
- data_test.shape[0])
- if metric == 'KL divergance estimate':
- kl_itr = 0.5 * (KLdivergence(data_test, shuffled_1) + KLdivergence(shuffled_1, data_test))
- kl_r_itr = 0.5 * (KLdivergence(shuffled_2, shuffled_1) + KLdivergence(shuffled_1, shuffled_2))
- elif metric == 'euclidean distance':
- kl_itr = euclidean_distance(data_test, shuffled_1)
- kl_r_itr = euclidean_distance(shuffled_2, shuffled_1)
- elif metric == 'epairs':
- kl_itr = epairs_metric(data_test, shuffled_1)
- kl_r_itr = epairs_metric(shuffled_2, shuffled_1)
- elif rnd_model == 'gaussian (tied)':
- shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
- np.diag(np.diag(np.cov(data_train.T))),
- data_train.shape[0])
- shuffled_2 = np.random.multivariate_normal(np.zeros(data_test.shape[1]),
- np.diag(np.diag(np.cov(data_test.T))),
- data_test.shape[0])
- if metric == 'KL divergance estimate':
- kl_itr = 0.5 * (KLdivergence(data_test, shuffled_1) + KLdivergence(shuffled_1, data_test))
- kl_r_itr = 0.5 * (KLdivergence(shuffled_2, shuffled_1) + KLdivergence(shuffled_1, shuffled_2))
- elif metric == 'euclidean distance':
- kl_itr = euclidean_distance(data_test, shuffled_1)
- kl_r_itr = euclidean_distance(shuffled_2, shuffled_1)
- elif metric == 'epairs':
- kl_itr = epairs_metric(data_test, shuffled_1)
- kl_r_itr = epairs_metric(shuffled_2, shuffled_1)
- KL[e, bootstrap] = kl_itr
- KL_r[e, bootstrap] = kl_r_itr
- if metric == 'KL divergance estimate':
- KL_opt[e, bootstrap] = 0.5 * (KLdivergence(s_opt, shuffled_1) + KLdivergence(shuffled_1, s_opt))
- elif metric == 'euclidean distance':
- KL_opt[e, bootstrap] = euclidean_distance(s_opt, shuffled_1)
- elif metric == 'epairs':
- KL_opt[e, bootstrap] = epairs_metric(s_opt, shuffled_1)
- if relative_dist:
- KL_r_avg = np.mean(KL_r, keepdims=True, axis=-1)
- KL -= KL_r_avg
- KL_r -= KL_r_avg
- KL_opt -= KL_r_avg
- KL_opt_avg = np.mean(KL_opt, keepdims=True, axis=-1)
- KL /= KL_opt_avg
- KL_r /= KL_opt_avg
- KL_opt /= KL_opt_avg
- p = 1 * (np.sum(KL_r >= np.mean(KL, axis=-1, keepdims=True), axis=-1) / n_bootstraps)
- if bon_correction:
- p = p * n_epochs
- p_vals.append(p)
- epoch = np.sort(list(range(1, n_epochs + 1)) * n_bootstraps)
- epoch_labels = np.concatenate([epoch, epoch, epoch])
- dist_label = np.concatenate([['observed'] * (n_bootstraps * n_epochs), [rnd_model] * (n_bootstraps * n_epochs),
- ['structured'] * (n_bootstraps * n_epochs)])
- KL_df = np.reshape(KL, KL.shape[0] * KL.shape[1])
- KL_r_df = np.reshape(KL_r, KL_r.shape[0] * KL_r.shape[1])
- KL_opt_df = np.reshape(KL_opt, KL_opt.shape[0] * KL_opt.shape[1])
- KL_all_df = np.concatenate([KL_df, KL_r_df, KL_opt_df])
- df = pd.DataFrame(np.array([KL_all_df, epoch_labels, dist_label]).T,
- columns=[metric, 'learning epoch', 'distribution'])
- df[metric] = df[metric].astype(float)
- df['model'] = model_names[model_n]
- dfs.append(df)
- df_all = pd.concat(dfs)
- df_all['divergence from'] = 'random selectivity'
- print(p_vals)
- return df_all, p_vals, [KL, KL_r, KL_opt]
- def dist_structured(epochs, n_bootstraps=1000, rnd_model='gaussian (spherical)',
- model_names=['cue + shape', 'cue + width'],
- design_model='0/1', bon_correction=False, metric='KL divergance estimate', relative_dist=False):
- n_epochs = len(epochs[0])
- dfs = []
- p_vals = []
- for model_n in range(len(model_names)):
- KL = np.zeros((n_epochs, n_bootstraps))
- KL_r = np.zeros((n_epochs, n_bootstraps))
- KL_opt = np.zeros((n_epochs, n_bootstraps))
- for e in range(n_epochs):
- for bootstrap in tqdm(range(n_bootstraps)):
- data_train = epochs[model_n][e][:, :, 0]
- data_test = epochs[model_n][e][:, :, 1]
- if design_model == '0/1':
- opt_cov1 = np.zeros((3, 3))
- m = np.mean(np.diag(np.cov(data_train.T)))
- opt_cov1[:2, :2] = 0.5 * m
- opt_cov1[2, 2] = 2 * m
- opt_cov1[2, :2] = -m;
- opt_cov1[:2, 2] = -m
- opt_cov2 = np.zeros((3, 3))
- m = np.mean(np.diag(np.cov(data_test.T)))
- opt_cov2[:2, :2] = 0.5 * m
- opt_cov2[2, 2] = 2 * m
- opt_cov2[2, :2] = -m;
- opt_cov2[:2, 2] = -m
- elif design_model == '+1/-1':
- opt_cov1 = np.diag([0, 0, np.cov(data_train.T)[2, 2]])
- opt_cov2 = np.diag([0, 0, np.cov(data_test.T)[2, 2]])
- s_opt_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
- opt_cov1,
- data_train.shape[0])
- s_opt_2 = np.random.multivariate_normal(np.zeros(data_test.shape[1]),
- opt_cov2,
- data_test.shape[0])
- if rnd_model == 'gaussian (tied)':
- shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
- np.diag(np.diag(np.cov(data_train.T))),
- data_train.shape[0])
- if metric == 'KL divergance estimate':
- KL[e, bootstrap] = 0.5 * (KLdivergence(data_test, s_opt_1) + KLdivergence(s_opt_1, data_test))
- KL_r[e, bootstrap] = 0.5 * (
- KLdivergence(shuffled_1, s_opt_1) + KLdivergence(s_opt_1, shuffled_1))
- KL_opt[e, bootstrap] = 0.5 * (KLdivergence(s_opt_2, s_opt_1) + KLdivergence(s_opt_1, s_opt_2))
- elif metric == 'euclidean distance':
- KL[e, bootstrap] = euclidean_distance(data_test, s_opt_1)
- KL_r[e, bootstrap] = euclidean_distance(shuffled_1, s_opt_1)
- KL_opt[e, bootstrap] = euclidean_distance(s_opt_2, s_opt_1)
- elif metric == 'epairs':
- KL[e, bootstrap] = epairs_metric(data_test, s_opt_1)
- KL_r[e, bootstrap] = epairs_metric(shuffled_1, s_opt_1)
- KL_opt[e, bootstrap] = epairs_metric(s_opt_2, s_opt_1)
- if rnd_model == 'gaussian (spherical)':
- m = np.mean(np.diag(np.cov(data_train.T)))
- shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
- np.diag([m, m, m]),
- data_train.shape[0])
- if metric == 'KL divergance estimate':
- KL[e, bootstrap] = 0.5 * (KLdivergence(data_test, s_opt_1) + KLdivergence(s_opt_1, data_test))
- KL_r[e, bootstrap] = 0.5 * (
- KLdivergence(shuffled_1, s_opt_1) + KLdivergence(s_opt_1, shuffled_1))
- KL_opt[e, bootstrap] = 0.5 * (KLdivergence(s_opt_2, s_opt_1) + KLdivergence(s_opt_1, s_opt_2))
- elif metric == 'euclidean distance':
- KL[e, bootstrap] = euclidean_distance(data_test, s_opt_1)
- KL_r[e, bootstrap] = euclidean_distance(shuffled_1, s_opt_1)
- KL_opt[e, bootstrap] = euclidean_distance(s_opt_2, s_opt_1)
- elif metric == 'epairs':
- KL[e, bootstrap] = epairs_metric(data_test, s_opt_1)
- KL_r[e, bootstrap] = epairs_metric(shuffled_1, s_opt_1)
- KL_opt[e, bootstrap] = epairs_metric(s_opt_2, s_opt_1)
- if relative_dist:
- KL_opt_avg = np.mean(KL_opt, keepdims=True, axis=-1)
- KL -= KL_opt_avg
- KL_r -= KL_opt_avg
- KL_opt -= KL_opt_avg
- KL_r_avg = np.mean(KL_r, keepdims=True, axis=-1)
- KL /= KL_r_avg
- KL_r /= KL_r_avg
- KL_opt /= KL_r_avg
- p = 1 * (np.sum(KL_r <= np.mean(KL, axis=-1, keepdims=True), axis=-1) / n_bootstraps)
- if bon_correction:
- p = p * n_epochs
- p_vals.append(p)
- epoch = np.sort(list(range(1, n_epochs + 1)) * n_bootstraps)
- epoch_labels = np.concatenate([epoch, epoch, epoch])
- dist_label = np.concatenate([['observed'] * (n_bootstraps * n_epochs),
- [rnd_model] * (n_bootstraps * n_epochs),
- ['structured'] * (n_bootstraps * n_epochs)])
- KL_df = np.reshape(KL, KL.shape[0] * KL.shape[1])
- KL_r_df = np.reshape(KL_r, KL_r.shape[0] * KL_r.shape[1])
- KL_opt_df = np.reshape(KL_opt, KL_opt.shape[0] * KL_opt.shape[1])
- KL_all_df = np.concatenate([KL_df, KL_r_df, KL_opt_df])
- df = pd.DataFrame(np.array([KL_all_df, epoch_labels, dist_label]).T,
- columns=[metric, 'learning epoch', 'distribution'])
- df[metric] = df[metric].astype(float)
- df['model'] = model_names[model_n]
- dfs.append(df)
- df_all = pd.concat(dfs)
- df_all['divergence from'] = 'structured selectivity'
- print(p_vals)
- return df_all, p_vals, [KL, KL_r, KL_opt]
- def fit_lines(data1, data2, if_xgen=False):
- r_best_n, r_opt_n = [], []
- for side in range(2):
- if side == 0:
- data_train = data1
- data_test = data2
- elif side == 1:
- data_train = data2
- data_test = data1
- N = data_train.shape[0]
- # colour as a regressor
- A1 = np.ones((N, 2)) # With intercept/constant
- A1[:, 1] = data_train[:, 0] # Use color selectivities as the regressor
- x_shape_c, _, _, _ = np.linalg.lstsq(A1, data_train[:, 1], rcond=-1) # Fit to the shape selectivies
- x_interaction_c, _, _, _ = np.linalg.lstsq(A1, data_train[:, 2], rcond=-1) # Fit to the interaction selectivies
- A1[:, 1] = data_test[:, 0] # Use color selectivities as the regressor
- pred_best1 = np.array([A1[:, 1] * x_shape_c[1], A1[:, 1] * x_interaction_c[1]]).T
- r_best1 = r2_score(data_test[:, 1:], pred_best1)
- # shape as a regressor
- A2 = np.ones((N, 2)) # With intercept/constant
- A2[:, 1] = data_train[:, 1] # Use shape selectivities as the regressor
- x_colour_s, _, _, _ = np.linalg.lstsq(A2, data_train[:, 0], rcond=-1) # Fit to the colour selectivies
- x_interaction_s, _, _, _ = np.linalg.lstsq(A2, data_train[:, 2], rcond=-1) # Fit to the interaction selectivies
- A2[:, 1] = data_test[:, 1] # Use color selectivities as the regressor
- pred_best2 = np.array([A2[:, 1] * x_colour_s[1], A2[:, 1] * x_interaction_s[1]]).T
- data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, -1][:, None]], axis=1)
- r_best2 = r2_score(data_compare, pred_best2)
- # interaction as a regressor
- A3 = np.ones((N, 2)) # With intercept/constant
- A3[:, 1] = data_train[:, 2] # Use interaction selectivities as the regressor
- x_colour_int, _, _, _ = np.linalg.lstsq(A3, data_train[:, 0], rcond=-1) # Fit to the colour selectivies
- x_shape_int, _, _, _ = np.linalg.lstsq(A3, data_train[:, 1], rcond=-1) # Fit to the shape selectivies
- A3[:, 1] = data_test[:, 2] # Use interaction selectivities as the regressor
- pred_best3 = np.array([A3[:, 1] * x_colour_int[1], A3[:, 1] * x_shape_int[1]]).T
- data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, 1][:, None]], axis=1)
- r_best3 = r2_score(data_compare, pred_best3)
- if if_xgen:
- pred_opt1 = np.zeros((data_train.shape[0], 2))
- pred_opt1[:, 0] = data_train[:, 0]
- pred_opt1[:, 1] = -2 * data_train[:, 0]
- r_opt1 = r2_score(data_test[:, 1:], pred_opt1)
- else:
- pred_opt1 = np.zeros((data_test.shape[0], 2))
- pred_opt1[:, 0] = data_test[:, 0]
- pred_opt1[:, 1] = -2 * data_test[:, 0]
- r_opt1 = r2_score(data_test[:, 1:], pred_opt1)
- if if_xgen:
- pred_opt2 = np.zeros((data_train.shape[0], 2))
- pred_opt2[:, 0] = data_train[:, 1]
- pred_opt2[:, 1] = -2 * data_train[:, 1]
- data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, -1][:, None]], axis=1)
- r_opt2 = r2_score(data_compare, pred_opt2)
- else:
- pred_opt2 = np.zeros((data_test.shape[0], 2))
- pred_opt2[:, 0] = data_test[:, 1]
- pred_opt2[:, 1] = -2 * data_test[:, 1]
- data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, -1][:, None]], axis=1)
- r_opt2 = r2_score(data_compare, pred_opt2)
- if if_xgen:
- pred_opt3 = np.zeros((data_train.shape[0], 2))
- pred_opt3[:, 0] = data_train[:, -1] / -2
- pred_opt3[:, 1] = data_train[:, -1] / -2
- data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, 1][:, None]], axis=1)
- r_opt3 = r2_score(data_compare, pred_opt3)
- else:
- pred_opt3 = np.zeros((data_test.shape[0], 2))
- pred_opt3[:, 0] = data_test[:, -1] / -2
- pred_opt3[:, 1] = data_test[:, -1] / -2
- data_compare = np.concatenate([data_test[:, 0][:, None], data_test[:, 1][:, None]], axis=1)
- r_opt3 = r2_score(data_compare, pred_opt3)
- r_best_n.append(np.mean([r_best1, r_best2, r_best3]))
- r_opt_n.append(np.mean([r_opt1, r_opt2, r_opt3]))
- r_best = np.mean(r_best_n)
- r_opt = np.mean(r_opt_n)
- return r_best, r_opt
- def r_squared_pop(epochs, rnd_model='gaussian (spherical)', n_bootstraps=1000,
- model_names=['cue + shape', 'cue + width'], break_axis=1.5, if_xgen=False):
- dfs = []
- n_epochs = len(epochs[0])
- for model_n in range(len(model_names)):
- fits_data = np.zeros((n_epochs, 2))
- fits_opt = np.zeros((n_bootstraps, n_epochs, 2))
- fits_rnd = np.zeros((n_bootstraps, n_epochs, 2))
- for e in range(n_epochs):
- s_train = epochs[model_n][e][:, :, 0]
- s_test = epochs[model_n][e][:, :, 1]
- # remove mean across neurons so we don't have to fit an intercept
- s_train -= np.mean(s_train, axis=0, keepdims=True) # length 3 vector
- s_test -= np.mean(s_test, axis=0, keepdims=True) # length 3 vector
- fits_data[e, :] = fit_lines(s_train, s_test, if_xgen=if_xgen)
- for bootstrap in tqdm(range(n_bootstraps)):
- opt_cov = np.eye(3)
- m = np.mean(np.diag(np.cov(s_train.T)))
- opt_cov[:2, :2] = 0.5 * m
- opt_cov[2, 2] = 2 * m
- opt_cov[2, :2] = -m
- opt_cov[:2, 2] = -m
- opt_cov = np.eye(3)
- m = np.mean(np.diag(np.cov(s_test.T)))
- opt_cov[:2, :2] = 0.5 * m
- opt_cov[2, 2] = 2 * m
- opt_cov[2, :2] = -m
- opt_cov[:2, 2] = -m
- s_train_opt = np.random.multivariate_normal(np.zeros(s_train.shape[1]),
- opt_cov,
- s_train.shape[0])
- s_test_opt = np.random.multivariate_normal(np.zeros(s_test.shape[1]),
- opt_cov,
- s_test.shape[0])
- fits_opt[bootstrap, e, :] = fit_lines(s_train_opt, s_test_opt)
- fits_opt[bootstrap, e, 0] += 0.05
- fits_opt[bootstrap, e, 1] -= 0.05
- if rnd_model == 'gaussian (spherical)':
- m = np.mean(np.diag(np.cov(s_train.T)))
- s_train_rnd = np.random.multivariate_normal(np.zeros(s_train.shape[1]),
- np.diag([m, m, m]),
- s_train.shape[0])
- m = np.mean(np.diag(np.cov(s_test.T)))
- s_test_rnd = np.random.multivariate_normal(np.zeros(s_test.shape[1]),
- np.diag([m, m, m]),
- s_test.shape[0])
- elif rnd_model == 'gaussian (tied)':
- s_train_rnd = np.random.multivariate_normal(np.zeros(s_train.shape[1]),
- np.diag(np.diag(np.cov(s_train.T))),
- s_train.shape[0])
- s_test_rnd = np.random.multivariate_normal(np.zeros(s_test.shape[1]),
- np.diag(np.diag(np.cov(s_test.T))),
- s_test.shape[0])
- fits_rnd[bootstrap, e, :] = fit_lines(s_train_rnd, s_test_rnd)
- if break_axis:
- fits_rnd[bootstrap, e, 1] += break_axis
- epoch_data = list(range(1, n_epochs + 1)) * 2
- epoch = np.concatenate([np.sort(list(range(1, n_epochs + 1)) * n_bootstraps)] * 2)
- epoch_labels = np.concatenate([epoch_data, epoch, epoch])
- gen_model = np.concatenate([['data'] * (n_epochs * 2),
- ['gaussian'] * (n_epochs * 2 * n_bootstraps),
- ['optimal'] * (n_epochs * 2 * n_bootstraps),
- ])
- dist_label = np.concatenate([['best fit line'] * n_epochs,
- ['optimal XOR line'] * n_epochs,
- ['best fit line'] * (n_epochs * n_bootstraps),
- ['optimal XOR line'] * (n_epochs * n_bootstraps),
- ['best fit line'] * (n_epochs * n_bootstraps),
- ['optimal XOR line'] * (n_epochs * n_bootstraps),
- ])
- data_df = np.reshape(fits_data, fits_data.shape[0] * fits_data.shape[1], order='F')
- rnd_df = np.reshape(fits_rnd, fits_rnd.shape[0] * fits_rnd.shape[1] * fits_rnd.shape[2], order='F')
- opt_df = np.reshape(fits_opt, fits_opt.shape[0] * fits_opt.shape[1] * fits_opt.shape[2], order='F')
- all_df = np.concatenate([data_df, rnd_df, opt_df])
- df = pd.DataFrame(np.array([all_df, epoch_labels, dist_label, gen_model]).T,
- columns=['r squared', 'learning epoch', 'fitted model', 'generative model'])
- df['r squared'] = df['r squared'].astype(float)
- df['model'] = model_names[model_n]
- dfs.append(df)
- return pd.concat(dfs)
- def xor_sim_noise(sig_min, sig_max, sig_n, n_neurons=400):
- s_list = np.linspace(sig_min, sig_max, sig_n)
- s_tot = s_list.shape[0]
- N = n_neurons
- n_coeffs = 3
- conds = np.ones((4, n_coeffs))
- conds[0, :2] = -1
- conds[1, [0, 2]] = -1
- conds[2, 1:] = -1
- targets = [1, -1, -1, 1] # XOR
- clf1 = linear_model.LogisticRegression() # 1 is no reg
- n = 100
- perf = np.zeros((s_tot, n, 2))
- print('XOr decoding performance as a function of noice (sigma) - running simulation ')
- for s_c, sig in tqdm(enumerate(s_list)):
- for i in range(n):
- # Random selectivity
- cov = np.eye(3)
- s = np.random.multivariate_normal(np.zeros(3), cov, N)
- r = s @ conds.T
- r_train = r.T + sig * np.random.normal(0, 1, (4, N))
- r_train = r_train - np.mean(r_train, axis=1, keepdims=True)
- clf1.fit(r_train, targets)
- r_test = r.T + sig * np.random.normal(0, 1, (4, N))
- r_test = r_test - np.mean(r_test, axis=1, keepdims=True)
- perf[s_c, i, 0] = clf1.score(r_test, targets)
- # Structured selectivity
- opt_cov = np.diag([0, 0, 3])
- s = np.random.multivariate_normal(np.zeros(3), opt_cov, N)
- r = s @ conds.T
- r_train = r.T + sig * np.random.normal(0, 1, (4, N))
- r_train = r_train - np.mean(r_train, axis=1, keepdims=True)
- clf1.fit(r_train, targets)
- r_test = r.T + sig * np.random.normal(0, 1, (4, N))
- r_test = r_test - np.mean(r_test, axis=1, keepdims=True)
- perf[s_c, i, 1] = clf1.score(r_test, targets)
- return perf, s_list
- def xor_sim_units(n_neurons, step_neurons, noise_sig):
- N_list = np.arange(1, n_neurons, step_neurons)
- N_tot = N_list.shape[0]
- conds = np.ones((4, 3))
- conds[0, :2] = -1
- conds[1, [0, 2]] = -1
- conds[2, 1:] = -1
- targets = [1, -1, -1, 1] # XOR
- clf1 = linear_model.LogisticRegression() # 1 is no reg
- sig = noise_sig
- n = 100
- perf_units = np.zeros((N_tot, n, 2))
- for N_c, N in enumerate(N_list):
- print(N, end=',')
- for i in range(n):
- s = np.random.normal(0, 1, (N, 3))
- r = s @ conds.T
- clf1.fit(r.T + sig * np.random.normal(0, 1, (4, N)), targets)
- perf_units[N_c, i, 0] = clf1.score(r.T + sig * np.random.normal(0, 1, (4, N)), targets)
- # Structured selectivity
- opt_cov = np.diag([0, 0, 3])
- s = np.random.multivariate_normal(np.zeros(3), opt_cov, N)
- r = s @ conds.T
- clf1.fit(r.T + sig * np.random.normal(0, 1, (4, N)), targets)
- perf_units[N_c, i, 1] = clf1.score(r.T + sig * np.random.normal(0, 1, (4, N)), targets)
- return perf_units, N_list
- def generate_boundary(sel, sig=0.9, n=100):
- conds = np.ones((4, 3))
- conds[0, :2] = -1
- conds[1, [0, 2]] = -1
- conds[2, 1:] = -1
- labels = [1, -1, -1, 1] # XOR
- neurons = sel[:2, :] # take the first two neurons
- rates = neurons @ conds.T
- rates = np.concatenate([rates] * n, axis=-1)
- rates = rates + np.random.normal(0, sig, (2, 4 * n))
- rates -= np.mean(rates, axis=1, keepdims=True)
- labels = np.concatenate([labels] * n, axis=0)
- clf = LDA()
- clf.fit(rates.T, labels)
- w = clf.coef_[0]
- a = -w[0] / w[1]
- xx = np.linspace(-10, 10)
- yy = a * xx - (clf.intercept_[0]) / w[1]
- return rates[:, :4], xx, yy
- def KL_optimal_2(epochs, n_bootstraps=1000, rnd_model='gaussian (tied)', model_names=['cue + shape', 'cue + width'],
- bon_correction=False):
- n_epochs = len(epochs[0])
- dfs = []
- p_vals = []
- for model_n in range(len(model_names)):
- KL = np.zeros((n_epochs, n_bootstraps))
- KL_r = np.zeros((n_epochs, n_bootstraps))
- for e in range(n_epochs):
- for bootstrap in tqdm(range(n_bootstraps)):
- data_train = epochs[model_n][e][:, :, 0]
- data_test = epochs[model_n][e][:, :, 1]
- s_opt = np.zeros(data_test.shape)
- s_opt[:, 0] = data_test[:, 0]
- s_opt[:, 1] = data_test[:, 0]
- s_opt[:, 2] = data_test[:, 0] * -2
- if rnd_model == 'gaussian (tied)':
- shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
- np.diag(np.diag(np.cov(data_train.T))),
- data_train.shape[0])
- KL[e, bootstrap] = 0.5 * (KLdivergence(data_test, s_opt) + KLdivergence(s_opt, data_test))
- KL_r[e, bootstrap] = 0.5 * (KLdivergence(shuffled_1, s_opt) + KLdivergence(s_opt, shuffled_1))
- if rnd_model == 'gaussian (spherical)':
- m = np.mean(np.diag(np.cov(data_train.T)))
- shuffled_1 = np.random.multivariate_normal(np.zeros(data_train.shape[1]),
- np.diag([m, m, m]),
- data_train.shape[0])
- KL[e, bootstrap] = 0.5 * (KLdivergence(data_train, s_opt) + KLdivergence(s_opt, data_train))
- KL_r[e, bootstrap] = 0.5 * (KLdivergence(shuffled_1, s_opt) + KLdivergence(s_opt, shuffled_1))
- p = 2 * (np.sum(KL_r <= np.mean(KL, axis=-1, keepdims=True), axis=-1) / n_bootstraps)
- if bon_correction:
- p = p * n_epochs
- p_vals.append(p)
- epoch = np.sort(list(range(1, n_epochs + 1)) * n_bootstraps)
- epoch_labels = np.concatenate([epoch, epoch])
- dist_label = np.concatenate([['observed'] * (n_bootstraps * n_epochs), [rnd_model] * (n_bootstraps * n_epochs)])
- KL_df = np.reshape(KL, KL.shape[0] * KL.shape[1])
- KL_r_df = np.reshape(KL_r, KL_r.shape[0] * KL_r.shape[1])
- KL_all_df = np.concatenate([KL_df, KL_r_df])
- df = pd.DataFrame(np.array([KL_all_df, epoch_labels, dist_label]).T,
- columns=['KL divergance estimate', 'learning epoch', 'distribution'])
- df['KL divergance estimate'] = df['KL divergance estimate'].astype(float)
- df['model'] = model_names[model_n]
- dfs.append(df)
- df_all = pd.concat(dfs)
- df_all['divergence from'] = 'optimal selectivity'
- print(p_vals)
- return df_all, p_vals
- def plot_pvals(df_data, p, ax, offset=0.0001, tail=1, n_bootstraps=1000, metric='KL divergance estimate'):
- y_vals = np.resize(df_data[df_data["distribution"] == 'observed'][metric].values,
- (len(p), n_bootstraps))
- y_avgs = y_vals.mean(-1)
- y_stds = y_vals.std(-1)
- if tail == 1:
- y_pos = y_avgs + y_stds + offset
- elif tail == -1:
- y_pos = y_avgs - y_stds - (offset * 12)
- for e in range(len(p)):
- if p[e] > 0.05:
- star = 'ns'
- size = 10
- elif (p[e] <= 0.05) & (p[e] > 0.01):
- star = '*'
- size = 14
- elif (p[e] <= 0.01) & (p[e] > 0.001):
- star = '**'
- size = 14
- elif p[e] <= 0.001:
- star = '***'
- size = 14
- ax.text(e, y_pos[e], star, ha='center', size=size)
- return
- def get_betas_cross_val_2(data, labels, condition_labels, normalisation='zscore', time_window=[140, 150],
- cross_val=False, if_xgen=False, task='task_1', shuffle=False):
- n_parts = len(data)
- betas_part = []
- for p in range(n_parts):
- betas = []
- for s in range(len(data[p])):
- firing_rates_ses = data[p][s]
- firing_rates_ses = firing_rates_ses[:, :, time_window[0]:time_window[1]]
- labels_ses = labels[p][s]
- if task == 'task_1':
- firing_rates_ses = firing_rates_ses[labels_ses < 9, :, :]
- labels_ses = labels_ses[labels_ses < 9]
- elif task == 'task_2':
- firing_rates_ses = firing_rates_ses[labels_ses > 8, :, :]
- labels_ses = labels_ses[labels_ses > 8]
- elif task == 'context_1':
- fac_cxt1 = np.array([1, 2, 3, 4, 9, 10, 11, 12])
- idc_cxt1 = np.isin(labels_ses, fac_cxt1)
- firing_rates_ses = firing_rates_ses[idc_cxt1, :, :]
- labels_ses = labels_ses[idc_cxt1]
- elif task == 'context_2':
- fac_cxt2 = np.array([5, 6, 7, 8, 13, 14, 15, 16])
- idc_cxt2 = np.isin(labels_ses, fac_cxt2)
- firing_rates_ses = firing_rates_ses[idc_cxt2, :, :]
- labels_ses = labels_ses[idc_cxt2]
- if cross_val:
- n_trls = firing_rates_ses.shape[0]
- idc_0 = np.zeros(n_trls // 2)
- idc_1 = np.ones(n_trls // 2)
- if (n_trls % 2) > 0:
- idc_1 = np.concatenate([idc_1, [1.0]])
- idc_rnd = np.concatenate([idc_0, idc_1])
- random.shuffle(idc_rnd)
- firing_rates_ses1 = firing_rates_ses[idc_rnd < 1, :, :]
- firing_rates_ses2 = firing_rates_ses[idc_rnd > 0, :, :]
- labels_ses1 = labels_ses[idc_rnd < 1]
- labels_ses2 = labels_ses[idc_rnd > 0]
- else:
- firing_rates_ses1 = firing_rates_ses
- firing_rates_ses2 = firing_rates_ses
- labels_ses1 = labels_ses
- labels_ses2 = labels_ses
- if if_xgen:
- firing_rates_ses1 = firing_rates_ses[labels_ses1 < 9, :, :]
- firing_rates_ses2 = firing_rates_ses[labels_ses2 > 8, :, :]
- labels_ses1 = labels_ses1[labels_ses1 < 9]
- labels_ses2 = labels_ses2[labels_ses2 > 8]
- firing_rates_ses = [firing_rates_ses1, firing_rates_ses2]
- labels_ses = [labels_ses1, labels_ses2]
- betas_split = []
- n_splits = len(labels_ses)
- for split in range(n_splits):
- firing_mean = np.mean(firing_rates_ses[split], axis=-1, keepdims=True)
- labels_ses_split = labels_ses[split]
- if shuffle:
- random.shuffle(labels_ses_split)
- if normalisation == 'zscore':
- firing_mean = (firing_mean - np.mean(firing_mean, axis=0, keepdims=True)) / (
- np.std(firing_mean, axis=0,
- keepdims=True))
- elif normalisation == 'mean_centred':
- firing_mean = firing_mean - np.mean(firing_mean, axis=0, keepdims=True)
- elif normalisation == 'none':
- firing_mean = firing_mean
- elif normalisation == 'soft':
- firing_mean = firing_mean / (
- (np.max(firing_mean, axis=0, keepdims=True) - np.min(firing_mean, axis=0,
- keepdims=True)) + 5)
- labels_cue = np.array(assign_lables(labels_ses_split, condition_labels[0]))
- labels_target = np.array(assign_lables(labels_ses_split, condition_labels[1]))
- labels_int1 = np.array(assign_lables(labels_ses_split, condition_labels[2]))
- labels_target2 = np.array(assign_lables(labels_ses_split, condition_labels[3]))
- labels_int2 = np.array(assign_lables(labels_ses_split, condition_labels[4]))
- constant = np.ones_like(labels_cue).astype(float)
- design_matrix = np.vstack(
- [constant, labels_cue, labels_target, labels_int1, labels_target2, labels_int2]).T
- if task == 'all':
- labels_context = np.array(assign_lables(labels_ses_split, condition_labels[0]))
- labels_shape = np.array(assign_lables(labels_ses_split, condition_labels[1]))
- labels_rew = np.array(assign_lables(labels_ses_split, condition_labels[2]))
- labels_task = np.array(assign_lables(labels_ses_split, condition_labels[3]))
- labels_width = np.array(assign_lables(labels_ses_split, condition_labels[4]))
- labels_int_irr = np.array(assign_lables(labels_ses_split, condition_labels[5]))
- constant = np.ones_like(labels_cue).astype(float)
- design_matrix = np.vstack(
- [constant, labels_context, labels_shape, labels_rew, labels_task, labels_width,
- labels_int_irr]).T
- design_matrix = design_matrix.astype(float)
- n_neurons = firing_mean.shape[1]
- n_times = firing_mean.shape[-1]
- n_models = design_matrix.shape[1]
- design_matrix_stacked = np.concatenate([design_matrix] * n_times, axis=0)
- firing_rates_ses_stacked = np.concatenate(np.array_split(firing_mean, n_times, axis=-1), axis=0)
- betas_ses = np.zeros((n_neurons, n_models, 1))
- for cell in range(n_neurons):
- betas_ses[cell, :, :] = \
- sp.linalg.lstsq(design_matrix_stacked, firing_rates_ses_stacked[:, cell, :])[0][:, :]
- betas_split.append(betas_ses)
- betas.append(np.array(betas_split)[:, :, :, 0])
- betas_part.append(betas)
- if task == 'all':
- epochs_rel = []
- for p in range(n_parts):
- epoch = np.concatenate(betas_part[p], axis=1)[:, :, 1:4].transpose((0, 2, 1)) # splits x coeffs x neurons
- epoch -= np.mean(epoch, axis=2, keepdims=True)
- epochs_rel.append(epoch.T)
- epochs_irrel = []
- for p in range(n_parts):
- epoch = np.concatenate(betas_part[p], axis=1)[:, :, 4:].transpose((0, 2, 1)) # splits x coeffs x neurons
- epoch -= np.mean(epoch, axis=2, keepdims=True)
- epochs_irrel.append(epoch.T)
- else:
- epochs_rel = []
- for p in range(n_parts):
- epoch = np.concatenate(betas_part[p], axis=1)[:, :, 1:4].transpose((0, 2, 1)) # splits x coeffs x neurons
- epoch -= np.mean(epoch, axis=2, keepdims=True)
- epochs_rel.append(epoch.T)
- epochs_irrel = []
- for p in range(n_parts):
- epoch_1 = np.concatenate(betas_part[p], axis=1)[:, :, 1][:, :, None].transpose(
- (0, 2, 1)) # splits x coeffs x neurons
- epoch_2 = np.concatenate(betas_part[p], axis=1)[:, :, 4:].transpose((0, 2, 1)) # splits x coeffs x neurons
- epoch = np.concatenate([epoch_1, epoch_2], axis=1)
- epoch -= np.mean(epoch, axis=2, keepdims=True)
- epochs_irrel.append(epoch.T)
- return epochs_rel, epochs_irrel
- def get_betas_cross_val_3(data, labels, condition_labels, normalisation='zscore', time_window=[140, 150],
- add_constant=True, n_splits=10):
- n_parts = len(data)
- betas_part = []
- for p in range(n_parts):
- betas_ses_all = []
- for s in tqdm(range(len(data[p]))):
- firing_rates_ses = data[p][s]
- firing_rates_ses = firing_rates_ses[:, :, time_window[0]:time_window[1]]
- labels_ses = labels[p][s]
- n_trl = firing_rates_ses.shape[0]
- firing_rates_ses_all = []
- labels_ses_all = []
- for split in range(n_splits):
- idc_1 = np.zeros(n_trl // 2)
- idc_2 = np.ones(n_trl // 2)
- if (n_trl % 2) > 0:
- idc_2 = np.append(idc_2, [1.0])
- idc_rnd = np.concatenate([idc_1, idc_2])
- random.shuffle(idc_rnd)
- firing_rates_ses1 = firing_rates_ses[idc_rnd < 1.0, :, :]
- firing_rates_ses2 = firing_rates_ses[idc_rnd > 0.0, :, :]
- labels_ses1 = labels_ses[idc_rnd < 1]
- labels_ses2 = labels_ses[idc_rnd > 0]
- firing_rates_ses_all.append([firing_rates_ses1, firing_rates_ses2])
- labels_ses_all.append([labels_ses1, labels_ses2])
- betas_split = []
- for split in range(n_splits):
- betas_run = []
- for run in range(2):
- firing_mean = np.mean(firing_rates_ses_all[split][run], axis=-1, keepdims=True)
- labels_ses_split = labels_ses_all[split][run]
- if normalisation == 'zscore':
- firing_mean = (firing_mean - np.mean(firing_mean, axis=0, keepdims=True)) / (
- np.std(firing_mean, axis=0,
- keepdims=True) + 1)
- elif normalisation == 'mean_centred':
- firing_mean = firing_mean - np.mean(firing_mean, axis=0, keepdims=True)
- elif normalisation == 'none':
- firing_mean = firing_mean
- elif normalisation == 'soft':
- firing_mean = firing_mean / (
- (np.max(firing_mean, axis=0, keepdims=True) - np.min(firing_mean, axis=0,
- keepdims=True)) + 5)
- labels_cue = np.array(assign_lables(labels_ses_split, condition_labels[0]))
- labels_target = np.array(assign_lables(labels_ses_split, condition_labels[1]))
- labels_int1 = np.array(assign_lables(labels_ses_split, condition_labels[2]))
- labels_cue2 = np.array(assign_lables(labels_ses_split, condition_labels[3]))
- labels_target2 = np.array(assign_lables(labels_ses_split, condition_labels[4]))
- labels_int2 = np.array(assign_lables(labels_ses_split, condition_labels[5]))
- if add_constant:
- constant = np.ones_like(labels_cue).astype(float)
- design_matrix = np.vstack(
- [constant, labels_cue, labels_target, labels_int1, labels_cue2, labels_target2,
- labels_int2]).T
- else:
- design_matrix = np.vstack(
- [labels_cue, labels_cue, labels_target, labels_int1, labels_cue2, labels_target2,
- labels_int2]).T
- design_matrix = design_matrix.astype(float)
- n_neurons = firing_mean.shape[1]
- n_times = firing_mean.shape[-1]
- n_models = design_matrix.shape[1]
- design_matrix_stacked = np.concatenate([design_matrix] * n_times, axis=0)
- firing_rates_ses_stacked = np.concatenate(np.array_split(firing_mean, n_times, axis=-1), axis=0)
- betas_ses = np.zeros((n_neurons, n_models, 1))
- for cell in range(n_neurons):
- betas_ses[cell, :, :] = \
- sp.linalg.lstsq(design_matrix_stacked, firing_rates_ses_stacked[:, cell, :])[0][:, :]
- betas_run.append(betas_ses)
- betas_split.append(np.mean(np.array(betas_run), axis=0)[:, :, 0])
- betas_ses_all.append(np.mean(np.array(betas_split), axis=0))
- betas_part.append(np.concatenate(betas_ses_all, axis=0))
- epochs_rel = []
- for p in range(n_parts):
- if add_constant:
- epoch = betas_part[p][:, 1:4]
- else:
- epoch = betas_part[p][:, :3]
- epoch = np.concatenate([epoch[:, :, None], epoch[:, :, None]], axis=-1)
- epoch -= np.mean(epoch, axis=0, keepdims=True)
- epochs_rel.append(epoch)
- epochs_irrel = []
- for p in range(n_parts):
- if add_constant:
- epoch = betas_part[p][:, 4:]
- else:
- epoch = betas_part[p][:, 3:]
- epoch = np.concatenate([epoch[:, :, None], epoch[:, :, None]], axis=-1)
- epoch -= np.mean(epoch, axis=0, keepdims=True)
- epochs_irrel.append(epoch)
- return epochs_rel, epochs_irrel
- def plot_pvals_2(data, ax, n_bootstraps, tail=1, offset=0.05, bon_correction=False, side=1,
- metric='KL divergance estimate'):
- df_data = data[data['distribution'] == 'observed']
- n_epochs = len(np.unique(df_data['learning epoch'].values))
- rel = df_data[df_data['model'] == 'cue + shape'][metric].values
- rel = np.resize(rel, (n_epochs, n_bootstraps))
- rel_avg = np.mean(rel, axis=-1, keepdims=True)
- irrel = df_data[df_data['model'] == 'cue + width'][metric].values
- irrel = np.resize(irrel, (n_epochs, n_bootstraps))
- if side == 1:
- side_fac = 1
- elif side == 2:
- side_fac = 2
- y_avgs = rel_avg[:, 0]
- y_stds = rel.std(-1)
- if tail == 1:
- p = side_fac * (np.sum(irrel >= rel_avg, axis=-1) / n_bootstraps)
- y_pos = y_avgs + y_stds + offset
- elif tail == -1:
- p = side_fac * (np.sum(irrel <= rel_avg, axis=-1) / n_bootstraps)
- y_pos = y_avgs - y_stds - (offset * 12)
- if bon_correction:
- p = p * n_epochs
- for e in range(n_epochs):
- if p[e] > 0.05:
- star = 'ns'
- size = 12
- elif (p[e] <= 0.05) & (p[e] > 0.01):
- star = '*'
- size = 20
- elif (p[e] <= 0.01) & (p[e] > 0.001):
- star = '**'
- size = 20
- elif p[e] <= 0.001:
- star = '***'
- size = 20
- ax.text(e, y_pos[e], star, ha='center', size=size)
- return
- def plot_pvals_3(y, x, p, ax, col, offset=0.05):
- if p > 0.05:
- star = 'ns'
- size = 10
- elif (p <= 0.05) & (p > 0.01):
- star = '*'
- size = 15
- elif (p <= 0.01) & (p > 0.001):
- star = '**'
- size = 15
- elif p <= 0.001:
- star = '***'
- size = 15
- ax.text(x, y + offset, star, ha='center', size=size, color=col)
- return
- def compare_r2s(data1, data2, n_bootstraps, tail=1, mode='diff'):
- n_units_1 = data1.shape[0]
- n_units_2 = data2.shape[0]
- data_all = np.concatenate([data1, data2], axis=0)
- fit_diff_rnd = np.zeros((n_bootstraps))
- fit_1_obs = fit_lines(data1[:, :, 0], data1[:, :, 1], if_xgen=True)
- fit_2_obs = fit_lines(data2[:, :, 0], data2[:, :, 1], if_xgen=True)
- if mode == 'diff':
- fit_diff1 = fit_1_obs[0] - fit_1_obs[1]
- fit_diff2 = fit_2_obs[0] - fit_2_obs[1]
- elif mode == 'best':
- fit_diff1 = fit_1_obs[0]
- fit_diff2 = fit_2_obs[0]
- elif mode == 'optimal':
- fit_diff1 = fit_1_obs[1]
- fit_diff2 = fit_2_obs[1]
- fit_diff = fit_diff1 - fit_diff2
- for n_bootstrap in range(n_bootstraps):
- rnd_idc1 = np.zeros(n_units_1)
- rnd_idc2 = np.ones(n_units_2)
- rnd_idx = np.concatenate([rnd_idc1, rnd_idc2])
- random.shuffle(rnd_idx)
- data1_rnd = data_all[rnd_idx < 1, :, :]
- data2_rnd = data_all[rnd_idx > 0, :, :]
- fit_1_rnd = fit_lines(data1_rnd[:, :, 0], data1_rnd[:, :, 1], if_xgen=True)
- fit_2_rnd = fit_lines(data2_rnd[:, :, 0], data2_rnd[:, :, 1], if_xgen=True)
- if mode == 'diff':
- fit_diff1_rnd = fit_1_rnd[0] - fit_1_rnd[1]
- fit_diff2_rnd = fit_2_rnd[0] - fit_2_rnd[1]
- elif mode == 'best':
- fit_diff1_rnd = fit_1_rnd[0]
- fit_diff2_rnd = fit_2_rnd[0]
- elif mode == 'optimal':
- fit_diff1_rnd = fit_1_rnd[1]
- fit_diff2_rnd = fit_2_rnd[1]
- fit_diff_rnd[n_bootstrap] = fit_diff1_rnd - fit_diff2_rnd
- if tail == -1:
- p = (np.sum(fit_diff_rnd >= fit_diff) / n_bootstraps)
- elif tail == 1:
- p = (np.sum(fit_diff_rnd <= fit_diff) / n_bootstraps)
- return p
- def data_handle(data, labels, min_n, do_rand=False):
- conditions = np.unique(labels)
- conditions_ = []
- for _ in range(len(conditions)):
- condition_dat = data[labels == conditions[_], :, :]
- if do_rand:
- idx_re = np.random.choice(np.array(list(range(0, condition_dat.shape[0]))), min_n, replace=False)
- conditions_.append(condition_dat[idx_re, :, :])
- else:
- conditions_.append(condition_dat[:min_n, :, :])
- conditions_ = np.array(conditions_)
- return np.concatenate(conditions_, axis=0)
- def euclidean_distance(data1, data2):
- return np.sqrt(np.sum((np.cov(data2.T) - np.cov(data1.T)) ** 2))
- def prepare_data(data, labels, which_trl='beginning', set_min_trl=None):
- n_trls_ses = []
- for sess in range(len(data)):
- n_trls_ses.append(data[sess].shape[0])
- n_trls = np.min(n_trls_ses)
- n_condi_min = []
- for sess in range(len(data)):
- if which_trl == 'beginning':
- freqs = get_freqs(labels[sess][:n_trls])
- elif which_trl == 'end':
- freqs = get_freqs(labels[sess][-n_trls:])
- elif which_trl == 'rnd':
- freqs = get_freqs(labels[sess])
- elif which_trl == 'middle':
- labels_sess = labels[sess]
- idc_half = int(len(labels_sess) / 2)
- # check whether even or odd number of trials
- if (len(labels_sess) % 2) > 0:
- labels_sess = labels_sess[idc_half - int(n_trls / 2):idc_half + int(n_trls / 2) + 1]
- else:
- labels_sess = labels_sess[idc_half - int(n_trls / 2):idc_half + int(n_trls / 2)]
- freqs = get_freqs(labels_sess)
- n_condi_min.append(freqs[min(freqs, key=freqs.get)])
- n_trl_condi = np.min(n_condi_min)
- if set_min_trl:
- n_trl_condi = set_min_trl
- dat_combined = []
- for sess in range(len(data)):
- if which_trl == 'beginning':
- data_sess = data[sess][:n_trls, :, :]
- labels_sess = labels[sess][:n_trls]
- dat_combined.append(data_handle(data_sess, labels_sess, n_trl_condi))
- elif which_trl == 'end':
- data_sess = data[sess][-n_trls:, :, :]
- labels_sess = labels[sess][-n_trls:]
- dat_combined.append(data_handle(data_sess, labels_sess, n_trl_condi))
- elif which_trl == 'rnd':
- data_sess = data_handle(data[sess], labels[sess], n_trl_condi, do_rand=False)
- dat_combined.append(data_sess)
- elif which_trl == 'middle':
- labels_sess = labels[sess]
- idc_half = int(len(labels_sess) / 2)
- # check whether even or odd number of trials
- if (len(labels_sess) % 2) > 0:
- labels_sess = labels_sess[idc_half - int(n_trls / 2):idc_half + int(n_trls / 2) + 1]
- data_sess = data[sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2) + 1, :, :]
- elif (len(labels_sess) % 2) == 0:
- labels_sess = labels_sess[idc_half - int(n_trls / 2):idc_half + int(n_trls / 2)]
- data_sess = data[sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2), :, :]
- dat_combined.append(data_handle(data_sess, labels_sess, n_trl_condi))
- dat_combined = np.concatenate(dat_combined, axis=1)
- labels_combined = np.sort(list(range(0, len(np.unique(labels[0])))) * n_trl_condi)
- return dat_combined, labels_combined
- def get_shattering_ids(cue, target, width, reward):
- n_condi = len(cue)
- all_combos = list(itertools.combinations(list(range(n_condi)), int(n_condi / 2)))
- n_combos = int(len(all_combos) / 2)
- decoding_targets = np.zeros((n_combos, n_condi))
- cue_tgt_width_reward_ids = np.zeros(int(n_condi / 2))
- for i in range(n_combos):
- decoding_targets[i, all_combos[i]] = 1
- # Find cue ids
- if np.sum(np.abs(decoding_targets[i, :] - cue)) == 0 or np.sum(np.abs(decoding_targets[i, :] - (1 - cue))) == 0:
- cue_tgt_width_reward_ids[0] = i
- # Find target ids
- if np.sum(np.abs(decoding_targets[i, :] - target)) == 0 or np.sum(
- np.abs(decoding_targets[i, :] - (1 - target))) == 0:
- cue_tgt_width_reward_ids[1] = i
- # Find width ids
- if np.sum(np.abs(decoding_targets[i, :] - width)) == 0 or np.sum(
- np.abs(decoding_targets[i, :] - (1 - width))) == 0:
- cue_tgt_width_reward_ids[2] = i
- # Find reward ids
- if np.sum(np.abs(decoding_targets[i, :] - reward)) == 0 or np.sum(
- np.abs(decoding_targets[i, :] - (1 - reward))) == 0:
- cue_tgt_width_reward_ids[3] = i
- cross_gen_decoding_train_ids = []
- cross_gen_decoding_test_ids = []
- for i in range(n_combos):
- current_zeros = np.squeeze(np.where(decoding_targets[i, :] == 0))
- current_combs1 = list(itertools.combinations(current_zeros.tolist(), 2))
- current_ones = np.squeeze(np.where(decoding_targets[i, :] == 1))
- current_combs2 = list(itertools.combinations(current_ones.tolist(), 2))
- n_combs = len(current_combs1)
- current_train = []
- current_test = []
- for j in range(n_combs):
- for k in range(n_combs):
- current_train.append([current_combs1[j], current_combs2[k]])
- current_test_zeros = list(np.setdiff1d(current_zeros, np.array(current_combs1[j])))
- current_test_ones = list(np.setdiff1d(current_ones, np.array(current_combs2[k])))
- current_test.append([current_test_zeros, current_test_ones])
- cross_gen_decoding_train_ids.append(current_train)
- cross_gen_decoding_test_ids.append(current_test)
- return decoding_targets, cue_tgt_width_reward_ids, cross_gen_decoding_train_ids, cross_gen_decoding_test_ids
- def decode(X, y, method='svm', n_inter=5, return_inter=False, n_jobs=None):
- y = np.array(y)
- scores_all = []
- for n in range(n_inter):
- if method == 'svm':
- clf1 = make_pipeline(StandardScaler(), SVM(C=5e-4))
- clf2 = make_pipeline(StandardScaler(), SVM(C=5e-4))
- if n_jobs != None:
- clf1 = SlidingEstimator(clf1, verbose='warning', n_jobs=n_jobs)
- clf2 = SlidingEstimator(clf2, verbose='warning', n_jobs=n_jobs)
- elif method == 'lda':
- clf1 = make_pipeline(StandardScaler(), LDA())
- if n_jobs != None:
- clf1 = SlidingEstimator(clf1, verbose='warning', n_jobs=n_jobs)
- clf2 = make_pipeline(StandardScaler(), LDA())
- if n_jobs != None:
- clf2 = SlidingEstimator(clf2, verbose='warning', n_jobs=n_jobs)
- if method == 'nonlin_svm':
- clf1 = make_pipeline(StandardScaler(), SVC(kernel='poly', C=1.0))
- clf2 = make_pipeline(StandardScaler(), SVC(kernel='poly', C=1.0))
- if n_jobs != None:
- clf1 = SlidingEstimator(clf1, verbose='warning', n_jobs=n_jobs)
- clf2 = SlidingEstimator(clf2, verbose='warning', n_jobs=n_jobs)
- n_trls = X.shape[0]
- idx_rnd = np.concatenate([np.zeros(n_trls // 2), np.ones(n_trls // 2)])
- if (n_trls % 2) > 0:
- idx_rnd = np.concatenate([idx_rnd, [1.0]])
- random.shuffle(idx_rnd)
- if n_jobs != None:
- X1 = X[idx_rnd < 1, :, :]
- y1 = y[idx_rnd < 1]
- X2 = X[idx_rnd > 0, :, :]
- y2 = y[idx_rnd > 0]
- else:
- X1 = X[idx_rnd < 1, :]
- y1 = y[idx_rnd < 1]
- X2 = X[idx_rnd > 0, :]
- y2 = y[idx_rnd > 0]
- clf1.fit(X1, y1)
- score1 = clf1.score(X2, y2)
- clf2.fit(X2, y2)
- score2 = clf2.score(X1, y1)
- scores_all.append(np.array([score1, score2]).mean(0))
- if return_inter:
- return scores_all
- else:
- return np.mean(scores_all)
- def get_decoding(data, labels, variables, time_window, method, n_jobs=None):
- decoding_targets, cue_tgt_width_reward_ids, _, _ = get_shattering_ids(variables[0], variables[1], variables[2],
- variables[3])
- n_combos = decoding_targets.shape[0]
- scores = []
- print('Cutting along all possible axes (shattering dimensionality)')
- for combo in tqdm(range(n_combos)):
- if n_jobs != None:
- X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
- else:
- X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=False)
- y = assign_lables(labels, decoding_targets[combo, :])
- scores.append(decode(X, y, method=method, n_jobs=n_jobs))
- score_cue = scores[int(cue_tgt_width_reward_ids[0])]
- score_target = scores[int(cue_tgt_width_reward_ids[1])]
- score_width = scores[int(cue_tgt_width_reward_ids[2])]
- score_rew = scores[int(cue_tgt_width_reward_ids[3])]
- score_shattering = scores
- var_ids = [int(cue_tgt_width_reward_ids) for cue_tgt_width_reward_ids in cue_tgt_width_reward_ids]
- scores_irrel = np.delete(scores, obj=var_ids)
- return [score_cue, score_target, score_width, score_rew, score_shattering], scores_irrel
- def get_decoding_null(data, labels, time_window, method, n_reps=100, n_jobs=None):
- if len(np.unique(labels)) < 9:
- fac = [0, 0, 0, 0, 1, 1, 1, 1]
- else:
- fac = [0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1]
- print('Generating the null decoding distribution...')
- scores = []
- for rep in tqdm(range(n_reps)):
- if n_jobs != None:
- X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
- else:
- X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=False)
- y = assign_lables(labels, factor=fac).copy()
- if len(np.unique(labels)) < 3:
- y = labels.copy()
- random.shuffle(y)
- scores.append(decode(X, y, method=method, n_jobs=n_jobs))
- return np.array(scores)
- def decode_xgen(X1, X2, y1, y2, method='svm', n_jobs=-1):
- # prepare a series of classifier applied at each time sample
- if method == 'svm':
- clf1 = make_pipeline(StandardScaler(), SVM(C=5e-4))
- if n_jobs != None:
- clf1 = SlidingEstimator(clf1, verbose=False, n_jobs=-1)
- clf2 = make_pipeline(StandardScaler(), SVM(C=5e-4))
- if n_jobs != None:
- clf2 = SlidingEstimator(clf2, verbose=False, n_jobs=-1)
- elif method == 'lda':
- clf1 = make_pipeline(StandardScaler(), LDA())
- if n_jobs != None:
- clf1 = SlidingEstimator(clf1, verbose=False, n_jobs=-1)
- clf2 = make_pipeline(StandardScaler(), LDA())
- if n_jobs != None:
- clf2 = SlidingEstimator(clf2, verbose=False, n_jobs=-1)
- clf1.fit(X1, y1)
- score1 = clf1.score(X2, y2)
- clf2.fit(X2, y2)
- score2 = clf2.score(X1, y1)
- return np.array([score1, score2]).mean(0)
- def get_x_gen_null(data, labels, method, n_reps=100, n_jobs=None):
- if len(np.unique(labels)) < 9:
- fac = [0, 0, 0, 0, 1, 1, 1, 1]
- else:
- fac = [0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1]
- print('Generating the null x-gen distribution...')
- scores = []
- for rep in tqdm(range(n_reps)):
- X = data
- y = assign_lables(labels, factor=fac).copy()
- if n_jobs != None:
- X1 = X[::2, :, :]
- X2 = X[1::2, :, :]
- else:
- X1 = X[::2, :]
- X2 = X[1::2, :]
- y1 = y[::2]
- y2 = y[1::2]
- random.shuffle(y1)
- random.shuffle(y2)
- scores.append(decode_xgen(X1, X2, y1, y2, method=method, n_jobs=n_jobs))
- return np.array(scores)
- def get_xgen(data, labels, variables, time_window, method='svm', mode='full', n_jobs=None, verbose=False):
- decoding_targets, cue_tgt_width_reward_ids, cross_gen_decoding_train_ids, cross_gen_decoding_test_ids = get_shattering_ids(
- variables[0], variables[1], variables[2], variables[3])
- var_ids = [int(cue_tgt_width_reward_ids) for cue_tgt_width_reward_ids in cue_tgt_width_reward_ids]
- n_combos = len(cross_gen_decoding_train_ids)
- if mode == 'only_rel':
- combo_list = var_ids
- elif mode == 'full':
- combo_list = list(range(n_combos))
- scores = []
- if verbose:
- print('Cutting along all possible axes (xgen)')
- for combo in tqdm(combo_list, disable=not verbose):
- xgen_train = cross_gen_decoding_train_ids[combo]
- xgen_test = cross_gen_decoding_test_ids[combo]
- if n_jobs != None:
- X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
- else:
- X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=False)
- n_axes = len(xgen_train)
- scores_axes = []
- for axis in range(n_axes):
- y1 = np.array(xgen_train[axis]).flatten()
- X1 = []
- for _ in range(len(y1)):
- if n_jobs != None:
- X1.append(X[labels == y1[_], :, :])
- else:
- X1.append(X[labels == y1[_], :])
- X1 = np.concatenate(X1, axis=0)
- y2 = np.array(xgen_test[axis]).flatten()
- X2 = []
- for _ in range(len(y1)):
- if n_jobs != None:
- X2.append(X[labels == y2[_], :, :])
- else:
- X2.append(X[labels == y2[_], :])
- X2 = np.concatenate(X2, axis=0)
- y = np.concatenate([np.zeros(int(X1.shape[0] / 2)), np.ones(int(X1.shape[0] / 2))])
- scores_axes.append(decode_xgen(X1, X2, y, y, method=method, n_jobs=n_jobs))
- scores.append(np.mean(scores_axes))
- if mode == 'full':
- scores_rel = [scores[var_ids[0]], scores[var_ids[1]], scores[var_ids[2]], scores[var_ids[3]]]
- scores_irrel = np.delete(scores, obj=var_ids)
- return scores_rel, scores_irrel
- elif mode == 'only_rel':
- return np.array(scores)
- def get_decoding_models(n_neurons=400, noise_std=1, n_trials=20, seed=42):
- np.random.seed(seed)
- conds = np.ones((4, 3))
- conds[0, :2] = -1
- conds[1, [0, 2]] = -1
- conds[2, 1:] = -1
- conds = np.array([conds] * n_trials)
- clfc = SVM(C=1e-5)
- clfs = SVM(C=1e-5)
- clfr = SVM(C=1e-5)
- N_std = [[n_neurons, noise_std]] # number of neurons and std of noise
- targets_r = [1, -1, -1, 1]
- targets_s = [-1, 1, -1, 1]
- targets_c = [-1, -1, 1, 1]
- decoding_tot = np.zeros((len(N_std), 100, 3, 2))
- cross_decoding_tot = np.zeros((len(N_std), 100, 3, 2))
- for counter, n_std in enumerate(N_std):
- N = n_std[0]
- noise_std = n_std[1]
- for model in range(100):
- for rand_opt in range(2):
- if rand_opt == 0:
- betas = np.random.normal(0, 1, (3, N))
- else:
- cov = np.diag([0, 0, 3])
- betas = np.random.multivariate_normal(np.zeros(3), cov, N).T
- x_train0 = conds @ betas + noise_std * np.random.normal(0, 1, (n_trials, 4, N))
- x_test0 = conds @ betas + noise_std * np.random.normal(0, 1, (n_trials, 4, N))
- x_train0_long = np.concatenate(np.array_split(x_train0, n_trials, axis=0), axis=1)[0, :, :]
- x_test0_long = np.concatenate(np.array_split(x_test0, n_trials, axis=0), axis=1)[0, :, :]
- clfc.fit(x_train0_long, np.concatenate([targets_c] * n_trials))
- clfs.fit(x_train0_long, np.concatenate([targets_s] * n_trials))
- clfr.fit(x_train0_long, np.concatenate([targets_r] * n_trials))
- decoding_tot[counter, model, 0, rand_opt] = clfc.score(x_test0_long,
- np.concatenate([targets_c] * n_trials))
- decoding_tot[counter, model, 1, rand_opt] = clfs.score(x_test0_long,
- np.concatenate([targets_s] * n_trials))
- decoding_tot[counter, model, 2, rand_opt] = clfr.score(x_test0_long,
- np.concatenate([targets_r] * n_trials))
- # Cross decoding
- r_train_id = [[0, 1], [0, 2], [3, 1], [3, 2]]
- r_test_id = [[3, 2], [3, 1], [0, 2], [0, 1]]
- s_train_id = [[0, 1], [0, 3], [2, 1], [2, 3]]
- s_test_id = [[2, 3], [2, 1], [0, 3], [0, 1]]
- c_train_id = [[0, 2], [0, 3], [1, 2], [1, 3]]
- c_test_id = [[1, 3], [1, 2], [0, 3], [0, 2]]
- for i in range(4):
- # color
- x_train1 = x_train0[:, c_train_id[i], :]
- x_test1 = x_test0[:, c_test_id[i], :]
- x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
- x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
- clfc.fit(x_train1, [0, 1] * n_trials)
- cross_decoding_tot[counter, model, 0, rand_opt] += 0.25 * clfc.score(x_test1, [0, 1] * n_trials)
- # shape
- x_train1 = x_train0[:, s_train_id[i], :]
- x_test1 = x_test0[:, s_test_id[i], :]
- x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
- x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
- clfs.fit(x_train1, [0, 1] * n_trials)
- cross_decoding_tot[counter, model, 1, rand_opt] += 0.25 * clfs.score(x_test1, [0, 1] * n_trials)
- # reward
- x_train1 = x_train0[:, r_train_id[i], :]
- x_test1 = x_test0[:, r_test_id[i], :]
- x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
- x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
- clfr.fit(x_train1, [0, 1] * n_trials)
- cross_decoding_tot[counter, model, 2, rand_opt] += 0.25 * clfr.score(x_test1, [0, 1] * n_trials)
- m_decode = np.mean(decoding_tot[0, :], axis=0)
- m_cross_gen = np.mean(cross_decoding_tot[0, :], axis=0)
- return m_decode, m_cross_gen
- def get_decoding_models(n_neurons=400, noise_std=1, n_trials=20, seed=42):
- np.random.seed(seed)
- conds = np.ones((4, 3))
- conds[0, :2] = -1
- conds[1, [0, 2]] = -1
- conds[2, 1:] = -1
- conds = np.array([conds] * n_trials)
- clfc = SVM(C=1e-5)
- clfs = SVM(C=1e-5)
- clfr = SVM(C=1e-5)
- N_std = [[n_neurons, noise_std]] # number of neurons and std of noise
- targets_r = [1, -1, -1, 1]
- targets_s = [-1, 1, -1, 1]
- targets_c = [-1, -1, 1, 1]
- decoding_tot = np.zeros((len(N_std), 100, 3, 2))
- cross_decoding_tot = np.zeros((len(N_std), 100, 3, 2))
- for counter, n_std in enumerate(N_std):
- N = n_std[0]
- noise_std = n_std[1]
- for model in range(100):
- for rand_opt in range(2):
- if rand_opt == 0:
- betas = np.random.normal(0, 1, (3, N))
- else:
- cov = np.diag([0, 0, 3])
- betas = np.random.multivariate_normal(np.zeros(3), cov, N).T
- x_train0 = conds @ betas + noise_std * np.random.normal(0, 1, (n_trials, 4, N))
- x_test0 = conds @ betas + noise_std * np.random.normal(0, 1, (n_trials, 4, N))
- x_train0_long = np.concatenate(np.array_split(x_train0, n_trials, axis=0), axis=1)[0, :, :]
- x_test0_long = np.concatenate(np.array_split(x_test0, n_trials, axis=0), axis=1)[0, :, :]
- clfc.fit(x_train0_long, np.concatenate([targets_c] * n_trials))
- clfs.fit(x_train0_long, np.concatenate([targets_s] * n_trials))
- clfr.fit(x_train0_long, np.concatenate([targets_r] * n_trials))
- decoding_tot[counter, model, 0, rand_opt] = clfc.score(x_test0_long,
- np.concatenate([targets_c] * n_trials))
- decoding_tot[counter, model, 1, rand_opt] = clfs.score(x_test0_long,
- np.concatenate([targets_s] * n_trials))
- decoding_tot[counter, model, 2, rand_opt] = clfr.score(x_test0_long,
- np.concatenate([targets_r] * n_trials))
- # Cross decoding
- r_train_id = [[0, 1], [0, 2], [3, 1], [3, 2]]
- r_test_id = [[3, 2], [3, 1], [0, 2], [0, 1]]
- s_train_id = [[0, 1], [0, 3], [2, 1], [2, 3]]
- s_test_id = [[2, 3], [2, 1], [0, 3], [0, 1]]
- c_train_id = [[0, 2], [0, 3], [1, 2], [1, 3]]
- c_test_id = [[1, 3], [1, 2], [0, 3], [0, 2]]
- for i in range(4):
- # color
- x_train1 = x_train0[:, c_train_id[i], :]
- x_test1 = x_test0[:, c_test_id[i], :]
- x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
- x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
- clfc.fit(x_train1, [0, 1] * n_trials)
- cross_decoding_tot[counter, model, 0, rand_opt] += 0.25 * clfc.score(x_test1, [0, 1] * n_trials)
- # shape
- x_train1 = x_train0[:, s_train_id[i], :]
- x_test1 = x_test0[:, s_test_id[i], :]
- x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
- x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
- clfs.fit(x_train1, [0, 1] * n_trials)
- cross_decoding_tot[counter, model, 1, rand_opt] += 0.25 * clfs.score(x_test1, [0, 1] * n_trials)
- # reward
- x_train1 = x_train0[:, r_train_id[i], :]
- x_test1 = x_test0[:, r_test_id[i], :]
- x_train1 = np.concatenate(np.array_split(x_train1, n_trials, axis=0), axis=1)[0, :, :]
- x_test1 = np.concatenate(np.array_split(x_test1, n_trials, axis=0), axis=1)[0, :, :]
- clfr.fit(x_train1, [0, 1] * n_trials)
- cross_decoding_tot[counter, model, 2, rand_opt] += 0.25 * clfr.score(x_test1, [0, 1] * n_trials)
- m_decode = np.mean(decoding_tot[0, :], axis=0)
- m_cross_gen = np.mean(cross_decoding_tot[0, :], axis=0)
- return m_decode, m_cross_gen
- def pearsonr(X, Y, axis=None, keepdims=False, strip_nans=False):
- """
- Pearson correlation across a specific axis.
- """
- if strip_nans:
- mymean = np.nanmean
- mystd = np.nanstd
- mysum = np.nansum
- else:
- mymean = np.mean
- mystd = np.std
- mysum = np.sum
- should_squeeze = axis is not None and not keepdims
- if axis is None:
- X = X.ravel()
- Y = Y.ravel()
- axis = 0
- xbar = mymean(X, axis=axis, keepdims=True)
- ybar = mymean(Y, axis=axis, keepdims=True)
- ssx = mysum((X - xbar) ** 2, axis=axis, keepdims=True)
- ssy = mysum((Y - ybar) ** 2, axis=axis, keepdims=True)
- # the following else-block is equivalent to:
- # num = np.sum( (X-xbar)*(Y-ybar) , axis=axis, keepdims=True)
- # but use MUCH less memory as they accumulate the sum in a loop.
- # the two approaches use the same amount of cputime
- if strip_nans:
- # use the memory-inefficient way in case we have to deal with nans
- num = mysum((X - xbar) * (Y - ybar), axis=axis, keepdims=True)
- else:
- tmpX = X.take(0, axis=axis).reshape(xbar.shape)
- tmpY = Y.take(0, axis=axis).reshape(ybar.shape)
- num = (tmpX - xbar) * (tmpY - ybar)
- for k in range(1, X.shape[axis]):
- tmpX = X.take(k, axis=axis).reshape(xbar.shape)
- tmpY = Y.take(k, axis=axis).reshape(ybar.shape)
- num += (tmpX - xbar) * (tmpY - ybar)
- denom = np.sqrt(ssx) * np.sqrt(ssy)
- r = num / denom
- if should_squeeze:
- s = list(r.shape)
- assert (s[axis] == 1) # the axis dimension should now be singleton
- s.pop(axis) # so remove it
- r = r.reshape(s)
- return r
- def dist_random_data2(const_coeffs, metric='euclidean distance', rnd_model='gaussian (spherical)', n_bootstraps=1000,
- relative_dist=True, bon_correction=False):
- n_epochs = len(const_coeffs)
- n_pairs = const_coeffs[0].shape[0]
- dist_data = np.zeros((n_epochs, n_pairs, n_bootstraps))
- dist_rnd = np.zeros((n_epochs, n_pairs, n_bootstraps))
- dist_str = np.zeros((n_epochs, n_pairs, n_bootstraps))
- for part in range(n_epochs):
- for n_bootstrap in tqdm(range(n_bootstraps)):
- for pair in range(n_pairs):
- data_coeffs = const_coeffs[part][pair, :, :]
- opt_cov = np.diag([0, 0, np.cov(data_coeffs.T)[2, 2]])
- s_opt = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
- opt_cov,
- data_coeffs.shape[0])
- m = np.mean(np.diag(np.cov(data_coeffs.T)))
- s_rnd_1 = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
- np.diag([m, m, m]),
- data_coeffs.shape[0])
- s_rnd_2 = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
- np.diag([m, m, m]),
- data_coeffs.shape[0])
- if metric == 'euclidean distance':
- dist_data[part, pair, n_bootstrap] = euclidean_distance(s_rnd_1, data_coeffs)
- dist_rnd[part, pair, n_bootstrap] = euclidean_distance(s_rnd_1, s_rnd_2)
- dist_str[part, pair, n_bootstrap] = euclidean_distance(s_rnd_1, s_opt)
- elif metric == 'KL divergance estimate':
- dist_data[part, pair, n_bootstrap] = 0.5 * (
- KLdivergence(s_rnd_1, data_coeffs) + KLdivergence(data_coeffs, s_rnd_1))
- dist_rnd[part, pair, n_bootstrap] = 0.5 * (
- KLdivergence(s_rnd_1, s_rnd_2) + KLdivergence(s_rnd_2, s_rnd_1))
- dist_str[part, pair, n_bootstrap] = 0.5 * (
- KLdivergence(s_rnd_1, s_opt) + KLdivergence(s_opt, s_rnd_1))
- dist_data = np.reshape(dist_data, (dist_data.shape[0], dist_data.shape[1] * dist_data.shape[2]))
- dist_rnd = np.reshape(dist_rnd, (dist_rnd.shape[0], dist_rnd.shape[1] * dist_rnd.shape[2]))
- dist_str = np.reshape(dist_str, (dist_str.shape[0], dist_str.shape[1] * dist_str.shape[2]))
- if relative_dist:
- dist_rnd_avg = np.mean(dist_rnd, keepdims=True, axis=-1)
- dist_str_avg = np.mean(dist_str, keepdims=True, axis=-1)
- dist_data -= dist_rnd_avg
- dist_rnd -= dist_rnd_avg
- dist_str -= dist_rnd_avg
- dist_data /= dist_str_avg
- dist_rnd /= dist_str_avg
- dist_str /= dist_str_avg
- p = 2 * (np.sum(dist_rnd >= np.mean(dist_data, axis=-1, keepdims=True), axis=-1) / n_bootstraps)
- if bon_correction:
- p = p * n_epochs
- print('p-values:')
- print(p)
- epoch = np.sort(list(range(1, n_epochs + 1)) * n_bootstraps * n_pairs)
- epoch_labels = np.concatenate([epoch, epoch, epoch])
- dist_label = np.concatenate([['observed'] * (n_bootstraps * n_epochs * n_pairs),
- [rnd_model] * (n_bootstraps * n_epochs * n_pairs),
- ['structured'] * (n_bootstraps * n_epochs * n_pairs)])
- data_df = np.reshape(dist_data, dist_data.shape[0] * dist_data.shape[1])
- rnd_df = np.reshape(dist_rnd, dist_rnd.shape[0] * dist_rnd.shape[1])
- str_df = np.reshape(dist_str, dist_str.shape[0] * dist_str.shape[1])
- all_df = np.concatenate([data_df, rnd_df, str_df])
- df = pd.DataFrame(np.array([all_df, epoch_labels, dist_label]).T,
- columns=[metric, 'learning epoch', 'distribution'])
- df[metric] = df[metric].astype(float)
- df['divergence from'] = 'random selectivity'
- return df, p, np.array([dist_data, dist_rnd, dist_str])
- def dist_structured_data2(const_coeffs, metric='euclidean distance', rnd_model='gaussian (spherical)',
- n_bootstraps=1000, relative_dist=True, bon_correction=False):
- n_epochs = len(const_coeffs)
- n_pairs = const_coeffs[0].shape[0]
- dist_data = np.zeros((n_epochs, n_pairs, n_bootstraps))
- dist_rnd = np.zeros((n_epochs, n_pairs, n_bootstraps))
- dist_str = np.zeros((n_epochs, n_pairs, n_bootstraps))
- for part in range(n_epochs):
- for n_bootstrap in tqdm(range(n_bootstraps)):
- for pair in range(n_pairs):
- data_coeffs = const_coeffs[part][pair, :, :]
- opt_cov = np.diag([0, 0, np.cov(data_coeffs.T)[2, 2]])
- s_opt_1 = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
- opt_cov,
- data_coeffs.shape[0])
- s_opt_2 = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
- opt_cov,
- data_coeffs.shape[0])
- m = np.mean(np.diag(np.cov(data_coeffs.T)))
- s_rnd = np.random.multivariate_normal(np.zeros(data_coeffs.shape[1]),
- np.diag([m, m, m]),
- data_coeffs.shape[0])
- if metric == 'euclidean distance':
- dist_data[part, pair, n_bootstrap] = euclidean_distance(s_opt_1, data_coeffs)
- dist_rnd[part, pair, n_bootstrap] = euclidean_distance(s_opt_1, s_rnd)
- dist_str[part, pair, n_bootstrap] = euclidean_distance(s_opt_1, s_opt_2)
- elif metric == 'KL divergance estimate':
- dist_data[part, pair, n_bootstrap] = 0.5 * (
- KLdivergence(s_opt_1, data_coeffs) + KLdivergence(data_coeffs, s_opt_1))
- dist_rnd[part, pair, n_bootstrap] = 0.5 * (
- KLdivergence(s_opt_1, s_rnd) + KLdivergence(s_rnd, s_opt_1))
- dist_str[part, pair, n_bootstrap] = 0.5 * (
- KLdivergence(s_opt_1, s_opt_2) + KLdivergence(s_opt_2, s_opt_1))
- dist_data = np.reshape(dist_data, (dist_data.shape[0], dist_data.shape[1] * dist_data.shape[2]))
- dist_rnd = np.reshape(dist_rnd, (dist_rnd.shape[0], dist_rnd.shape[1] * dist_rnd.shape[2]))
- dist_str = np.reshape(dist_str, (dist_str.shape[0], dist_str.shape[1] * dist_str.shape[2]))
- if relative_dist:
- dist_str_avg = np.mean(dist_str, keepdims=True, axis=-1)
- dist_rnd_avg = np.mean(dist_rnd, keepdims=True, axis=-1)
- dist_data -= dist_str_avg
- dist_rnd -= dist_str_avg
- dist_str -= dist_str_avg
- dist_data /= dist_rnd_avg
- dist_rnd /= dist_rnd_avg
- dist_str /= dist_rnd_avg
- p = 2 * (np.sum(dist_rnd <= np.mean(dist_data, axis=-1, keepdims=True), axis=-1) / n_bootstraps)
- if bon_correction:
- p = p * n_epochs
- print('p-values:')
- print(p)
- epoch = np.sort(list(range(1, n_epochs + 1)) * n_bootstraps * n_pairs)
- epoch_labels = np.concatenate([epoch, epoch, epoch])
- dist_label = np.concatenate([['observed'] * (n_bootstraps * n_epochs * n_pairs),
- [rnd_model] * (n_bootstraps * n_epochs * n_pairs),
- ['structured'] * (n_bootstraps * n_epochs * n_pairs)])
- data_df = np.reshape(dist_data, dist_data.shape[0] * dist_data.shape[1])
- rnd_df = np.reshape(dist_rnd, dist_rnd.shape[0] * dist_rnd.shape[1])
- str_df = np.reshape(dist_str, dist_str.shape[0] * dist_str.shape[1])
- all_df = np.concatenate([data_df, rnd_df, str_df])
- df = pd.DataFrame(np.array([all_df, epoch_labels, dist_label]).T,
- columns=[metric, 'learning epoch', 'distribution'])
- df[metric] = df[metric].astype(float)
- df['divergence from'] = 'structured selectivity'
- return df, p, np.array([dist_data, dist_rnd, dist_str])
- def get_decoding_exp2(data, labels, variables, time_window, method):
- n_combos = len(variables)
- scores = []
- print('Decoding task variables')
- for combo in tqdm(range(n_combos)):
- X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
- y = assign_lables(labels, variables[combo])
- scores.append(decode(X, y, method=method, return_inter=True))
- return np.array(scores)[:, :, 0]
- def shattering_dim_rel(data, labels, fac_new, time_window, method='svm'):
- labels_new = assign_lables(labels, factor=fac_new)
- n_condi = len(np.unique(fac_new))
- all_combos = list(itertools.combinations(list(range(n_condi)), int(n_condi / 2)))
- n_combos = int(len(all_combos) / 2)
- combo_targets = np.ones((n_combos, n_condi))
- for combo in range(n_combos):
- combo_targets[combo, all_combos[combo][0]] = 0
- combo_targets[combo, all_combos[combo][1]] = 0
- scores = []
- print('Cutting along all possible axes (shattering dimensionality)')
- for combo in tqdm(range(n_combos)):
- X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
- y = assign_lables(labels_new, combo_targets[combo, :])
- scores.append(decode(X, y, method=method))
- return np.array(scores).mean()
- def cross_decoding_1axis(data, labels, target_ax, splitting_ax, method='svm', n_jobs=-1, if_rnd=False):
- labels_splitting_ax = np.array(assign_lables(labels, splitting_ax))
- labels_target_ax = np.array(assign_lables(labels, target_ax))
- if if_rnd:
- random.shuffle(labels_target_ax)
- X1 = data[labels_splitting_ax < 1, :, :]
- y1 = labels_target_ax[labels_splitting_ax < 1]
- X2 = data[labels_splitting_ax > 0, :, :]
- y2 = labels_target_ax[labels_splitting_ax > 0]
- score = decode_xgen(X1, X2, y1, y2, method=method, n_jobs=n_jobs)
- return score
- def get_cell_frate(dat, lab):
- cells_lis_cue = []
- cells_lis_shape = []
- cells_lis_xor = []
- cue = np.array([0, 0, 0, 0, 1, 1, 1, 1])
- shape = np.array([0, 0, 1, 1, 0, 0, 1, 1])
- width = np.array([0, 1, 0, 1, 0, 1, 0, 1])
- xor = np.array([1, 1, 0, 0, 0, 0, 1, 1])
- for _ in range(len(dat)):
- lables_fac = assign_lables(lab[_], factor=cue)
- cells_lis_cue.append(condi_avg(dat[_], lables_fac))
- lables_fac = assign_lables(lab[_], factor=shape)
- cells_lis_shape.append(condi_avg(dat[_], lables_fac))
- lables_fac = assign_lables(lab[_], factor=xor)
- cells_lis_xor.append(condi_avg(dat[_], lables_fac))
- cells_arr_cue = np.concatenate(cells_lis_cue, axis=1)
- cells_lis_shape = np.concatenate(cells_lis_shape, axis=1)
- cells_lis_xor = np.concatenate(cells_lis_xor, axis=1)
- return cells_arr_cue, cells_lis_shape, cells_lis_xor
- def decode_epoch(data, labels, n_reps=100, method='svm', n_inter=1, n_jobs=-1):
- obs = np.array(decode(data, labels, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
- rnd = np.zeros((n_reps, data.shape[-1]))
- for _ in range(n_reps):
- n_trls = data.shape[0]
- idx_rnd = np.concatenate([np.zeros(n_trls // 2), np.ones(n_trls // 2)])
- if (n_trls % 2) > 0:
- idx_rnd = np.concatenate([idx_rnd, [1.0]])
- random.shuffle(idx_rnd)
- rnd[_, :] = np.array(
- decode(data, idx_rnd, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
- return obs, rnd
- def decode_epoch_diff(data1, data2, labels, n_reps=100, method='svm', n_inter=10, tail=1, n_jobs=-1):
- n_cells_1 = data1.shape[1]
- n_cells_2 = data2.shape[1]
- data_all = np.concatenate([data1, data2], axis=1)
- obs1 = np.array(decode(data1, labels, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
- obs2 = np.array(decode(data2, labels, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
- if tail == 1:
- obs = obs1 - obs2
- elif tail == -1:
- obs = obs2 - obs1
- rnd = np.zeros((n_reps, data1.shape[-1]))
- for _ in range(n_reps):
- cell_idx = np.concatenate([np.zeros(n_cells_1), np.ones(n_cells_2)])
- random.shuffle(cell_idx)
- data1_rnd = data_all[:, cell_idx == 0, :]
- data2_rnd = data_all[:, cell_idx == 1, :]
- rnd_1 = np.array(
- decode(data1_rnd, labels, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
- rnd_2 = np.array(
- decode(data2_rnd, labels, method=method, n_inter=n_inter, return_inter=True, n_jobs=n_jobs)).mean(0)
- if tail == 1:
- rnd[_, :] = rnd_1 - rnd_2
- elif tail == -1:
- rnd[_, :] = rnd_2 - rnd_1
- return obs, rnd
- def smooth(x, window_len=5, window='hanning'):
- s = np.r_[x[window_len - 1:0:-1], x, x[-2:-window_len - 1:-1]]
- if window == 'flat': # moving average
- w = np.ones(window_len, 'd')
- else:
- w = eval('np.' + window + '(window_len)')
- y = np.convolve(w / w.sum(), s, mode='valid')
- return y[int(window_len / 2):-int((window_len / 2))]
- def compute_perm_stats(obs_parts, rnd_parts, tails, if_smooth=False):
- clu_times, clu_lables = [], []
- for _ in range(obs_parts.shape[0]):
- if if_smooth:
- obs = smooth(obs_parts[_, :])
- rnd = np.array([smooth(rnd_parts[_, rep, :]) for rep in range(rnd_parts.shape[1])])
- else:
- obs = obs_parts[_, :]
- rnd = rnd_parts[_, :, :]
- clu_t, clu_lable = permutation_test(obs, rnd.T, tail=tails[_])
- clu_times.append(clu_t)
- clu_lables.append(clu_lable)
- return clu_times, clu_lables
- def melt_data(data, models, times):
- dat_models = []
- for a, m in enumerate(models):
- dat = pd.DataFrame(data[a, :].T)
- # melt data into a long format
- dat['time (s)'] = times
- dat['Model'] = m # add the model label before melting
- dat_models.append(pd.melt(dat, id_vars=['time (s)', 'Model']))
- # combine every model into long format
- data_melt = pd.concat(dat_models)
- return data_melt
- def get_slices(clu_t, clu_lable, times):
- clusters = np.unique(clu_lable)
- clusters = clusters[clusters > 0]
- result = np.zeros((len(clusters), 3))
- for _ in range(len(clusters)):
- idc = clu_lable == _ + 1
- p_value = np.unique(clu_t[idc])[0]
- slices = times[idc]
- result[_, :] = p_value, slices[0], slices[-1]
- return result
- def plot_perm_results(obs, clt_times, clt_labels, model_names, part_names, times=np.linspace(-0.5, 2, 250),
- threshold=0.050001):
- data_all_parts = []
- for m in range(len(model_names)):
- data_melted = melt_data(obs[m], models=part_names, times=times)
- data_melted['variable'] = model_names[m]
- data_all_parts.append(data_melted)
- df = pd.concat(data_all_parts)
- df['decoding accuracy'] = df['value']
- sns.set_style("ticks")
- sns.set_context("notebook", rc={"lines.linewidth": 3})
- cols = sns.cubehelix_palette(len(clt_times[0]), rot=-.25, light=.7)
- g = sns.FacetGrid(df, col="variable", hue="Model", palette=cols, legend_out=True)
- g.map(sns.lineplot, "time (s)", "decoding accuracy")
- for m in range(len(model_names)):
- g.axes[0][m].axhline(0, linestyle='--', linewidth=0.8, color='black')
- g.axes[0][m].axhline(0.5, linestyle='--', linewidth=0.8, color='black')
- g.axes[0][m].axvline(0, linestyle='--', linewidth=0.8, color='black')
- g.axes[0][m].axvline(0.5, linestyle='--', linewidth=0.8, color='black')
- g.axes[0][m].axvline(1., linestyle='--', linewidth=0.8, color='black')
- for p in range(len(clt_times[0])):
- slices = get_slices(clt_times[m][p], clt_labels[m][p], times)
- for s_i in range(slices.shape[0]):
- if slices[s_i, 0] <= threshold:
- g.axes[0][m].hlines(xmin=slices[s_i, 1], xmax=slices[s_i, 2], colors=cols[p],
- y=0.45 - (p / 30),
- linewidth=4)
- # g.set_titles(row_template='', col_template='')
- # plt.ylim([-10, 10]
- g.set_titles("{col_name}")
- # g.axes[0][-1].legend(loc='lower left')
- return
- def plot_clusters(ax, clt_times, clt_labels, times, epoch_names, colour_lis, p_threshold=0.05, plot_chance_lvl=0.45):
- for p in range(len(epoch_names)):
- slices = get_slices(clt_times[p], clt_labels[p], times)
- for s_i in range(slices.shape[0]):
- if slices[s_i, 0] <= p_threshold:
- if p == 2:
- ax.hlines(xmin=slices[s_i, 1], xmax=slices[s_i, 2], colors=colour_lis[p],
- y=plot_chance_lvl - (p / 80), linestyles=(0, (1, 0.5)),
- linewidth=2)
- else:
- ax.hlines(xmin=slices[s_i, 1], xmax=slices[s_i, 2], colors=colour_lis[p],
- y=plot_chance_lvl - (p / 80),
- linewidth=2)
- return
- def p_into_stars(p_val):
- if p_val <= 0.001:
- stars = "***"
- elif p_val <= 0.01:
- stars = "**"
- elif p_val <= 0.05:
- stars = "*"
- elif (p_val > 0.05) & (p_val <= 0.1):
- stars = '†'
- elif p_val > 0.1:
- stars = 'ns'
- return stars
- def reg_par_recovery(n_neurons=100, max_noise=5, n_trials=101, fit_to_noise=False, n_bootstraps=100):
- noise_levels = np.linspace(0.0001, max_noise, 9)
- corrs_min = np.zeros((2, len(noise_levels), n_bootstraps))
- r2_min = np.zeros((2, len(noise_levels), n_bootstraps))
- corrs_rnd = np.zeros((2, len(noise_levels), n_bootstraps))
- r2_rnd = np.zeros((2, len(noise_levels), n_bootstraps))
- coefs_min_lis = []
- coefs_rnd_lis = []
- print('Recovering the underlying covariance matrix')
- for sig in tqdm(range(len(noise_levels))):
- for n_bootstrap in range(n_bootstraps):
- # Structured selectivity
- cov = np.zeros((3, 3))
- cov[2, 2] = 1
- # simulate neuronal selectivity profiles of minimal
- s_opt = np.random.multivariate_normal(np.zeros(3), cov, n_neurons)
- # simulate neuronal selectivity profiles of random
- s_rnd = np.random.multivariate_normal(np.zeros(3), np.diag([1 / 3, 1 / 3, 1 / 3]), n_neurons)
- design = np.zeros((4, 3))
- design[0, 0] = -0.5
- design[0, 1] = 0.5
- design[0, 2] = -0.5
- design[1, 0] = 0.5
- design[1, 1] = -0.5
- design[1, 2] = -0.5
- design[2, :] = -0.5
- design[2, 2] = 0.5
- design[3, :] = 0.5 # rewarded
- # construct design matrix populated with orthogonalised coefficients
- design_stack = np.tile(design.T, n_trials).T
- # prepare the regression models
- clf1 = linear_model.LinearRegression(fit_intercept=True)
- clf2 = linear_model.LinearRegression(fit_intercept=True)
- # generate firing rate for the optimal model using orthogonalised coefficients
- r_min = s_opt @ design_stack.T
- r_rnd = s_rnd @ design_stack.T
- # add noise to the firing rates
- r_train_min = r_min.T + np.random.normal(0, noise_levels[sig], (4 * n_trials, n_neurons))
- r_train_rnd = r_rnd.T + np.random.normal(0, noise_levels[sig], (4 * n_trials, n_neurons))
- # try to recover the model when only noise was supplied
- if fit_to_noise:
- r_train_min = np.random.normal(0, noise_levels[sig], (4 * n_trials, n_neurons))
- r_train_rnd = np.random.normal(0, noise_levels[sig], (4 * n_trials, n_neurons))
- design = np.zeros((4, 3))
- design[0, 0] = -0.5;
- design[0, 1] = 0.5;
- design[0, 2] = -0.5
- design[1, 0] = 0.5;
- design[1, 1] = -0.5;
- design[1, 2] = -0.5
- design[2, :] = -0.5;
- design[2, 2] = 0.5
- design[3, :] = 0.5 # rewarded
- design_stack = np.tile(design.T, n_trials).T
- # fit regression to data (underlying structured selectivity) using orthogonalised coefficients
- clf1.fit(design_stack, r_train_min)
- clf2.fit(design_stack, r_train_rnd)
- # get and mean-center orthogonalised coefficients
- coefs_min = clf1.coef_
- coefs_min = coefs_min - np.mean(coefs_min, axis=0, keepdims=True)
- coefs_rnd = clf2.coef_
- coefs_rnd = coefs_rnd - np.mean(coefs_rnd, axis=0, keepdims=True)
- # get true covariance under min generative data
- true_min_min = np.zeros((3, 3))
- true_min_min[2, 2] = np.cov(coefs_min.T)[2, 2]
- true_min_min = true_min_min.flatten()
- m_rnd = np.mean(np.diag(np.cov(coefs_min.T)))
- true_min_rnd = np.diag([m_rnd, m_rnd, m_rnd]).flatten()
- r2_min[0, sig, n_bootstrap] = r2_score(true_min_min, np.cov(coefs_min.T).flatten())
- corrs_min[0, sig, n_bootstrap] = pearsonr(true_min_min, np.cov(coefs_min.T).flatten())[0]
- r2_min[1, sig, n_bootstrap] = r2_score(true_min_rnd, np.cov(coefs_min.T).flatten())
- corrs_min[1, sig, n_bootstrap] = pearsonr(true_min_rnd, np.cov(coefs_min.T).flatten())[0]
- # get true covariance under rnd generative data
- true_rnd_min = np.zeros((3, 3))
- true_rnd_min[2, 2] = np.cov(coefs_rnd.T)[2, 2]
- true_rnd_min = true_rnd_min.flatten()
- m_rnd = np.mean(np.diag(np.cov(coefs_rnd.T)))
- true_rnd_rnd = np.diag([m_rnd, m_rnd, m_rnd]).flatten()
- r2_rnd[0, sig, n_bootstrap] = r2_score(true_rnd_min, np.cov(coefs_rnd.T).flatten())
- corrs_rnd[0, sig, n_bootstrap] = pearsonr(true_rnd_min, np.cov(coefs_rnd.T).flatten())[0]
- r2_rnd[1, sig, n_bootstrap] = r2_score(true_rnd_rnd, np.cov(coefs_rnd.T).flatten())
- corrs_rnd[1, sig, n_bootstrap] = pearsonr(true_rnd_rnd, np.cov(coefs_rnd.T).flatten())[0]
- coefs_min_lis.append(coefs_min)
- coefs_rnd_lis.append(coefs_rnd)
- return corrs_min, r2_min, corrs_rnd, r2_rnd, coefs_min_lis, coefs_rnd_lis
- def grab_variables(Name):
- with open(Name + '.txt', "rb") as f:
- data = pickle.load(f)
- total_cost_over_time = data[1]
- betas_final = data[0][0]
- w_final = data[0][1]
- w_b_final = data[0][2]
- perf_cost_final = data[0][3]
- reg_cost_final = data[0][4]
- del data
- return betas_final, w_final, w_b_final, perf_cost_final, reg_cost_final, total_cost_over_time
- def cos_sim(v1, v2):
- return np.dot(v1, v2) / (norm(v1) * norm(v2))
- def epairs_metric(weights1, weights2, l=20):
- def epairs(weights, l=5):
- N = weights.shape[0]
- angles = np.zeros((N))
- for n in range(N):
- cosine = weights[n, :] @ weights.T / (
- np.linalg.norm(weights[n, :]) *
- np.linalg.norm(weights, axis=1))
- # cosine = np.abs(cosine)
- cosine = np.delete(cosine, n)
- cosine.sort()
- angles[n] = np.median(np.arccos(cosine[-l:]))
- return angles
- return np.abs(np.mean(epairs(weights1, l=l)) - np.mean(epairs(weights2, l=l)))
- def split_data(data, labels1, n_splits=10, min_trl=100, n_condi=8):
- data1_splits = np.zeros((n_splits, int((min_trl / n_splits) * n_condi), data.shape[1], 250))
- for c in range(n_condi):
- data1 = data[labels1 == c, :, :]
- trl_idc = np.repeat(list(range(n_splits)), int(min_trl / n_splits))
- random.shuffle(trl_idc)
- for split in range(n_splits):
- data1_splits[split, c * 10:c * 10 + 10, :, :] = data1[trl_idc == split, :, :]
- return data1_splits
- def get_sd_dimensions(data, labels, time_window, method='svm', n_jobs=None, n_inter=1, n_condi=16):
- all_combos = list(itertools.combinations(list(range(n_condi)), int(n_condi / 2)))
- n_combos = int(len(all_combos) / 2)
- decoding_targets = np.zeros((n_combos, n_condi))
- for i in range(n_combos):
- decoding_targets[i, all_combos[i]] = 1
- n_combos = decoding_targets.shape[0]
- scores = []
- print('Cutting along all possible axes (shattering dimensionality)')
- for combo in tqdm(range(n_combos)):
- if n_jobs != None:
- X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=True)
- else:
- X = np.mean(data[:, :, time_window[0]:time_window[1]], axis=-1, keepdims=False)
- y = assign_lables(labels, decoding_targets[combo, :])
- scores.append(decode(X, y, method=method, n_jobs=n_jobs, n_inter=n_inter))
- return np.array(scores)
- def get_sd_min_axis(data, mode='momentum'):
- from scipy.optimize import curve_fit
- def func_sigmoid(x, L, x0, k, b):
- return L / (1 + np.exp(-k * (x - x0))) + b
- axes_min = np.zeros(data.shape[0])
- for rep in range(data.shape[0]):
- if mode == 'momentum':
- data_sorted = np.array(sorted(data[rep, :], reverse=True))
- axes_min[rep] = next((i for i, j in enumerate(data_sorted < 0) if j), None)
- elif mode == 'sigmoid':
- data_sorted = np.array(sorted(data[rep, :], reverse=False))
- x = data_sorted
- y = np.array(list(range(1, data_sorted.shape[0] + 1)))
- p0 = [max(y), np.median(x), 1, min(y)]
- axes_min[rep] = curve_fit(func_sigmoid, x, y, p0, maxfev=5000)[0][1]
- return axes_min
- def get_xgen_exp2(data, labels, condi_fac, target_var_fac, method='svm'):
- condi_lis1 = condi_fac[target_var_fac == 0]
- condi_lis2 = condi_fac[target_var_fac == 1]
- combos = []
- for x in condi_lis1:
- for y in condi_lis2:
- combos.append([x, y])
- train_test_combos_unfltr = list(itertools.combinations(combos, 2))
- # filter for duplicates
- train_test_combos = []
- for _ in range(len(train_test_combos_unfltr)):
- lis = list(np.array(train_test_combos_unfltr[_]).flatten())
- has_dubs = len(lis) != len(np.unique(lis))
- if not has_dubs:
- train_test_combos.append(train_test_combos_unfltr[_])
- n_combos = len(train_test_combos)
- scores = []
- print('Cutting along all possible axes (xgen)')
- for combo in tqdm(range(n_combos)):
- train_condi = train_test_combos[combo][0]
- test_condi = train_test_combos[combo][1]
- idc_train = np.in1d(labels, train_condi)
- idc_test = np.in1d(labels, test_condi)
- X1 = data[idc_train, :]
- X2 = data[idc_test, :]
- y = np.concatenate([np.zeros(int(X1.shape[0] / 2)), np.ones(int(X1.shape[0] / 2))])
- scores.append(decode_xgen(X1, X2, y, y, method=method, n_jobs=None))
- return np.mean(scores)
- def cross_val_pca(data, labels, factor, n_comps, n_splits=10):
- # zscore the data
- data = (data - np.mean(data, axis=0, keepdims=True)) / (np.std(data, axis=0, keepdims=True) + 1)
- labels = np.array(assign_lables(labels, factor))
- n_trls = data.shape[0]
- idx_rnd = np.concatenate([np.zeros(n_trls // 2), np.ones(n_trls // 2)])
- if (n_trls % 2) > 0:
- idx_rnd = np.concatenate([idx_rnd, [1.0]])
- var_ratios_splits = []
- for split in range(n_splits):
- random.shuffle(idx_rnd)
- data1 = data[idx_rnd == 0, :, :]
- data2 = data[idx_rnd == 1, :, :]
- labels1 = labels[idx_rnd == 0]
- labels2 = labels[idx_rnd == 1]
- data1 = condi_avg(data1, labels1)
- data2 = condi_avg(data2, labels2)
- dat1 = data1.mean(-1)
- dat2 = data2.mean(-1)
- pca1 = PCA(n_components=n_comps, random_state=42)
- pca2 = PCA(n_components=n_comps, random_state=42)
- pca1.fit(dat1)
- comps = pca1.transform(dat2)
- var_ratio1 = pca1.explained_variance_ratio_
- pca2.fit(dat2)
- comps = pca2.transform(dat1)
- var_ratio2 = pca2.explained_variance_ratio_
- var_ratios_splits.append(np.array([var_ratio1, var_ratio2]).mean(0))
- var_ratio = np.array(var_ratios_splits).mean(0)
- return var_ratio
- def set_seed(seed=None):
- """
- Function that controls randomness. NumPy and random modules must be imported.
- Args:
- seed : Integer
- A non-negative integer that defines the random state. Default is `None`.
- Returns:
- Nothing.
- """
- if seed is None:
- seed = np.random.choice(2 ** 32)
- random.seed(seed)
- np.random.seed(seed)
- print(f'Random seed {seed} has been set.')
- def equalise_data_witihn_session(data, labels, which_trl='end', n_splits=4):
- trl_min_ses = np.min([data[i].shape[0] for i in range(len(data))])
- if which_trl == 'beginning':
- data_re = [data[_][:trl_min_ses, :, :] for _ in range(len(data))]
- labels_re = [labels[_][:trl_min_ses] for _ in range(len(data))]
- elif which_trl == 'end':
- data_re = [data[_][-trl_min_ses:, :, :] for _ in range(len(data))]
- labels_re = [labels[_][-trl_min_ses:] for _ in range(len(data))]
- elif which_trl == 'rnd':
- idx_re = np.random.choice(np.array(list(range(0, data.shape[0]))), trl_min_ses, replace=False)
- data_re = [data[_][idx_re, :, :] for _ in range(len(data))]
- labels_re = [labels[_][idx_re] for _ in range(len(data))]
- elif which_trl == 'middle':
- labels_re = []
- data_re = []
- for i_sess in range(len(data)):
- idc_half = int(len(labels[i_sess]) / 2)
- n_trls = len(labels[i_sess])
- if (n_trls % 2) > 0:
- labels_re.append(labels[i_sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2) + 1])
- data_re.append(data[i_sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2) + 1, :, :])
- else:
- labels_re.append(labels[i_sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2)])
- data_re.append(data[i_sess][idc_half - int(n_trls / 2):idc_half + int(n_trls / 2), :, :])
- min_condi = int(
- np.min([np.min(list(get_freqs(labels_re[i]).values())) for i in range(len(labels_re))]) / n_splits) - 1
- n_trl_bck = int(trl_min_ses / n_splits)
- trial_start = 0
- trial_stop = n_trl_bck
- labels_all_blcks = []
- data_all_blcks = []
- for split in range(n_splits):
- data_block = [data_re[i][trial_start:trial_stop, :, :] for i in range(len(data_re))]
- label_block = [labels_re[i][trial_start:trial_stop] for i in range(len(data_re))]
- data_blck, label_blck = prepare_data(data_block, label_block, which_trl='beginning', set_min_trl=min_condi)
- trial_start += n_trl_bck
- trial_stop += n_trl_bck
- data_all_blcks.append(data_blck)
- labels_all_blcks.append(label_blck)
- min_trl_blck_post = np.min([data_all_blcks[i].shape[0] for i in range(n_splits)])
- data_eq = np.array([data_all_blcks[i][:min_trl_blck_post, :, :] for i in range(n_splits)])
- labels_eq = np.array([labels_all_blcks[i][:min_trl_blck_post] for i in range(n_splits)])
- return data_eq, labels_eq
- def compute_95ci(data, n_bootstraps=10000):
- # Bootstrap resampling
- bootstrap_means = np.empty(n_bootstraps)
- for i in range(n_bootstraps):
- bootstrap_sample = np.random.choice(data, size=len(data), replace=True)
- bootstrap_means[i] = np.mean(bootstrap_sample)
- # Compute the 95% confidence interval
- lower_bound = np.percentile(bootstrap_means, 2.5)
- upper_bound = np.percentile(bootstrap_means, 97.5)
- return lower_bound, upper_bound
- def get_xgen_cross_set_null(X1, X2, labels1, labels2, n_reps=100, method='svm', n_jobs=None):
- xgen_null = np.zeros(n_reps)
- print('Computing null distribution (cross-variable generalisation)')
- for _ in tqdm(range(n_reps)):
- labels2_rnd = labels2.copy()
- np.random.shuffle(labels2_rnd)
- xgen_null[_] = decode_xgen_within_ses(X1, X2, labels1, labels2_rnd, method=method, n_jobs=n_jobs)
- return xgen_null
- def compute_p_value(obs1, obs2, rnd1, rnd2, tail='greater'):
- diff = obs2 - obs1
- diff_rnd = rnd2 - rnd1
- if tail == 'smaller':
- p_val = np.sum(diff_rnd > diff) / diff_rnd.shape[0]
- elif tail == 'greater':
- p_val = np.sum(diff_rnd < diff) / diff_rnd.shape[0]
- elif tail == 'two':
- p_val = np.sum(np.abs(diff_rnd) >= abs(diff)) / diff_rnd.shape[0]
- else:
- raise ValueError('tail must be greater, smaller or two')
- print(' Stats: M1 = ', str(round(obs1, 3)), ', M2 = ', str(round(obs2, 3)), ' | p-value = ', str(round(p_val, 3)))
- return p_val
- def run_within_session_decoding(X, X_times, labels, variables, N_SPLITS, N_REPS, TIME_WINDOW, splitting_factor):
- n_variables = len(variables)
- splitting_labels = np.array(assign_lables(labels[0, :], factor=splitting_factor))
- X_set1 = X[:, splitting_labels == 0, :]
- labels_set1 = labels[:, splitting_labels == 0]
- X_set2 = X[:, splitting_labels == 1, :]
- labels_set2 = labels[:, splitting_labels == 1]
- decoding = np.zeros((n_variables, 2, N_SPLITS))
- decoding_nulls = np.zeros((n_variables, 2, N_SPLITS, N_REPS))
- for i_var in range(n_variables):
- for i_block in range(N_SPLITS):
- decoding[i_var, 0, i_block] = decode(X[i_block, :, :],
- assign_lables(labels[i_block], factor=variables[i_var]), method='svm',
- n_inter=40)
- decoding_nulls[i_var, 0, i_block, :] = get_decoding_null(X_times[i_block, :, :, :],
- assign_lables(labels[i_block],
- factor=variables[i_var]),
- TIME_WINDOW,
- n_jobs=None, method='svm', n_reps=N_REPS)
- decoding[i_var, 1, i_block] = decode_xgen_within_ses(X_set1[i_block, :, :], X_set2[i_block, :, :],
- assign_lables(labels_set1[i_block],
- factor=variables[i_var][
- splitting_factor == 0]),
- assign_lables(labels_set2[i_block],
- factor=variables[i_var][
- splitting_factor == 0]),
- method='svm', n_jobs=None)
- decoding_nulls[i_var, 1, i_block, :] = get_xgen_cross_set_null(X_set1[i_block, :, :], X_set2[i_block, :, :],
- assign_lables(labels_set1[i_block],
- factor=variables[i_var][
- splitting_factor == 0]),
- assign_lables(labels_set2[i_block],
- factor=variables[i_var][
- splitting_factor == 0]),
- method='svm', n_jobs=None, n_reps=N_REPS)
- return decoding, decoding_nulls
- def plot_within_session_decoding(decoding, decoding_nulls, variable_names, learning_stages, trails, ylim=[0.4, 0.8],
- plot_p=True):
- fig, ax = plt.subplots(1, len(variable_names), figsize=(2.5 * len(variable_names), 2.5))
- for i_var, variables in enumerate(variable_names):
- ax[i_var].plot(learning_stages, decoding[i_var, 0, :], label='decoding', color='black')
- ax[i_var].plot(learning_stages, decoding[i_var, 1, :], label='cross-gen.\ndecoding', color='grey', zorder=-5)
- ax[i_var].scatter(learning_stages, decoding[i_var, 0, :], color='black')
- ax[i_var].scatter(learning_stages, decoding[i_var, 1, :], color='grey', zorder=-5)
- ax[i_var].set_title(variables)
- ax[i_var].set_ylim(ylim)
- ax[i_var].set_xlabel('learning stage')
- ax[i_var].set_ylabel('accuracy')
- ax[i_var].axhline(0.5, color='black', linestyle='--')
- sns.despine(top=True, right=True)
- # plot p-values for decoding and cross-gen decoding
- if plot_p:
- p_val_dec = compute_p_value(decoding[i_var, 0, -1], decoding[i_var, 0, 0], decoding_nulls[i_var, 0, -1],
- decoding_nulls[i_var, 0, 1], tail=trails[i_var])
- p_val_cross = compute_p_value(decoding[i_var, 1, -1], decoding[i_var, 1, 0], decoding_nulls[i_var, 1, -1],
- decoding_nulls[i_var, 1, 1], tail=trails[i_var])
- ax[i_var].text(0.5, 0.75, 'p = ' + str(np.round(p_val_dec, 3)), fontsize=8, transform=ax[i_var].transAxes,
- ha='center', color='black')
- ax[i_var].text(0.5, 0.65, 'p = ' + str(np.round(p_val_cross, 3)), fontsize=8, transform=ax[i_var].transAxes,
- ha='center', color='grey')
- ax[0].legend()
- plt.tight_layout()
- plt.show()
- def decode_xgen_within_ses(X1, X2, y1, y2, method='svm', n_jobs=-1, n_iter=10):
- # prepare a series of classifier applied at each time sample
- if method == 'svm':
- clf1 = make_pipeline(StandardScaler(), SVM(C=5e-4))
- if n_jobs != None:
- clf1 = SlidingEstimator(clf1, verbose=False, n_jobs=-1)
- clf2 = make_pipeline(StandardScaler(), SVM(C=5e-4))
- if n_jobs != None:
- clf2 = SlidingEstimator(clf2, verbose=False, n_jobs=-1)
- elif method == 'lda':
- clf1 = make_pipeline(StandardScaler(), LDA())
- if n_jobs != None:
- clf1 = SlidingEstimator(clf1, verbose=False, n_jobs=-1)
- clf2 = make_pipeline(StandardScaler(), LDA())
- if n_jobs != None:
- clf2 = SlidingEstimator(clf2, verbose=False, n_jobs=-1)
- scores = np.zeros((n_iter, 2))
- for _ in range(n_iter):
- idc1 = np.arange(X1.shape[0])
- idc2 = np.arange(X2.shape[0])
- idx_rnd1 = np.random.choice(idc1, X1.shape[0], replace=True)
- idx_rnd2 = np.random.choice(idc2, X2.shape[0], replace=True)
- X1_rnd = X1[idx_rnd1, :]
- y1_rnd = np.array(y1)[idx_rnd1]
- X2_rnd = X2[idx_rnd2, :]
- y2_rnd = np.array(y2)[idx_rnd2]
- clf1.fit(X1_rnd, y1_rnd)
- scores[_, 0] = clf1.score(X2_rnd, y2_rnd)
- clf2.fit(X2_rnd, y2_rnd)
- scores[_, 1] = clf2.score(X1_rnd, y1_rnd)
- return scores.mean((1, 0))
- def split_into_blocks(dat, labels, n_splits):
- n_trl_blck = int(dat.shape[0] / n_splits)
- i_start = 0
- i_stop = n_trl_blck
- data_blocks = []
- labels_blocks = []
- for i_block in range(n_splits):
- data_blocks.append(dat[i_start:i_stop, :, :])
- labels_blocks.append(labels[i_start:i_stop])
- i_start += n_trl_blck
- i_stop += n_trl_blck
- return data_blocks, labels_blocks
- def split_vector(input_vector, k, d):
- n = len(input_vector) # Total number of elements in the input vector
- output_vectors = []
- # Calculate the step between the starts of each output vector to distribute elements evenly
- step = max(1, (n - d) // (k - 1)) if k > 1 else 0
- # Generate the output vectors
- for i in range(k):
- start_index = i * step
- end_index = start_index + d
- # Adjust the end index if it goes beyond the input vector length
- if end_index > n:
- start_index = max(0, n - d) # Move back to fit the last vector
- end_index = n
- output_vector = input_vector[start_index:end_index]
- output_vectors.append(output_vector)
- if end_index == n: # Stop if the last vector reaches the end of the input vector
- break
- return output_vectors
- def split_data_blocks_moveavg(data, labels, N_SPLITS, N_WINDOWS):
- dat_split, labels_split = [], []
- for i_sess in range(len(data)):
- dat_split_ses, labels_split_ses = split_into_blocks(data[i_sess], labels[i_sess], N_SPLITS)
- dat_split.append(dat_split_ses)
- labels_split.append(labels_split_ses)
- dat_blocks, labs_blocks = [], []
- for i_block in range(N_SPLITS):
- dat_blocks.append([dat_split[_][i_block] for _ in range(len(data))])
- labs_blocks.append([labels_split[_][i_block] for _ in range(len(data))])
- data_blocks, labels_blocks = [], []
- for i_block in range(N_SPLITS):
- dat = dat_blocks[i_block]
- labs = labs_blocks[i_block]
- trl_number_min = np.min([dat[_].shape[0] for _ in range(len(dat))])
- dat_block_windows = []
- labels_block_windows = []
- for i_sess in range(len(dat)):
- dat_ses = dat[i_sess]
- labs_ses = labs[i_sess]
- idc_trials = np.arange(dat_ses.shape[0])
- idc_windows = split_vector(idc_trials, N_WINDOWS, trl_number_min - N_WINDOWS)
- dat_block_windows.append([dat_ses[idc_windows[_], :, :] for _ in range(N_WINDOWS)])
- labels_block_windows.append([labs_ses[idc_windows[_]] for _ in range(N_WINDOWS)])
- data_pseudopop_wins, labs_pseudopop_wins = [], []
- for i_window in range(N_WINDOWS):
- dat_window = [dat_block_windows[_][i_window] for _ in range(len(dat))]
- labels_window = [labels_block_windows[_][i_window] for _ in range(len(dat))]
- data_pseudopop, labs_pseudopop = prepare_data(dat_window, labels_window, which_trl="end")
- data_pseudopop_wins.append(data_pseudopop)
- labs_pseudopop_wins.append(labs_pseudopop)
- data_pseudopop_wins = np.array(data_pseudopop_wins)
- labs_pseudopop_wins = np.array(labs_pseudopop_wins)
- data_blocks.append(data_pseudopop_wins)
- labels_blocks.append(labs_pseudopop_wins)
- return data_blocks, labels_blocks
- def split_data_blocks_moveavg_sel(data, labels, N_SPLITS):
- dat_split, labels_split = [], []
- for i_sess in range(len(data)):
- dat_split_ses, labels_split_ses = split_into_blocks(data[i_sess], labels[i_sess], N_SPLITS)
- dat_split.append(dat_split_ses)
- labels_split.append(labels_split_ses)
- dat_blocks, labs_blocks = [], []
- for i_block in range(N_SPLITS):
- dat_blocks.append([dat_split[_][i_block] for _ in range(len(data))])
- labs_blocks.append([labels_split[_][i_block] for _ in range(len(data))])
- return dat_blocks, labs_blocks
- def compute_p_val_learning(data_first, data_last, n_perm=10000, tail='greater'):
- mean_diff = np.mean(data_last) - np.mean(data_first)
- data_all = np.concatenate((data_first, data_last))
- mean_diff_rnd = []
- for _ in range(n_perm):
- idc_rnd = np.random.permutation(len(data_all))
- data_first_rnd = data_all[idc_rnd[:len(data_first)]]
- data_last_rnd = data_all[idc_rnd[len(data_first):]]
- mean_diff_rnd.append(np.mean(data_last_rnd) - np.mean(data_first_rnd))
- if tail == 'greater':
- p_value = np.sum(mean_diff_rnd > mean_diff) / len(mean_diff_rnd)
- elif tail == 'less':
- p_value = np.sum(mean_diff_rnd < mean_diff) / len(mean_diff_rnd)
- elif tail == 'two-sided':
- p_value = np.sum(np.abs(mean_diff_rnd) > np.abs(mean_diff)) / len(mean_diff_rnd)
- return p_value
- def plot_fixation_breaks(reward_prop, animal):
- plt.figure()
- sessions = list(range(1, len(reward_prop) + 1))
- plt.plot(sessions, reward_prop, label='No reward - reward')
- plt.legend()
- plt.xlabel('Session')
- plt.ylabel('Proportion of fixation breaks')
- plt.title(animal)
- plt.show()
- def sum_values(trial_types, keys):
- return sum([trial_types.get(key, 0) for key in keys])
- def compute_p_val_learning(data_first, data_last, n_perm=10000, tail='greater'):
- mean_diff = np.mean(data_last) - np.mean(data_first)
- data_all = np.concatenate((data_first, data_last))
- mean_diff_rnd = []
- for _ in range(n_perm):
- idc_rnd = np.random.permutation(len(data_all))
- data_first_rnd = data_all[idc_rnd[:len(data_first)]]
- data_last_rnd = data_all[idc_rnd[len(data_first):]]
- mean_diff_rnd.append(np.mean(data_last_rnd) - np.mean(data_first_rnd))
- if tail == 'greater':
- p_value = np.sum(mean_diff_rnd > mean_diff) / len(mean_diff_rnd)
- elif tail == 'less':
- p_value = np.sum(mean_diff_rnd < mean_diff) / len(mean_diff_rnd)
- elif tail == 'two-sided':
- p_value = np.sum(np.abs(mean_diff_rnd) > np.abs(mean_diff)) / len(mean_diff_rnd)
- return p_value
- def plot_mean_and_ci(reward, no_reward, n_perm=10000, tail='greater'):
- plt.figure(figsize=(4, 3))
- x = list(range(1, len(reward) + 1))
- y1 = np.array([np.mean(reward[i], axis=0) for i in range(len(reward))])
- y2 = np.array([np.mean(no_reward[i], axis=0) for i in range(len(no_reward))])
- # plote 95% confidence interval
- y1_95ci = np.array([1.96 * np.std(reward[i], axis=0) / np.sqrt(len(reward[i])) for i in range(len(reward))])
- y2_95ci = np.array(
- [1.96 * np.std(no_reward[i], axis=0) / np.sqrt(len(no_reward[i])) for i in range(len(no_reward))])
- plt.plot(x, y1, label='reward', color='black')
- plt.errorbar(x, y1, yerr=y1_95ci, fmt='o', color='black')
- # plt.fill_between(x, y1 - y1_sem, y1 + y1_sem, alpha=0.5)
- plt.plot(x, y2, label='no reward', color='grey')
- plt.errorbar(x, y2, yerr=y2_95ci, fmt='o', color='grey')
- # plt.fill_between(x, y2 - y2_sem, y2 + y2_sem, alpha=0.5)
- p_value_norew = compute_p_val_learning(no_reward[0], no_reward[-1], n_perm=n_perm, tail=tail)
- p_value_rew = compute_p_val_learning(reward[0], reward[-1], n_perm=n_perm, tail=tail)
- # annotate the plot with p-value
- plt.text(0.1, 0.9, f'p-value no reward = {round(p_value_norew, 3)}\np-value reward = {round(p_value_rew, 3)}',
- horizontalalignment='left', verticalalignment='center', transform=plt.gca().transAxes)
- plt.legend()
- plt.xlabel('learning stage')
- plt.xticks([1, 2, 3, 4])
- plt.ylabel('fixation breaks (%)')
- plt.legend(['reward', 'no reward'], loc='lower left')
- sns.despine(top=True, right=True)
- plt.tight_layout()
- plt.show()
- 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):
- ax = fig.add_subplot(gs)
- x = list(range(1, len(reward_prop) + 1))
- y1 = np.array([np.mean(reward_prop[i], axis=0) for i in range(len(reward_prop))])
- y1_95ci = np.array(
- [1.96 * np.std(reward_prop[i], axis=0) / np.sqrt(len(reward_prop[i])) for i in range(len(reward_prop))])
- # plot individual datapoints as empty circles with jitter
- for i, stage_data in enumerate(reward_prop):
- jitter = np.random.uniform(-jitter_width, jitter_width, size=len(stage_data))
- ax.scatter(np.full(len(stage_data), x[i]) + jitter, stage_data,
- facecolors='none', edgecolors='black', s=20, zorder=1, alpha=0.5)
- ax.plot(x, y1, color='black', zorder=2)
- ax.errorbar(x, y1, yerr=y1_95ci, fmt='o', color='black', zorder=3)
- if stat == 'reg':
- y_flat = list(itertools.chain.from_iterable(reward_prop))
- x_flat = [[i+1] * len(reward_prop[i]) for i in range(len(reward_prop))]
- x_flat = list(itertools.chain.from_iterable(x_flat))
- slope, intercept, p_value, r_value = compute_lin_reg(x_flat, y_flat, n_perm=n_perm, tail=tail)
- ax.text(0.1, 0.9, f'p-value = {round(p_value, 3)}\nr = {round(r_value, 3)}',
- horizontalalignment='left', verticalalignment='center', transform=plt.gca().transAxes)
- elif stat == 'ttest':
- p_value_rew_prop = compute_p_val_learning(reward_prop[0], reward_prop[-1], n_perm=n_perm, tail=tail)
- ax.text(0.1, 0.9, f'p-value = {round(p_value_rew_prop, 3)}',
- horizontalalignment='left', verticalalignment='center', transform=plt.gca().transAxes)
- if plot_comp_dat != False:
- y2 = np.array([np.mean(plot_comp_dat[i], axis=0) for i in range(len(plot_comp_dat))])
- y2_95ci = np.array(
- [1.96 * np.std(plot_comp_dat[i], axis=0) / np.sqrt(len(plot_comp_dat[i])) for i in
- range(len(plot_comp_dat))])
- # plot individual datapoints for comparison data
- for i, stage_data in enumerate(plot_comp_dat):
- jitter = np.random.uniform(-jitter_width, jitter_width, size=len(stage_data))
- ax.scatter(np.full(len(stage_data), x[i]) + jitter, stage_data,
- facecolors='none', edgecolors='grey', s=20, zorder=1, alpha=0.5)
- ax.plot(x, y2, color='grey', zorder=2)
- ax.errorbar(x, y2, yerr=y2_95ci, fmt='o', color='grey', zorder=3)
- ax.set_xlabel('learning stage')
- ax.set_xticks(x)
- ax.set_ylabel(y_label)
- ax.set_ylim([vmin, vmax])
- sns.despine(top=True, right=True)
- ax.axhline(y=baseline_val, color='black', linestyle='--', linewidth=0.8)
- ax.set_title(title)
- def get_fixation_breaks(sessions_animals, experiment_label='exp1'):
- with open('config.yml', 'r') as f:
- configs = yaml.safe_load(f)
- code_labels = configs['TRIGGER_CODES']
- fix_breaks_animals = []
- for animal in range(len(sessions_animals)):
- sessions = sessions_animals[animal]
- fix_numers = np.zeros((len(sessions), 5))
- for s in range(len(sessions)):
- print('Session: ', s + 1)
- beh = io.loadmat(configs['PATHS']['in_template_beh'].format(sessions[s]))
- codes = beh['uecode']
- codes = list(codes[0, :])
- times = beh['timingms']
- times = list(times[0, :])
- size = len(codes)
- idx_list = [idx for idx, val in
- enumerate(codes) if val == 6116]
- trials = [codes[i: j] for i, j in
- zip([0] + idx_list, idx_list +
- ([size] if idx_list[-1] != size else []))]
- trials = trials[1:]
- conditions = []
- for _ in range(len(trials)):
- if code_labels['CUE1_ON'] in trials[_] and code_labels['TARGET1_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(1)
- elif code_labels['CUE1_ON'] in trials[_] and code_labels['TARGET2_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(2)
- elif code_labels['CUE1_ON'] in trials[_] and code_labels['TARGET3_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(3)
- elif code_labels['CUE1_ON'] in trials[_] and code_labels['TARGET4_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(4)
- elif code_labels['CUE2_ON'] in trials[_] and code_labels['TARGET1_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(5)
- elif code_labels['CUE2_ON'] in trials[_] and code_labels['TARGET2_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(6)
- elif code_labels['CUE2_ON'] in trials[_] and code_labels['TARGET3_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(7)
- elif code_labels['CUE2_ON'] in trials[_] and code_labels['TARGET4_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(8)
- elif code_labels['CUE3_ON'] in trials[_] and code_labels['TARGET1_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(9)
- elif code_labels['CUE3_ON'] in trials[_] and code_labels['TARGET2_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(10)
- elif code_labels['CUE3_ON'] in trials[_] and code_labels['TARGET3_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(11)
- elif code_labels['CUE3_ON'] in trials[_] and code_labels['TARGET4_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(12)
- elif code_labels['CUE4_ON'] in trials[_] and code_labels['TARGET1_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(13)
- elif code_labels['CUE4_ON'] in trials[_] and code_labels['TARGET2_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(14)
- elif code_labels['CUE4_ON'] in trials[_] and code_labels['TARGET3_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(15)
- elif code_labels['CUE4_ON'] in trials[_] and code_labels['TARGET4_ON'] in trials[_] and (
- code_labels['BREAK_TARGET_ERROR'] in trials[_]):
- conditions.append(16)
- if code_labels['BREAK_CUE_ERROR'] in trials[_] or code_labels['BREAK_TARGET_ERROR'] in trials[_] or (
- code_labels['BREAK_ERROR'] in trials[_]) or (code_labels['FIXATION_ERROR'] in trials[_]):
- conditions.append(99)
- if code_labels['BREAK_CUE_ERROR'] in trials[_]:
- conditions.append(999)
- all_trls = []
- for _ in range(len(trials)):
- if code_labels['CUE1_ON'] in trials[_]:
- all_trls.append(1)
- elif code_labels['CUE2_ON'] in trials[_]:
- all_trls.append(1)
- elif code_labels['CUE3_ON'] in trials[_]:
- all_trls.append(2)
- elif code_labels['CUE4_ON'] in trials[_]:
- all_trls.append(2)
- trial_types = {value: len(list(freq)) for value, freq in groupby(sorted(conditions))}
- trial_all = {value: len(list(freq)) for value, freq in groupby(sorted(all_trls))}
- if experiment_label == 'exp1':
- fix_numers[s, 0] = len(trials)
- fix_numers[s, 1] = sum_values(trial_types,
- [1, 2, 7, 8]) # number of fixation breaks in reward trials
- fix_numers[s, 2] = sum_values(trial_types,
- [3, 4, 5, 6]) # number of fixation breaks in no reward trials
- fix_numers[s, 3] = trial_types[99] # number of trials with all fixation breaks
- fix_numers[s, 4] = trial_types[999] # number of trials with cue fixation breaks
- elif experiment_label == 'exp2':
- fix_numers[s, 0] = len(trials)
- fix_numers[s, 1] = sum_values(trial_types,
- [9, 10, 15, 16]) # number of fixation breaks in reward trials
- fix_numers[s, 2] = sum_values(trial_types,
- [11, 12, 13, 14]) # number of fixation breaks in no reward trials
- fix_numers[s, 3] = trial_types[99] # number of trials with all fixation breaks
- fix_numers[s, 4] = trial_types[999] # number of trials with cue fixation breaks
- fix_breaks_animals.append(fix_numers)
- return fix_breaks_animals
- def get_data_stages(observe_or_run='observe', file_name=None, return_data=False, session_list=None):
- with open('config.yml', 'r') as file:
- configs = yaml.safe_load(file)
- if observe_or_run == 'run':
- data_all_parts, labels = get_data(session_list=session_list,
- path_spikes=configs['PATHS']['out_template_spks'],
- path_meta=configs['PATHS']['out_template_meta'],
- window=[0, 250],
- cut_off=None
- )
- data, _, _ = exclude_neurons(data=data_all_parts,
- session_list=session_list,
- path_locations=configs['PATHS']['out_template_loc'],
- path_sel_exclude=configs['PATHS']['out_template_sel_list'],
- loc=configs['ANALYSIS_PARAMS']['SAMPLED_AREAS']
- )
- save_data([data, labels], ['data', 'labels'], configs['PATHS']['output_path'] + file_name + '.pickle')
- elif observe_or_run == 'observe':
- return_data = True
- obj_loaded = load_data(configs['PATHS']['output_path'] + file_name + '.pickle')
- data = obj_loaded['data']
- labels = obj_loaded['labels']
- if return_data:
- return data, labels
- def create_sliding_windows_adaptive(sessions, n_stages=5, min_window_ratio=0.25):
- """
- Create n overlapping windows with adaptive overlap, ensuring continuous coverage.
- """
- total = len(sessions)
- # Window size based on minimum ratio
- window_size = max(2, int(np.ceil(total * min_window_ratio)))
- window_size = min(window_size, total)
- if n_stages == 1:
- return [sessions]
- # Calculate step size to evenly distribute windows across all sessions
- # This ensures the last window ends at total while maintaining overlap
- step = (total - window_size) / (n_stages - 1)
- windows = []
- for i in range(n_stages):
- start = int(i * step)
- end = min(start + window_size, total)
- # Ensure we don't go past the end
- if end > total:
- end = total
- start = max(0, end - window_size)
- windows.append(sessions[start:end])
- return windows
- def combine_session_lists(mode='time', which_exp='exp1', combine_all=True):
- with open('config.yml', 'r') as file:
- configs = yaml.safe_load(file)
- if which_exp == 'exp1':
- animal1_ses_labels = configs['SESSION_NAMES']['sessions_womble_1']
- animal2_ses_labels = configs['SESSION_NAMES']['sessions_wilfred_1']
- elif which_exp == 'exp2':
- animal1_ses_labels = configs['SESSION_NAMES']['sessions_womble_2']
- animal2_ses_labels = configs['SESSION_NAMES']['sessions_wilfred_2']
- if mode == 'time':
- min_num_sessions = min(len(animal1_ses_labels), len(animal2_ses_labels))
- sesions_split_animal1 = np.array_split(animal1_ses_labels, min_num_sessions)
- sesions_split_animal2 = np.array_split(animal2_ses_labels, min_num_sessions)
- ses_com = []
- for _ in range(min_num_sessions):
- ses_com.append(list(sesions_split_animal1[_]))
- ses_com.append(list(sesions_split_animal2[_]))
- if combine_all:
- ses_com = list(chain.from_iterable(ses_com))
- elif mode == 'fix_bias':
- session_labels = [animal1_ses_labels, animal2_ses_labels]
- fix_breaks_animals = get_fixation_breaks(session_labels, experiment_label=which_exp)
- sessions_sorted = []
- for i_anim in range(2):
- no_reward_fix = fix_breaks_animals[i_anim][:, 2] / fix_breaks_animals[i_anim][:, 3]
- reward_fix = fix_breaks_animals[i_anim][:, 1] / fix_breaks_animals[i_anim][:, 3]
- dat = no_reward_fix / reward_fix
- # sort data_prop_s and get the indices of the sorting
- sort_idx = np.argsort(dat)
- sessions_sorted.append(np.array(session_labels[i_anim])[sort_idx])
- ses_wom = np.array_split(sessions_sorted[0], configs['ANALYSIS_PARAMS']['N_STAGES'])
- ses_wil = np.array_split(sessions_sorted[1], configs['ANALYSIS_PARAMS']['N_STAGES'])
- ses_com = []
- for _ in range(configs['ANALYSIS_PARAMS']['N_STAGES']):
- ses_com.append(list(ses_wom[_]) + list(ses_wil[_]))
- if combine_all:
- ses_com = list(chain.from_iterable(ses_com))
- elif mode == 'stages':
- ses_wom = np.array_split(animal1_ses_labels, configs['ANALYSIS_PARAMS']['N_STAGES'])
- ses_wil = np.array_split(animal2_ses_labels, configs['ANALYSIS_PARAMS']['N_STAGES'])
- ses_com = []
- for _ in range(configs['ANALYSIS_PARAMS']['N_STAGES']):
- ses_com.append(list(ses_wom[_]) + list(ses_wil[_]))
- if combine_all:
- ses_com = list(chain.from_iterable(ses_com))
- elif mode == 'cxt_cost':
- session_labels = [animal1_ses_labels, animal2_ses_labels]
- _, _, switch_costs_cxt = get_switch_costs(session_labels)
- sessions_sorted = []
- for i_anim in range(2):
- dat = switch_costs_cxt[i_anim]
- sort_idx = np.argsort(dat)
- sessions_sorted.append(np.array(session_labels[i_anim])[sort_idx])
- ses_wom = np.array_split(sessions_sorted[0], configs['ANALYSIS_PARAMS']['N_STAGES'])
- ses_wil = np.array_split(sessions_sorted[1], configs['ANALYSIS_PARAMS']['N_STAGES'])
- ses_com = []
- for _ in range(configs['ANALYSIS_PARAMS']['N_STAGES']):
- ses_com.append(list(ses_wom[_]) + list(ses_wil[_]))
- if combine_all:
- ses_com = list(chain.from_iterable(ses_com))
- elif mode == 'colour_cost':
- session_labels = [animal1_ses_labels, animal2_ses_labels]
- switch_costs_col, _, _ = get_switch_costs(session_labels)
- sessions_sorted = []
- for i_anim in range(2):
- dat = switch_costs_col[i_anim]
- sort_idx = np.argsort(dat)
- sessions_sorted.append(np.array(session_labels[i_anim])[sort_idx])
- ses_wom = np.array_split(sessions_sorted[0], configs['ANALYSIS_PARAMS']['N_STAGES'])
- ses_wil = np.array_split(sessions_sorted[1], configs['ANALYSIS_PARAMS']['N_STAGES'])
- ses_com = []
- for _ in range(configs['ANALYSIS_PARAMS']['N_STAGES']):
- ses_com.append(list(ses_wom[_]) + list(ses_wil[_]))
- if combine_all:
- ses_com = list(chain.from_iterable(ses_com))
- elif mode == 'sliding_window':
- parts_animal1 = create_sliding_windows_adaptive(animal1_ses_labels, n_stages=configs['ANALYSIS_PARAMS']['N_STAGES'], min_window_ratio=0.25)
- parts_animal2 = create_sliding_windows_adaptive(animal2_ses_labels, n_stages=configs['ANALYSIS_PARAMS']['N_STAGES'], min_window_ratio=0.25)
- ses_com = [pw + pwi for pw, pwi in zip(parts_animal1, parts_animal2)]
- if combine_all:
- ses_com = list(chain.from_iterable(ses_com))
- return ses_com
- def plot_reg_prop(reward_prop, n_perm=10000, tail='greater'):
- plt.figure(figsize=(3, 2))
- # flatten reward_prop and construct list with stage labels
- y_flat = list(itertools.chain.from_iterable(reward_prop))
- x_flat = [1] * len(reward_prop[0]) + [2] * len(reward_prop[1]) + [3] * len(reward_prop[2]) + [4] * len(
- reward_prop[3])
- slope, intercept, p_value, r_value = compute_lin_reg(x_flat, y_flat, n_perm=n_perm, tail=tail)
- # plot regression
- sns.regplot(x=x_flat, y=y_flat, color='black', scatter_kws={'color': 'black'},
- line_kws={'color': 'black'})
- # annotate the plot with p-value
- plt.text(0.1, 0.9, f'p-value = {round(p_value, 3)}\nr = {round(r_value, 3)}',
- horizontalalignment='left', verticalalignment='center', transform=plt.gca().transAxes)
- plt.xlabel('learning stage')
- plt.xticks([1, 2, 3, 4])
- plt.ylabel('no reward/reward\nproporiton')
- sns.despine(top=True, right=True)
- plt.axhline(y=1, color='black', linestyle='--')
- plt.title('fixation breaks')
- plt.tight_layout()
- plt.show()
- def compute_lin_reg(x, y, n_perm=1000, tail='greater'):
- A = np.vstack([x, np.ones(len(x))]).T
- results = np.linalg.lstsq(A, y, rcond=None)
- slope = results[0][0]
- intercept = results[0][1]
- null_slopes = []
- for _ in range(n_perm):
- x_perm = np.random.permutation(x)
- results = np.linalg.lstsq(np.vstack([x_perm, np.ones(len(x))]).T, y, rcond=None)
- null_slopes.append(results[0][0])
- if tail == 'greater':
- p_value = np.sum(null_slopes > slope) / len(null_slopes)
- elif tail == 'less':
- p_value = np.sum(null_slopes < slope) / len(null_slopes)
- elif tail == 'two-sided':
- p_value = np.sum(np.abs(null_slopes) > np.abs(slope)) / len(null_slopes)
- # compute the r value of the regression
- r_value = spearmanr(x, y)[0]
- return slope, intercept, p_value, r_value
- def plot_regression(df, n_perm=100000):
- # Unique animals
- animals = df['animal'].unique()
- # Create subplots
- fig, axes = plt.subplots(1, len(animals), figsize=(5 * len(animals), 5), sharey=False)
- # Check if there is only one subplot (axis) and make it iterable
- if len(animals) == 1:
- axes = [axes]
- # Iterate over each animal and its corresponding axis
- for animal, ax in zip(animals, axes.flatten()):
- # Filter the DataFrame for the current animal
- animal_df = df[df['animal'] == animal]
- # Plot the regression for the current animal
- sns.regplot(x='x', y='y', data=animal_df, ax=ax, color='black', scatter_kws={'color': 'black'},
- line_kws={'color': 'black'})
- # Compute the regression for the current animal
- slope, intercept, p_value, r_value = compute_lin_reg(animal_df['x'], animal_df['y'], n_perm=n_perm,
- tail='greater')
- # Calculate buffer for x and y limits
- x_buffer = (animal_df['x'].max() - animal_df['x'].min()) * 0.1
- y_buffer = (animal_df['y'].max() - animal_df['y'].min()) * 0.1
- # Set individual x and y limits with buffer
- ax.set_xlim(animal_df['x'].min() - x_buffer, animal_df['x'].max() + x_buffer)
- ax.set_ylim(animal_df['y'].min() - y_buffer, animal_df['y'].max() + y_buffer)
- # Annotate the plot with p-value and r-value
- ax.text(0.5, 0.9, f'p-value = {round(p_value, 3)}\nr = {round(r_value, 3)}',
- horizontalalignment='center', verticalalignment='center', transform=ax.transAxes)
- # Set title
- ax.set_title(animal)
- ax.set_xlabel('Session')
- ax.set_ylabel('Proportion of fixation breaks')
- sns.despine(ax=ax, top=True, right=True)
- # Set common labels
- plt.tight_layout()
- plt.show()
- def creat_plot_grid(n_rows, n_cols, size, width_ratios=None):
- fig = plt.figure(figsize=(size * n_cols, (size * n_rows) * 0.8))
- gs = gridspec.GridSpec(n_rows, n_cols, figure=fig, width_ratios=width_ratios)
- return fig, gs
- def split_data_stages_moveavg(data, labels, n_stages, n_windows, trl_min=None, if_rnd=False):
- # collapse list of lists into a single list
- data_lis = list(chain.from_iterable(data))
- trl_number_min = np.min([data_lis[_].shape[0] for _ in range(len(data_lis))])
- data_stages, labels_stages = [], []
- for i_stage in range(n_stages):
- dat = data[i_stage]
- labs = labels[i_stage]
- dat_stage_windows = []
- labels_stage_windows = []
- for i_sess in range(len(dat)):
- dat_ses = dat[i_sess]
- labs_ses = labs[i_sess]
- idc_trials = np.arange(dat_ses.shape[0])
- idc_windows = split_vector(idc_trials, n_windows, trl_number_min - n_windows)
- dat_stage_windows.append([dat_ses[idc_windows[i_win], :, :] for i_win in range(n_windows)])
- labels_stage_windows.append([labs_ses[idc_windows[i_win]] for i_win in range(n_windows)])
- data_pseudopop_wins, labs_pseudopop_wins = [], []
- for i_window in range(n_windows):
- dat_window = [dat_stage_windows[_][i_window] for _ in range(len(dat))]
- labels_window = [labels_stage_windows[_][i_window] for _ in range(len(dat))]
- if trl_min is None:
- data_pseudopop, labs_pseudopop = prepare_data(dat_window, labels_window, which_trl="end")
- else:
- if if_rnd:
- data_pseudopop, labs_pseudopop = prepare_data(dat_window, labels_window, which_trl="rnd",
- set_min_trl=trl_min)
- else:
- data_pseudopop, labs_pseudopop = prepare_data(dat_window, labels_window, which_trl="end",
- set_min_trl=trl_min)
- data_pseudopop_wins.append(data_pseudopop)
- labs_pseudopop_wins.append(labs_pseudopop)
- data_pseudopop_wins = np.array(data_pseudopop_wins)
- labs_pseudopop_wins = np.array(labs_pseudopop_wins)
- data_stages.append(data_pseudopop_wins)
- labels_stages.append(labs_pseudopop_wins)
- return data_stages, labels_stages
- def run_moving_window_decoding(data_eq, labels_eq, variables, time_window, method='svm', n_jobs=None, if_xgen=True,
- if_null=False, n_reps=None):
- n_windows = data_eq[0].shape[0]
- n_stages = len(data_eq)
- shattering_dim = np.zeros((n_stages, 31, n_windows))
- decoding = np.zeros((n_stages, 4, n_windows))
- if if_xgen:
- xgen_decoding = np.zeros((n_stages, 4, n_windows))
- for i_stage in range(n_stages):
- for i_window in range(n_windows):
- X_stage = data_eq[i_stage][i_window, :, :, :]
- y_stage = labels_eq[i_stage][i_window, :]
- decoding_sess, decoding_dich_sess = get_decoding(data=X_stage,
- labels=y_stage,
- variables=variables,
- time_window=time_window,
- method=method,
- n_jobs=n_jobs)
- if if_xgen:
- decoding_xgen_sess = get_xgen(data=X_stage,
- labels=y_stage,
- variables=variables,
- time_window=time_window,
- method=method,
- mode='only_rel',
- n_jobs=n_jobs,
- )
- xgen_decoding[i_stage, :, i_window] = decoding_xgen_sess
- shattering_dim[i_stage, :, i_window] = decoding_dich_sess
- decoding[i_stage, :, i_window] = decoding_sess[:4]
- shattering_dim = shattering_dim.mean(-1)
- decoding = decoding.mean(-1)
- if if_xgen:
- xgen_decoding = xgen_decoding.mean(-1)
- if if_null:
- decoding_null = np.zeros((n_stages, n_reps, n_windows))
- if if_xgen:
- xgen_null = np.zeros((n_stages, n_reps, 4, n_windows))
- for i_rep in tqdm(range(n_reps)):
- for i_stage in range(n_stages):
- X_stage = data_eq[i_stage]
- y_stage = labels_eq[i_stage]
- idc_rnd = np.arange(X_stage.shape[1])
- random.shuffle(idc_rnd)
- y_stage_rnd = y_stage[:, idc_rnd]
- for i_window in range(n_windows):
- X_stage_win = X_stage[i_window, :, :, :]
- y_stage_win = y_stage_rnd[i_window, :]
- labels_random = assign_lables(y_stage_win, factor=[0, 0, 0, 0, 1, 1, 1, 1])
- decoding_null[i_stage, i_rep, i_window] = decode(
- X_stage_win[:, :, time_window[0]: time_window[1]].mean(-1), labels_random, method=method,
- n_jobs=n_jobs)
- if if_xgen:
- xgen_null[i_stage, i_rep, :, i_window] = get_xgen(data=X_stage_win,
- labels=y_stage_win,
- variables=variables,
- time_window=time_window,
- method=method,
- mode='only_rel',
- n_jobs=n_jobs,
- verbose=False
- )
- decoding_null = decoding_null.mean(-1)
- if if_xgen:
- xgen_null = xgen_null.mean(-1)
- if if_xgen:
- if if_null:
- return shattering_dim, decoding, xgen_decoding, decoding_null, xgen_null
- else:
- return shattering_dim, decoding, xgen_decoding
- else:
- if if_null:
- return shattering_dim, decoding, decoding_null
- else:
- return shattering_dim, decoding
- def shuffle_stages(stage1, stage4, label1, seed=None):
- if seed is None:
- random.seed()
- else:
- random.seed(seed)
- data_all = np.concatenate((stage1, stage4), axis=2)
- n_cells_stage1 = stage1.shape[2]
- idc_cells = np.arange(data_all.shape[2])
- random.shuffle(idc_cells) # Python’s shuffle
- dat_stage1_shuffled = data_all[:, :, idc_cells[:n_cells_stage1], :]
- dat_stage4_shuffled = data_all[:, :, idc_cells[n_cells_stage1:], :]
- return [dat_stage1_shuffled, dat_stage4_shuffled], [label1, label1]
- def run_moving_window_decoding_ler_null(data1, data2, labels1, variables, time_window, n_reps=100, if_xgen=True, method='svm'):
- if time_window is not None:
- shattering_dim_rnd = np.zeros((n_reps, 2, 31))
- decoding_rnd = np.zeros((n_reps, 2, 4))
- xgen_decoding_rnd = np.zeros((n_reps, 2, 4))
- for rep in range(n_reps):
- print("Rep: ", rep + 1, " of ", n_reps)
- data_epochs1and4, labels_epochs1and4 = shuffle_stages(data1, data2, labels1)
- if if_xgen:
- shattering_dim_rep, decoding_rep, xgen_decoding_rep = run_moving_window_decoding(data_epochs1and4,
- labels_epochs1and4,
- variables, time_window,
- if_xgen=if_xgen, method=method)
- xgen_decoding_rnd[rep, :, :] = xgen_decoding_rep
- else:
- shattering_dim_rep, decoding_rep = run_moving_window_decoding(data_epochs1and4, labels_epochs1and4,
- variables, time_window, if_xgen=if_xgen, method=method)
- shattering_dim_rnd[rep, :, :] = shattering_dim_rep
- decoding_rnd[rep, :, :] = decoding_rep
- if if_xgen:
- return shattering_dim_rnd, decoding_rnd, xgen_decoding_rnd
- else:
- return shattering_dim_rnd, decoding_rnd
- if time_window is None:
- n_times = data1[0].shape[-1] - 80
- shattering_dim_rnd = np.zeros((n_times, n_reps, 2, 31))
- decoding_rnd = np.zeros((n_times, n_reps, 2, 4))
- for i_rep in range(n_reps):
- print("Rep: ", i_rep + 1, " of ", n_reps)
- data_epochs1and4, labels_epochs1and4 = shuffle_stages(data1, data2, labels1)
- for i_time in range(30, 200):
- time_sliding = [i_time, i_time + 1]
- print('Time point: ', i_time + 1, ' of ', data1[0].shape[-1] - 50)
- shattering_dim_t, decoding_t = run_moving_window_decoding(data_epochs1and4, labels_epochs1and4,
- variables, time_sliding, if_xgen=False, method=method)
- shattering_dim_rnd[i_time-30, i_rep, :, :] = shattering_dim_t
- decoding_rnd[i_time-30, i_rep, :, :] = decoding_t
- return shattering_dim_rnd, decoding_rnd
- def save_data(data_list, names_list, path):
- data_dict = {}
- for i, name in enumerate(names_list):
- data_dict[name] = data_list[i]
- with open(path, 'wb') as f:
- pickle.dump(data_dict, f)
- def load_data(path):
- with open(path, 'rb') as f:
- data_dict = pickle.load(f)
- return data_dict
- def line_plot_timevar(gs, fig, x, y, color, xlabel, ylabel, title, ylim, xlim, xticks=None, xticklabels=None,
- baseline_line=None,
- patch_pars=None, if_title=True, if_sem=False, xaxese_booo=False):
- ax = fig.add_subplot(gs)
- [ax.plot(x, y.mean(0)[_, :], color=color[_]) for _ in range(y.shape[1])]
- sns.despine(top=True, right=True)
- # plot the shaded area using standard error of the mean
- if if_sem:
- [ax.fill_between(x, y.mean(0)[_] - y.std(0)[_] / np.sqrt(y.shape[0]),
- y.mean(0)[_] + y.std(0)[_] / np.sqrt(y.shape[0]), color=color[_], alpha=0.3) for _ in
- range(y.shape[1])]
- ax.set_ylabel(ylabel)
- if xaxese_booo:
- ax.set_xlabel(xlabel)
- ax.set_ylim(ylim)
- ax.set_xlim(xlim)
- if if_title:
- ax.set_title(title)
- if xticks is not None:
- ax.set_xticks(xticks)
- [ax.axvline(xticks[_], linewidth=0.8, color='black', linestyle='--') for _ in range(len(xticks))]
- if xticklabels is not None:
- ax.set_xticklabels(xticklabels)
- if baseline_line is not None:
- ax.axhline(baseline_line, color='black', linestyle='--', linewidth=0.8)
- if patch_pars is not None:
- rect = patches.Rectangle(patch_pars['xy'], patch_pars['width'], patch_pars['height'], edgecolor=None,
- facecolor='peachpuff',
- zorder=-8)
- ax.add_patch(rect)
- return ax
- def plot_sig_bars(ax, dat_obs, dat_rnd, times, tails_lis, colour_lis, p_threshold=0.05, plot_chance_lvl=0.5,
- if_smooth=False, variable_name='[name here]', time_window=(-0.2, 1.4)):
- """
- Plot significance bars for cluster-corrected permutation tests.
- Parameters:
- -----------
- ax : matplotlib axis
- Axis to plot on
- dat_obs : ndarray
- Observed data
- dat_rnd : ndarray
- Randomization data
- times : ndarray
- Time points
- tails_lis : list
- Tail directions for tests
- colour_lis : list
- Colors for each condition
- p_threshold : float
- P-value threshold for significance
- plot_chance_lvl : float
- Y-position for significance bars
- if_smooth : bool
- Whether to smooth data
- variable_name : str
- Name for printing results
- time_window : tuple or None
- (start_time, end_time) to restrict analysis to specific time window.
- If None, uses entire time range.
- """
- # Restrict to time window if specified
- if time_window is not None:
- start_time, end_time = time_window
- # Find indices corresponding to time window
- time_mask = (times >= start_time) & (times <= end_time)
- time_indices = np.where(time_mask)[0]
- # Subset data and times
- dat_obs_subset = dat_obs[:, time_indices]
- dat_rnd_subset = dat_rnd[:, :, time_indices]
- times_subset = times[time_indices]
- print(f'Restricting analysis to time window: {start_time:.3f}s to {end_time:.3f}s')
- else:
- dat_obs_subset = dat_obs
- dat_rnd_subset = dat_rnd
- times_subset = times
- # Run permutation test on subset
- clt_times, clt_labels = compute_perm_stats(dat_obs_subset,
- dat_rnd_subset,
- tails=tails_lis,
- if_smooth=if_smooth
- )
- print('Cluster perm. test: ' + variable_name)
- for p in range(dat_obs_subset.shape[0]):
- slices = get_slices(clt_times[p], clt_labels[p], times_subset)
- for s_i in range(slices.shape[0]):
- if slices[s_i, 0] <= p_threshold:
- start_time_sig = slices[s_i, 1]
- end_time_sig = slices[s_i, 2]
- p_value = slices[s_i, 0]
- # Print time window for each significant cluster
- print(
- 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})")
- if p == 2:
- ax.hlines(xmin=start_time_sig, xmax=end_time_sig, colors=colour_lis[p],
- y=plot_chance_lvl - (p / 80), linestyles=(0, (1, 0.5)),
- linewidth=2)
- else:
- ax.hlines(xmin=start_time_sig, xmax=end_time_sig, colors=colour_lis[p],
- y=plot_chance_lvl - (p / 80),
- linewidth=2)
- return
- def plot_scatter(gs, fig, x, y, scale, yaxis_label, xaxis_label, offset_axis=20, overlay_model='contour', title=None,
- 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):
- from scipy.stats import gaussian_kde
- ax = fig.add_subplot(gs)
- ax.scatter(x, y, color=color, s=dot_size, zorder=-5,
- edgecolor=edgecolor)
- # move the left spine (y axis) to the right
- ax.spines['left'].set_position(('axes', 0.5))
- # move the bottom spine (x axis) up
- ax.spines['bottom'].set_position(('axes', 0.5))
- # turn off the right and top spines
- ax.spines['right'].set_visible(False)
- ax.spines['top'].set_visible(False)
- ax.set_ylim([-scale, scale])
- ax.set_xlim([-scale, scale])
- ax.set_yticks([-scale, scale])
- ax.set_xticks([-scale,
fun_lib.py at commit 48ada80, no license · at the source
Overview
- Centre for Neural Circuits and Behaviour, Department of Physiology, Anatomy and Genetics, University of Oxford,Oxford, UK
- Department of Experimental Psychology, University of Oxford,Oxford, UK
- Department of Engineering, University of Cambridge,Cambridge, UK
- Coherence Neuro Global,San Francisco, CA USA
- MRC Cognition and Brain Sciences Unit, University of Cambridge,Cambridge, UK
- Centre for Neurotechnology, Neuromodulation and Neurotherapeutics, University of Nottingham,Nottingham, UK
- Department of Psychiatry, University of Oxford,Oxford, UK
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
48ada8054940f6a7ac26e8e83d150357a9f249d2, 27 July 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
13 files
- figure_1.py, Python, 93 lines
- figure_2.py, Python, 342 lines
- figure_3.py, Python, 339 lines
- figure_4.py, Python, 500 lines
- fun_lib.py, Python, 4,746 lines, 5 matches
- supp_fig_1.py, Python, 453 lines
- supp_fig_2.py, Python, 431 lines
- supp_fig_3.py, Python, 374 lines, 1 match
- supp_fig_4.py, Python, 261 lines
- supp_fig_5.py, Python, 99 lines
- supp_fig_6.py, Python, 255 lines
- supp_fig_7.py, Python, 293 lines
- README.md, Text, 285 lines
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://
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
- doi:10.5061/
dryad.c2fqz61kb , at Dryad; found in DataCite - doi:10.6080/
k0zw1hvd , at the source; found in the references
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://
BibTeX
@article{wojcik2026learn
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/
url = {https://
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/
VL - 29
IS - 8
SP - 1966
EP - 1975
SN - 1097-6256
PB - Nature Portfolio
DO - 10.1038/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1038/
"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":
"volume": "29",
"issue": "8",
"page": "1966-1975",
"DOI": "10.1038/
"PMID": "42350815",
"PMCID": "PMC13433252",
"ISSN": "1097-6256",
"publisher": "Nature Portfolio",
"URL": "https://
"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 communicationsIn 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 communicationsIn 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 biologyIn 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 neuroscienceIn 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: iScienceIn 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: eLifeIn 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 communicationsIn 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 communicationsIn 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 communicationsIn 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 oneIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 12 scripts, and 6 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:64684de214f00138…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
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.
