OSCR

Transformations of the spatial activity manifold convey aversive information in CA3.

A correction to this paper has been published: the notice, 42766758, from Europe PMC.

Code ↔ Paper

10 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 10 matches
  1. [1] § Materials and Methods › Position Decoding. ↔ aversive_scripts_2025/processing_functions.py, lines 848–910 · score 0.71 · XGBoost, cross validation, SVR, circle, angle, Wiener
  2. [2] § Results › Common Spatial Manifolds Exist Across Sessions. ↔ aversive_scripts_2025/simulation_controls.py, lines 766–809 · score 0.59 · periodic spiral, Noisy, incompatible, double, twisted, noise
  3. [3] § Materials and Methods › mCCA. ↔ aversive_scripts_2025/mCCA.py, lines 386–444 · score 0.58 · canonical space, cross predictions, position prediction, unaligned, alignment
  4. [4] § Materials and Methods › mCCA. ↔ aversive_scripts_2025/simulation_controls.py, lines 494–554 · score 0.57 · cross predictions, mCCA, position prediction, warped, quantified, unaligned
  5. [5] § Materials and Methods › mCCA. › TCA. ↔ aversive_scripts_2025/main_figures.py, lines 2887–3029 · score 0.56 · ncp hals, dimensionality reduction, tensor, TCA, Components, bins
  6. [6] § Results › Aversive Information Is Evenly Distributed Across Dimensions and Cell Types. ↔ aversive_scripts_2025/main_figures.py, lines 49–101 · score 0.54 · decoding weight, TCA dimensions, position prediction, quantified, variances, PCA
  7. [7] § Results › Affective Information Is Encoded Within the Spatial Activity Manifold. ↔ aversive_scripts_2025/main_figures.py, lines 2887–3029 · score 0.54 · TCA repetitions, TCA factor, trial shuffle, tensor, LDA, periodic
  8. [8] § Materials and Methods › Air Puff Decoding. ↔ aversive_scripts_2025/main_figures.py, lines 4054–4134 · score 0.54 · way ANOVA, F1 scores, variable, class, AP, shuffle
  9. [9] § Results › Affective Information Is Encoded Within the Spatial Activity Manifold. ↔ aversive_scripts_2025/project_parameters.py, lines 170–200 · score 0.53 · LDA decoding, TCA factor, trial shuffle, repetitions, belt, periodic
  10. [10] § Materials and Methods › mCCA. › TCA. ↔ aversive_scripts_2025/APdecoding_funs.py, lines 30–65 · score 0.53 · ncp hals, TCA dimensions, tensor, Components, bins

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

  1. # -*- coding: utf-8 -*-
  2. """
  3. Created on Thu Oct 5 03:08:39 2023
  4. @author: Albert
  5. Scripts to create the figures in "Transformations of the spatial activity manifold convey aversive information in CA3" (2025)
  6. """
  7. import time
  8. import scipy
  9. import scipy.io
  10. import numpy as np
  11. import matplotlib as mpl
  12. import matplotlib.pyplot as plt
  13. from matplotlib.ticker import MaxNLocator
  14. from matplotlib.patches import Rectangle
  15. import h5py
  16. import pandas as pd
  17. from sklearn import decomposition
  18. # from decoders import WienerFilterRegression, WienerCascadeRegression, KalmanFilterRegression, SVRegression, NaiveBayesRegression,XGBoostRegression
  19. # from sklearn.model_selection import KFold, StratifiedKFold, GroupKFold, LeaveOneGroupOut
  20. from sklearn.discriminant_analysis import LinearDiscriminantAnalysis
  21. import statsmodels.api as sm
  22. from statsmodels.formula.api import ols
  23. #Global project variables
  24. import project_parameters as pparam
  25. from project_parameters import (FAT_CLUSTER_PATH, OUTPUT_PATH, SESSION_NAMES, SESSION_REPEATS, MOUSE_TYPE_LABELS)
  26. #Project scripts
  27. import processing_functions as pf
  28. import mCCA as mCCA_funs
  29. import APdecoding_funs as APfuns
  30. plt.rcParams['font.family'] = 'Arial'
  31. def main():
  32. # save_pca_data(np.arange(8), np.arange(9)) #Uncomment this once to pre-compute PCA data.
  33. ''' Figure 1 plots '''
  34. # count_neurons()
  35. # plot_fig1_C_firing_rates()
  36. # plot_fig1_D_pcas()
  37. # plot_fig1_E_position_prediction_example()
  38. # plot_fig1_F_prediction_across_sessions()
  39. ''' Figure 1 SI plots '''
  40. # plot_fig1SI_A_traces()
  41. # plot_fig1SI_B_C_variance_explained()
  42. # plot_fig1SI_D_prediction_vs_dimensionality()
  43. ''' Figure 2 SI plots '''
  44. # fig2_CCA_example()
  45. # fig2_CCA_aligned_pcas(shuffle=False)
  46. # fig2_D_E_F_CCA_quantification()
  47. ''' Figure 2 SI plots '''
  48. # fig2_CCA_aligned_pcas(shuffle=True)
  49. ''' Figure 3 plots '''
  50. # fig3_A_and_fig3SI_A_TCA_factors()
  51. # fig3_B_LDA_on_TCA()
  52. # fig3_C_D_f1_plots_and_fig3SI_B_C_accuracy_plots()
  53. # fig3_F_and_3SI_A_B_belt_restriction_plots()
  54. ''' Figure 3 SI plots '''
  55. # fig3SI_C()
  56. # fig3SI_D()
  57. # fig3SI_F1_by_CCA_and_TCA_dimension_and_mouse()
  58. ''' Figure 4 plots '''
  59. # fig4_A_session_comparisons()
  60. fig4_C_D_E_distance_measures()
  61. ''' Figure 4 SI plots '''
  62. # fig4SI_G_H_I_distance_measures()
  63. ''' Figure 5 plots '''
  64. # fig5_A_F1_by_TCA_dimension()
  65. # fig5_B_C_D_and_fig5SI_A_B_C_D_E_decoding_weight_plots()
  66. ''' Figure 5 SI plots '''
  67. # fig5SI_F_decoding_weight_control()
  68. # <>
  69. ######## Functions related to saving and retrieving data ########
  70. def compare_parameter_dictionaries(dict1, dict2):
  71. ''' Checks if two dictionaries are equal, return "False" if they are not '''
  72. are_dictionaries_equal = False
  73. try:
  74. np.testing.assert_equal(dict1, dict2)
  75. are_dictionaries_equal = True
  76. except AssertionError:
  77. are_dictionaries_equal = False
  78. return are_dictionaries_equal
  79. def get_analysis_dict_filename(analysis_name):
  80. ''' name is string (no .npy or previous path)
  81. if "name" corresponds to an analysis name, give the default name from pparams
  82. otherwise treat it as a custom name
  83. '''
  84. if analysis_name in pparam.ANALYSIS_NAME_LIST:
  85. filename = pparam.default_param_dicts_names[analysis_name]
  86. else:
  87. filename = analysis_name
  88. analysis_dict_filename = OUTPUT_PATH + filename + ".npy"
  89. return analysis_dict_filename
  90. def load_previous_analysis(analysis_name):
  91. '''
  92. Parameters
  93. ----------
  94. analysis_name: 'preprocessing', 'alignment', 'APdecoding'
  95. if neither of these, it must be the name to a different analysis dict filename (assumed to be in OUTPUT_PATH, without the .npy)
  96. Returns
  97. -------
  98. Previous analysis dictionary
  99. '''
  100. analysis_dict_filename = get_analysis_dict_filename(analysis_name)
  101. try:
  102. analysis_dict = np.load(analysis_dict_filename, allow_pickle=True)[()]
  103. except FileNotFoundError:
  104. analysis_dict = None
  105. return analysis_dict
  106. def process_input_parameter_dict(param_dict_new, analysis_name, custom_filename = None):
  107. ''' Checks incoming parameter dict:
  108. - Adds missing parameters from default dictionary
  109. - Compares with old dictionary and returns "are_different" = False if they are different
  110. analysis_name: 'preprocessing', 'alignment', 'APdecoding'
  111. (from pparam.ANALYSIS_NAME_LIST[0])
  112. '''
  113. if custom_filename is None:
  114. analysis_filename = analysis_name
  115. else:
  116. analysis_filename = custom_filename
  117. analysis_dict_old = load_previous_analysis(analysis_filename)
  118. param_dict_default = pparam.default_param_dicts[analysis_name]
  119. #Substitute default parameters if missing
  120. for k,v in param_dict_default.items():
  121. if k not in param_dict_new:
  122. param_dict_new[k] = v
  123. #If no previous analysis was found, indicate that it must be repeated
  124. if analysis_dict_old is None:
  125. print('No previous analysis found, starting from scratch')
  126. return False, param_dict_new, {}
  127. #If previous analysis existss, check if it has the same parameters
  128. try:
  129. param_dict_old = analysis_dict_old['param_dict']
  130. except KeyError: #There's not parameter dict saved in the dictionary
  131. print('Warning: no previous parameter dict found in the analysis dictionary for "%s"'%analysis_name)
  132. have_parameters_been_used = False
  133. else:
  134. have_parameters_been_used = compare_parameter_dictionaries(param_dict_old, param_dict_new)
  135. return have_parameters_been_used, param_dict_new, analysis_dict_old
  136. def save_figure(fig, fig_name):
  137. print(pparam.FIGURES_PATH + fig_name + ".svg")
  138. fig.savefig(pparam.FIGURES_PATH + fig_name + ".svg")
  139. def save_pca_data(mouse_list, session_list):
  140. ''' Convenience function to pre-compute PCA information '''
  141. ## Get PCA weights ##
  142. components_dict = {} #(mnum, snum) = (pca_dims, num_neurons)
  143. variance_explained_dict = {} #(mnum, snum) = (pca_dims, num_neurons)
  144. PCA_dict = {} #(mnum, snum) = (position, pca_data) with size(position) = timepoints and size(pca_data) = (pca_dims X timepoints)
  145. input_data_dict = {} #(mnum, snum) = (position, preprocessed_spikes) with size(position) = timepoints and size(preprocessed_spikes) = (num_neurons X timepoints)
  146. time_bin_size = 1 # Number of elements to average over, each dt should be ~65ms
  147. distance_bin_size = 1 # mm, track is 1500mm, data is in mm
  148. gaussian_size = 25 # Why not
  149. data_used = 'amplitudes'
  150. running = True
  151. eliminate_v_zeros = False
  152. pos_max = 1500
  153. for mnum in mouse_list:
  154. for snum in session_list:
  155. #Get data
  156. data_dict = pf.read_and_preprocess_data(fat_cluster, mnum, snum, gaussian_size, time_bin_size, distance_bin_size,
  157. only_running=running, eliminate_v_zeros=eliminate_v_zeros, pos_max=pos_max)
  158. pca_input_data, position, times = pf.get_data_from_datadict(data_dict, data_used)
  159. # place_cell_bool = pf.load_place_cell_boolean(mnum, snum, criteria='dombeck').astype(bool)
  160. # place_cell_idxs = np.where(place_cell_bool)[0]
  161. # nplace_cell_idxs = np.where(np.invert(place_cell_bool))[0]
  162. # pca_input_data = pca_input_data[nplace_cell_idxs]
  163. # print(mnum, snum, pca_input_data.shape)
  164. # if pca_input_data.shape[0] < 3:
  165. # continue
  166. # print(mnum, snum, pca_input_data.shape)
  167. input_data_dict[mnum, snum] = (position, pca_input_data)
  168. num_neurons = pca_input_data.shape[0]
  169. #PCA
  170. pca = decomposition.PCA(n_components=num_neurons)
  171. pca.fit(pca_input_data.T)
  172. #PCA over time
  173. pca_data = pf.project_spikes_PCA(pca_input_data, pca_instance = pca, num_components = num_neurons)
  174. if eliminate_v_zeros == True:
  175. position, pca_data, _ = pf.compute_velocity_and_eliminate_zeros(position, pca_data, pos_max = pos_max)
  176. print(pca_data.shape)
  177. PCA_dict[mnum, snum] = (position, pca_data)
  178. #Components
  179. components = pca.components_
  180. components_dict[mnum, snum] = components
  181. #Variance explained
  182. variance_explained = pca.explained_variance_ratio_
  183. variance_explained_dict[mnum, snum] = variance_explained
  184. np.save(OUTPUT_PATH + "input_data_dict.npy", input_data_dict)
  185. np.save(OUTPUT_PATH + "pca_components_dict.npy", components_dict)
  186. np.save(OUTPUT_PATH + "variance_explained_dict.npy", variance_explained_dict)
  187. np.save(OUTPUT_PATH + "PCA_dict.npy", PCA_dict)
  188. def perform_pca_on_multiple_mice_param_dict(pca_param_dict):
  189. have_parameters_been_used, pca_param_dict, pca_analysis_dict_previous = process_input_parameter_dict(pca_param_dict, pparam.ANALYSIS_NAME_LIST[0])
  190. print("Are pca parameters repeated?", have_parameters_been_used)
  191. if have_parameters_been_used == True:
  192. return pca_analysis_dict_previous
  193. mouse_list = pca_param_dict['mouse_list']
  194. session_list = pca_param_dict['session_list']
  195. #Preprocessing
  196. time_bin_size = pca_param_dict['time_bin_size'] # Number of elements to average over, each dt should be ~65ms
  197. distance_bin_size = pca_param_dict['distance_bin_size'] # mm, track is 1500mm, data is in mm
  198. gaussian_size = pca_param_dict['gaussian_size'] # Why not
  199. data_used = pca_param_dict['data_used']
  200. running = pca_param_dict['running']
  201. eliminate_v_zeros = pca_param_dict['eliminate_v_zeros']
  202. #PCA
  203. num_components = pca_param_dict['num_components']
  204. ## Store PCA weights ##
  205. components_dict = {} #(mnum, snum) = (pca_dims, num_neurons)
  206. variance_explained_dict = {} #(mnum, snum) = (pca_dims, num_neurons)
  207. PCA_dict = {} #(mnum, snum) = (position, pca_data) with size(position) = timepoints and size(pca_data) = (pca_dims X timepoints)
  208. input_data_dict = {} #(mnum, snum) = (position, preprocessed_spikes) with size(position) = timepoints and size(preprocessed_spikes) = (num_neurons X timepoints)
  209. for mnum in mouse_list:
  210. for snum in session_list:
  211. #Get data
  212. data_dict = pf.read_and_preprocess_data(fat_cluster, mnum, snum, gaussian_size, time_bin_size, distance_bin_size,
  213. only_running=running, eliminate_v_zeros=eliminate_v_zeros, pos_max=pparam.MAX_POS)
  214. pca_input_data, position, times = pf.get_data_from_datadict(data_dict, data_used)
  215. input_data_dict[mnum, snum] = (position, pca_input_data)
  216. num_neurons = pca_input_data.shape[0]
  217. #PCA
  218. pca = decomposition.PCA(n_components=num_neurons)
  219. pca.fit(pca_input_data.T)
  220. #PCA over time
  221. pca_data = pf.project_spikes_PCA(pca_input_data, pca_instance = pca, num_components = num_components)
  222. if eliminate_v_zeros == True:
  223. position, pca_data, _ = pf.compute_velocity_and_eliminate_zeros(position, pca_data, pos_max = pparam.MAX_POS)
  224. PCA_dict[mnum, snum] = (position, pca_data)
  225. #Components
  226. components = pca.components_
  227. components_dict[mnum, snum] = components
  228. #Variance explained
  229. variance_explained = pca.explained_variance_ratio_
  230. variance_explained_dict[mnum, snum] = variance_explained
  231. PCA_analysis_dict = {
  232. 'param_dict':pca_param_dict,
  233. 'mouse_list':mouse_list,
  234. 'session_list':session_list,
  235. 'input_data_dict':input_data_dict,
  236. 'PCA_dict':PCA_dict,
  237. 'components_dict':components_dict,
  238. 'variance_explained_dict':variance_explained_dict
  239. }
  240. np.save(OUTPUT_PATH + pparam.preprocessing_dict_name + ".npy", PCA_analysis_dict)
  241. return PCA_analysis_dict
  242. def filter_repeated_sessions(mnum, session_list):
  243. ''' Given a mouse number and a list of session numbers, return the list with non-repeated sessions.
  244. If two sessions are repeats, keep only the first instance of the copy.
  245. Uses "SESSION_REPEATS" from parameter file
  246. '''
  247. session_list = list(np.copy(session_list))
  248. if mnum in SESSION_REPEATS:
  249. for repeated_pair in SESSION_REPEATS[mnum]:
  250. if repeated_pair[0] in session_list and repeated_pair[1] in session_list:
  251. session_list.remove(repeated_pair[1])
  252. return session_list
  253. def pca_and_pos_from_dict_to_lists(pca_dict, mnum, session_list):
  254. ''' Given "pca_dict" from "save_pca_data", return a list of positions and pcas for each session'''
  255. pos_list = []
  256. pca_list = []
  257. for snum in session_list:
  258. pos, pca = pca_dict[mnum, snum]
  259. # pca, _, _ = pf.normalize_data(pca, axis=1)
  260. pos_list.append(pos)
  261. pca_list.append(pca)
  262. return pos_list, pca_list
  263. #~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ Fig 1 ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
  264. def count_neurons():
  265. ########## PARAMETERS ###########
  266. mlist = list(range(8))
  267. # mlist = [2,3,5,6]
  268. # slist = list(range(1,9))
  269. slist = list(range(9))
  270. gaussian_size = 25
  271. time_bin_size = 1
  272. distance_bin_size = 1
  273. running = True
  274. eliminate_v_zeros = False
  275. data_used = 'amplitudes'
  276. num_neurons_list = []
  277. num_trials_list = []
  278. for midx, mnum in enumerate(mlist):
  279. for sidx, snum in enumerate(slist):
  280. data_dict = pf.read_and_preprocess_data(fat_cluster, mnum, snum, gaussian_size, time_bin_size, distance_bin_size,
  281. only_running=running, eliminate_v_zeros=eliminate_v_zeros, pos_max=pparam.MAX_POS)
  282. pca_input_data, position, times = pf.get_data_from_datadict(data_dict, data_used)
  283. num_neurons_list.append(pca_input_data.shape[0])
  284. num_trials_list.append(data_dict['num_trials'])
  285. print(num_neurons_list)
  286. print("Average neurons per session: %.1f" %np.average(num_neurons_list))
  287. print(num_trials_list)
  288. print("Average number of trials per session: %.1f" %np.average(num_trials_list))
  289. def plot_fig1SI_A_traces(fig_num = None):
  290. '''
  291. Plots example Ca2+ traces
  292. '''
  293. if fig_num is None:
  294. fig_num = plt.gcf().number + 1
  295. mnum = 6
  296. snum = 3
  297. time_bin_size = 1 # Number of elements to average over, each dt should be ~65ms
  298. distance_bin_size = 1 # mm, track is 1500mm, data is in mm
  299. gaussian_size = 0 # Why not
  300. running = False
  301. eliminate_v_zeros = False
  302. ## Plot params ##
  303. fs = 15
  304. #Get raw data
  305. data_dict = pf.read_and_preprocess_data(fat_cluster, mnum, snum, gaussian_size, time_bin_size, distance_bin_size,
  306. only_running=running, eliminate_v_zeros=eliminate_v_zeros, pos_max=pparam.MAX_POS)
  307. amplitudes = data_dict['amplitudes_binned_normalized']
  308. data = amplitudes
  309. fig = plt.figure(fig_num, figsize=(7,7)); fig_num += 1
  310. ax = plt.gca()
  311. neurons_to_plot = [0,2,3,4,6,7,8,9,10,11,15,20]
  312. for nidx, n in enumerate(neurons_to_plot):
  313. trace = data[n]
  314. trace = (trace - np.min(trace))/(np.max(trace) - np.min(trace))
  315. trace = 0.8*trace
  316. yy = trace + nidx
  317. xx = np.arange(len(yy))/(60*10)
  318. ax.plot(xx, yy, color='black')
  319. ax.set_xlabel('Time (min)', fontsize=fs+4)
  320. ax.set_yticks(np.arange(len(neurons_to_plot)), neurons_to_plot)
  321. ax.tick_params(axis='x', labelsize=fs)
  322. ax.tick_params(axis='y', labelsize=fs)
  323. ax.spines[['right', 'top']].set_visible(False)
  324. for axis in ['top','bottom','left','right']:
  325. ax.spines[axis].set_linewidth(3)
  326. ax.set_ylabel('Axon #', fontsize=fs+4)
  327. ax.set_title('$\Delta$ F / F', fontsize=fs+4)
  328. fig_name = "fig1SI_traces"
  329. save_figure(fig,fig_name)
  330. return fig_num
  331. def plot_fig1_C_firing_rates(fig_num = None):
  332. '''
  333. Plots example cell rates
  334. '''
  335. if fig_num is None:
  336. fig_num = plt.gcf().number + 1
  337. mnum = 6
  338. snum = 2
  339. time_bin_size = 1 # Number of elements to average over, each dt should be ~65ms
  340. distance_bin_size = 1 # mm, track is 1500mm, data is in mm
  341. gaussian_size = 15 # Why not
  342. running = False
  343. eliminate_v_zeros = False
  344. ## Plot params ##
  345. fs = 20
  346. bins = 30
  347. spine_width = 3
  348. tick_length = 10
  349. #Get raw data
  350. data_dict = pf.read_and_preprocess_data(fat_cluster, mnum, snum, gaussian_size, time_bin_size, distance_bin_size,
  351. only_running=running, eliminate_v_zeros=eliminate_v_zeros, pos_max=pparam.MAX_POS)
  352. amplitudes = data_dict['amplitudes_binned_normalized']
  353. amplitudes, _, _ = pf.normalize_data(amplitudes)
  354. pos = data_dict['distance']
  355. pos, data, _ = pf.warping(pos, amplitudes, bins, max_pos = pparam.MAX_POS, warp_sampling_type = 'interpolation', warp_based_on = 'position', return_flattened = False)
  356. num_neurons, num_bins, num_trials = data.shape
  357. neurons_to_plot = [0,2,9]
  358. colors = ['royalblue', 'darkblue', 'mediumslateblue']
  359. ### Non-overlapped ###
  360. fig, axs = plt.subplots(len(neurons_to_plot), 1, num=fig_num, figsize=(7,6)); fig_num += 1
  361. # axs = plt.gca()
  362. for nidx, n in enumerate(neurons_to_plot):
  363. ax = axs.ravel()[nidx]
  364. rate = data[n]
  365. avg = np.mean(rate, axis=1)
  366. std = np.std(rate, axis=1)/np.sqrt(num_trials)
  367. xx = pos[:,0]
  368. ax.plot(xx, avg, color=colors[nidx], linewidth=5, label='Neuron %d'%n)
  369. ax.fill_between(xx, avg-std, avg+std, color='royalblue', alpha=0.3)
  370. if nidx == 0:
  371. neuron_list_string = [str(n) for n in neurons_to_plot]
  372. neuron_list_string = ','.join(neuron_list_string)
  373. ax.set_title('Animal %d / %s / Neurons %s'%(mnum, pparam.SESSION_NAMES[snum], neuron_list_string), fontsize=fs+4, pad=20)
  374. #Set y label on middle neuron
  375. if nidx == len(neurons_to_plot)//2:
  376. ax.set_ylabel('$\Delta F / F$', fontsize=fs+8)
  377. if nidx == len(neurons_to_plot)-1:
  378. #X axis for bottom plot
  379. ax.set_xlabel('Position (mm)', fontsize=fs+4)
  380. ax.spines.bottom.set_position(('outward', 10))
  381. ax.tick_params(axis='x', labelsize=fs+4)
  382. else:
  383. #Eliminate X axis for the rest
  384. ax.get_xaxis().set_visible(False)
  385. ax.spines[['bottom']].set_visible(False)
  386. #Set the Y axis ticks
  387. ymin, ymax = np.min(avg), np.max(avg)
  388. ymin, ymax = ax.get_ylim()
  389. ax.set_ylim([ymin, ymax])
  390. ax.set_yticks([np.around(ymin, 1), np.around(ymax, 1)])
  391. ax.xaxis.set_tick_params(width=spine_width, length=tick_length)
  392. ax.yaxis.set_tick_params(width=spine_width, length=tick_length)
  393. ax.tick_params(axis='y', labelsize=fs+4)
  394. #Align all the X axis and move them apart
  395. ax.set_xlim([0,pparam.MAX_POS])
  396. ax.spines.left.set_position(('outward', 10))
  397. ax.spines[['right', 'top']].set_visible(False)
  398. #Spine width
  399. for axis in ['top','bottom','left','right']:
  400. ax.spines[axis].set_linewidth(spine_width)
  401. fig.tight_layout()
  402. fig_name = "fig1_firing_rates"
  403. save_figure(fig,fig_name)
  404. return fig_num
  405. def plot_fig1_D_pcas(fig_num = None):
  406. ''' Plot PCA examples for figure 1 '''
  407. if fig_num is None:
  408. fig_num = plt.gcf().number + 1
  409. mlist = [6,6,6]
  410. slist = [2,3,8]
  411. angle_list = [75,50,-90]
  412. angle_azim_list = [-90, None, -100]
  413. rows = 2 #One row per session
  414. cols = len(slist) #raw data and trial averaged
  415. fig, axs = plt.subplots(rows, cols, subplot_kw={"projection": "3d"}, figsize=(9,7))
  416. fig_num = fig_num+1
  417. time_bin_size = 1 # Number of elements to average over, each dt should be ~65ms
  418. distance_bin_size = 1 # mm, track is 1500mm, data is in mm
  419. gaussian_size = 25 # Why not
  420. data_used = 'amplitudes'
  421. running = True
  422. eliminate_v_zeros = True
  423. ## Plot params ##
  424. fs = 15
  425. pca_plot_bin_size = 30
  426. max_pos = pparam.MAX_POS
  427. for idx, snum in enumerate(slist):
  428. mnum = mlist[idx]
  429. angle = angle_list[idx]
  430. angle_azim = angle_azim_list[idx]
  431. #Get PCA
  432. data_dict = pf.read_and_preprocess_data(fat_cluster, mnum, snum, gaussian_size, time_bin_size, distance_bin_size,
  433. only_running=running, eliminate_v_zeros=eliminate_v_zeros, pos_max=pparam.MAX_POS)
  434. pca_input_data, position, times = pf.get_data_from_datadict(data_dict, data_used)
  435. pca = pf.project_spikes_PCA(pca_input_data, num_components = 3)
  436. position, pca, _ = pf.warping(position, pca, 200, max_pos=pparam.MAX_POS,
  437. warp_sampling_type = 'interpolation', warp_based_on = 'time', return_flattened=True)
  438. # print(pca.shape)
  439. # fig = plt.figure()
  440. ax = axs[0,idx]
  441. pf.plot_pca_with_position(pca, position, ax=ax, max_pos = pparam.MAX_POS, cmap_name = pparam.PCA_CMAP, fs=15, scatter=True, cbar=False, cbar_label='Position (mm)',
  442. alpha=1, angle=angle, angle_azim=angle_azim, axis = 'off', show_axis_labels=False, axis_label=None,
  443. ms = 10, lw=3)
  444. # fig_num += 1
  445. ax = axs[1, idx]
  446. position_unique, pca_average, pca_std = pf.compute_average_data_by_position(pca, position, position_bin_size=pca_plot_bin_size, max_pos=max_pos)
  447. pf.plot_pca_with_position(pca_average, position_unique, ax=ax, max_pos = pparam.MAX_POS, cmap_name = pparam.PCA_CMAP, fs=15, cbar=False, cbar_label='Position (mm)',
  448. alpha=1, angle=angle, angle_azim=angle_azim, axis = 'off', show_axis_labels=False, axis_label=None,
  449. scatter=False, ms = 250, lw=6)
  450. slist_names = [pparam.SESSION_NAMES[snum] for snum in slist]
  451. fig.suptitle('M%d %s, M%d %s, M%d %s' %(mlist[0], slist_names[0], mlist[1], slist_names[1], mlist[2], slist_names[2]),
  452. fontsize=30)
  453. fig.subplots_adjust(wspace=0, hspace=0)
  454. fig.tight_layout()
  455. fig_name = "fig1_pcas"
  456. save_figure(fig,fig_name)
  457. plt.figure(fig_num); fig_num += 1
  458. fig = plt.gcf()
  459. cbar = pf.add_distance_cbar(fig, pparam.PCA_CMAP, vmin = 0, vmax = pparam.MAX_POS, fs=fs,
  460. cbar_label = '',
  461. cbar_kwargs = {'fraction':0.055, 'pad':0.04, 'aspect':10})
  462. cbar.ax.tick_params(axis='y', labelsize=25)
  463. # #Put everything on left
  464. cbar.ax.set_ylabel('Position (mm)', fontsize=30)
  465. cbar.ax.yaxis.set_label_position('left')
  466. cbar.ax.yaxis.set_ticks_position('left')
  467. fig.tight_layout()
  468. ax = plt.gca()
  469. ax.set_axis_off()
  470. fig_name = "fig1_pca_colorbar"
  471. save_figure(fig,fig_name)
  472. return fig_num
  473. def plot_fig1_E_position_prediction_example(fig_num = None):
  474. if fig_num is None:
  475. fig_num = plt.gcf().number + 1
  476. ########## PARAMETERS ###########
  477. mnum = 6
  478. snum = 3
  479. time_bin_size = 1 # Number of elements to average over, each dt should be ~65ms
  480. distance_bin_size = 1 # mm, track is 1500mm, data is in mm
  481. gaussian_size = 25 # Why not
  482. data_used = 'amplitudes'
  483. running = True
  484. eliminate_v_zeros = False
  485. ## Predictor parameters ##
  486. cv_folds = 5
  487. predictor_name = 'Wiener'
  488. error_type = 'sse'
  489. ## Plot params ##
  490. fs = 20
  491. markersize=50
  492. #Get PCA
  493. data_dict = pf.read_and_preprocess_data(fat_cluster, mnum, snum, gaussian_size, time_bin_size, distance_bin_size,
  494. only_running=running, eliminate_v_zeros=eliminate_v_zeros, pos_max=pparam.MAX_POS)
  495. pca_input_data, position, times = pf.get_data_from_datadict(data_dict, data_used)
  496. pca = pf.project_spikes_PCA(pca_input_data, num_components = 3)
  497. # position, pca, _ = pf.warping(position, pca, 200, max_pos=pparam.MAX_POS,
  498. # warp_sampling_type = 'interpolation', warp_based_on = 'time', return_flattened=True)
  499. pos_pred, error, predictor = pf.predict_position_CV(pca, position, n_splits=cv_folds, shuffle=False, periodic=True, pmin=0, pmax=pparam.MAX_POS,
  500. predictor_name=predictor_name, predictor_default=None, return_error=error_type)
  501. ###### OPTIONAL: ADD SHUFFLE ######
  502. np.random.seed(100)
  503. timepoints = len(position)
  504. idxs = list(np.arange(timepoints))
  505. shift = np.random.randint(-timepoints, timepoints)
  506. idxs_shifted = idxs[shift:] + idxs[:shift]
  507. idxs_shifted = np.random.choice(idxs_shifted, size=len(idxs_shifted), replace=False)
  508. position_shuffled = position[idxs_shifted]
  509. pca_shuffled = np.copy(pca)
  510. pos_pred_shuffle, error_shuffle, _ = pf.predict_position_CV(pca_shuffled, position_shuffled, n_splits=cv_folds, shuffle=False, periodic=True, pmin=0, pmax=pparam.MAX_POS,
  511. predictor_name=predictor_name, predictor_default=None, return_error=error_type)
  512. fig = plt.figure(fig_num, figsize=(7,4)); fig_num += 1
  513. ax = plt.gca()
  514. timesteps = np.arange(len(position))
  515. for plot_idx, p in enumerate([position, pos_pred]):
  516. label = pparam.PREDICTION_LABELS[plot_idx]
  517. color = pparam.PREDICTION_COLORS[label]
  518. ax.scatter(timesteps, p, s=markersize, color=color, label=label)
  519. # ax.scatter(timesteps, pos_pred_shuffle, s=markersize, color='red', label='shuffle')
  520. ax.set_xlabel('Timestep', fontsize=fs+4)
  521. # ax.set_yticks(np.arange(len(neurons_to_plot)), neurons_to_plot)
  522. ax.tick_params(axis='x', labelsize=fs)
  523. ax.tick_params(axis='y', labelsize=fs)
  524. ax.spines[['right', 'top']].set_visible(False)
  525. for axis in ['top','bottom','left','right']:
  526. ax.spines[axis].set_linewidth(3)
  527. ax.set_ylabel('Position (mm)', fontsize=fs+4)
  528. ax.set_title('Animal %d, session %s, SSE=%.1f (cm)'%(mnum, pparam.SESSION_NAMES[snum], error),
  529. pad=20, fontsize=fs+4)
  530. ax.legend(fontsize=fs-4, loc='upper right', frameon=False)
  531. fig.tight_layout()
  532. fig_name = "fig1_prediction_example"
  533. save_figure(fig,fig_name)
  534. return fig_num
  535. def plot_fig1_F_prediction_across_sessions(fig_num = None):
  536. if fig_num is None:
  537. fig_num = plt.gcf().number + 1
  538. ########## PARAMETERS ###########
  539. mlist = list(range(8))
  540. # mlist = [2,3,5,6]
  541. # slist = list(range(1,9))
  542. slist = list(range(9))
  543. # slist = [2,3,6]
  544. preprocessing_param_dict = {
  545. #session params
  546. 'mouse_list':mlist,
  547. 'session_list':slist,
  548. #Preprocessing parameters
  549. 'time_bin_size':1,
  550. 'distance_bin_size':1,
  551. 'gaussian_size':25,
  552. 'data_used':'amplitudes',
  553. 'running':True,
  554. 'eliminate_v_zeros':True,
  555. 'num_components':'.9',
  556. }
  557. ## Predictor parameters ##
  558. cv_folds = 5
  559. predictor_name = 'Wiener'
  560. error_type = 'sse'
  561. shuffle_reps = 10
  562. ## Plot params ##
  563. fs = 20
  564. num_mice = len(mlist)
  565. num_sessions = len(slist)
  566. error_array = np.ones((num_mice, num_sessions))*-1
  567. error_array_shuffle = np.ones((num_mice, num_sessions))*-1
  568. PCA_analysis_dict = perform_pca_on_multiple_mice_param_dict(preprocessing_param_dict)
  569. PCA_dict = PCA_analysis_dict['PCA_dict']
  570. for midx, mnum in enumerate(mlist):
  571. for sidx, snum in enumerate(slist):
  572. # #Get PCA
  573. position, pca = PCA_dict[mnum, snum]
  574. pos_pred, error, predictor = pf.predict_position_CV(pca, position, n_splits=cv_folds, shuffle=False, periodic=True, pmin=0, pmax=pparam.MAX_POS,
  575. predictor_name=predictor_name, predictor_default=None, return_error=error_type)
  576. error_array[midx, sidx] = error
  577. if shuffle_reps != 0:
  578. timepoints = len(position)
  579. idxs = list(range(timepoints))
  580. error_shuffle_list = []
  581. for rep in range(shuffle_reps):
  582. shift = np.random.randint(-timepoints, timepoints)
  583. idxs_shifted = idxs[shift:] + idxs[:shift]
  584. idxs_shifted = np.random.choice(idxs_shifted, size=len(idxs_shifted), replace=False)
  585. position_shuffled = position[idxs_shifted]
  586. pca_shuffled = np.copy(pca)
  587. # position_shuffled = np.copy(position)
  588. # pca_input_data_shifted = pca_input_data[:, idxs_shifted]
  589. pos_pred_shuffle, error_shuffle, _ = pf.predict_position_CV(pca_shuffled, position_shuffled, n_splits=cv_folds, shuffle=False, periodic=True, pmin=0, pmax=pparam.MAX_POS,
  590. predictor_name=predictor_name, predictor_default=None, return_error=error_type)
  591. error_shuffle_list.append(error_shuffle)
  592. error_array_shuffle[midx, sidx] = np.average(error_shuffle_list)
  593. fig = plt.figure(fig_num, figsize=(7,4)); fig_num += 1
  594. ax = plt.gca()
  595. mice_types = pparam.MOUSE_TYPE_LABELS
  596. for mtype_idx, mtype in enumerate(pparam.MOUSE_TYPE_LABELS):
  597. midx_list_type = [mlist.index(mnum) for mnum in mlist if mnum in pparam.MOUSE_TYPE_INDEXES[mice_types[mtype_idx]]]
  598. error_array_type = error_array[midx_list_type]
  599. num_mice_type, num_sessions_type = error_array_type.shape
  600. _, xx_scatter = np.mgrid[:num_mice_type, :num_sessions_type]
  601. xx_scatter = xx_scatter.ravel()
  602. error_scatter = error_array_type.ravel()
  603. color = pparam.MOUSE_TYPE_COLORS[mtype_idx]
  604. ax.scatter(xx_scatter, error_scatter, color=color)
  605. avg = np.average(error_array_type, axis=0)
  606. std = np.std(error_array_type, axis=0)/np.sqrt(error_array_type.shape[0])
  607. # print(error_array.shape)
  608. xx = np.arange(num_sessions)
  609. ax.plot(xx, avg, lw=3, color=color, label=mtype)
  610. ax.fill_between(xx, avg-std, avg+std, color=color, alpha=0.5)
  611. # #Significance testing (session by session)
  612. # for sidx, snum in enumerate(slist):
  613. # id_vals = [error_array[mnum, sidx] for mnum in mlist if mnum in pparam.MOUSE_TYPE_INDEXES[mice_types[0]]]
  614. # dd_vals = [error_array[mnum, sidx] for mnum in mlist if mnum in pparam.MOUSE_TYPE_INDEXES[mice_types[1]]]
  615. # tstat, pval = scipy.stats.ttest_ind(id_vals, dd_vals, equal_var=True, permutations=None, alternative='two-sided')
  616. # pval_label = pf.get_significance_label(np.abs(pval), [0.05], asterisk = True, ns=False)
  617. # print(pval, pval_label)
  618. # if pval_label != 'ns':
  619. # max_val = np.maximum(np.max(id_vals), np.max(dd_vals))
  620. # ax.text(sidx-0.1, max_val + 3, pval_label, fontsize=fs+10, style='normal', color='black')
  621. if shuffle_reps != 0:
  622. avg = np.average(error_array_shuffle, axis=0)
  623. std = np.std(error_array_shuffle, axis=0)/np.sqrt(error_array_shuffle.shape[0])
  624. xx = np.arange(num_sessions)
  625. ax.plot(xx, avg, lw=3, color=pparam.SHUFFLE_DEFAULT_COLOR, label='Shuffle')
  626. ax.fill_between(xx, avg-std, avg+std, color=pparam.SHUFFLE_DEFAULT_COLOR, alpha=0.5)
  627. #Significance testing (overall)
  628. # id_vals = error_array[]
  629. midx_list_type = [[mlist.index(mnum) for mnum in mlist if mnum in pparam.MOUSE_TYPE_INDEXES[mtype]] for mtype in pparam.MOUSE_TYPE_LABELS]
  630. id_vals = error_array[midx_list_type[0]].ravel()
  631. dd_vals = error_array[midx_list_type[1]].ravel()
  632. shuffle_vals = error_array_shuffle.ravel()
  633. # tstat, pval_id_shuffle = scipy.stats.ttest_ind(id_vals, shuffle_vals, equal_var=False, permutations=None, alternative='two-sided')
  634. # tstat, pval_dd_shuffle = scipy.stats.ttest_ind(dd_vals, shuffle_vals, equal_var=False, permutations=None, alternative='two-sided')
  635. # tstat, pval_id_dd = scipy.stats.ttest_ind(id_vals, dd_vals, equal_var=False, permutations=None, alternative='two-sided')
  636. tstat, pval_id_shuffle = scipy.stats.mannwhitneyu(id_vals, shuffle_vals, use_continuity=False, alternative='two-sided')
  637. tstat, pval_dd_shuffle = scipy.stats.mannwhitneyu(dd_vals, shuffle_vals, use_continuity=False, alternative='two-sided')
  638. tstat, pval_id_dd = scipy.stats.mannwhitneyu(id_vals, dd_vals, use_continuity=False, alternative='two-sided')
  639. print("Pval id against shuffle", pval_id_shuffle)
  640. print("Pval dd against shuffle", pval_dd_shuffle)
  641. print("Pval id against dd", pval_id_dd)
  642. #X axis
  643. ax.set_xlabel('Session', fontsize=fs+4)
  644. # ax.set_yticks(np.arange(len(neurons_to_plot)), neurons_to_plot)
  645. ax.set_xticks(np.arange(num_sessions), [pparam.SESSION_NAMES[snum] for snum in slist])
  646. ax.tick_params(axis='x', labelsize=fs)
  647. #Y axis
  648. ax.tick_params(axis='y', labelsize=fs)
  649. ax.set_ylabel('Error (cm)', fontsize=fs+4)
  650. #Both axis
  651. ax.spines[['right', 'top']].set_visible(False)
  652. for axis in ['top','bottom','left','right']:
  653. ax.spines[axis].set_linewidth(3)
  654. #Legend
  655. ax.legend(fontsize=fs-4, loc='upper right', frameon=False)
  656. #Figure params
  657. fig.tight_layout()
  658. fig_name = "fig1_prediction_across_sessions"
  659. save_figure(fig,fig_name)
  660. return fig_num
  661. def plot_fig1SI_B_C_variance_explained(fig_num = None):
  662. if fig_num is None:
  663. fig_num = plt.gcf().number + 1
  664. mlist = np.arange(8)
  665. # mlist = np.arange(4)
  666. num_mice = len(mlist)
  667. slist = np.arange(9)
  668. # slist = np.arange(3)
  669. num_sessions = len(slist)
  670. #PCA params
  671. variance_to_explain = 0.8
  672. fig_num = 1
  673. variance_explained_dict = np.load(OUTPUT_PATH +"variance_explained_dict.npy", allow_pickle=True)[()]
  674. PCA_dict = np.load(OUTPUT_PATH +"PCA_dict.npy", allow_pickle=True)[()]
  675. dimensions_array = np.zeros((num_mice, num_sessions))
  676. dimensions_array_abs = np.zeros((num_mice, num_sessions))
  677. dimension_bins = np.arange(0,105,5)
  678. var_explained_by_dimension_bin = [[] for i in range(len(dimension_bins))]
  679. cumvar_explained_by_dimension_bin = [[] for i in range(len(dimension_bins))]
  680. for midx, mnum in enumerate(mlist):
  681. for sidx, snum in enumerate(slist):
  682. pos, pca_data = PCA_dict[mnum, snum]
  683. num_neurons = pca_data.shape[0]
  684. variance_explained = variance_explained_dict[mnum, snum]
  685. variance_explained_cum = np.cumsum(variance_explained)
  686. dimensions_for_x = pf.dimensions_to_explain_variance(variance_explained, variance_to_explain)
  687. dimensions_array[midx, sidx] = int(np.around(100 * dimensions_for_x/float(num_neurons)))
  688. dimensions_array_abs[midx, sidx] = dimensions_for_x
  689. dimensions_percentage = np.array(100 * np.arange(num_neurons)/num_neurons, dtype=int)
  690. for d_idx,d in enumerate(dimensions_percentage):
  691. dim_bin = np.argmax(d <= dimension_bins)
  692. var_explained_by_dimension_bin[dim_bin].append(variance_explained[d_idx])
  693. cumvar_explained_by_dimension_bin[dim_bin].append(variance_explained_cum[d_idx])
  694. fig = plt.figure(fig_num); fig_num += 1
  695. ax = plt.gca()
  696. fs = 18
  697. color1 = 'royalblue'
  698. color2 = 'indianred'
  699. #Var explained ratio
  700. color1 = 'royalblue'
  701. var_explained_avg = np.array([np.average(varlist) for varlist in var_explained_by_dimension_bin])
  702. var_explained_std = np.array([np.std(varlist) for varlist in var_explained_by_dimension_bin])
  703. # var_explained_std = [np.std(varlist)/np.sqrt(len(varlist)) for varlist in var_explained_by_dimension_bin]
  704. ax.plot(dimension_bins, var_explained_avg, color=color1, lw=3)
  705. ax.fill_between(dimension_bins, var_explained_avg - var_explained_std, var_explained_avg + var_explained_std,
  706. color=color1, alpha=0.5)
  707. ax.set_xlabel('Dimensions (%)', fontsize=fs)
  708. ax.set_ylabel('Var. explained (ratio)', color=color1, fontsize=fs)
  709. ax.tick_params(axis='x', labelsize=fs)
  710. ax.tick_params(axis='y', labelcolor=color1, labelsize=fs)
  711. #Cumulative variance explained
  712. ax2 = ax.twinx()
  713. cumvar_explained_avg = np.array([np.average(varlist) for varlist in cumvar_explained_by_dimension_bin])
  714. cumvar_explained_std = np.array([np.std(varlist) for varlist in cumvar_explained_by_dimension_bin])
  715. # var_explained_std = [np.std(varlist)/np.sqrt(len(varlist)) for varlist in var_explained_by_dimension_bin]
  716. ax2.plot(dimension_bins, cumvar_explained_avg, color=color2, lw=3)
  717. ax2.fill_between(dimension_bins, cumvar_explained_avg - cumvar_explained_std, cumvar_explained_avg + cumvar_explained_std,
  718. color=color2, alpha=0.5)
  719. ax2.set_ylabel('Cumulative var. explained', color=color2, fontsize=fs)
  720. ax2.tick_params(axis='y', labelcolor=color2, labelsize=fs)
  721. xlims = ax.get_xlim()
  722. ylims = ax2.get_ylim()
  723. #Lines at cumulative variance explained
  724. avg_var_explained_idx = np.argmax(variance_to_explain <= cumvar_explained_avg)
  725. dim_for_avg_var_explained = dimension_bins[avg_var_explained_idx]
  726. yy = np.linspace(ax.get_ylim()[0], 0.8)
  727. ax2.plot([dim_for_avg_var_explained]*len(yy), yy, '--', alpha=1, color='gray')
  728. xx = np.linspace(dim_for_avg_var_explained, xlims[1])
  729. ax2.plot(xx, [0.8]*len(xx), '--', alpha=1, color='gray')
  730. ax.set_xlim(xlims)
  731. ax2.set_xlim(xlims)
  732. ax2.set_ylim(ylims)
  733. ax.spines[['right', 'top']].set_visible(False)
  734. for axis in ['top','bottom','left','right']:
  735. ax.spines[axis].set_linewidth(3)
  736. ax2.spines[['left', 'top']].set_visible(False)
  737. ax2.spines['right'].set_linewidth(3)
  738. fig.tight_layout()
  739. save_figure(fig, 'fig1SI_B_variance_explained')
  740. print(r'Average dimension %% to explain %d of the variance: %.1f +- %.1f, median: %d'%(100*variance_to_explain, np.average(dimensions_array), np.std(dimensions_array), np.median(dimensions_array)))
  741. print(r'ID Average dimension %% to explain %d of the variance: %.1f +- %.1f, median: %d'%(100*variance_to_explain, np.average(dimensions_array[:4]), np.std(dimensions_array[:4]), np.median(dimensions_array[:4])))
  742. print(r'Average dimension %% to explain %d of the variance: %.1f +- %.1f, median: %d'%(100*variance_to_explain, np.average(dimensions_array[4:]), np.std(dimensions_array[4:]), np.median(dimensions_array[4:])))
  743. # Plot
  744. figsize= (7, 6)
  745. fs = 15
  746. condition_color = ['forestgreen', 'khaki']
  747. fig, axs = plt.subplots(2,1, figsize=figsize); fig_num += 1
  748. for sidx, snum in enumerate(slist):
  749. ax = axs[0]
  750. dim_list = dimensions_array[:4, sidx].ravel()
  751. dim_mean = np.mean(dim_list)
  752. dim_std = np.std(dim_list)
  753. ax.bar([SESSION_NAMES[sidx]], dim_mean, yerr=dim_std, width=.7, zorder=1, color=condition_color[0])
  754. ax.scatter([sidx] * len(dim_list), dim_list, s=12, zorder=2, color='dimgray')
  755. ax = axs[1]
  756. dim_list = dimensions_array[4:8, sidx].ravel()
  757. dim_mean = np.mean(dim_list)
  758. dim_std = np.std(dim_list)
  759. ax.bar([SESSION_NAMES[sidx]], dim_mean, yerr=dim_std, width=.7, zorder=1, color=condition_color[1])
  760. ax.scatter([sidx] * len(dim_list), dim_list, s=12, zorder=2, color='dimgray')
  761. fig.suptitle('Dimensions to explain %d%% of the variance' %int(variance_to_explain*100), fontsize=fs)
  762. # fig.suptitle('Dimensions to explain % of the variance', fontsize=fs)
  763. axs[0].set_title('%s'%pparam.MOUSE_TYPE_LABELS[0], fontsize=fs+3)
  764. axs[0].set_ylabel('Dimensions (%)', fontsize=fs)
  765. axs[0].tick_params(axis='x', labelsize=fs)
  766. axs[0].tick_params(axis='y', labelsize=fs)
  767. axs[0].spines[['right', 'top']].set_visible(False)
  768. for axis in ['top','bottom','left','right']:
  769. axs[0].spines[axis].set_linewidth(3)
  770. axs[1].set_title('%s'%pparam.MOUSE_TYPE_LABELS[0], fontsize=fs+3)
  771. axs[1].set_ylabel('Dimensions (%)', fontsize=fs)
  772. axs[1].tick_params(axis='x', labelsize=fs)
  773. axs[1].tick_params(axis='y', labelsize=fs)
  774. axs[1].spines[['right', 'top']].set_visible(False)
  775. for axis in ['top','bottom','left','right']:
  776. axs[1].spines[axis].set_linewidth(3)
  777. fig.tight_layout()
  778. save_figure(fig, 'fig1SI_C_prop_dimensions_to_explain_%d_of_the_variance'%int(variance_to_explain*100))
  779. figsize= (7, 6)
  780. fs = 15
  781. condition_color = ['forestgreen', 'khaki']
  782. fig, axs = plt.subplots(2,1, figsize=figsize); fig_num += 1
  783. for sidx, snum in enumerate(slist):
  784. ax = axs[0]
  785. dim_list = dimensions_array_abs[:4, sidx].ravel()
  786. dim_mean = np.mean(dim_list)
  787. dim_std = np.std(dim_list)
  788. ax.bar([SESSION_NAMES[sidx]], dim_mean, yerr=dim_std, width=.7, zorder=1, color=condition_color[0])
  789. ax.scatter([sidx] * len(dim_list), dim_list, s=12, zorder=2, color='dimgray')
  790. ax = axs[1]
  791. dim_list = dimensions_array_abs[4:8, sidx].ravel()
  792. dim_mean = np.mean(dim_list)
  793. dim_std = np.std(dim_list)
  794. ax.bar([SESSION_NAMES[sidx]], dim_mean, yerr=dim_std, width=.7, zorder=1, color=condition_color[1])
  795. ax.scatter([sidx] * len(dim_list), dim_list, s=12, zorder=2, color='dimgray')
  796. fig.suptitle('Dimensions to explain %d%% of the variance' %int(variance_to_explain*100), fontsize=fs)
  797. # fig.suptitle('Dimensions to explain % of the variance', fontsize=fs)
  798. axs[0].set_title('%s'%pparam.MOUSE_TYPE_LABELS[0], fontsize=fs+3)
  799. axs[0].set_ylabel('Dimensions', fontsize=fs)
  800. axs[0].tick_params(axis='x', labelsize=fs)
  801. axs[0].tick_params(axis='y', labelsize=fs)
  802. axs[0].spines[['right', 'top']].set_visible(False)
  803. for axis in ['top','bottom','left','right']:
  804. axs[0].spines[axis].set_linewidth(3)
  805. axs[1].set_title('%s'%pparam.MOUSE_TYPE_LABELS[0], fontsize=fs+3)
  806. axs[1].set_ylabel('Dimensions', fontsize=fs)
  807. axs[1].tick_params(axis='x', labelsize=fs)
  808. axs[1].tick_params(axis='y', labelsize=fs)
  809. axs[1].spines[['right', 'top']].set_visible(False)
  810. for axis in ['top','bottom','left','right']:
  811. axs[1].spines[axis].set_linewidth(3)
  812. fig.tight_layout()
  813. save_figure(fig, 'fig1SI_D_raw_dimensions_to_explain_%d_of_the_variance'%int(variance_to_explain*100))
  814. def plot_fig1SI_D_prediction_vs_dimensionality(fig_num = None):
  815. if fig_num is None:
  816. fig_num = plt.gcf().number + 1
  817. ########## PARAMETERS ###########
  818. mlist = list(range(8))
  819. # mlist = [2,3,5,6]
  820. # slist = list(range(1,9))
  821. slist = list(range(9))
  822. # slist = [2,3,6]
  823. preprocessing_param_dict = {
  824. #session params
  825. 'mouse_list':mlist,
  826. 'session_list':slist,
  827. #Preprocessing parameters
  828. 'time_bin_size':1,
  829. 'distance_bin_size':1,
  830. 'gaussian_size':25,
  831. 'data_used':'amplitudes',
  832. 'running':True,
  833. 'eliminate_v_zeros':True,
  834. 'num_components':None, #will be varied
  835. }
  836. components_list = [0.1, 0.15, 0.25, 0.5, 0.8, 0.9, 0.95, 0.98, 0.99]
  837. # components_list = [0.1, 0.2, 0.8]
  838. components_list = ["%.3f"%ncomp for ncomp in components_list]
  839. components_list_len = len(components_list)
  840. print(components_list)
  841. ## Predictor parameters ##
  842. cv_folds = 5
  843. predictor_name = 'Wiener'
  844. error_type = 'sse'
  845. shuffle_reps = 10
  846. ## Plot params ##
  847. fs = 20
  848. error_by_dim = {(mtype, ncomp):[] for mtype in pparam.MOUSE_TYPE_LABELS for ncomp in components_list} # {("normal" or "shuffle", ncomp) : [list of errors]}
  849. error_by_dim_shuffle = {(mtype, ncomp):[] for mtype in pparam.MOUSE_TYPE_LABELS for ncomp in components_list} # {("normal" or "shuffle", ncomp) : [list of errors]}
  850. # error_by_dim_shuffle = {(mtype, ncomp):[] for mtype in pparam.MOUSE_TYPE_LABELS for ncomp in components_list} # {("normal" or "shuffle", ncomp) : [list of errors]}
  851. for ncomp_idx, ncomp in enumerate(components_list):
  852. preprocessing_param_dict['num_components'] = ncomp
  853. PCA_analysis_dict = perform_pca_on_multiple_mice_param_dict(preprocessing_param_dict)
  854. PCA_dict = PCA_analysis_dict['PCA_dict']
  855. for midx, mnum in enumerate(mlist):
  856. mtype = pparam.MOUSE_TYPE_LABEL_BY_MOUSE[mnum]
  857. for sidx, snum in enumerate(slist):
  858. # #Get PCA
  859. position, pca = PCA_dict[mnum, snum]
  860. pos_pred, error, predictor = pf.predict_position_CV(pca, position, n_splits=cv_folds, shuffle=False, periodic=True, pmin=0, pmax=pparam.MAX_POS,
  861. predictor_name=predictor_name, predictor_default=None, return_error=error_type)
  862. error_by_dim[mtype, ncomp].append(error)
  863. if shuffle_reps != 0:
  864. timepoints = len(position)
  865. idxs = list(range(timepoints))
  866. error_shuffle_list = []
  867. for rep in range(shuffle_reps):
  868. shift = np.random.randint(-timepoints, timepoints)
  869. idxs_shifted = idxs[shift:] + idxs[:shift]
  870. idxs_shifted = np.random.choice(idxs_shifted, size=len(idxs_shifted), replace=False)
  871. position_shuffled = position[idxs_shifted]
  872. pca_shuffled = np.copy(pca)
  873. pos_pred_shuffle, error_shuffle, _ = pf.predict_position_CV(pca_shuffled, position_shuffled, n_splits=cv_folds, shuffle=False, periodic=True, pmin=0, pmax=pparam.MAX_POS,
  874. predictor_name=predictor_name, predictor_default=None, return_error=error_type)
  875. error_shuffle_list.append(error_shuffle)
  876. error_shuffle_avg = np.average(error_shuffle_list)
  877. error_by_dim_shuffle[mtype, ncomp].append(error_shuffle_avg)
  878. #Compute averages to plot
  879. error_by_dim_avg = np.zeros((len(pparam.MOUSE_TYPE_LABELS), components_list_len))
  880. error_by_dim_std = np.zeros((len(pparam.MOUSE_TYPE_LABELS), components_list_len))
  881. error_by_dim_shuffle_avg = np.zeros(components_list_len)
  882. error_by_dim_shuffle_std = np.zeros(components_list_len)
  883. for ncomp_idx, ncomp in enumerate(components_list):
  884. for mtype_idx, mtype in enumerate(pparam.MOUSE_TYPE_LABELS):
  885. errors = error_by_dim[mtype, ncomp]
  886. avg = np.average(errors)
  887. std = np.std(errors)/np.sqrt(len(errors))
  888. error_by_dim_avg[mtype_idx, ncomp_idx] = avg
  889. error_by_dim_std[mtype_idx, ncomp_idx] = std
  890. if shuffle_reps != 0:
  891. errors = error_by_dim_shuffle[pparam.MOUSE_TYPE_LABELS[0], ncomp] + error_by_dim_shuffle[pparam.MOUSE_TYPE_LABELS[1], ncomp]
  892. avg = np.average(errors)
  893. std = np.std(errors)/np.sqrt(len(errors))
  894. error_by_dim_shuffle_avg[ncomp_idx] = avg
  895. error_by_dim_shuffle_std[ncomp_idx] = std
  896. ## Plot results ##
  897. fig = plt.figure(fig_num, figsize=(7,4)); fig_num += 1
  898. ax = plt.gca()
  899. # xx = np.array(components_list).astype(float) * 100
  900. # ax.set_xscale('log')
  901. xx = np.arange(len(components_list))
  902. for mtype_idx, mtype in enumerate(pparam.MOUSE_TYPE_LABELS):
  903. avgs = error_by_dim_avg[mtype_idx]
  904. stds = error_by_dim_std[mtype_idx]
  905. color = pparam.MOUSE_TYPE_COLORS[mtype_idx]
  906. ax.plot(xx, avgs, lw=3, color=color, label=mtype)
  907. ax.fill_between(xx, avgs-stds, avgs+stds, color=color, alpha=0.5)
  908. if shuffle_reps != 0:
  909. ax.plot(xx, error_by_dim_shuffle_avg, lw=3, color=pparam.SHUFFLE_DEFAULT_COLOR, label='Shuffle')
  910. ax.fill_between(xx, error_by_dim_shuffle_avg-error_by_dim_shuffle_std, error_by_dim_shuffle_avg+error_by_dim_shuffle_std,
  911. color=pparam.SHUFFLE_DEFAULT_COLOR, alpha=0.5)
  912. #X axis
  913. ax.set_xlabel('Variance explained (%)', fontsize=fs)
  914. ax.set_xticks([], []) #Delete log ticks
  915. xlabels = [int(100*float(ncomp)) for ncomp in components_list]
  916. ax.set_xticks(xx, xlabels) #Add the ticks I want
  917. ax.tick_params(axis='x', labelsize=fs)
  918. ax.minorticks_off()
  919. #Y axis
  920. ax.tick_params(axis='y', labelsize=fs)
  921. ax.set_ylabel('Error (cm)', fontsize=fs)
  922. #Both axis
  923. ax.spines[['right', 'top']].set_visible(False)
  924. for axis in ['top','bottom','left','right']:
  925. ax.spines[axis].set_linewidth(3)
  926. #Legend
  927. ax.legend(fontsize=fs-4, loc='upper right', frameon=False)
  928. #Figure params
  929. fig.tight_layout()
  930. fig_name = "fig1SI_prediction_by_dimensionality"
  931. save_figure(fig,fig_name)
  932. return
  933. #~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ Fig 2 CCA ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
  934. def perform_mCCA_on_pca_dict_param_dict(pca_analysis_dict, cca_param_dict, force_recalculation=True):
  935. '''Given a PCA analysis dictionary (from "save_pca_data"), perform mCCA on it
  936. Input:
  937. PCA_analysis_dict: output of "save_pca_data"
  938. sessions_to_align: either 'all' or a list of sessions that need to be aligned
  939. pca_dim: int or percentage. Number of PCA dimensions to keep.
  940. return_warped_data: bool, if True, returns data with task-normalized bins
  941. return_trimmed_data: bool, if True, returns dataset trimmed so each session has the same number of laps.
  942. If False, some sessions may include laps not used in the alignment process
  943. verbose: bool, if True says which sessions it's aligning
  944. plot: bool, if True plots the result
  945. Returns:
  946. aligned_data_dict: dictionary with all relevant CCA data
  947. '''
  948. #First check if the cca parameters are the same
  949. analysis_name = pparam.ANALYSIS_NAME_LIST[1]
  950. if cca_param_dict['shuffle'] == True:
  951. custom_filename = pparam.default_param_dicts_names[analysis_name] + "_shuffle"
  952. else:
  953. custom_filename = analysis_name
  954. are_cca_params_equal, cca_param_dict, aligned_data_dict_previous = process_input_parameter_dict(cca_param_dict, analysis_name, custom_filename)
  955. #Then check if the pca parameters were the same as well
  956. preprocessing_param_dict = pca_analysis_dict['param_dict']
  957. try:
  958. preprocessing_param_dict_old = aligned_data_dict_previous[pparam.ANALYSIS_NAME_LIST[0] + '_param_dict']
  959. are_preprocessing_params_equal = compare_parameter_dictionaries(preprocessing_param_dict, preprocessing_param_dict_old)
  960. except KeyError:
  961. #No properly parsed previous analysis found
  962. are_preprocessing_params_equal = False
  963. # print('Are CCA parameters repeated?', are_preprocessing_params_equal,
  964. # 'And PCA?', are_cca_params_equal)
  965. #If both the PCA and CCA analysis were the same, return previous result
  966. if force_recalculation == False and are_cca_params_equal == True and are_preprocessing_params_equal == True and cca_param_dict['shuffle'] == False: #Always recalculate for shuffle!
  967. return aligned_data_dict_previous
  968. #Path to save results
  969. analysis_dict_filename = get_analysis_dict_filename(custom_filename)
  970. #Warping parameters
  971. warping_bins = 150
  972. warp_based_on = 'position'
  973. # warp_based_on = 'position'
  974. #Prediction parameters
  975. error_type = 'sse'
  976. n_splits = 5
  977. predictor_name = 'Wiener'
  978. # predictor_name = 'SVR'
  979. ## CCA params ##
  980. CCA_dim = cca_param_dict['CCA_dim']
  981. return_warped_data = cca_param_dict['return_warped_data']
  982. return_trimmed_data = cca_param_dict['return_trimmed_data']
  983. sessions_to_align = cca_param_dict['sessions_to_align']
  984. shuffle = cca_param_dict['shuffle']
  985. skip_alignment = False
  986. if 'skip_alignment' in cca_param_dict:
  987. skip_alignment = cca_param_dict['skip_alignment']
  988. #Data parameters
  989. max_pos = pparam.MAX_POS
  990. mouse_list = pca_analysis_dict['mouse_list']
  991. pca_data_dict = pca_analysis_dict['PCA_dict']
  992. variance_explained_dict = pca_analysis_dict['variance_explained_dict']
  993. if sessions_to_align == 'all':
  994. session_list_original = pca_analysis_dict['session_list']
  995. else:
  996. session_list_original = sessions_to_align
  997. aligned_data_dict = {}
  998. for midx, mnum in enumerate(mouse_list):
  999. session_list = filter_repeated_sessions(mnum, session_list_original)
  1000. pos_list, pca_list = pca_and_pos_from_dict_to_lists(pca_data_dict, mnum, session_list)
  1001. M = len(session_list)
  1002. #Set PCA dimension
  1003. pca_list = mCCA_funs.set_dimension_of_pca_list(pca_list, CCA_dim, variance_explained_list = [variance_explained_dict[mnum, snum] for snum in session_list])
  1004. #Perform mCCA
  1005. pos_list_aligned, pca_dict_aligned, mCCA = mCCA_funs.perform_warped_mCCA(pos_list, pca_list, max_pos, warping_bins, warp_based_on,
  1006. return_warped_data, return_trimmed_data, shuffle, skip_alignment)
  1007. #Normalize PCA after alignment changes
  1008. pca_dict_aligned = mCCA_funs.normalize_pca_dict_aligned(pca_dict_aligned, mCCA)
  1009. #Find space with best alignment
  1010. best_space = mCCA_funs.return_best_mCCA_space(pos_list_aligned, pca_dict_aligned, max_pos=1500, verbose=False)
  1011. pca_list_aligned = pca_dict_aligned[best_space]
  1012. pca_list = [pca_dict_aligned[m][m] for m in range(M)]
  1013. unaligned_error_array, aligned_error_array = mCCA_funs.get_cross_prediction_errors(pos_list_aligned, pca_list, pos_list_aligned, pca_dict_aligned,
  1014. max_pos, n_splits, error_type, predictor_name)
  1015. # print('WARNING: IGNORING CCA IN APdecoding_pipeline AS A TEST, REVERT THIS CHANGE')
  1016. # pca_list_aligned = pca_list
  1017. # aligned_error_array = unaligned_error_array
  1018. # print('WARNING: IGNORING CCA IN APdecoding_pipeline AS A TEST, REVERT THIS CHANGE')
  1019. aligned_data_dict[mnum, 'pos'] = pos_list_aligned
  1020. aligned_data_dict[mnum, 'pca'] = pca_list_aligned
  1021. aligned_data_dict[mnum, 'pca_unaligned'] = pca_list
  1022. aligned_data_dict[mnum, 'pca_dict_aligned'] = pca_dict_aligned
  1023. aligned_data_dict[mnum, 'best_space'] = best_space
  1024. aligned_data_dict[mnum, 'session_list'] = session_list
  1025. aligned_data_dict[mnum, 'unaligned_error_array'] = unaligned_error_array
  1026. aligned_data_dict[mnum, 'aligned_error_array'] = aligned_error_array
  1027. aligned_data_dict[mnum, 'mCCA_instance'] = mCCA
  1028. aligned_data_dict['mouse_list'] = mouse_list
  1029. aligned_data_dict['num_bins'] = warping_bins
  1030. aligned_data_dict['param_dict'] = cca_param_dict
  1031. aligned_data_dict[pparam.ANALYSIS_NAME_LIST[0] + '_param_dict'] = pca_analysis_dict['param_dict']
  1032. # np.save(OUTPUT_PATH + pparam.cca_dict_name + ".npy", aligned_data_dict)
  1033. np.save(analysis_dict_filename, aligned_data_dict)
  1034. return aligned_data_dict
  1035. def perform_mCCA_on_pca_dict(
  1036. pca_analysis_dict,
  1037. sessions_to_align = 'all', #[0,1,2,3,4,5,6,7,8] [0,1,2,3,4,5,7,8]
  1038. pca_dim = 12, #6, '85%'
  1039. return_warped_data = False,
  1040. return_trimmed_data = False,
  1041. plot=True,
  1042. shuffle=False
  1043. ):
  1044. '''
  1045. Given a PCA analysis dictionary (from "save_pca_data"), perform mCCA on it
  1046. Input:
  1047. PCA_analysis_dict: output of "save_pca_data"
  1048. sessions_to_align: either 'all' or a list of sessions that need to be aligned
  1049. pca_dim: int or percentage. Number of PCA dimensions to keep.
  1050. return_warped_data: bool, if True, returns data with task-normalized bins
  1051. return_trimmed_data: bool, if True, returns dataset trimmed so each session has the same number of laps.
  1052. If False, some sessions may include laps not used in the alignment process
  1053. verbose: bool, if True says which sessions it's aligning
  1054. plot: bool, if True plots the result
  1055. Returns:
  1056. aligned_data_dict: dictionary with all relevant CCA data
  1057. '''
  1058. #Warping parameters
  1059. warping_bins = 150
  1060. warp_based_on = 'position'
  1061. #Prediction parameters
  1062. error_type = 'sse'
  1063. n_splits = 5
  1064. predictor_name = 'Wiener'
  1065. #Data parameters
  1066. max_pos = pparam.MAX_POS
  1067. mouse_list = pca_analysis_dict['mouse_list']
  1068. pca_data_dict = pca_analysis_dict['PCA_dict']
  1069. variance_explained_dict = pca_analysis_dict['variance_explained_dict']
  1070. if sessions_to_align == 'all':
  1071. session_list_original = pca_analysis_dict['session_list']
  1072. else:
  1073. session_list_original = sessions_to_align
  1074. aligned_data_dict = {}
  1075. for midx, mnum in enumerate(mouse_list):
  1076. session_list = filter_repeated_sessions(mnum, session_list_original)
  1077. pos_list, pca_list = pca_and_pos_from_dict_to_lists(pca_data_dict, mnum, session_list)
  1078. M = len(session_list)
  1079. #Set PCA dimension
  1080. pca_list = mCCA_funs.set_dimension_of_pca_list(pca_list, pca_dim, variance_explained_list = [variance_explained_dict[mnum, snum] for snum in session_list])
  1081. #Perform mCCA
  1082. pos_list_aligned, pca_dict_aligned, mCCA = mCCA_funs.perform_warped_mCCA(pos_list, pca_list, max_pos, warping_bins, warp_based_on, return_warped_data, return_trimmed_data, shuffle)
  1083. #Normalize PCA after alignment changes
  1084. pca_dict_aligned = mCCA_funs.normalize_pca_dict_aligned(pca_dict_aligned, mCCA)
  1085. #Find space with best alignment
  1086. best_space = mCCA_funs.return_best_mCCA_space(pos_list_aligned, pca_dict_aligned, max_pos=1500, verbose=False)
  1087. pca_list_aligned = pca_dict_aligned[best_space]
  1088. pca_list = [pca_dict_aligned[m][m] for m in range(M)]
  1089. unaligned_error_array, aligned_error_array = mCCA_funs.get_cross_prediction_errors(pos_list_aligned, pca_list, pos_list_aligned, pca_dict_aligned, max_pos, n_splits, error_type, predictor_name)
  1090. aligned_data_dict[mnum, 'pos'] = pos_list_aligned
  1091. aligned_data_dict[mnum, 'pca'] = pca_list_aligned
  1092. aligned_data_dict[mnum, 'pca_unaligned'] = pca_list
  1093. aligned_data_dict[mnum, 'pca_dict_aligned'] = pca_dict_aligned
  1094. aligned_data_dict[mnum, 'best_space'] = best_space
  1095. aligned_data_dict[mnum, 'session_list'] = session_list
  1096. aligned_data_dict[mnum, 'unaligned_error_array'] = unaligned_error_array
  1097. aligned_data_dict[mnum, 'aligned_error_array'] = aligned_error_array
  1098. aligned_data_dict['mouse_list'] = mouse_list
  1099. aligned_data_dict['num_bins'] = warping_bins
  1100. return aligned_data_dict
  1101. def get_cca_filename_full_path(CCA_dim, return_trimmed_data, return_warped_data):
  1102. cca_filename = 'aligned_data_dict'
  1103. cca_filename += '_%s'%str(CCA_dim)
  1104. cca_filename += ['_untrimmed','_trimmed'][return_trimmed_data]
  1105. cca_filename += ['_unwarped','_warped'][return_warped_data]
  1106. cca_filename_full_path = OUTPUT_PATH + cca_filename + '.npy'
  1107. return cca_filename_full_path
  1108. def fig2_CCA_example():
  1109. ''' Plots the result of alignment on a single '''
  1110. mnum = 6
  1111. snum1 = 1
  1112. snum2 = 3
  1113. mouse_list = [mnum]
  1114. session_list = range(9)
  1115. figsize = (8,6)
  1116. preprocessing_param_dict = {
  1117. #session params
  1118. 'mouse_list':mouse_list,
  1119. 'session_list':session_list,
  1120. #Preprocessing parameters
  1121. 'time_bin_size':1,
  1122. 'distance_bin_size':1,
  1123. 'gaussian_size':25,
  1124. 'data_used':'amplitudes',
  1125. 'running':True,
  1126. 'eliminate_v_zeros':False,
  1127. 'num_components':'all',
  1128. }
  1129. ## CCA params ##
  1130. cca_param_dict = {
  1131. 'CCA_dim':12,
  1132. 'return_warped_data':False,
  1133. 'return_trimmed_data':False,
  1134. 'sessions_to_align':'all',
  1135. 'shuffle':False
  1136. }
  1137. ## Plot params ##
  1138. fig_num = 1
  1139. fs = 17
  1140. pca_plot_bins = 50
  1141. max_pos = pparam.MAX_POS
  1142. predictor_name_example_plot = 'XGBoost'
  1143. #PCA
  1144. PCA_analysis_dict = perform_pca_on_multiple_mice_param_dict(preprocessing_param_dict)
  1145. #CCA
  1146. aligned_data_dict = perform_mCCA_on_pca_dict_param_dict(PCA_analysis_dict, cca_param_dict)
  1147. rows=3
  1148. cols=2
  1149. fig = plt.figure(fig_num, figsize=figsize); fig_num += 1
  1150. pca_list_unaligned = aligned_data_dict[mnum, 'pca_unaligned']
  1151. pca_list_aligned = aligned_data_dict[mnum, 'pca']
  1152. pos_list = aligned_data_dict[mnum, 'pos']
  1153. #Force them to be the same size
  1154. min_size = np.min([np.size(pos_list[snum]) for snum in [snum1, snum2]])
  1155. pos_list = [pos[:min_size] for pos in pos_list]
  1156. pca_list_unaligned = [pca[:,:min_size] for pca in pca_list_unaligned]
  1157. pca_list_aligned = [pca[:,:min_size] for pca in pca_list_aligned]
  1158. subplot_counter = 0
  1159. for row in range(rows):
  1160. if row==0:
  1161. snums = [snum1]
  1162. elif row ==1:
  1163. snums = [snum2]
  1164. elif row == 2:
  1165. snums = [snum1, snum2]
  1166. #Plot position prediction
  1167. subplot_counter += 1
  1168. ax = fig.add_subplot(rows, cols, subplot_counter, aspect=0.4)
  1169. snum = snums[-1]
  1170. pos = pos_list[snum]
  1171. if row != 2:
  1172. pca = pca_list_unaligned[snum]
  1173. else:
  1174. pca = pca_list_aligned[snum]
  1175. if row == 0:
  1176. pos_pred, error, _ = pf.predict_position_CV(pca, pos, n_splits=5, pmax=max_pos, predictor_name=predictor_name_example_plot)
  1177. _,_, predictor = pf.predict_position_CV(pca, pos, n_splits=0, pmax=max_pos, predictor_name=predictor_name_example_plot)
  1178. elif row == 1:
  1179. pos_pred, error = pf.predict_position_from_predictor_object(pca, pos, predictor, periodic=True, pmax=max_pos)
  1180. elif row == 2:
  1181. _,_, predictor = pf.predict_position_CV(pca_list_aligned[snum1], pos_list[snum1], n_splits=0, pmax=max_pos, predictor_name=predictor_name_example_plot)
  1182. pos_pred, error = pf.predict_position_from_predictor_object(pca_list_aligned[snum2], pos_list[snum2], predictor, periodic=True, pmax=max_pos)
  1183. time_steps = np.arange(pos.size)
  1184. ax.scatter(time_steps, pos, s=4, alpha=0.6, lw=2, color=pparam.PREDICTION_COLORS[pparam.PREDICTION_LABELS[0]])
  1185. ax.scatter(time_steps, pos_pred, alpha=1, s=4,lw=2, color=pparam.PREDICTION_COLORS[pparam.PREDICTION_LABELS[1]])
  1186. #Ticks
  1187. tickw, tickl = 4, 7
  1188. ax.set_yticks([0, 750, 1500])
  1189. ax.set_xticks(np.arange(0, np.max(time_steps), 400).astype(int))
  1190. ax.tick_params(axis='x', labelsize=fs, width=tickw, length=tickl)
  1191. ax.tick_params(axis='y', labelsize=fs, width=tickw, length=tickl)
  1192. ax.set_ylim([-0.1*pparam.MAX_POS, ax.get_ylim()[-1]])
  1193. if row == 2:
  1194. ax.set_xlabel('Time step', fontsize=fs)
  1195. # ax[1].tick_params(width=3, length=6)
  1196. #Axis
  1197. ax.spines[['right', 'top']].set_visible(False)
  1198. for axis in ['bottom','left']:
  1199. ax.spines[axis].set_linewidth(3)
  1200. if row == 0:
  1201. ax.plot([0],[0], label=pparam.PREDICTION_LABELS[0], color=pparam.PREDICTION_COLORS[pparam.PREDICTION_LABELS[0]])
  1202. ax.plot([0],[0], label=pparam.PREDICTION_LABELS[1], color=pparam.PREDICTION_COLORS[pparam.PREDICTION_LABELS[1]])
  1203. ax.legend(fontsize=fs-5, loc='upper left', frameon=False)
  1204. if row == 1:
  1205. ax.set_ylabel('Position (mm)', fontsize=fs)
  1206. #Plot PCA
  1207. subplot_counter += 1
  1208. ax = fig.add_subplot(rows, cols, subplot_counter, projection='3d')
  1209. for snum_idx, snum in enumerate(snums):
  1210. pos = pos_list[snum]
  1211. if row != 2:
  1212. pca = pca_list_unaligned[snum]
  1213. else:
  1214. pca = pca_list_aligned[snum]
  1215. # cbar = [None, True][row==0]
  1216. cbar = None
  1217. pos_bins, pca_avg, _ = pf.compute_average_data_by_position(pca, pos, position_bin_size=pca_plot_bins)
  1218. pf.plot_pca_with_position(pca_avg, pos_bins, ax=ax, scatter=False, cbar=cbar)
  1219. fig.tight_layout()
  1220. fig_name = "fig2_cca_example"
  1221. save_figure(fig,fig_name)
  1222. def fig2_CCA_aligned_pcas(shuffle=False):
  1223. mouse_list = np.arange(8)
  1224. session_list = np.arange(9)
  1225. mnum1 = 1
  1226. mnum2 = 6
  1227. # time_bin_size = 1 # Number of elements to average over, each dt should be ~65ms
  1228. # distance_bin_size = 1 # mm, track is 1500mm, data is in mm
  1229. # gaussian_size = 25 # Why not
  1230. # data_used = 'amplitudes'
  1231. # running = True
  1232. # eliminate_v_zeros = False
  1233. preprocessing_param_dict = {
  1234. #session params
  1235. 'mouse_list':mouse_list,
  1236. 'session_list':session_list,
  1237. #Preprocessing parameters
  1238. 'time_bin_size':1,
  1239. 'distance_bin_size':1,
  1240. 'gaussian_size':25,
  1241. 'data_used':'amplitudes',
  1242. 'running':True,
  1243. 'eliminate_v_zeros':False,
  1244. 'num_components':3,
  1245. }
  1246. ## CCA params ##
  1247. return_warped_data = True
  1248. return_trimmed_data = True
  1249. CCA_dim = 12
  1250. ## Plot params ##
  1251. fig_num = 1
  1252. pca_plot_bins = 50
  1253. figsize1 = (5,5)
  1254. angle_list1 = [40, 30, 80, 50]
  1255. angle_azim_list1 = [-170, -25, -120, None]
  1256. figsize2 = (8,6)
  1257. angle_list2 = [80, -170, -90, 50, 50, -50]
  1258. angle_azim_list2 = [-50, -90, 50, -150, None, -27]
  1259. PCA_analysis_dict = perform_pca_on_multiple_mice_param_dict(preprocessing_param_dict)
  1260. cca_filename_full_path = get_cca_filename_full_path(CCA_dim, return_trimmed_data, return_warped_data)
  1261. #Can be commented after being run once
  1262. aligned_data_dict = perform_mCCA_on_pca_dict(
  1263. PCA_analysis_dict,
  1264. sessions_to_align = 'all', #[0,1,2,3,4,5,6,7,8] [0,1,2,3,4,5,7,8]
  1265. pca_dim = CCA_dim, #6, '85%'
  1266. return_warped_data = return_warped_data,
  1267. return_trimmed_data = return_trimmed_data,
  1268. plot=False,
  1269. shuffle=shuffle
  1270. )
  1271. np.save(cca_filename_full_path, aligned_data_dict, allow_pickle=True)
  1272. #Can be commented after being run once
  1273. aligned_data_dict = np.load(cca_filename_full_path, allow_pickle=True)[()]
  1274. #FIG 2b: UNALIGNED VS ALIGNED
  1275. fig_CCAcomp, axs_CCAcomp = plt.subplots(nrows=2, ncols = 2, squeeze=False, figsize=figsize1, num=fig_num, subplot_kw={'projection':'3d'})
  1276. # fig_CCAcomp, axs_CCAcomp = plt.subplots(nrows=2, ncols = 2, squeeze=False, figsize=subfig_b_size, num=fig_num)
  1277. fig_num += 1
  1278. subplot_counter = 0
  1279. for midx, mnum in enumerate([mnum1, mnum2]):
  1280. for alignment_idx in range(2): #0 is unaligned, 1 is aligned
  1281. ax = axs_CCAcomp[midx, alignment_idx]
  1282. if alignment_idx == 0:
  1283. pca_list = aligned_data_dict[mnum, 'pca_unaligned']
  1284. elif alignment_idx == 1:
  1285. pca_list = aligned_data_dict[mnum, 'pca']
  1286. pos_list = aligned_data_dict[mnum, 'pos']
  1287. num_sessions = len(pca_list)
  1288. for sidx in range(num_sessions):
  1289. pca = pca_list[sidx]
  1290. pos = pos_list[sidx]
  1291. pos_bins, pca_avg, _ = pf.compute_average_data_by_position(pca, pos, position_bin_size=pca_plot_bins)
  1292. pf.plot_pca_with_position(pca_avg, pos_bins, ax=ax, scatter=False, cbar=None,
  1293. angle = angle_list1[subplot_counter], angle_azim = angle_azim_list1[subplot_counter])
  1294. subplot_counter += 1
  1295. fig_CCAcomp.tight_layout()
  1296. fig_name = "fig2_cca_aligned_pcas"
  1297. save_figure(fig_CCAcomp,fig_name)
  1298. # return
  1299. #FIG 2c: ALL ALIGNED
  1300. fig_CCAaligned, axs_CCAaligned = plt.subplots(nrows=2, ncols = 3, squeeze=False, figsize=figsize2, num=fig_num, subplot_kw={'projection':'3d'})
  1301. fig_num += 1
  1302. #Ignore plots from previous plot
  1303. mlist_for_all_aligned_plot = [mnum for mnum in range(8) if mnum not in [mnum1,mnum2] ]
  1304. subplot_counter = 0
  1305. for midx, mnum in enumerate(mlist_for_all_aligned_plot):
  1306. ax = axs_CCAaligned[midx//3, midx%3]
  1307. if mnum in mouse_list:
  1308. pca_list = aligned_data_dict[mnum, 'pca']
  1309. pos_list = aligned_data_dict[mnum, 'pos']
  1310. num_sessions = len(pca_list)
  1311. for sidx in range(num_sessions):
  1312. pca = pca_list[sidx]
  1313. pos = pos_list[sidx]
  1314. pos_bins, pca_avg, _ = pf.compute_average_data_by_position(pca, pos, position_bin_size=pca_plot_bins)
  1315. pf.plot_pca_with_position(pca_avg, pos_bins, ax=ax, scatter=False, cbar=None,
  1316. angle = angle_list2[subplot_counter], angle_azim = angle_azim_list2[subplot_counter])
  1317. # ax.set_title('M%d'%mnum, fontsize=25)
  1318. subplot_counter += 1
  1319. # fig_CCAaligned.subplots_adjust(wspace=-200, hspace=0)
  1320. fig_CCAaligned.tight_layout()
  1321. fig_name = "fig2_cca_aligned_pcas_remaining"
  1322. save_figure(fig_CCAaligned,fig_name)
  1323. def fig2_D_E_F_CCA_quantification():
  1324. ''' Plot figures related to quantifying CCA alignemnt '''
  1325. mouse_list = np.arange(8)
  1326. # mouse_list = [2,7]
  1327. session_list = np.arange(9)
  1328. preprocessing_param_dict = {
  1329. #session params
  1330. 'mouse_list':mouse_list,
  1331. 'session_list':session_list,
  1332. #Preprocessing parameters
  1333. 'time_bin_size':1,
  1334. 'distance_bin_size':1,
  1335. 'gaussian_size':25,
  1336. 'data_used':'amplitudes',
  1337. 'running':True,
  1338. 'eliminate_v_zeros':True,
  1339. 'num_components':'all',
  1340. }
  1341. ## CCA params ##
  1342. cca_param_dict = {
  1343. 'CCA_dim':'.9',
  1344. 'return_warped_data':False,
  1345. 'return_trimmed_data':False,
  1346. 'sessions_to_align':'all',
  1347. 'shuffle':False,
  1348. 'skip_alignment':False
  1349. }
  1350. ## CCA params ##
  1351. add_shuffle = True
  1352. ## Plot params ##
  1353. fig_num = 1
  1354. fs = 15
  1355. CCAmap_figsize = (7,7)
  1356. hist_figsize = (7,7)
  1357. CCAavg_figsize = (5,4)
  1358. figsize_sessions = (6,5)
  1359. num_sessions = len(session_list)
  1360. ########## STEP 1 - PCA ###########
  1361. PCA_analysis_dict = perform_pca_on_multiple_mice_param_dict(preprocessing_param_dict)
  1362. aligned_data_dict = perform_mCCA_on_pca_dict_param_dict(PCA_analysis_dict, cca_param_dict)
  1363. if add_shuffle:
  1364. cca_param_dict_shuffle = {k:v for k,v in cca_param_dict.items()}
  1365. cca_param_dict_shuffle['shuffle'] = True
  1366. aligned_data_dict_shuffle = perform_mCCA_on_pca_dict_param_dict(PCA_analysis_dict, cca_param_dict_shuffle)
  1367. ### ADD SHUFFLE ###
  1368. #Get prediction error by alignment kind
  1369. error_by_alignment = {} #(mouse_type_label, CCA_label): list of errors
  1370. for mtype_idx, mouse_type_label in enumerate(pparam.MOUSE_TYPE_LABELS):
  1371. #Gather together mice of same type
  1372. mlist = pparam.MOUSE_TYPE_INDEXES[mouse_type_label]
  1373. for label_idx, label in enumerate(pparam.CCA_LABELS[:4]):
  1374. if label_idx == 3 and add_shuffle == False:
  1375. continue
  1376. error_list = []
  1377. for mnum in mlist:
  1378. if label_idx == 0:
  1379. errors = aligned_data_dict[mnum, 'unaligned_error_array'].diagonal()
  1380. else:
  1381. if label_idx == 1:
  1382. errors = aligned_data_dict[mnum, 'unaligned_error_array']
  1383. elif label_idx == 2:
  1384. errors = aligned_data_dict[mnum, 'aligned_error_array']
  1385. elif label_idx == 3:
  1386. errors = aligned_data_dict_shuffle[mnum, 'aligned_error_array']
  1387. errors = errors[np.where(~np.eye(errors.shape[0], dtype=bool))]
  1388. error_list.extend(errors)
  1389. error_by_alignment[mouse_type_label, label] = error_list
  1390. #Get *relative* errors by alignment kind
  1391. errordiff_by_alignment = {} #(mouse_type_label, CCA_label)
  1392. errornorm_by_alignment = {} #(mouse_type_label, CCA_label)
  1393. for mtype_idx, mouse_type_label in enumerate(pparam.MOUSE_TYPE_LABELS):
  1394. #Gather together mice of same type
  1395. mlist = pparam.MOUSE_TYPE_INDEXES[mouse_type_label]
  1396. for midx, mnum in enumerate(mlist):
  1397. #Get reference error
  1398. errors_self = aligned_data_dict[mnum, 'unaligned_error_array'].diagonal().astype(int)
  1399. for label_idx, label in enumerate(pparam.CCA_LABELS):
  1400. if label == pparam.CCA_LABELS[0]: #Self, skip
  1401. continue
  1402. if label_idx == 3 and add_shuffle == False:
  1403. continue
  1404. if label == pparam.CCA_LABELS[1]: #Unaligned
  1405. errors = aligned_data_dict[mnum, 'unaligned_error_array']
  1406. elif label == pparam.CCA_LABELS[2]: #Aligned
  1407. errors = aligned_data_dict[mnum, 'aligned_error_array']
  1408. elif label == pparam.CCA_LABELS[3]:
  1409. errors = aligned_data_dict_shuffle[mnum, 'aligned_error_array']
  1410. #Error difference
  1411. error_difference = errors - errors_self[:, np.newaxis]
  1412. error_difference = error_difference[np.where(~np.eye(error_difference.shape[0], dtype=bool))]
  1413. if midx == 0:
  1414. errordiff_by_alignment[mouse_type_label, label] = []
  1415. errordiff_by_alignment[mouse_type_label, label].extend(error_difference)
  1416. #Normalized error
  1417. error_norm = errors / errors_self[:, np.newaxis]
  1418. error_norm = error_norm[np.where(~np.eye(error_norm.shape[0], dtype=bool))]
  1419. if midx == 0:
  1420. errornorm_by_alignment[mouse_type_label, label] = []
  1421. errornorm_by_alignment[mouse_type_label, label].extend(error_norm)
  1422. #Fig 2d: CCA colormaps of cross prediction errors
  1423. fig_CCAmap, axs_CCAmap = plt.subplots(2, 2, num=fig_num, figsize=CCAmap_figsize); fig_num += 1
  1424. #Average them
  1425. mouse_list = np.array(mouse_list)
  1426. mouse_list_by_kind = [mouse_list[mouse_list< 4], mouse_list[mouse_list >=4]]
  1427. overall_max_error = 0
  1428. overall_min_error = np.inf
  1429. for mlist_idx, mlist in enumerate(mouse_list_by_kind):
  1430. num_sessions_max = len(session_list)
  1431. alignment_names = ['unaligned_error_array', 'aligned_error_array']
  1432. alignment_error_arrays = {name:np.zeros((num_sessions_max, num_sessions_max)) for name in alignment_names}
  1433. alignment_counter = {name:np.zeros((num_sessions_max, num_sessions_max)) for name in alignment_names}
  1434. for mnum in mlist:
  1435. for name in alignment_names:
  1436. error_array = aligned_data_dict[mnum, name]
  1437. current_session_list = aligned_data_dict[mnum, 'session_list']
  1438. for scounter1, snum1 in enumerate(current_session_list):
  1439. sidx1 = np.where(np.array(session_list) == snum1)[0]
  1440. for scounter2, snum2 in enumerate(current_session_list):
  1441. sidx2 = np.where(np.array(session_list) == snum2)[0]
  1442. alignment_error_arrays[name][sidx1, sidx2] += error_array[scounter1, scounter2]
  1443. alignment_counter[name][sidx1, sidx2] += 1
  1444. alignment_error_arrays = {name:alignment_error_arrays[name]/alignment_counter[name] for name in alignment_names}
  1445. min_error = np.min([np.min(alignment_error_arrays[name]) for name in alignment_names])
  1446. max_error = np.max([np.max(alignment_error_arrays[name]) for name in alignment_names])
  1447. overall_max_error = np.maximum(overall_max_error, max_error)
  1448. overall_min_error = np.minimum(overall_min_error, min_error)
  1449. unaligned_error_array = alignment_error_arrays['unaligned_error_array']
  1450. aligned_error_array = alignment_error_arrays['aligned_error_array']
  1451. plot_data_list = [unaligned_error_array, aligned_error_array]
  1452. title_list = ['Unaligned', 'Aligned']
  1453. fs = 15
  1454. for alignment_idx in range(2):
  1455. ax = axs_CCAmap[mlist_idx, alignment_idx]
  1456. data = plot_data_list[alignment_idx]
  1457. ax.imshow(data, cmap=pparam.ERROR_CMAP, interpolation='nearest', vmin=min_error, vmax=max_error)
  1458. snames = [SESSION_NAMES[s] for s in session_list]
  1459. ax.set_xticks(range(num_sessions), snames, fontsize=fs)
  1460. ax.set_yticks(range(num_sessions), snames, fontsize=fs)
  1461. ax.set_title(title_list[alignment_idx], fontsize = fs+6, pad=10)
  1462. for m in range(num_sessions):
  1463. ax.add_patch(Rectangle((m-0.5, m-0.5), 1, 1, fill=False, edgecolor='indianred', lw=3))
  1464. # axs_CCAmap[0,0].annotate('CA3 i-D', xy=(0, 0.5), xytext=(170, 180), fontsize=20, xycoords='axes points')
  1465. # axs_CCAmap[1,0].annotate('CA3 D-D', xy=(0, 0.5), xytext=(170, 175), fontsize=20, xycoords='axes points')
  1466. axs_CCAmap[0,0].set_ylabel('Trained on', fontsize=fs+4)
  1467. axs_CCAmap[1,0].set_xlabel('Predicted on', fontsize=fs+4)
  1468. fig_CCAmap.subplots_adjust(right=0.90)
  1469. cbar_ax = fig_CCAmap.add_axes([1.05, 0.17, 0.035, 0.7])
  1470. norm = mpl.colors.Normalize(vmin=0, vmax=overall_max_error)
  1471. cbar = fig_CCAmap.colorbar(mpl.cm.ScalarMappable(norm=norm, cmap=pparam.ERROR_CMAP), cax=cbar_ax, orientation='vertical', fraction=0.01)
  1472. cbar.ax.set_ylabel('Prediction error (cm)', fontsize=fs+7, rotation=270, labelpad=25)
  1473. cbar.ax.tick_params(axis='both', which='major', labelsize=fs+5)
  1474. fig_CCAmap.subplots_adjust(wspace=0, hspace=0)
  1475. fig_CCAmap.tight_layout()
  1476. #Create separate colorbar
  1477. plt.figure(fig_num); fig_num += 1
  1478. fig_CCAcbar = plt.gcf()
  1479. cbar = pf.add_distance_cbar(fig_CCAcbar, pparam.ERROR_CMAP, vmin = 0, vmax = overall_max_error, fs=fs,
  1480. cbar_label = '',
  1481. cbar_kwargs = {'fraction':0.555, 'pad':0.04, 'aspect':15})
  1482. # cbar_ticks = 5 * (np.linspace(0, 5 * (overall_max_error//5), num=4)//5) #Round up to the closest multiple of 5
  1483. cbar_ticks = np.linspace(0, 5*(overall_max_error//5), num=4).astype(int)
  1484. cbar.ax.set_yticks(cbar_ticks)
  1485. cbar.ax.tick_params(axis='y', labelsize=25)
  1486. cbar.ax.set_ylabel('Prediction error (cm)', fontsize=fs+7, rotation=270, labelpad=25)
  1487. #Get prediction error by alignment and session
  1488. error_by_session = {} #(mouse_type_label, CCA_label, snum): list of errors
  1489. error_by_session = {(mtype, CCA_label, snum):[] for mtype in pparam.MOUSE_TYPE_LABELS for CCA_label in pparam.CCA_LABELS[:3] for snum in session_list}
  1490. error_across_sessions = {(mtype, CCA_label):[] for mtype in pparam.MOUSE_TYPE_LABELS for CCA_label in pparam.CCA_LABELS[:3]}
  1491. for mtype_idx, mouse_type_label in enumerate(pparam.MOUSE_TYPE_LABELS):
  1492. #Gather together mice of same type
  1493. mlist = mouse_list_by_kind[mtype_idx]
  1494. for label_idx, cca_label in enumerate(pparam.CCA_LABELS[:3]):
  1495. #To be plotted at the end
  1496. for mnum in mlist:
  1497. current_session_list = aligned_data_dict[mnum, 'session_list']
  1498. for sidx, snum in enumerate(current_session_list):
  1499. if label_idx == 0:
  1500. errors = aligned_data_dict[mnum, 'unaligned_error_array'].diagonal()
  1501. errors = [errors[sidx]]
  1502. elif label_idx == 1:
  1503. errors = aligned_data_dict[mnum, 'unaligned_error_array']
  1504. errors = errors[:, sidx] #Only when predicting on current session
  1505. non_current_session_idxs = np.where(~(np.arange(len(current_session_list))==sidx))
  1506. errors = errors[non_current_session_idxs] #Exclude self prediction
  1507. elif label_idx == 2:
  1508. errors = aligned_data_dict[mnum, 'aligned_error_array']
  1509. errors = errors[:, sidx] #Only when predicting on current session
  1510. non_current_session_idxs = np.where(~(np.arange(len(current_session_list))==sidx))
  1511. errors = errors[non_current_session_idxs] #Exclude self prediction
  1512. error_by_session[mouse_type_label, cca_label, snum].extend(list(errors))
  1513. # error_across_sessions[mouse_type_label, cca_label].extend(list(errors))
  1514. error_across_sessions[mouse_type_label, cca_label].extend([np.average(errors)])
  1515. #Plot error across sessionf for id and dd
  1516. def plot_error_across_sessions(fig_num, add_session_averages=False):
  1517. ''' This function is NOT usable outside fig2_quantification!
  1518. Used to create small variations on the same data
  1519. '''
  1520. fig_across_sessions, axs_session = plt.subplots(2, 1, num=fig_num, figsize=figsize_sessions); fig_num += 1
  1521. ax = plt.gca()
  1522. # mice_types = pparam.MOUSE_TYPE_LABELS
  1523. num_sessions = len(session_list)
  1524. for mtype_idx, mouse_type_label in enumerate(pparam.MOUSE_TYPE_LABELS):
  1525. ax = axs_session[mtype_idx]
  1526. for cca_label_idx, cca_label in enumerate(pparam.CCA_LABELS[:3]):
  1527. color = pparam.CCA_COLORS[cca_label_idx]
  1528. avg_list = []
  1529. std_list = []
  1530. for sidx, snum in enumerate(session_list):
  1531. errors = error_by_session[mouse_type_label, cca_label, snum]
  1532. avg_list.append(np.average(errors))
  1533. std_list.append(scipy.stats.sem(errors))
  1534. # ax.scatter([sidx]*len(errors), errors, color=color)
  1535. xx = np.arange(num_sessions)
  1536. avg_list = np.array(avg_list)
  1537. std_list = np.array(std_list)
  1538. ax.plot(xx, avg_list, '-', lw=3, color=color, label=cca_label)
  1539. ax.fill_between(xx, avg_list-std_list, avg_list+std_list, color=color, alpha=0.5)
  1540. if add_session_averages == True:
  1541. #Plot the across-session avg and std
  1542. errors = error_across_sessions[mouse_type_label, cca_label]
  1543. last_xpos = xx[-1] + 1 + 0.5*cca_label_idx
  1544. # #Errorbar
  1545. # err = np.std(errors)/np.sqrt(len(errors))
  1546. # # err = np.std(errors)
  1547. # ax.errorbar([last_xpos], np.average(errors), yerr = err, fmt='o', markersize=6,
  1548. # markeredgewidth=3, elinewidth = 3, zorder=2, color=color)
  1549. #Boxplot
  1550. bplot = ax.boxplot([errors], positions = [last_xpos],
  1551. showfliers=False,
  1552. vert=True,
  1553. patch_artist=True,
  1554. widths = 0.3)
  1555. for box in bplot['boxes']:
  1556. box.set_facecolor(color)
  1557. box.set_alpha(0.5)
  1558. #X axis
  1559. if mtype_idx == 1:
  1560. ax.set_xlabel('Session', fontsize=fs+4)
  1561. ax.set_xticks(np.arange(num_sessions), [pparam.SESSION_NAMES[snum] for snum in session_list])
  1562. ax.tick_params(axis='x', labelsize=fs+4)
  1563. #Y axis
  1564. ax.set_yticks(np.linspace(10 * (ax.get_ylim()[0]//10), 10 * (ax.get_ylim()[1]//10), num=3))
  1565. ax.tick_params(axis='y', labelsize=fs+6)
  1566. if mtype_idx == 0:
  1567. ax.set_ylabel('Error (cm)', fontsize=fs+6)
  1568. #Both axis
  1569. ax.spines[['right', 'top']].set_visible(False)
  1570. for axis in ['top','bottom','left','right']:
  1571. ax.spines[axis].set_linewidth(3)
  1572. if mtype_idx == 0:
  1573. #Legend
  1574. ax.legend(fontsize=fs, loc='upper right', frameon=False)
  1575. fig_across_sessions.subplots_adjust(hspace=35)
  1576. #Figure params
  1577. fig_across_sessions.tight_layout()
  1578. return fig_num, fig_across_sessions
  1579. fig_num, fig_across_sessions = plot_error_across_sessions(fig_num, add_session_averages=False)
  1580. fig_num, fig_across_sessions_boxplot = plot_error_across_sessions(fig_num, add_session_averages=True)
  1581. #Fig 2d&e: CCA histograms
  1582. fs=20
  1583. ### Plot error results
  1584. fig_hist, axs_hist = plt.subplots(2, 1, num=fig_num, figsize=hist_figsize); fig_num += 1
  1585. num_bins_cca_hist = 15
  1586. bin_min = overall_min_error
  1587. bin_max = overall_max_error
  1588. # bin_list = np.linspace(bin_min, bin_max, int((bin_max-bin_min)/2)) #Around 2 error units covered by each bin
  1589. bin_list = np.linspace(bin_min, bin_max, num_bins_cca_hist) #Around 2 error units covered by each bin
  1590. for mtype_idx, mouse_type_label in enumerate(pparam.MOUSE_TYPE_LABELS):
  1591. ax = axs_hist.ravel()[mtype_idx]
  1592. #Plot bars
  1593. for label_idx, label in enumerate(pparam.CCA_LABELS[:3]):
  1594. errors = error_by_alignment[mouse_type_label, label]
  1595. if label_idx == 0:
  1596. errors = errors * (len(session_list)-1)
  1597. ax.hist(errors, bins=bin_list, density=True, stacked=False, alpha=0.6,
  1598. label=label, color=pparam.CCA_COLORS[label_idx])
  1599. #Plot average dots
  1600. ymax = ax.get_ylim()[1] * 1.1
  1601. for label_idx, label in enumerate(pparam.CCA_LABELS[:3]):
  1602. errors = error_by_alignment[mouse_type_label, label]
  1603. #Plot average
  1604. ax.errorbar(np.average(errors), ymax, xerr = np.std(errors)/np.sqrt(len(errors)), fmt='o', markersize=6,
  1605. markeredgewidth=3, elinewidth = 3, zorder=2, color=pparam.CCA_COLORS[label_idx])
  1606. #Labels, fontsize
  1607. if mtype_idx == 0:
  1608. ax.legend(fontsize=fs-5, frameon=False)
  1609. # ax.set_ylim([0, ymax+3])
  1610. ax.set_ylabel('Normalized density', fontsize=fs)
  1611. elif mtype_idx == 1:
  1612. ax.set_xlabel('Error (cm)', fontsize=fs)
  1613. ax.set_title('%s Mice' %mouse_type_label, fontsize=fs)
  1614. ax.tick_params(axis='both', which='major', labelsize=fs)
  1615. fig_hist.tight_layout()
  1616. #Plot averages
  1617. def plot_cca_summary(fig_num,
  1618. error_by_alignment_dict,
  1619. plot_type = 'boxplot'):
  1620. ''' This function is NOT usable outside fig2_quantification!
  1621. Used to create small variations on the same data
  1622. plot_type: 'bar' or 'boxplot'
  1623. '''
  1624. fig_CCAavg = plt.figure(num=fig_num, figsize=CCAavg_figsize); fig_num += 1
  1625. ax = plt.gca()
  1626. all_bars_width = 0.5
  1627. num_of_bars = 3
  1628. barwidth = all_bars_width/(num_of_bars)
  1629. xlabels_pos_list = [0, 1]
  1630. # if not add_shuffle:
  1631. # cca_labels = pparam.CCA_LABELS[:3]
  1632. # else:
  1633. # cca_labels = pparam.CCA_LABELS[:4]
  1634. cca_labels_by_alignment = [k[1] for k in error_by_alignment_dict.keys()]
  1635. cca_labels_in_dict = [l for l in pparam.CCA_LABELS if l in cca_labels_by_alignment]
  1636. for mtype_idx, mouse_type_label in enumerate(pparam.MOUSE_TYPE_LABELS):
  1637. xpos_ref = xlabels_pos_list[mtype_idx]
  1638. errors_list = []
  1639. xpos_list = []
  1640. for label_idx, label in enumerate(cca_labels_in_dict):
  1641. errors = error_by_alignment_dict[mouse_type_label, label]
  1642. errors_list.append(errors)
  1643. if plot_type == 'bar':
  1644. xpos = xpos_ref - all_bars_width*0.5 + barwidth * (0.5+label_idx)
  1645. avg = np.average(errors)
  1646. # err = np.std(errors) ## Change also height in significance!
  1647. err = scipy.stats.sem(errors) ## Change also height in significance!
  1648. pltlabel = [None, label][mtype_idx==0]
  1649. ax.bar(xpos, avg, width=barwidth, alpha=0.7, edgecolor=None, color=pparam.CCA_COLORS[label_idx], label=pltlabel)
  1650. ax.errorbar([xpos], avg, yerr=err, fmt='', markersize=35, markeredgewidth=5, elinewidth = 5, zorder=2, color=pparam.CCA_COLORS[label_idx], alpha=0.6)
  1651. elif plot_type == 'boxplot':
  1652. xpos = xpos_ref - all_bars_width*0.5 + barwidth * (0.5+label_idx)
  1653. bplot = ax.boxplot([errors], positions = [xpos],
  1654. showfliers=False,
  1655. vert=True,
  1656. patch_artist=True,
  1657. widths = barwidth,
  1658. labels=[pparam.CCA_COLORS[label_idx]])
  1659. for box in bplot['boxes']:
  1660. box.set_facecolor(pparam.CCA_COLORS[label_idx])
  1661. box.set_alpha(0.5)
  1662. xpos_list.append(xpos)
  1663. #Add significance with unaligned
  1664. # errors_list_height = [np.average(e)+np.std(e) for e in errors_list]
  1665. errors_list_height = [np.average(e)+np.std(e)/np.sqrt(len(e)) for e in errors_list]
  1666. max_height = np.max(errors_list_height) + 3
  1667. extra_height = 0
  1668. d0 = 0.
  1669. dp = -1
  1670. label_padding = -0.025
  1671. pval_fs = 20
  1672. ref_labels = ['Unaligned', 'Aligned (shift)']
  1673. ref_idxs = [cca_labels_in_dict.index(l) for l in cca_labels_in_dict if l in ref_labels]
  1674. for ref_idx_idx, ref_idx in enumerate(ref_idxs):
  1675. ref_label = ref_labels[ref_idx_idx]
  1676. # _, pval_self = scipy.stats.ttest_ind(errors_list[0], errors_list[ref_idx], equal_var=False, permutations=None, alternative='two-sided')
  1677. # _, pval_aligned = scipy.stats.ttest_rel(errors_list[2], errors_list[ref_idx], alternative='two-sided')
  1678. tstat, pval_self = scipy.stats.mannwhitneyu(errors_list[0], errors_list[ref_idx], use_continuity=False, alternative='two-sided')
  1679. tstat, pval_aligned = scipy.stats.wilcoxon(errors_list[2], errors_list[ref_idx], zero_method='wilcox', correction=False, alternative='two-sided')
  1680. if ref_idx_idx == 0: #Ignore self to aligned (shift)
  1681. print('self to unaligned', pval_self)
  1682. pf.draw_significance(ax, pval_self, max_height + extra_height, xpos_list[0], xpos_list[ref_idx], d0, dp, orientation='top', thresholds = [0.01], fs=pval_fs, label_padding=label_padding)
  1683. extra_height += 3
  1684. print('aligned to %s'%ref_label, pval_aligned)
  1685. pf.draw_significance(ax, pval_aligned, max_height + extra_height, xpos_list[2], xpos_list[ref_idx], d0, dp, orientation='top', thresholds = [0.01], fs=pval_fs, label_padding=label_padding)
  1686. extra_height += 3
  1687. # #Significance self with aligned
  1688. # _, pval_self_aligned = scipy.stats.ttest_ind(errors_list[0], errors_list[2], equal_var=True, permutations=None, alternative='two-sided')
  1689. # pf.draw_significance(ax, pval_self_aligned, max_height*1.01 + extra_height, xpos_list[0], xpos_list[2], d0, dp, orientation='top', thresholds = [0.01], fs=pval_fs, label_padding=label_padding)
  1690. xlabels = pparam.MOUSE_TYPE_LABELS
  1691. ax.set_xticks(xlabels_pos_list, xlabels, fontsize=fs)
  1692. ax.set_ylabel('Error (cm)', fontsize=fs)
  1693. ax.tick_params(axis='y', which='major', labelsize=fs)
  1694. ax.legend(fontsize=fs-7, frameon=False)
  1695. ax.spines[['right', 'top']].set_visible(False)
  1696. for axis in ['top','bottom','left','right']:
  1697. ax.spines[axis].set_linewidth(3)
  1698. fig_CCAavg.tight_layout()
  1699. return fig_num, fig_CCAavg, ax
  1700. fig_num, fig_CCAavg, ax_CCAavg = plot_cca_summary(fig_num,
  1701. error_by_alignment,
  1702. plot_type='bar')
  1703. fig_to_save = [(fig_across_sessions, "fig2_D_prediction_across_sessions"),
  1704. (fig_across_sessions_boxplot, "fig2_D_prediction_across_sessions_boxplot"),
  1705. (fig_CCAmap, 'fig2_C_cross_prediction_errors'),
  1706. (fig_CCAcbar, 'fig2_C_cross_prediction_errors_colorbar'),
  1707. (fig_hist, 'fig2_histogram'),
  1708. (fig_CCAavg, 'fig2_F_error_avg')
  1709. ]
  1710. for fig, fig_name in fig_to_save:
  1711. save_figure(fig, fig_name)
  1712. return
  1713. #~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ FIG 3 - TCA ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
  1714. def perform_APdecoding_on_cca_param_dict(CCA_analysis_dict, ap_decoding_param_dict):
  1715. ''' Function that takes the output of CCA, takes the trial factors, and predicts the required label
  1716. See "ap_decoding_param_dict_default" in project_parameters for an explanation of each parameter
  1717. '''
  1718. LDA_components = 1
  1719. max_tca_on_lda_attempts = 100 #TCA on LDA is repeated until the required number of repetitions is reached, but in case it never converges, this will stop it
  1720. max_shuffle_LDA_attempts = 25 #Shuffle is less likely to converge when using LDA, so we might need many tries
  1721. mouse_list = CCA_analysis_dict['mouse_list']
  1722. TCA_dims_param = ap_decoding_param_dict['TCA_factors']
  1723. # #LDA parameters
  1724. LDA_imbalance_prop = ap_decoding_param_dict['LDA_imbalance_prop']
  1725. LDA_imbalance_repetitions = ap_decoding_param_dict['LDA_imbalance_repetitions']
  1726. LDA_trial_shuffles = ap_decoding_param_dict['LDA_trial_shuffles']
  1727. LDA_session_shuffles = ap_decoding_param_dict['LDA_session_shuffles']
  1728. num_bins = CCA_analysis_dict['num_bins']
  1729. APdecoding_dict = {}
  1730. for mnum in mouse_list:
  1731. session_list = CCA_analysis_dict[mnum, 'session_list']
  1732. pos_list = CCA_analysis_dict[mnum, 'pos']
  1733. pca_list = CCA_analysis_dict[mnum, 'pca']
  1734. # pca_list = CCA_analysis_dict[mnum, 'pca_unaligned']
  1735. data_by_trial, pos_by_trial, snum_by_trial = pf.reshape_pca_list_by_trial(pca_list, pos_list, num_bins, session_list)
  1736. num_features, num_bins, total_trials = data_by_trial.shape
  1737. num_CCA_dims = num_features
  1738. if TCA_dims_param == 'max':
  1739. num_TCA_dims = num_CCA_dims
  1740. else:
  1741. num_TCA_dims = int(TCA_dims_param)
  1742. print('Performing TCA+LDA on M%d // CCA dim: %d, // TCA dim: %d' %(mnum, num_CCA_dims, num_TCA_dims))
  1743. #Limit position (if indicated)
  1744. if ap_decoding_param_dict['exclude_positions'] == True:
  1745. positions = pos_by_trial[:,0] #Assumes all trials are binned using the same positions
  1746. pos_bool_filtered_out = pf.get_idxs_in_periodic_interval(positions, ap_decoding_param_dict['pos_to_exclude_from'],
  1747. ap_decoding_param_dict['pos_to_exclude_to'], pparam.MAX_POS)
  1748. pos_bool_selected = np.invert(pos_bool_filtered_out)
  1749. data_by_trial = data_by_trial[:, pos_bool_selected, :]
  1750. pos_by_trial = pos_by_trial[pos_bool_selected, :]
  1751. num_bins_current = data_by_trial.shape[1]
  1752. #Selecting trials to decode
  1753. trials_to_keep, label_by_trial = APfuns.get_trials_to_keep_and_labels(snum_by_trial, ap_decoding_param_dict['session_comparisons'])
  1754. num_trials_to_keep = len(label_by_trial)
  1755. f1_best = -1
  1756. f1_array = np.zeros((0, 2)) #1st axis is TCA repetition, 2nd is class 0 or 1
  1757. # accuracy_array = []
  1758. accuracy_array = np.zeros((0, 2)) #1st axis is TCA repetition, 2nd is class 0 or 1
  1759. LDA_prob_array = np.zeros((0, num_trials_to_keep)) #1st axis is TCA repetition, 2nd is trial
  1760. LDA_projection_array = np.zeros((0, num_trials_to_keep)) #1st axis is TCA repetition, 2nd is trial
  1761. label_by_trial_predicted = np.zeros((0, num_trials_to_keep), dtype=int) #1st axis is TCA repetition, 2nd is trial
  1762. feature_factors = np.zeros((0, num_CCA_dims, num_TCA_dims)) #1st axis is TCA repetition, 2nd is latent input dimension, 3rd is TCA dim
  1763. time_factors = np.zeros((0, num_bins_current, num_TCA_dims)) #1st axis is TCA repetitions, 2nd is time bin, 3rd is TCA dim
  1764. trial_factors = np.zeros((0, num_trials_to_keep, num_TCA_dims)) #1st axis is TCA repetition, 2nd is trial dim, 3rd is TCA dim
  1765. APdecoding_weights = np.zeros((0, data_by_trial.shape[0])) #1st axis is TCA repetition, 2nd is input data feature. Gets the weight of each input dimension for AP decoding
  1766. TCA_on_LDA_counter = 0
  1767. while TCA_on_LDA_counter < ap_decoding_param_dict['TCA_on_LDA_repetitions'] and TCA_on_LDA_counter < max_tca_on_lda_attempts:
  1768. # Step 4: TCA
  1769. KTensor = APfuns.perform_TCA(data_by_trial, num_TCA_dims, ap_decoding_param_dict['TCA_replicates'],
  1770. ap_decoding_param_dict['TCA_method'], ap_decoding_param_dict['TCA_convergence_attempts'])
  1771. feature_factors_temp, time_factors_temp, trial_factors_temp = KTensor
  1772. LDA_input = trial_factors_temp[trials_to_keep]
  1773. #LDA on TCA factors
  1774. try:
  1775. LDA_results_dict = APfuns.perform_LDA(LDA_input, label_by_trial, LDA_components, LDA_imbalance_prop, LDA_imbalance_repetitions)
  1776. except np.linalg.LinAlgError:
  1777. continue
  1778. #Update arrays
  1779. f1 = LDA_results_dict['f1']
  1780. # print(f1)
  1781. f1_array = np.vstack((f1_array, f1))
  1782. print(f1_array)
  1783. # accuracy_array.append([LDA_results_dict['accuracy']])
  1784. accuracy_array = np.vstack((accuracy_array, LDA_results_dict['accuracy']))
  1785. LDA_prob_array = np.vstack((LDA_prob_array, LDA_results_dict['LDA_prob']))
  1786. LDA_projection_array = np.vstack((LDA_projection_array, LDA_results_dict['LDA_projection'].squeeze()))
  1787. label_by_trial_predicted = np.vstack((label_by_trial_predicted, LDA_results_dict['label_predicted']))
  1788. feature_factors = np.vstack((feature_factors, feature_factors_temp[np.newaxis]))
  1789. time_factors = np.vstack((time_factors, time_factors_temp[np.newaxis]))
  1790. trial_factors = np.vstack((trial_factors, LDA_input[np.newaxis]))
  1791. #Compute dimensional weights
  1792. weights = np.abs(LDA_results_dict['weights'].reshape(num_TCA_dims, 1))
  1793. # feature_factors = KTensor[0] #TCA weights for each data dimension
  1794. APdecoding_weights_rep = np.dot(feature_factors_temp, weights)
  1795. APdecoding_weights = np.vstack((APdecoding_weights, np.squeeze(APdecoding_weights_rep)))
  1796. if np.min(f1) > np.min(f1_best): #Get the best one
  1797. best_LDA_input = LDA_input #Used as reference for shuffle
  1798. f1_best = f1
  1799. TCA_on_LDA_counter += 1
  1800. if f1_array.shape[0] == 0:
  1801. print('WARNING: LDA ON TCA FAILED TO CONVERGE AFTER %d TRIES'%max_tca_on_lda_attempts)
  1802. #Random shuffle (done on best TCA run)
  1803. f1_array_shuffle = np.zeros((LDA_trial_shuffles,2))
  1804. accuracy_array_shuffle = np.zeros((LDA_trial_shuffles, 2))
  1805. LDA_prob_shuffle = np.zeros((LDA_trial_shuffles))
  1806. LDA_projection_array_shuffle = np.zeros((LDA_trial_shuffles, num_trials_to_keep)) #1st axis is TCA repetition, 2nd is trial
  1807. num_trials = len(label_by_trial)
  1808. if LDA_trial_shuffles != 0 or LDA_session_shuffles != 0:
  1809. print('Starting shuffle')
  1810. for randidx in range(LDA_trial_shuffles):
  1811. shuffle_attempt_counter = 0
  1812. while shuffle_attempt_counter < max_shuffle_LDA_attempts: #We repeat until it converges
  1813. shuffle_attempt_counter += 1
  1814. shuffled_idxs = np.random.choice(np.arange(num_trials), size=num_trials, replace=False)
  1815. label_by_trial_shuffled = label_by_trial[shuffled_idxs]
  1816. try:
  1817. LDA_results_dict_shuffle = APfuns.perform_LDA(best_LDA_input, label_by_trial_shuffled, LDA_components, LDA_imbalance_prop, LDA_imbalance_repetitions)
  1818. except np.linalg.LinAlgError:
  1819. continue
  1820. f1_array_shuffle[randidx] = LDA_results_dict_shuffle['f1']
  1821. accuracy_array_shuffle[randidx] = LDA_results_dict_shuffle['accuracy']
  1822. LDA_prob_shuffle[randidx] = np.average(LDA_results_dict_shuffle['LDA_prob'])
  1823. LDA_projection_array_shuffle = np.vstack((LDA_projection_array_shuffle, LDA_results_dict_shuffle['LDA_projection'].squeeze()))
  1824. break
  1825. else:
  1826. print('we are in trouble! random shuffle max attempts reached!')
  1827. #SESSION SHUFFLE (NOT USED)
  1828. f1_array_session_shuffle = np.zeros((LDA_session_shuffles,2))
  1829. LDA_prob_session_shuffle = np.zeros((LDA_session_shuffles))
  1830. LDA_projection_array_session_shuffle = np.zeros((LDA_trial_shuffles, num_trials_to_keep)) #1st axis is TCA repetition, 2nd is trial
  1831. #Step 1: which sessions were selected?
  1832. # print(trials_to_keep)
  1833. snum_by_trial_selected = snum_by_trial[trials_to_keep]
  1834. selected_sessions = np.unique(snum_by_trial_selected)
  1835. num_selected_sessions = len(selected_sessions)
  1836. #Step 2: get label assigned to each session
  1837. first_trial_per_snum = {snum:np.argmax(snum_by_trial_selected==snum) for snum in selected_sessions}
  1838. label_list = np.array([label_by_trial[first_trial_per_snum[snum]] for snum in selected_sessions])
  1839. for randidx in range(LDA_session_shuffles):
  1840. session_shuffle_attempt_counter = 0
  1841. while session_shuffle_attempt_counter < max_shuffle_LDA_attempts:
  1842. session_shuffle_attempt_counter += 1
  1843. # #Shuffle indexes at the session level, making sure they are different from original
  1844. while True:
  1845. shuffled_session_idxs = np.random.choice(np.arange(num_selected_sessions), size=num_selected_sessions, replace=False)
  1846. label_list_shuffled = label_list[shuffled_session_idxs]
  1847. label_by_trial_session_shuffled = np.array([label_list_shuffled[list(selected_sessions).index(snum)] for snum in snum_by_trial_selected])
  1848. if np.allclose(label_list, label_list_shuffled) == False: #All labels must not be equal to the original
  1849. break
  1850. try:
  1851. LDA_results_dict_session_shuffle = APfuns.perform_LDA(best_LDA_input, label_by_trial_session_shuffled, LDA_components, LDA_imbalance_prop, LDA_imbalance_repetitions)
  1852. except np.linalg.LinAlgError:
  1853. continue
  1854. f1_array_session_shuffle[randidx] = LDA_results_dict_session_shuffle['f1']
  1855. LDA_prob_session_shuffle[randidx] = np.average(LDA_results_dict_session_shuffle['LDA_prob'])
  1856. LDA_projection_array_session_shuffle = np.vstack((LDA_projection_array_session_shuffle, LDA_results_dict_session_shuffle['LDA_projection'].squeeze()))
  1857. break
  1858. else:
  1859. print('we are in trouble! random session shuffle max attempts reached!')
  1860. snum_by_trial = snum_by_trial[trials_to_keep]
  1861. #Objective dictionary
  1862. APdecoding_dict.update({
  1863. (mnum, 'snum_by_trial'):snum_by_trial,
  1864. (mnum, 'selected_trials'):trials_to_keep,
  1865. (mnum, 'label_by_trial'):label_by_trial,
  1866. (mnum, 'label_by_trial_predicted'):label_by_trial_predicted,
  1867. (mnum, 'LDA_projection_array'):LDA_projection_array,
  1868. (mnum, 'feature_factors'):feature_factors,
  1869. (mnum, 'time_factors'):time_factors,
  1870. (mnum, 'trial_factors'):trial_factors,
  1871. (mnum, 'APdecoding_weights'):APdecoding_weights,
  1872. (mnum, 'LDA_prob_correct'):LDA_prob_array,
  1873. (mnum, 'f1_array'):f1_array,
  1874. (mnum, 'accuracy_array'):np.array(accuracy_array),
  1875. (mnum, 'LDA_prob_correct_shuffle'):LDA_prob_shuffle,
  1876. (mnum, 'f1_array_shuffle'):f1_array_shuffle,
  1877. (mnum, 'accuracy_array_shuffle'):accuracy_array_shuffle,
  1878. (mnum, 'LDA_prob_correct_session_shuffle'):LDA_prob_session_shuffle,
  1879. (mnum, 'f1_array_session_shuffle'):f1_array_session_shuffle,
  1880. })
  1881. return APdecoding_dict
  1882. def APdecoding_pipeline(preprocessing_param_dict, cca_param_dict, ap_decoding_param_dict, force_recalculation=False):
  1883. ''' Given a parameter dict, peform PCA, CCA, and TCA+LDA on it '''
  1884. ########## STEP 1 - PCA ###########
  1885. PCA_analysis_dict = perform_pca_on_multiple_mice_param_dict(preprocessing_param_dict)
  1886. ############## STEP 2: mCCA ############
  1887. CCA_analysis_dict = perform_mCCA_on_pca_dict_param_dict(PCA_analysis_dict, cca_param_dict)
  1888. ############# STEP 3: TCA + LDA ############
  1889. APdecoding_dict = perform_APdecoding_on_cca_param_dict(CCA_analysis_dict, ap_decoding_param_dict)
  1890. pipeline_output_dict = {
  1891. 'PCA_analysis_dict':PCA_analysis_dict,
  1892. 'CCA_analysis_dict':CCA_analysis_dict,
  1893. 'APdecoding_dict':APdecoding_dict,
  1894. }
  1895. np.save(OUTPUT_PATH + "pipeline_output_dict.npy", pipeline_output_dict, allow_pickle=True)
  1896. return pipeline_output_dict
  1897. def fig3_A_and_fig3SI_A_TCA_factors():
  1898. '''
  1899. Plot LDA projection from TCA factors
  1900. '''
  1901. mouse_list = [6]
  1902. preprocessing_param_dict = {
  1903. #session params
  1904. 'mouse_list':mouse_list,
  1905. 'session_list':np.arange(7),
  1906. #Preprocessing parameters
  1907. 'time_bin_size':1,
  1908. 'distance_bin_size':1,
  1909. 'gaussian_size':25,
  1910. 'data_used':'amplitudes',
  1911. 'running':True,
  1912. 'eliminate_v_zeros':True,
  1913. 'num_components':'all'
  1914. }
  1915. cca_param_dict = {
  1916. 'CCA_dim':'.9', #11, '.9'
  1917. 'return_warped_data':True,
  1918. 'return_trimmed_data':False,
  1919. 'sessions_to_align':'all',
  1920. 'shuffle':False
  1921. }
  1922. ap_decoding_param_dict = {
  1923. 'exclude_positions':False,
  1924. 'pos_to_exclude_from':200,
  1925. 'pos_to_exclude_to':1300,
  1926. ## TCA params ##
  1927. 'TCA_method': "ncp_hals", #"cp_als", "mcp_als", "ncp_bcd", "ncp_hals"
  1928. 'TCA_factors':'max', #int, or 'max' to get the maximum possible (determined by CCA)
  1929. 'TCA_replicates':10,
  1930. 'TCA_convergence_attempts':10, #Number of times TCA can fail before giving up
  1931. 'TCA_on_LDA_repetitions':2,
  1932. ## LDA params ##
  1933. 'LDA_imbalance_prop':.51,
  1934. 'LDA_imbalance_repetitions':10,
  1935. 'LDA_trial_shuffles':0,
  1936. 'LDA_session_shuffles':0,
  1937. 'session_comparisons':'BT' #'airpuff', 'BT', 'TP', 'BP'
  1938. }
  1939. fig_num = plt.gcf().number + 1
  1940. fs = 25
  1941. fs_title = 35
  1942. lw = 5
  1943. trial_smoothing_size = 3 #trial smoothing in "trial" number
  1944. ########## STEP 1 - PCA ###########
  1945. PCA_analysis_dict = perform_pca_on_multiple_mice_param_dict(preprocessing_param_dict)
  1946. ############## STEP 2: mCCA ############
  1947. CCA_analysis_dict = perform_mCCA_on_pca_dict_param_dict(PCA_analysis_dict, cca_param_dict)
  1948. # print('WARNING: IGNORING CCA IN THE AP DECODING PIPELINE AS A TEST, REVERT THIS CHANGE')
  1949. # CCA_analysis_dict = PCA_analysis_dict
  1950. ############# STEP 3: TCA + LDA ############
  1951. mouse_list = CCA_analysis_dict['mouse_list']
  1952. TCA_dims_param = ap_decoding_param_dict['TCA_factors']
  1953. # #LDA parameters
  1954. num_bins = CCA_analysis_dict['num_bins']
  1955. for mnum in mouse_list:
  1956. session_list = CCA_analysis_dict[mnum, 'session_list']
  1957. pos_list = CCA_analysis_dict[mnum, 'pos']
  1958. pca_list = CCA_analysis_dict[mnum, 'pca']
  1959. # pca_list = CCA_analysis_dict[mnum, 'pca_unaligned']
  1960. data_by_trial, pos_by_trial, snum_by_trial = pf.reshape_pca_list_by_trial(pca_list, pos_list, num_bins, session_list)
  1961. num_features, num_bins, total_trials = data_by_trial.shape
  1962. num_CCA_dims = num_features
  1963. if TCA_dims_param == 'max':
  1964. num_TCA_dims = num_CCA_dims
  1965. else:
  1966. num_TCA_dims = int(TCA_dims_param)
  1967. print('Performing TCA+LDA on M%d // CCA dim: %d, // TCA dim: %d' %(mnum, num_CCA_dims, num_TCA_dims))
  1968. #Limit position (if indicated)
  1969. if ap_decoding_param_dict['exclude_positions'] == True:
  1970. positions = pos_by_trial[:,0] #Assumes all trials are binned using the same positions
  1971. pos_bool_filtered_out = pf.get_idxs_in_periodic_interval(positions, ap_decoding_param_dict['pos_to_exclude_from'],
  1972. ap_decoding_param_dict['pos_to_exclude_to'], pparam.MAX_POS)
  1973. pos_bool_selected = np.invert(pos_bool_filtered_out)
  1974. data_by_trial = data_by_trial[:, pos_bool_selected, :]
  1975. pos_by_trial = pos_by_trial[pos_bool_selected, :]
  1976. #Selecting trials to decode
  1977. trials_to_keep, label_by_trial = APfuns.get_trials_to_keep_and_labels(snum_by_trial, ap_decoding_param_dict['session_comparisons'])
  1978. # snum_by_trial = snum_by_trial[trials_to_keep]
  1979. # data_by_trial = data_by_trial[:, :, trials_to_keep]
  1980. plotted_factors = False
  1981. while plotted_factors == False:
  1982. # Step 4: TCA
  1983. KTensor, TCA_ensemble = APfuns.perform_TCA(data_by_trial, num_TCA_dims, ap_decoding_param_dict['TCA_replicates'],
  1984. ap_decoding_param_dict['TCA_method'], ap_decoding_param_dict['TCA_convergence_attempts'],
  1985. return_ensemble = True)
  1986. feature_factors, time_factors, trial_factors = KTensor
  1987. plotted_factors = True
  1988. ''' Plots TCA results for trials of an aversive task following Negar's experimental design.
  1989. Assumes the trials are concatenated across sessions, computed for PCA.
  1990. TCA_ensemble: output from tensortools' TCA method
  1991. session_list: list of session numbers used
  1992. num_trials_by_snum: number of trials per session number
  1993. '''
  1994. TCA_dim = feature_factors.shape[1]
  1995. tca_factors_figsize = (3 * 4, TCA_dim*2)
  1996. ncols = 3
  1997. nrows = feature_factors.shape[0]
  1998. fig, axs = plt.subplots(nrows=nrows, ncols=ncols, squeeze=False, figsize=tca_factors_figsize, num=fig_num); fig_num += 1
  1999. for factor in range(TCA_dim):
  2000. for factor_type in range(3):
  2001. ax = axs[factor, factor_type]
  2002. vals = KTensor[factor_type][:, factor]
  2003. xx = np.arange(0, len(vals))
  2004. if factor_type == 2:
  2005. vals = np.convolve(vals, np.ones(trial_smoothing_size), mode='same')/trial_smoothing_size
  2006. color = 'black' #'tab:blue'
  2007. ax.plot(xx, vals, color=color, lw=lw)
  2008. minval, maxval = np.min(vals), np.max(vals)
  2009. ax.set_ylim([minval, maxval])
  2010. ax.set_yticks([minval, maxval], [np.around(minval, decimals=2), np.around(maxval, decimals=2)], fontsize=fs)
  2011. if factor_type == 0: #PCA visualizations
  2012. pca_dim_list = range(vals.shape[0])
  2013. ax.set_xticks(pca_dim_list, pca_dim_list, fontsize=fs)
  2014. # ax.set_yticks([minval, maxval], [np.around(minval, decimals=2), np.around(maxval, decimals=2)], fontsize=fs)
  2015. for pca_dim_idx in pca_dim_list:
  2016. yy = np.linspace(minval, maxval)
  2017. xx = [pca_dim_idx] * len(yy)
  2018. ax.plot(xx, yy, '--', color='gray', alpha=0.5)
  2019. ax.xaxis.set_major_locator(MaxNLocator(integer=True))
  2020. if factor == TCA_dim-1:
  2021. ax.set_xlabel('PCA dim', fontsize=fs_title)
  2022. if factor == 0:
  2023. ax.set_title('TCA factors (AU)', fontsize=fs_title)
  2024. if factor_type == 1: #Time/position factor visualization
  2025. yy = np.linspace(minval, maxval)
  2026. pos_landmarks = [0, 50, 100, 149]
  2027. pos_landmarks_names = [0, 500, 1000, 1500]
  2028. colors_landmarks = ['black', 'indianred', 'forestgreen', 'black']
  2029. ax.set_xticks(pos_landmarks, pos_landmarks_names, fontsize=fs)
  2030. ax.xaxis.set_major_locator(MaxNLocator(integer=True))
  2031. for landmark in range(len(pos_landmarks)):
  2032. xx = [pos_landmarks[landmark]] * len(yy)
  2033. ax.plot(xx, yy, '--', color=colors_landmarks[landmark], lw=lw, alpha=0.6)
  2034. if factor == TCA_dim-1:
  2035. ax.set_xlabel('Position (mm)', fontsize=fs_title)
  2036. if factor_type == 2: #Trial factors
  2037. APfuns.add_session_delimiters_to_trial_plot(ax, snum_by_trial, fs=fs)
  2038. if factor == TCA_dim-1:
  2039. ax.set_xlabel('Trial number', fontsize=fs_title)
  2040. ax.tick_params(axis='x', which='major', labelsize=fs)
  2041. ax.tick_params(axis='y', which='major', labelsize=fs)
  2042. ax.spines[['right', 'top']].set_visible(False)
  2043. for axis in ['top','bottom','left','right']:
  2044. ax.spines[axis].set_linewidth(lw)
  2045. fig.tight_layout()
  2046. fig_name = 'fig3_A_TCA_factors_M%d'%mnum
  2047. save_figure(fig, fig_name)
  2048. return
  2049. def fig3_B_LDA_on_TCA():
  2050. '''
  2051. Plot LDA projection from TCA factors
  2052. '''
  2053. mouse_list = [2,6]
  2054. preprocessing_param_dict = {
  2055. #session params
  2056. 'mouse_list':mouse_list,
  2057. 'session_list':np.arange(7),
  2058. #Preprocessing parameters
  2059. 'time_bin_size':1,
  2060. 'distance_bin_size':1,
  2061. 'gaussian_size':25,
  2062. 'data_used':'amplitudes',
  2063. 'running':True,
  2064. 'eliminate_v_zeros':True,
  2065. 'num_components':'all'
  2066. }
  2067. cca_param_dict = {
  2068. 'CCA_dim':'.9', #11, '.9'
  2069. 'return_warped_data':True,
  2070. 'return_trimmed_data':False,
  2071. 'sessions_to_align':'all',
  2072. 'shuffle':False
  2073. }
  2074. ap_decoding_param_dict = {
  2075. 'exclude_positions':False,
  2076. 'pos_to_exclude_from':200,
  2077. 'pos_to_exclude_to':1300,
  2078. ## TCA params ##
  2079. 'TCA_method': "ncp_hals", #"cp_als", "mcp_als", "ncp_bcd", "ncp_hals"
  2080. 'TCA_factors':'max', #int, or 'max' to get the maximum possible (determined by CCA)
  2081. 'TCA_replicates':10,
  2082. 'TCA_convergence_attempts':10, #Number of times TCA can fail before giving up
  2083. 'TCA_on_LDA_repetitions':20,
  2084. ## LDA params ##
  2085. 'LDA_imbalance_prop':.51,
  2086. 'LDA_imbalance_repetitions':10,
  2087. 'LDA_trial_shuffles':0,
  2088. 'LDA_session_shuffles':0,
  2089. 'session_comparisons':'BT' #'airpuff', 'BT', 'TP', 'BP'
  2090. }
  2091. ########## STEP 1 - PCA ###########
  2092. PCA_analysis_dict = perform_pca_on_multiple_mice_param_dict(preprocessing_param_dict)
  2093. ############## STEP 2: mCCA ############
  2094. CCA_analysis_dict = perform_mCCA_on_pca_dict_param_dict(PCA_analysis_dict, cca_param_dict)
  2095. ############# STEP 3: TCA + LDA ############
  2096. LDA_components = 1
  2097. mouse_list = CCA_analysis_dict['mouse_list']
  2098. TCA_dims_param = ap_decoding_param_dict['TCA_factors']
  2099. # #LDA parameters
  2100. num_bins = CCA_analysis_dict['num_bins']
  2101. for mnum in mouse_list:
  2102. session_list = CCA_analysis_dict[mnum, 'session_list']
  2103. pos_list = CCA_analysis_dict[mnum, 'pos']
  2104. pca_list = CCA_analysis_dict[mnum, 'pca']
  2105. # pca_list = CCA_analysis_dict[mnum, 'pca_unaligned']
  2106. data_by_trial, pos_by_trial, snum_by_trial = pf.reshape_pca_list_by_trial(pca_list, pos_list, num_bins, session_list)
  2107. num_features, num_bins, total_trials = data_by_trial.shape
  2108. num_CCA_dims = num_features
  2109. if TCA_dims_param == 'max':
  2110. num_TCA_dims = num_CCA_dims
  2111. else:
  2112. num_TCA_dims = int(TCA_dims_param)
  2113. print('Performing TCA+LDA on M%d // CCA dim: %d, // TCA dim: %d' %(mnum, num_CCA_dims, num_TCA_dims))
  2114. #Limit position (if indicated)
  2115. if ap_decoding_param_dict['exclude_positions'] == True:
  2116. positions = pos_by_trial[:,0] #Assumes all trials are binned using the same positions
  2117. pos_bool_filtered_out = pf.get_idxs_in_periodic_interval(positions, ap_decoding_param_dict['pos_to_exclude_from'],
  2118. ap_decoding_param_dict['pos_to_exclude_to'], pparam.MAX_POS)
  2119. pos_bool_selected = np.invert(pos_bool_filtered_out)
  2120. data_by_trial = data_by_trial[:, pos_bool_selected, :]
  2121. pos_by_trial = pos_by_trial[pos_bool_selected, :]
  2122. #Selecting trials to decode
  2123. trials_to_keep, label_by_trial = APfuns.get_trials_to_keep_and_labels(snum_by_trial, ap_decoding_param_dict['session_comparisons'])
  2124. num_trials_to_keep = len(label_by_trial)
  2125. label_by_trial_predicted = np.zeros((0, num_trials_to_keep), dtype=int) #1st axis is TCA repetition, 2nd is trial
  2126. plotted_factors = False
  2127. while plotted_factors == False:
  2128. # Step 4: TCA
  2129. KTensor = APfuns.perform_TCA(data_by_trial, num_TCA_dims, ap_decoding_param_dict['TCA_replicates'],
  2130. ap_decoding_param_dict['TCA_method'], ap_decoding_param_dict['TCA_convergence_attempts'])
  2131. feature_factors_temp, time_factors_temp, trial_factors_temp = KTensor
  2132. LDA_input = trial_factors_temp[trials_to_keep]
  2133. #LDA on TCA factors
  2134. # LDA_results_dict = APfuns.perform_LDA(LDA_input, label_by_trial, LDA_components, LDA_imbalance_prop, LDA_imbalance_repetitions)
  2135. LDA = LinearDiscriminantAnalysis(solver="eigen", #svd, lsqr, eigen
  2136. shrinkage=0, #None, "auto", float 0-1
  2137. n_components=LDA_components, #Dimensionality reduction
  2138. store_covariance=False #Only useful for svd, which doesn't automatically calculate it
  2139. )
  2140. try:
  2141. LDA.fit(LDA_input, label_by_trial)
  2142. except np.linalg.LinAlgError:
  2143. continue
  2144. label_by_trial_predicted = (LDA.predict(LDA_input))
  2145. f1 = np.average(pf.multiclass_f1(label_by_trial, label_by_trial_predicted))
  2146. trial_factors_LDA_projection = np.squeeze(LDA.transform(LDA_input))
  2147. fig, ax = APfuns.plot_LDA_projection(trial_factors_LDA_projection, label_by_trial, snum_by_trial=snum_by_trial, plot_legend=True, ax=None)
  2148. ax.set_title('M%d, F1 = %.1f'%(mnum, f1), fontsize=15)
  2149. fig_name = 'fig3_B_LDA_example_M%d'%mnum
  2150. save_figure(fig, fig_name)
  2151. plotted_factors = True
  2152. return
  2153. def fig3_C_D_f1_plots_and_fig3SI_B_C_accuracy_plots():
  2154. ''' Performs TCA across sessions for an animal.
  2155. Step 1: perform PCA, limit to minimum possible of dimensions
  2156. Step 2: align through CCA, so every dimension represents something similar about the data
  2157. Step 3: split into trials through warping
  2158. Step 4: TCA!
  2159. '''
  2160. mouse_list = np.arange(8)
  2161. # mouse_list = [6]
  2162. # mouse_list = [2,6]
  2163. preprocessing_param_dict = {
  2164. #session params
  2165. 'mouse_list':mouse_list,
  2166. 'session_list':np.arange(9),
  2167. #Preprocessing parameters
  2168. 'time_bin_size':1,
  2169. 'distance_bin_size':1,
  2170. 'gaussian_size':25,
  2171. 'data_used':'amplitudes',
  2172. 'running':True,
  2173. 'eliminate_v_zeros':True,
  2174. 'num_components':'all'
  2175. }
  2176. cca_param_dict = {
  2177. 'CCA_dim':'.9', #11, '.9'
  2178. 'return_warped_data':True,
  2179. 'return_trimmed_data':False,
  2180. 'sessions_to_align':'all',
  2181. 'shuffle':False,
  2182. 'skip_alignment':True
  2183. }
  2184. ap_decoding_param_dict = {
  2185. 'exclude_positions':False,
  2186. 'pos_to_exclude_from':200,
  2187. 'pos_to_exclude_to':1300,
  2188. ## TCA params ##
  2189. 'TCA_method': "ncp_hals", #"cp_als", "mcp_als", "ncp_bcd", "ncp_hals"
  2190. 'TCA_factors':'max', #int, or 'max' to get the maximum possible (determined by CCA)
  2191. 'TCA_replicates':10,
  2192. 'TCA_convergence_attempts':10, #Number of times TCA can fail before giving up
  2193. 'TCA_on_LDA_repetitions':25,
  2194. ## LDA params ##
  2195. 'LDA_imbalance_prop':.51,
  2196. 'LDA_imbalance_repetitions':10,
  2197. 'LDA_trial_shuffles':25,
  2198. 'LDA_session_shuffles':0,
  2199. 'session_comparisons':'BT' #'airpuff', 'BT', 'TP', 'BP'
  2200. }
  2201. CCA_random_shifts = 25 #Number of random shifts WARNING: LDA PROB NOT CALCULATED FOR IT!!!
  2202. if cca_param_dict['skip_alignment'] == True:
  2203. print('WARNING: SKIPPING ALIGNMENT DURING AP DECODING!')
  2204. print('STARTING')
  2205. print('Sessions: %d'%len(preprocessing_param_dict['session_list']))
  2206. print('V zeros: %s'%(str(preprocessing_param_dict['eliminate_v_zeros'])))
  2207. print('CCA dim: %s'%str(cca_param_dict['CCA_dim']))
  2208. print('TCA dim: %s'%str(ap_decoding_param_dict['TCA_factors']))
  2209. #Do analysis
  2210. pipeline_output_dict = APdecoding_pipeline(preprocessing_param_dict, cca_param_dict, ap_decoding_param_dict)
  2211. if CCA_random_shifts > 0:
  2212. cca_param_dict_shift = {k:v for k,v in cca_param_dict.items()}
  2213. ap_decoding_param_dict_shift = {k:v for k,v in ap_decoding_param_dict.items()}
  2214. cca_param_dict_shift['shuffle'] = True
  2215. cca_param_dict_shift['skip_alignment'] = False
  2216. ap_decoding_param_dict_shift['TCA_on_LDA_repetitions'] = 5
  2217. ap_decoding_param_dict_shift['LDA_trial_shuffles'] = 0
  2218. ap_decoding_param_dict_shift['LDA_session_shuffles'] = 0
  2219. f1_array_shift = np.zeros((len(mouse_list), CCA_random_shifts, 2))
  2220. accuracy_array_shift = np.zeros((len(mouse_list), CCA_random_shifts, 2))
  2221. for shift_run in range(CCA_random_shifts):
  2222. pipeline_output_dict_shift = APdecoding_pipeline(preprocessing_param_dict, cca_param_dict_shift, ap_decoding_param_dict_shift)
  2223. AP_decoding_shift_current = pipeline_output_dict_shift['APdecoding_dict']
  2224. for midx, mnum in enumerate(mouse_list):
  2225. f1 = AP_decoding_shift_current[mnum, 'f1_array']
  2226. f1_array_shift[midx, shift_run] = np.average(f1, axis=0)
  2227. acc = AP_decoding_shift_current[mnum, 'accuracy_array']
  2228. accuracy_array_shift[midx, shift_run] = np.average(acc, axis=0)
  2229. #PUTTING DATA IN ARRAYS
  2230. mtype_by_mouse = np.array([pparam.MOUSE_TYPE_LABEL_BY_MOUSE[mnum] for mnum in mouse_list])
  2231. num_mice = len(mouse_list)
  2232. LDA_trial_shuffles = ap_decoding_param_dict['LDA_trial_shuffles']
  2233. LDA_session_shuffles = ap_decoding_param_dict['LDA_session_shuffles']
  2234. APdecoding_dict = pipeline_output_dict['APdecoding_dict']
  2235. best_lda_by_mouse = np.zeros(len(mouse_list), dtype=int)
  2236. best_f1_by_mouse = np.zeros((len(mouse_list), 2))
  2237. f1_avg_by_mouse = np.zeros((len(mouse_list), 2))
  2238. f1_std_by_mouse = np.zeros((len(mouse_list), 2))
  2239. accuracy_avg_by_mouse = np.zeros((len(mouse_list), 2))
  2240. accuracy_std_by_mouse = np.zeros((len(mouse_list), 2))
  2241. for midx, mnum in enumerate(mouse_list):
  2242. f1_array = APdecoding_dict[mnum, 'f1_array']
  2243. min_f1 = np.min(f1_array, axis=1)
  2244. best_lda_idx = np.argmax(min_f1)
  2245. best_lda_by_mouse[midx] = best_lda_idx
  2246. best_f1 = f1_array[best_lda_idx]
  2247. best_f1_by_mouse[midx] = best_f1
  2248. #Avg F1
  2249. f1_avg_by_mouse[midx] = np.average(f1_array, axis=0)
  2250. f1_std_by_mouse[midx] = np.std(f1_array, axis=0)
  2251. # f1_std_by_mouse[midx] = np.std(f1_array, axis=0)/np.sqrt(f1_array.shape[0])
  2252. #Accuracy
  2253. acc = APdecoding_dict[mnum, 'accuracy_array']
  2254. accuracy_avg_by_mouse[midx] = np.average(acc, axis=0)
  2255. accuracy_std_by_mouse[midx] = np.std(acc, axis=0)
  2256. #Take the indicated shuffles
  2257. control_names = []
  2258. control_label_idxs = []
  2259. if LDA_trial_shuffles > 0:
  2260. control_names.append('shuffle')
  2261. control_label_idxs.append(2)
  2262. if LDA_session_shuffles > 0:
  2263. control_names.append('session_shuffle')
  2264. control_label_idxs.append(3)
  2265. if CCA_random_shifts > 0:
  2266. control_names.append('CCA_shift')
  2267. control_label_idxs.append(4)
  2268. control_data = {}
  2269. for control in control_names:
  2270. f1_avg_control_by_mouse = np.zeros((len(mouse_list), 2))
  2271. f1_std_control_by_mouse = np.zeros((len(mouse_list), 2))
  2272. accuracy_avg_control_by_mouse = np.zeros((len(mouse_list), 2))
  2273. accuracy_std_control_by_mouse = np.zeros((len(mouse_list), 2))
  2274. for midx, mnum in enumerate(mouse_list):
  2275. if 'shuffle' in control:
  2276. f1_control = APdecoding_dict[mnum, 'f1_array_' + control]
  2277. accuracy_control = APdecoding_dict[mnum, 'accuracy_array_' + control]
  2278. elif 'CCA' in control:
  2279. f1_control = f1_array_shift[midx]
  2280. accuracy_control = accuracy_array_shift[midx]
  2281. f1_avg_control_by_mouse[midx] = np.average(f1_control, axis=0)
  2282. f1_std_control_by_mouse[midx] = np.std(f1_control, axis=0)
  2283. # f1_std_control_by_mouse[midx] = np.std(f1_control, axis=0)/np.sqrt(f1_control.shape[0])
  2284. # f1_std_control_by_mouse[midx] = np.average(np.std(f1_control, axis=0))
  2285. accuracy_avg_control_by_mouse[midx] = np.average(accuracy_control, axis=0)
  2286. accuracy_std_control_by_mouse[midx] = np.std(accuracy_control, axis=0)
  2287. control_data[control, 'f1_avg'] = f1_avg_control_by_mouse
  2288. control_data[control, 'f1_std'] = f1_std_control_by_mouse
  2289. control_data[control, 'accuracy_avg'] = accuracy_avg_control_by_mouse
  2290. control_data[control, 'accuracy_std'] = accuracy_std_control_by_mouse
  2291. #Plot parameters
  2292. fig_num = 1
  2293. pca_plot_bins = 50
  2294. nrows = 2
  2295. ncols = num_mice // 2 + num_mice % 2
  2296. if num_mice == 1:
  2297. nrows=1
  2298. #### PLOT PCAS ####
  2299. fig_pca, axs_pca = plt.subplots(nrows=nrows, ncols=ncols, squeeze=False, figsize=(4 * ncols, 3 * nrows), num=fig_num, subplot_kw={'projection':'3d'}); fig_num += 1
  2300. for midx, mnum in enumerate(mouse_list):
  2301. CCA_analysis_dict = pipeline_output_dict['CCA_analysis_dict']
  2302. pca_list, pos_list = CCA_analysis_dict[mnum, 'pca'], CCA_analysis_dict[mnum, 'pos']
  2303. ax = axs_pca.ravel()[midx]; ax.set_title("M%d" %mnum, fontsize=15)
  2304. pf.plot_pca_overlapped(pca_list, pos_list, average_by_session=True, pca_plot_bins=pca_plot_bins, ax=ax)
  2305. #### DECODING PERFORMANCE PLOTS ####
  2306. fs=15
  2307. ymin = 0.2
  2308. figsize_decoding_by_mouse = (6,3)
  2309. figsize_decoding_summary = (4,4)
  2310. decoding_avg_std_pairs = [(f1_avg_by_mouse, f1_std_by_mouse), (accuracy_avg_by_mouse, accuracy_std_by_mouse)]
  2311. decoding_key_list = ['f1', 'accuracy']
  2312. decoding_key_list = ['f1', 'accuracy']
  2313. for decoding_pair_idx, (avg_by_mouse, std_by_mouse) in enumerate(decoding_avg_std_pairs):
  2314. decoding_key = decoding_key_list[decoding_pair_idx]
  2315. decoding_label = decoding_key[0].upper() + decoding_key[1:] #Make first letter uppercase
  2316. fig, ax = plt.subplots(nrows=1, ncols=1, squeeze=False, figsize=figsize_decoding_by_mouse, num=fig_num)
  2317. ax = ax[0,0]; fig_num += 1
  2318. xpos_list = np.arange(num_mice)
  2319. shift = 0.15
  2320. for class_idx, class_label in enumerate(range(2)):
  2321. label=pparam.AP_DECODING_LABELS[class_idx]
  2322. color=pparam.AP_DECODING_COLORS[class_idx]
  2323. perf = avg_by_mouse[:, class_idx]
  2324. perfstd = std_by_mouse[:, class_idx]
  2325. xx = xpos_list
  2326. xx = xpos_list - shift * (1-2*class_idx)
  2327. ax.errorbar(xx, perf, yerr=perfstd, fmt='_', color=color, alpha=0.9, label=label,
  2328. markersize=20, markeredgewidth=3, elinewidth = 3, zorder=2)
  2329. for idx, control in enumerate(control_names):
  2330. control_idx = control_label_idxs[idx]
  2331. for class_idx, class_label in enumerate(range(2)):
  2332. if class_idx == 0:
  2333. label=pparam.AP_DECODING_LABELS[control_idx]
  2334. else:
  2335. label=None
  2336. color=pparam.AP_DECODING_COLORS[control_idx]
  2337. perf = control_data[control, decoding_key + '_avg'][:, class_idx]
  2338. perfstd = control_data[control, decoding_key + '_std'][:, class_idx]
  2339. xx = xpos_list
  2340. xx = xpos_list - shift * (1-2*class_idx)
  2341. ax.errorbar(xx, perf, yerr=perfstd, fmt='_', color=color, alpha=0.9, label=label,
  2342. markersize=20, markeredgewidth=3, elinewidth = 3, zorder=2)
  2343. ## Plot significances with controls
  2344. ymax = ax.get_ylim()[1]
  2345. if len(control_names) > 0:
  2346. for midx, mnum in enumerate(mouse_list):
  2347. ## Pval for each individually. Done for each mouse and class against each control
  2348. significance = True
  2349. for class_idx, class_label in enumerate(range(2)):
  2350. if decoding_key == 'f1':
  2351. perf = APdecoding_dict[mnum, 'f1_array'][:, class_idx]
  2352. elif decoding_key == 'accuracy':
  2353. perf = APdecoding_dict[mnum, 'accuracy_array'][:, class_idx]
  2354. for control_idx, control in enumerate(control_names):
  2355. if 'shuffle' in control:
  2356. perf_control = APdecoding_dict[mnum, decoding_key + '_array_' + control][:, class_idx]
  2357. elif 'CCA' in control:
  2358. if decoding_key == 'f1':
  2359. perf_control = f1_array_shift[midx][:, class_idx]
  2360. elif decoding_key == 'accuracy':
  2361. perf_control = accuracy_array_shift[midx][:, class_idx]
  2362. # tstat, pval = scipy.stats.ttest_ind(perf, perf_control, equal_var=True, permutations=None, alternative='greater')
  2363. tstat, pval = scipy.stats.mannwhitneyu(perf, perf_control, use_continuity=False, alternative='greater')
  2364. #If even a single class vs control test is non-significant, mark it as non-significant
  2365. if pval > 0.05:
  2366. significance = False
  2367. if significance == True:
  2368. maxval = np.max(avg_by_mouse[midx])
  2369. xp = xpos_list[midx]-0.05
  2370. yp = maxval + 0.05
  2371. ax.text(xp, yp, '*', fontsize = 25, style='italic')
  2372. ymax = np.maximum(ymax, yp)
  2373. ## Two way anova, using variables "class type" (AP or No AP) and "analysis type" (shuffle, normal)
  2374. class_type_label = []
  2375. analysis_type_label = []
  2376. value_list = [] #f1 or acc
  2377. for class_idx, class_label in enumerate(range(2)):
  2378. if decoding_key == 'f1':
  2379. perf = APdecoding_dict[mnum, 'f1_array'][:, class_idx]
  2380. elif decoding_key == 'accuracy':
  2381. perf = APdecoding_dict[mnum, 'accuracy_array'][:, class_idx]
  2382. class_type_label.extend([class_label]*len(perf))
  2383. analysis_type_label.extend(["normal"]*len(perf))
  2384. value_list.extend(perf)
  2385. for control_idx, control in enumerate(control_names):
  2386. if 'shuffle' in control:
  2387. perf_control = APdecoding_dict[mnum, decoding_key + '_array_' + control][:, class_idx]
  2388. elif 'CCA' in control:
  2389. if decoding_key == 'f1':
  2390. perf_control = f1_array_shift[midx][:, class_idx]
  2391. elif decoding_key == 'accuracy':
  2392. perf_control = accuracy_array_shift[midx][:, class_idx]
  2393. class_type_label.extend([class_label]*len(perf_control))
  2394. analysis_type_label.extend([control]*len(perf_control))
  2395. value_list.extend(perf_control)
  2396. # all_data = np.hstack((shuffle_id, shuffle_dd, shuffle2_id, shuffle2_dd, f1_id, f1_dd))
  2397. df = pd.DataFrame({'ctype':class_type_label,
  2398. 'atype':analysis_type_label,
  2399. "value":value_list})
  2400. model = ols('value ~ C(ctype) + C(atype) + C(ctype):C(atype)', data=df).fit()
  2401. result = sm.stats.anova_lm(model, typ=2)
  2402. print(result)
  2403. print("M%d"%mnum, " // Controls pval: ", result.loc[["C(atype)"]]['PR(>F)'].values[0],
  2404. " // Class pval: ", result.loc[["C(ctype)"]]['PR(>F)'].values[0])
  2405. ax.legend(fontsize=15, loc = 'upper right', frameon=False)
  2406. ax.set_xlim([-shift-0.1, num_mice-(1-shift-0.1)])
  2407. # ax.set_ylim([0, ax.get_ylim()[1]])
  2408. yminlim = np.minimum(ax.get_ylim()[0], ymin)
  2409. ymaxlim = np.maximum(ax.get_ylim()[1], ymax)
  2410. ax.set_ylim([yminlim, ymaxlim])
  2411. xpos_list = np.arange(num_mice)
  2412. xlabels = ["M%d"%mnum for mnum in mouse_list]
  2413. ax.set_xticks(xpos_list, xlabels, fontsize=fs)
  2414. ax.set_ylabel("%s score"%decoding_label, fontsize=fs)
  2415. ax.tick_params(axis='y', labelsize=fs)
  2416. xlims = ax.get_xlim()
  2417. ax.plot(xlims, [0.5, 0.5], '--k', alpha=0.5)
  2418. ax.spines[['right', 'top']].set_visible(False)
  2419. for axis in ['top','bottom','left','right']:
  2420. ax.spines[axis].set_linewidth(3)
  2421. fig.tight_layout()
  2422. fig_name = 'fig3_%s_by_mouse'%decoding_label
  2423. save_figure(fig, fig_name)
  2424. #VD vs DD summary plot
  2425. fig, ax = plt.subplots(nrows=1, ncols=1, squeeze=False, figsize=figsize_decoding_summary, num=fig_num)
  2426. ax = ax[0,0]; fig_num += 1
  2427. xpos_list = np.arange(2)
  2428. shift = 0.1
  2429. ymax = 0.5
  2430. for mouse_type_idx, mouse_type_label in enumerate(pparam.MOUSE_TYPE_LABELS):
  2431. m_idxs = mtype_by_mouse == mouse_type_label
  2432. if np.sum(m_idxs) == 0:
  2433. continue
  2434. for class_idx, class_label in enumerate(range(2)):
  2435. perf_avg = np.average(avg_by_mouse[m_idxs, class_idx])
  2436. perf_std = np.average(std_by_mouse[m_idxs, class_idx])
  2437. ymax = np.maximum(ymax, perf_avg)
  2438. xx = xpos_list[mouse_type_idx] - shift * (1-2*class_idx)
  2439. label=pparam.AP_DECODING_LABELS[class_idx]
  2440. if mouse_type_idx == 1:
  2441. label=None
  2442. color = pparam.AP_DECODING_COLORS[class_idx]
  2443. ax.errorbar(xx, perf_avg, yerr=perf_std, fmt='_', color=color, alpha=0.9, label=label,
  2444. markersize=20, markeredgewidth=3, elinewidth = 3, zorder=2)
  2445. if decoding_pair_idx == 0:
  2446. print(mouse_type_label, class_label, perf_avg)
  2447. for idx, control in enumerate(control_names):
  2448. for class_idx, class_label in enumerate(range(2)):
  2449. perf_avg = np.average(control_data[control, decoding_key + '_avg'][m_idxs, class_idx])
  2450. perf_std = np.average(control_data[control, decoding_key + '_std'][m_idxs, class_idx])
  2451. # print(mouse_type_label, control, class_label, perf_avg)
  2452. control_idx = control_label_idxs[idx]
  2453. label =pparam.AP_DECODING_LABELS[control_idx]
  2454. if mouse_type_idx != 0 or class_idx != 0:
  2455. label=None
  2456. color = pparam.AP_DECODING_COLORS[control_idx]
  2457. xx = xpos_list[mouse_type_idx] - shift * (1-2*class_idx)
  2458. ax.errorbar(xx, perf_avg, yerr=perf_std, fmt='_', color=color, alpha=0.9, label=label,
  2459. markersize=20, markeredgewidth=3, elinewidth = 3, zorder=2)
  2460. if decoding_pair_idx == 0:
  2461. print(mouse_type_label, class_label, perf_avg)
  2462. #Plot significances with controls. Done for each mouse class against each control
  2463. if len(control_names) > 0 and decoding_key == 'f1':
  2464. for mouse_type_idx, mouse_type_label in enumerate(pparam.MOUSE_TYPE_LABELS):
  2465. m_idxs = mtype_by_mouse == mouse_type_label
  2466. if np.sum(m_idxs) == 0:
  2467. continue
  2468. significance = True
  2469. for class_idx, class_label in enumerate(range(2)):
  2470. perf = avg_by_mouse[m_idxs, class_idx]
  2471. for control_idx, control in enumerate(control_names):
  2472. perf_control = control_data[control, decoding_key + '_avg'][m_idxs, class_idx]
  2473. # tstat, pval = scipy.stats.ttest_ind(perf, perf_control, equal_var=True, permutations=None, alternative='greater')
  2474. tstat, pval = scipy.stats.mannwhitneyu(perf, perf_control, use_continuity=False, alternative='greater')
  2475. #If even a single class vs control test is non-significant, mark it as non-significant
  2476. if pval > 0.05:
  2477. significance = False
  2478. if significance == True:
  2479. maxval = np.max([np.average(avg_by_mouse[m_idxs, class_idx]) for class_idx in range(2)])
  2480. xp = xpos_list[mouse_type_idx]-0.05
  2481. yp = maxval + 0.05
  2482. ax.text(xp, yp, '*', fontsize = 25, style='italic')
  2483. ymax = np.maximum(ymax, yp)
  2484. #Plot ID vs DD significances (merge both class labels together)
  2485. perf_by_mtype = []
  2486. for mouse_type_idx, mouse_type_label in enumerate(pparam.MOUSE_TYPE_LABELS):
  2487. m_idxs = mtype_by_mouse == mouse_type_label
  2488. perf_by_mtype.append(avg_by_mouse[m_idxs, :].ravel())
  2489. # tstat, pval = scipy.stats.ttest_ind(perf_by_mtype[0], perf_by_mtype[1], equal_var=True, permutations=None, alternative='two-sided')
  2490. tstat, pval = scipy.stats.wilcoxon(perf_by_mtype[0], perf_by_mtype[1], zero_method='wilcox', correction=False, alternative='two-sided')
  2491. print("ID vs DD, Wilcoxon pval: %.4f"%pval)
  2492. p00 = ymax + 0.05
  2493. p11 = xpos_list[0]
  2494. p10 = xpos_list[1]
  2495. d0 = 0
  2496. dp = 0.05
  2497. ymax = np.maximum(ymax, p00)
  2498. pf.draw_significance(ax, pval, p00, p10, p11, d0, dp, orientation='top', label_padding=0, thresholds = [0.05], fs=20)
  2499. ## Two way anova, using variables "class type" (AP or No AP) and "analysis type" (shuffle, normal)
  2500. if len(control_names) > 0 and decoding_key == 'f1':
  2501. axon_type_label = []
  2502. analysis_type_label = []
  2503. value_list = []
  2504. for mouse_type_idx, mouse_type_label in enumerate(pparam.MOUSE_TYPE_LABELS):
  2505. m_idxs = mtype_by_mouse == mouse_type_label
  2506. if np.sum(m_idxs) == 0:
  2507. continue
  2508. for class_idx, class_label in enumerate(range(2)):
  2509. perf = avg_by_mouse[m_idxs, class_idx]
  2510. axon_type_label.extend([mouse_type_label]*len(perf))
  2511. analysis_type_label.extend(["normal"]*len(perf))
  2512. value_list.extend(perf)
  2513. for control_idx, control in enumerate(control_names):
  2514. perf_control = control_data[control, decoding_key + '_avg'][m_idxs, class_idx]
  2515. axon_type_label.extend([mouse_type_label]*len(perf_control))
  2516. analysis_type_label.extend([control]*len(perf_control))
  2517. value_list.extend(perf_control)
  2518. df = pd.DataFrame({'axontype':axon_type_label,
  2519. 'atype':analysis_type_label,
  2520. "value":value_list})
  2521. model = ols('value ~ C(axontype) + C(atype) + C(axontype):C(atype)', data=df).fit()
  2522. result = sm.stats.anova_lm(model, typ=2)
  2523. print(result)
  2524. print("ID vs DD ANOVA // Controls pval: ", result.loc[["C(atype)"]]['PR(>F)'].values[0],
  2525. " // Class pval: ", result.loc[["C(axontype)"]]['PR(>F)'].values[0])
  2526. ax.legend(fontsize=15, frameon=False)
  2527. # ax.set_xlim([-shift-0.1, num_mice-(1-shift-0.1)])
  2528. yminlim = np.minimum(ax.get_ylim()[0], ymin)
  2529. ymaxlim = np.maximum(ax.get_ylim()[1], ymax)
  2530. ax.set_ylim([yminlim, ymaxlim])
  2531. ax.legend(fontsize=15, loc = 'upper right', frameon=False)
  2532. xlabels = pparam.MOUSE_TYPE_LABELS
  2533. ax.set_xticks(xpos_list, xlabels, fontsize=fs)
  2534. ax.set_ylabel("%s score"%decoding_label, fontsize=fs)
  2535. ax.tick_params(axis='y', labelsize=fs)
  2536. xlims = ax.get_xlim()
  2537. ax.plot(xlims, [0.5, 0.5], '--k', alpha=0.5)
  2538. ax.spines[['right', 'top']].set_visible(False)
  2539. for axis in ['top','bottom','left','right']:
  2540. ax.spines[axis].set_linewidth(3)
  2541. fig.tight_layout()
  2542. fig_name = 'fig3_%s_summary'%decoding_label
  2543. save_figure(fig, fig_name)
  2544. return
  2545. def fig3_F_and_3SI_A_B_belt_restriction_plots():
  2546. exclusion_center_list = [250, 750, 1250]
  2547. exclusion_interval_size = 500
  2548. fig_name = 'fig3_belt_restriction'
  2549. belt_restriction_plots(exclusion_center_list, exclusion_interval_size, fig_name)
  2550. def fig3SI_C():
  2551. exclusion_center_list = [0, 500, 1000]
  2552. exclusion_interval_size = 500
  2553. fig_name = 'fig3SI_C'
  2554. belt_restriction_plots(exclusion_center_list, exclusion_interval_size, fig_name)
  2555. def fig3SI_D():
  2556. exclusion_center_list = [0, 500, 1000]
  2557. exclusion_interval_size = 1000
  2558. fig_name = 'fig3SI_D'
  2559. belt_restriction_plots(exclusion_center_list, exclusion_interval_size, fig_name)
  2560. def belt_restriction_plots(
  2561. exclusion_center_list = [250, 750, 1250],
  2562. exclusion_interval_size = 500,
  2563. fig_name = 'fig3_belt_restriction'
  2564. ):
  2565. '''
  2566. Performs air puff decoding when part of the belt is excluded from the TCA analysis.
  2567. exclusion_center_list: centers of the segment that is going to be excluded.
  2568. exclusion_interval_size: total length of the excluded segment.
  2569. e.g. if 250 is the center and 500 the interval size, a section from x=0 to x=500 will be excluded.
  2570. '''
  2571. mouse_list = np.arange(8)
  2572. # mouse_list = [2,6]
  2573. # mouse_list = [0,2,6,7]
  2574. # mouse_list = [0,3,6,7]
  2575. # mouse_list = [6]
  2576. preprocessing_param_dict = {
  2577. #session params
  2578. 'mouse_list':mouse_list,
  2579. 'session_list':np.arange(9),
  2580. #Preprocessing parameters
  2581. 'time_bin_size':1,
  2582. 'distance_bin_size':1,
  2583. 'gaussian_size':25,
  2584. 'data_used':'amplitudes',
  2585. 'running':True,
  2586. 'eliminate_v_zeros':False,
  2587. 'num_components':'all'
  2588. }
  2589. cca_param_dict = {
  2590. 'CCA_dim':'.9',
  2591. 'return_warped_data':True,
  2592. 'return_trimmed_data':False,
  2593. 'sessions_to_align':'all',
  2594. 'shuffle':False
  2595. }
  2596. ap_decoding_param_dict = {
  2597. ## TCA params ##
  2598. 'TCA_method': "ncp_hals", #"cp_als", "mcp_als", "ncp_bcd", "ncp_hals"
  2599. 'TCA_factors':'max',
  2600. 'TCA_replicates':10,
  2601. 'TCA_convergence_attempts':10, #Number of times TCA can fail before giving up
  2602. 'TCA_on_LDA_repetitions':20,
  2603. ## LDA params ##
  2604. 'LDA_imbalance_prop':.51,
  2605. 'LDA_imbalance_repetitions':10,
  2606. 'LDA_trial_shuffles':20,
  2607. 'LDA_session_shuffles':0,
  2608. 'session_comparisons':'BT' #'airpuff', 'BT', 'TP', 'BP'
  2609. }
  2610. fig_num = 1
  2611. fs=15
  2612. plot_cut_pca = True
  2613. exclusion_left_list = [(exclusion_center_list[i]-exclusion_interval_size/2)%pparam.MAX_POS for i in range(len(exclusion_center_list))]
  2614. exclusion_right_list = [(exclusion_center_list[i]+exclusion_interval_size/2)%pparam.MAX_POS for i in range(len(exclusion_center_list))]
  2615. #### ANALYZE FOR EACH EXCLUSION INTERVAL ####
  2616. shuffle_num = ap_decoding_param_dict['LDA_trial_shuffles']
  2617. f1_dict_by_center_and_mouse = {}
  2618. f1_shuffle_dict_by_center_and_mouse = {}
  2619. for center_idx, center in enumerate(exclusion_center_list):
  2620. start, end = exclusion_left_list[center_idx], exclusion_right_list[center_idx]
  2621. ap_decoding_param_dict['exclude_positions'] = True
  2622. ap_decoding_param_dict['pos_to_exclude_from'] = start
  2623. ap_decoding_param_dict['pos_to_exclude_to'] = end
  2624. # #Do analysis
  2625. pipeline_output_dict = APdecoding_pipeline(preprocessing_param_dict, cca_param_dict, ap_decoding_param_dict)
  2626. # #Put data in arrays
  2627. APdecoding_dict = pipeline_output_dict['APdecoding_dict']
  2628. for midx, mnum in enumerate(mouse_list):
  2629. f1_array = APdecoding_dict[mnum, 'f1_array']
  2630. f1_dict_by_center_and_mouse[center, mnum] = f1_array
  2631. f1_array = APdecoding_dict[mnum, 'f1_array_shuffle']
  2632. f1_shuffle_dict_by_center_and_mouse[center, mnum] = f1_array
  2633. ## OPTIONAL: PLOT CUT DATA (only for the first center and mouse) ##
  2634. if plot_cut_pca == True:
  2635. CCA_analysis_dict = pipeline_output_dict['CCA_analysis_dict']
  2636. num_bins = CCA_analysis_dict['num_bins']
  2637. APdecoding_dict = {}
  2638. for midx,mnum in enumerate(mouse_list):
  2639. if midx != 0:
  2640. continue
  2641. fig = plt.figure(fig_num, figsize=(7,7)); fig_num += 1
  2642. ax = plt.subplot(projection='3d')
  2643. session_list = CCA_analysis_dict[mnum, 'session_list']
  2644. pos_list = CCA_analysis_dict[mnum, 'pos']
  2645. pca_list = CCA_analysis_dict[mnum, 'pca']
  2646. data_by_trial, pos_by_trial, snum_by_trial = pf.reshape_pca_list_by_trial(pca_list, pos_list, num_bins, session_list)
  2647. num_features, num_bins, total_trials = data_by_trial.shape
  2648. #Limit position (if indicated)
  2649. if ap_decoding_param_dict['exclude_positions'] == True:
  2650. positions = pos_by_trial[:,0] #Assumes all trials are binned using the same positions
  2651. pos_bool_filtered_out = pf.get_idxs_in_periodic_interval(positions, ap_decoding_param_dict['pos_to_exclude_from'],
  2652. ap_decoding_param_dict['pos_to_exclude_to'], pparam.MAX_POS)
  2653. pos_bool_selected = np.invert(pos_bool_filtered_out)
  2654. data_by_trial = data_by_trial[:, pos_bool_selected, :]
  2655. pos_by_trial = pos_by_trial[pos_bool_selected, :]
  2656. for sidx, snum in enumerate(np.unique(snum_by_trial)):
  2657. # print(data_by_trial.shape)
  2658. trials = snum_by_trial == snum
  2659. pca = pf.flatten_warped_data(data_by_trial[:, :, trials])
  2660. pos = pf.flatten_warped_data(pos_by_trial[:, trials])
  2661. pos, pca, _ = pf.compute_average_data_by_position(pca, pos, position_bin_size=5)
  2662. cbar = [None, True][midx == 0 and sidx == 0]
  2663. pf.plot_pca_with_position(pca, pos, ax=ax, scatter=True, cbar=cbar, ms=100)
  2664. ## OPTIONAL: PLOT CUT DATA ##
  2665. #Get mouse-averaged results
  2666. f1_avg_list = []; f1_std_list = []
  2667. f1_shuffle_avg_list = []; f1_shuffle_std_list = []
  2668. for center_idx, center in enumerate(exclusion_center_list):
  2669. f1_all = np.array([f1_dict_by_center_and_mouse[center, mnum].ravel() for mnum in mouse_list]).ravel()
  2670. f1_avg_list.append(np.average(f1_all))
  2671. f1_all_std = np.array([np.std(f1_dict_by_center_and_mouse[center, mnum].ravel()) for mnum in mouse_list]).ravel()
  2672. f1_std_list.append(np.average(f1_all_std))
  2673. f1_all = np.array([f1_shuffle_dict_by_center_and_mouse[center, mnum].ravel() for mnum in mouse_list]).ravel()
  2674. f1_shuffle_avg_list.append(np.average(f1_all))
  2675. # f1_shuffle_std_list.append(np.std(f1_all))
  2676. f1_shuffle_all_std = np.array([np.std(f1_shuffle_dict_by_center_and_mouse[center, mnum].ravel()) for mnum in mouse_list]).ravel()
  2677. f1_shuffle_std_list.append(np.average(f1_shuffle_all_std))
  2678. #### ANALYZE WITHOUT EXCLUSION, AS COMPARISON ####
  2679. ap_decoding_param_dict['exclude_positions'] = False
  2680. pipeline_output_dict = APdecoding_pipeline(preprocessing_param_dict, cca_param_dict, ap_decoding_param_dict)
  2681. #Store reference results
  2682. #The standard deviation is the average of the standard deviation of each mouse
  2683. APdecoding_dict_ref = pipeline_output_dict['APdecoding_dict']
  2684. f1_ref_list = []
  2685. f1_ref_std_list = []
  2686. for midx, mnum in enumerate(mouse_list):
  2687. f1_array = APdecoding_dict_ref[mnum, 'f1_array']
  2688. f1_ref_list.extend(f1_array.ravel())
  2689. f1_ref_std_list.append(np.std(f1_array.ravel()))
  2690. f1_ref_avg = np.average(f1_ref_list)
  2691. f1_ref_std = np.average(f1_ref_std_list)
  2692. ###### PLOT RESULTS ######
  2693. #If the exclusion interval is larger than half the belt, plot the center of the inclusion interval instead
  2694. #EG: if interval size is 1400 and period is 1500, the result only takes a small 100mm window into account. Use that as reference instead.
  2695. if exclusion_interval_size < 750:
  2696. center_list = exclusion_center_list
  2697. plot_interval_size = exclusion_interval_size
  2698. center_type = 'Exclusion'
  2699. else:
  2700. center_list = ((np.array(exclusion_center_list)+pparam.MAX_POS/2)%pparam.MAX_POS).astype(int)
  2701. plot_interval_size = int(pparam.MAX_POS - exclusion_interval_size)
  2702. center_type = 'Inclusion'
  2703. #SI: results for each mouse separately
  2704. fig, axs = plt.subplots(2, 4, figsize=(4 * 4, 3 * 2)); fig_num += 1
  2705. for midx, mnum in enumerate(mouse_list):
  2706. ax = axs.ravel()[mnum]
  2707. #Plot main results
  2708. f1_avg_mouse = [np.average(f1_dict_by_center_and_mouse[center, mnum]) for center in exclusion_center_list]
  2709. f1_std_mouse = [np.std(f1_dict_by_center_and_mouse[center, mnum]) for center in exclusion_center_list]
  2710. ax.errorbar(center_list, f1_avg_mouse, yerr=f1_std_mouse, fmt='_', color='black', alpha=0.9, label='Restricted',
  2711. markersize=20, markeredgewidth=3, elinewidth = 3, zorder=2)
  2712. #Plot shuffle
  2713. if shuffle_num > 0:
  2714. f1_shuffle_avg_mouse = [np.average(f1_shuffle_dict_by_center_and_mouse[center, mnum]) for center in exclusion_center_list]
  2715. f1_shuffle_std_mouse = [np.std(f1_shuffle_dict_by_center_and_mouse[center, mnum]) for center in exclusion_center_list]
  2716. ax.errorbar(center_list, f1_shuffle_avg_mouse, yerr=f1_shuffle_std_mouse, fmt='_', color=pparam.AP_DECODING_COLORS[2], alpha=0.9,
  2717. label=pparam.AP_DECODING_LABELS[2], markersize=20, markeredgewidth=3, elinewidth = 3, zorder=2)
  2718. #Plot reference
  2719. f1_avg_ref_mouse = np.average(APdecoding_dict_ref[mnum, 'f1_array'])
  2720. f1_std_ref_mouse = np.std(APdecoding_dict_ref[mnum, 'f1_array'])
  2721. #Draw F1 reference average
  2722. xlims = ax.get_xlim()
  2723. xx = np.linspace(xlims[0], xlims[1])
  2724. ax.plot(xx, [f1_avg_ref_mouse]*len(xx), '--', color='black', lw=3, alpha=0.5, label='Full')
  2725. ax.fill_between(xx, [f1_avg_ref_mouse-f1_std_ref_mouse]*len(xx), [f1_avg_ref_mouse+f1_std_ref_mouse]*len(xx), color='gray', alpha=0.3)
  2726. ax.set_xlim(xlims)
  2727. #Plot significance with shuffle for each segment center
  2728. ymax = ax.get_ylim()[1]
  2729. if shuffle_num > 0:
  2730. for center_idx, center in enumerate(exclusion_center_list):
  2731. f1_all = np.array(f1_dict_by_center_and_mouse[center, mnum]).ravel()
  2732. f1_all_shuffle = np.array(f1_shuffle_dict_by_center_and_mouse[center, mnum]).ravel()
  2733. # tstat, pval = scipy.stats.ttest_ind(f1_all, f1_all_shuffle, equal_var=True, permutations=None, alternative='two-sided')
  2734. tstat, pval = scipy.stats.mannwhitneyu(f1_all, f1_all_shuffle, use_continuity=False, alternative='two-sided')
  2735. if 0.05 > pval:
  2736. xp = center_list[center_idx]-50
  2737. yp = np.max(f1_avg_ref_mouse + f1_std_ref_mouse) + 0.05
  2738. ax.text(xp, yp, '*', fontsize = 25, style='italic')
  2739. ymax = np.maximum(ymax, yp)
  2740. #X axis
  2741. ax.set_xticks(center_list, center_list)
  2742. ax.tick_params(axis='x', labelsize=fs)
  2743. ax.set_xlim([np.min(center_list)-150, np.max(center_list)+150])
  2744. #Y axis
  2745. ax.tick_params(axis='y', labelsize=fs+4)
  2746. ax.set_ylim([0.2, ymax+0.1])
  2747. #Both axis
  2748. ax.spines[['right', 'top']].set_visible(False)
  2749. for axis in ['top','bottom','left','right']:
  2750. ax.spines[axis].set_linewidth(3)
  2751. if midx == 0:
  2752. ax.set_xlabel('%s segment center'%center_type, fontsize=fs+4)
  2753. ax.set_ylabel('Avg F1 score', fontsize=fs+4)
  2754. ax.legend(fontsize=fs-6, frameon=False)
  2755. ax.set_title('M%d'%mnum, fontsize=fs+2)
  2756. fig.suptitle(r'%s range:$\pm$ %d (mm)'%(center_type, plot_interval_size/2), fontsize=fs+3)
  2757. fig.tight_layout()
  2758. fig_name_by_mouse = '%s_by_mouse'%fig_name
  2759. save_figure(fig, fig_name_by_mouse)
  2760. #Fig4 subplot: mice-averaged results
  2761. fig = plt.figure(fig_num, figsize=(5,5)); fig_num += 1
  2762. ax = plt.gca()
  2763. fs = 18
  2764. #Plot exclusion results
  2765. ax.errorbar(center_list, f1_avg_list, yerr=f1_std_list, fmt='_', color='black', alpha=0.9, label='Restricted',
  2766. markersize=20, markeredgewidth=3, elinewidth = 3, zorder=2)
  2767. #Plot shuffle results
  2768. if shuffle_num > 0:
  2769. ax.errorbar(center_list, f1_shuffle_avg_list, yerr=f1_shuffle_std_list, fmt='_', color=pparam.AP_DECODING_COLORS[2], alpha=0.9,
  2770. label=pparam.AP_DECODING_LABELS[2], markersize=20, markeredgewidth=3, elinewidth = 3, zorder=2)
  2771. #Plot reference
  2772. xlims = ax.get_xlim()
  2773. xx = np.linspace(xlims[0], xlims[1])
  2774. ax.plot(xx, [f1_ref_avg]*len(xx), '--', color='black', lw=3, alpha=0.5, label='Full')
  2775. ax.fill_between(xx, [f1_ref_avg-f1_ref_std]*len(xx), [f1_ref_avg+f1_ref_std]*len(xx), color='gray', alpha=0.3)
  2776. ax.set_xlim(xlims)
  2777. #Plot significance with shuffle for each segment center
  2778. ymax = ax.get_ylim()[1]
  2779. if shuffle_num > 0:
  2780. for center_idx, center in enumerate(exclusion_center_list):
  2781. f1_all = np.array([f1_dict_by_center_and_mouse[center, mnum].ravel() for mnum in mouse_list]).ravel()
  2782. f1_all_shuffle = np.array([f1_shuffle_dict_by_center_and_mouse[center, mnum].ravel() for mnum in mouse_list]).ravel()
  2783. # tstat, pval = scipy.stats.ttest_ind(f1_all, f1_all_shuffle, equal_var=True, permutations=None, alternative='two-sided')
  2784. tstat, pval = scipy.stats.mannwhitneyu(f1_all, f1_all_shuffle, use_continuity=False, alternative='two-sided')
  2785. if 0.05 > pval:
  2786. xp = center_list[center_idx]-25
  2787. yp = np.max(f1_avg_ref_mouse + f1_std_ref_mouse) + 0.02
  2788. ax.text(xp, yp, '*', fontsize = 25, style='italic')
  2789. ymax = np.maximum(ymax, yp)
  2790. #X axis
  2791. ax.set_xlabel('%s segment center'%center_type, fontsize=fs)
  2792. ax.set_xticks(center_list, center_list)
  2793. ax.tick_params(axis='x', labelsize=fs)
  2794. ax.set_xlim([np.min(center_list)-150, np.max(center_list)+150])
  2795. #Y axis
  2796. ax.tick_params(axis='y', labelsize=fs)
  2797. ax.set_ylabel('Avg F1 score', fontsize=fs)
  2798. ax.set_ylim([0.2, ymax+0.05])
  2799. #Both axis
  2800. ax.spines[['right', 'top']].set_visible(False)
  2801. for axis in ['top','bottom','left','right']:
  2802. ax.spines[axis].set_linewidth(3)
  2803. #Title
  2804. ax.set_title(r'%s range:$\pm$ %d (mm)'%(center_type, plot_interval_size/2), fontsize=fs)
  2805. ax.legend(fontsize=fs, frameon=False)
  2806. #Figure params
  2807. fig.tight_layout()
  2808. fig_name_summary = '%s_summary'%fig_name
  2809. save_figure(fig, fig_name_summary)
  2810. def fig4_A_session_comparisons():
  2811. ''' Performs TCA across sessions for an animal.
  2812. Step 1: perform PCA, limit to minimum possible of dimensions
  2813. Step 2: align through CCA, so every dimension represents something similar about the data
  2814. Step 3: split into trials through warping
  2815. Step 4: TCA!
  2816. '''
  2817. mouse_list = np.arange(8)
  2818. # mouse_list = [6]
  2819. # mouse_list = [0,2,6,7]
  2820. # mouse_list = [0,3,6,7]
  2821. # mouse_list = [2,6]
  2822. preprocessing_param_dict = {
  2823. #session params
  2824. 'mouse_list':mouse_list,
  2825. 'session_list':np.arange(9),
  2826. #Preprocessing parameters
  2827. 'time_bin_size':1,
  2828. 'distance_bin_size':1,
  2829. 'gaussian_size':25,
  2830. 'data_used':'amplitudes',
  2831. 'running':True,
  2832. 'eliminate_v_zeros':True,
  2833. 'num_components':'all'
  2834. }
  2835. cca_param_dict = {
  2836. 'CCA_dim':'.9', #'.9'
  2837. 'return_warped_data':True,
  2838. 'return_trimmed_data':False,
  2839. 'sessions_to_align':'all',
  2840. 'shuffle':False
  2841. }
  2842. ap_decoding_param_dict = {
  2843. 'exclude_positions':False,
  2844. 'pos_to_exclude_from':1000,
  2845. 'pos_to_exclude_to':1500,
  2846. ## TCA params ##
  2847. 'TCA_method': "ncp_hals", #"cp_als", "mcp_als", "ncp_bcd", "ncp_hals"
  2848. 'TCA_factors':'max',
  2849. 'TCA_replicates':10,
  2850. 'TCA_convergence_attempts':10, #Number of times TCA can fail before giving up
  2851. 'TCA_on_LDA_repetitions':20,
  2852. ## LDA params ##
  2853. 'LDA_imbalance_prop':.51,
  2854. 'LDA_imbalance_repetitions':10,
  2855. 'LDA_trial_shuffles':20,
  2856. 'LDA_session_shuffles':0,
  2857. 'session_comparisons':'BT' #'airpuff', 'BT', 'TP', 'BP'
  2858. }
  2859. fig_num = 1
  2860. fs = 15
  2861. # # figsize_f1 = (6,3)
  2862. # # figsize_f1_summary = (3,4)
  2863. #Do analysis
  2864. session_comparisons_list = ['airpuff', 'BT', 'BP', 'TP']
  2865. # session_comparisons_list = ['BT', 'TP']
  2866. session_comparisons_label_dict = {'airpuff':'BP-T', 'BT':'B-T', 'BP': 'B-P', 'TP':'T-P'}
  2867. f1_array_list = []
  2868. f1_array_shuffle_list = []
  2869. for comp_idx, session_comparison in enumerate(session_comparisons_list):
  2870. ap_decoding_param_dict['session_comparisons'] = session_comparison
  2871. pipeline_output_dict = APdecoding_pipeline(preprocessing_param_dict, cca_param_dict, ap_decoding_param_dict)
  2872. #Put data in arrays
  2873. num_mice = len(mouse_list)
  2874. LDA_trial_shuffles = ap_decoding_param_dict['LDA_trial_shuffles']
  2875. APdecoding_dict = pipeline_output_dict['APdecoding_dict']
  2876. f1_array = np.zeros((num_mice, 2)) #num mice X num classes
  2877. # f1_array_std = np.zeros((num_mice, 2)) #num mice X num classes
  2878. f1_array_shuffle = np.zeros((num_mice, LDA_trial_shuffles, 2))
  2879. # f1_array_shuffle_std = np.zeros((num_mice, LDA_trial_shuffles, 2))
  2880. for midx, mnum in enumerate(mouse_list):
  2881. f1_array[midx, :] = np.average(APdecoding_dict[mnum, 'f1_array'], axis=0)
  2882. f1_array_shuffle[midx, :] = np.average(APdecoding_dict[mnum, 'f1_array_shuffle'], axis=0)
  2883. f1_array_list.append(f1_array)
  2884. f1_array_shuffle_list.append(f1_array_shuffle)
  2885. #Mouse types
  2886. mtype_by_mouse = np.array([pparam.MOUSE_TYPE_LABEL_BY_MOUSE[mnum] for mnum in mouse_list])
  2887. mouse_types = pparam.MOUSE_TYPE_LABELS
  2888. labels = [session_comparisons_label_dict[scomp] for scomp in session_comparisons_list]
  2889. fig = plt.figure(fig_num, figsize=(5,3)); fig_num += 1
  2890. ax = plt.gca()
  2891. xpos_list = np.arange(len(session_comparisons_list))
  2892. xgap = 0.2
  2893. for mtype_idx, mtype in enumerate(mouse_types):
  2894. midxs = np.where(mtype_by_mouse == mtype)[0]
  2895. f1avg = [np.average(f1[midxs]) for f1 in f1_array_list]
  2896. f1std = [np.std(f1[midxs]) for f1 in f1_array_list]
  2897. ax.errorbar(xpos_list + xgap * (2*mtype_idx-1), f1avg, yerr=f1std, label=mtype, color=pparam.MOUSE_TYPE_COLORS[mtype_idx],
  2898. fmt='_', markersize=20, markeredgewidth=3, elinewidth = 3, zorder=2)
  2899. if LDA_trial_shuffles > 0:
  2900. # for mtype_idx, mtype in enumerate(mouse_types):
  2901. midxs = np.where(mtype_by_mouse == mtype)[0]
  2902. f1avg = [np.average(f1[midxs]) for f1 in f1_array_shuffle_list]
  2903. f1std = [np.std(f1[midxs]) for f1 in f1_array_shuffle_list]
  2904. label = [None, pparam.AP_DECODING_LABELS[2]][mtype_idx == 0]
  2905. ax.errorbar(xpos_list + xgap * (2*mtype_idx-1), f1avg, yerr=f1std, label=label, color=pparam.SHUFFLE_DEFAULT_COLOR,
  2906. fmt='_', markersize=20, markeredgewidth=3, elinewidth = 3, zorder=2)
  2907. ## Plot significances with controls. Done for each mouse type against their respective shuffle control
  2908. overall_max_val = ax.get_ylim()[1]
  2909. for session_comparison_idx, session_comparison in enumerate(session_comparisons_list):
  2910. #If one of the two axon types is non-significant, don't plot the asterisk!
  2911. significance = True
  2912. session_comparison_maxval = 0
  2913. for mtype_idx, mtype in enumerate(mouse_types):
  2914. midxs = np.where(mtype_by_mouse == mtype)[0]
  2915. f1 = f1_array_list[session_comparison_idx][midxs].ravel()
  2916. f1_control = f1_array_shuffle_list[session_comparison_idx][midxs]
  2917. f1_control = np.average(f1_control, axis=1).ravel()
  2918. # tstat, pval = scipy.stats.ttest_ind(f1, f1_control, equal_var=True, permutations=None, alternative='greater')
  2919. tstat, pval = scipy.stats.mannwhitneyu(f1, f1_control, use_continuity=False, alternative='greater')
  2920. if pval > 0.05:
  2921. significance = False
  2922. session_comparison_maxval = np.maximum(session_comparison_maxval, np.average(f1) + np.std(f1))
  2923. if significance == True:
  2924. xp = xpos_list[session_comparison_idx]-0.1
  2925. yp = session_comparison_maxval + 0.01
  2926. ax.text(xp, yp, '*', fontsize = 25, style='italic')
  2927. overall_max_val = np.maximum(overall_max_val, yp)
  2928. ## Two way anova, using variables "class type" (AP or No AP) and "analysis type" (shuffle, normal)
  2929. axon_type_label = []
  2930. analysis_type_label = []
  2931. value_list = []
  2932. for session_comparison_idx, session_comparison in enumerate(session_comparisons_list):
  2933. #If one of the two axon types is non-significant, don't plot the asterisk!
  2934. significance = True
  2935. session_comparison_maxval = 0
  2936. for mtype_idx, mtype in enumerate(mouse_types):
  2937. midxs = np.where(mtype_by_mouse == mtype)[0]
  2938. f1 = f1_array_list[session_comparison_idx][midxs].ravel()
  2939. axon_type_label.extend([mtype]*f1.size)
  2940. analysis_type_label.extend(['normal']*f1.size)
  2941. value_list.extend(f1)
  2942. f1_control = f1_array_shuffle_list[session_comparison_idx][midxs].ravel()
  2943. axon_type_label.extend([mtype]*f1_control.size)
  2944. analysis_type_label.extend(['shuffle']*f1_control.size)
  2945. value_list.extend(f1_control)
  2946. df = pd.DataFrame({'axontype':axon_type_label,
  2947. 'atype':analysis_type_label,
  2948. "value":value_list})
  2949. model = ols('value ~ C(axontype) + C(atype) + C(axontype):C(atype)', data=df).fit()
  2950. result = sm.stats.anova_lm(model, typ=2)
  2951. print(result)
  2952. print("%s ANOVA // Controls pval: "%session_comparison, result.loc[["C(atype)"]]['PR(>F)'].values[0],
  2953. " // Class pval: ", result.loc[["C(axontype)"]]['PR(>F)'].values[0])
  2954. xlabels = labels
  2955. ax.set_xticks(xpos_list, xlabels, fontsize=fs)
  2956. ax.set_ylabel("Avg F1 score", fontsize=fs)
  2957. ax.tick_params(axis='y', labelsize=fs)
  2958. ax.tick_params(axis='x', labelsize=fs)
  2959. ymin = 0.2
  2960. ymin = np.minimum(ymin, ax.get_ylim()[0])
  2961. ax.set_ylim([ymin, overall_max_val])
  2962. xlims = ax.get_xlim()
  2963. ax.plot(xlims, [0.5, 0.5], '--k', alpha=0.5)
  2964. ax.spines[['right', 'top']].set_visible(False)
  2965. for axis in ['top','bottom','left','right']:
  2966. ax.spines[axis].set_linewidth(3)
  2967. ax.legend(fontsize=15, loc='lower right', frameon=False)
  2968. save_figure(fig, 'fig4_session_comparisons')
  2969. return
  2970. def fig4_C_D_E_distance_measures():
  2971. compute_and_plot_trial_factor_distances(distance_comparison = 'session_type')
  2972. def fig4SI_G_H_I_distance_measures():
  2973. compute_and_plot_trial_factor_distances(distance_comparison = 'airpuff')
  2974. def compute_and_plot_trial_factor_distances(distance_comparison = 'session_type'):
  2975. ''' Calculate distances between trial factors '''
  2976. mouse_list = np.arange(8)
  2977. # mouse_list = [2,6]
  2978. # mouse_list = [2,3,5,6]
  2979. # mouse_list = [0,1,2,3]
  2980. # mouse_list = [4,5,6,7]
  2981. # mouse_list = [0,6]
  2982. #Distance parameters
  2983. # distance_comparison = 'session_type' #'session_type', 'airpuff', 'single_session'
  2984. distance_shuffles = 10
  2985. preprocessing_param_dict = {
  2986. #session params
  2987. 'mouse_list':mouse_list,
  2988. 'session_list':np.arange(9),
  2989. #Preprocessing parameters
  2990. 'time_bin_size':1,
  2991. 'distance_bin_size':1,
  2992. 'gaussian_size':25,
  2993. 'data_used':'amplitudes',
  2994. 'running':True,
  2995. 'eliminate_v_zeros':False,
  2996. 'num_components':'all'
  2997. }
  2998. cca_param_dict = {
  2999. 'CCA_dim':11,
  3000. 'return_warped_data':True,
  3001. 'return_trimmed_data':False,
  3002. 'sessions_to_align':'all',
  3003. 'shuffle':False
  3004. }
  3005. ap_decoding_param_dict = {
  3006. 'exclude_positions':False,
  3007. 'pos_to_exclude_from':1000,
  3008. 'pos_to_exclude_to':1500,
  3009. ## TCA params ##
  3010. 'TCA_method': "ncp_hals", #"cp_als", "mcp_als", "ncp_bcd", "ncp_hals"
  3011. 'TCA_factors':'max',
  3012. 'TCA_replicates':10,
  3013. 'TCA_convergence_attempts':10, #Number of times TCA can fail before giving up
  3014. 'TCA_on_LDA_repetitions':15,
  3015. ## LDA params ##
  3016. 'LDA_imbalance_prop':.6,
  3017. 'LDA_imbalance_repetitions':1,
  3018. 'LDA_trial_shuffles':0,
  3019. 'LDA_session_shuffles':0,
  3020. 'session_comparisons':'BT'
  3021. }
  3022. mouse_list = preprocessing_param_dict['mouse_list']
  3023. session_list = preprocessing_param_dict['session_list']
  3024. ## TCA params ##
  3025. # TCA parameters
  3026. TCA_factors = ap_decoding_param_dict['TCA_factors']
  3027. ########## STEP 1 - PCA ###########
  3028. PCA_analysis_dict = perform_pca_on_multiple_mice_param_dict(preprocessing_param_dict)
  3029. ############## STEP 2: mCCA ############
  3030. CCA_analysis_dict = perform_mCCA_on_pca_dict_param_dict(PCA_analysis_dict, cca_param_dict)
  3031. num_bins = CCA_analysis_dict['num_bins']
  3032. mouse_list = CCA_analysis_dict['mouse_list']
  3033. num_bins = CCA_analysis_dict['num_bins']
  3034. #Prepare distance matrices
  3035. if distance_comparison == 'session_type':
  3036. distance_type_names = pparam.SESSION_TYPE_LABELS
  3037. distance_type_colors = pparam.SESSION_TYPE_COLORS
  3038. distance_type_num = len(distance_type_names)
  3039. distance_type_names_short = [n[0] for n in distance_type_names]
  3040. elif distance_comparison == 'airpuff':
  3041. distance_type_names = ['No AP', 'AP']
  3042. distance_type_colors = pparam.SESSION_TYPE_COLORS[:2]
  3043. distance_type_num = len(distance_type_names)
  3044. distance_type_names_short = distance_type_names
  3045. elif distance_comparison == 'single_session':
  3046. distance_type_names = pparam.SESSION_NAMES
  3047. cmap = mpl.cm.get_cmap('Set2')
  3048. distance_type_colors = [cmap(num) for num in np.linspace(0,1,9)]
  3049. distance_type_num = len(distance_type_names)
  3050. distance_type_names_short = distance_type_names
  3051. distance_dict_by_mouse_and_type = {(mnum, type1_idx, type2_idx):[] for mnum in mouse_list for type1_idx in range(distance_type_num) for type2_idx in range(distance_type_num)}
  3052. distance_dict_by_mouse_and_type_shuffle = {(mnum, type1_idx, type2_idx):[] for mnum in mouse_list for type1_idx in range(distance_type_num) for type2_idx in range(distance_type_num)}
  3053. for midx, mnum in enumerate(mouse_list):
  3054. print('Performing TCA+LDA on M%d'%mnum)
  3055. session_list = CCA_analysis_dict[mnum, 'session_list']
  3056. pos_list = CCA_analysis_dict[mnum, 'pos']
  3057. pca_list = CCA_analysis_dict[mnum, 'pca']
  3058. # pca_list = CCA_analysis_dict[mnum, 'pca_unaligned']
  3059. # session_list_to_decode = [snum for snum in session_list if snum in ap_decoding_param_dict['sessions_to_decode']]
  3060. data_by_trial, pos_by_trial, snum_by_trial = pf.reshape_pca_list_by_trial(pca_list, pos_list, num_bins, session_list)
  3061. print(mnum, data_by_trial.shape)
  3062. #Limit position (if indicated)
  3063. if ap_decoding_param_dict['exclude_positions'] == True:
  3064. positions = pos_by_trial[:,0] #Assumes all trials are binned using the same positions
  3065. pos_bool = np.invert(np.bitwise_and(positions > ap_decoding_param_dict['pos_to_exclude_from'], positions < ap_decoding_param_dict['pos_to_exclude_to']))
  3066. data_by_trial = data_by_trial[:, pos_bool, :]
  3067. pos_by_trial = pos_by_trial[pos_bool, :]
  3068. #Selecting trials to decode
  3069. # ap_decoding_param_dict['session_comparisons'] = 'airpuff'
  3070. trials_to_keep, label_by_trial = APfuns.get_trials_to_keep_and_labels(snum_by_trial, ap_decoding_param_dict['session_comparisons'])
  3071. if TCA_factors == 'max':
  3072. TCA_dimensions = data_by_trial.shape[0]
  3073. TCA_counter_total = 0
  3074. for TCA_on_LDA_counter in range(ap_decoding_param_dict['TCA_on_LDA_repetitions']):
  3075. # #Step 4: TCA
  3076. KTensor = APfuns.perform_TCA(data_by_trial, TCA_dimensions, ap_decoding_param_dict['TCA_replicates'],
  3077. ap_decoding_param_dict['TCA_method'], ap_decoding_param_dict['TCA_convergence_attempts'])
  3078. trial_factors = KTensor[2]
  3079. TCA_counter_total += 1
  3080. ############## STARTING THE DISTANCE MEASUREMENTS ###################
  3081. fig_num = 1
  3082. fs = 15
  3083. trial_type_by_trial = np.array([pparam.SESSION_TYPE_LABELS.index(pparam.SESSION_LABEL_BY_SNUM[snum]) for snum in snum_by_trial])
  3084. AP_by_trial = pparam.get_AP_labels_from_snum_by_trial(snum_by_trial)
  3085. # stype_by_trial = stype_by_trial[trials_to_keep]
  3086. if distance_comparison == 'session_type':
  3087. label_by_trial = trial_type_by_trial
  3088. elif distance_comparison == 'airpuff':
  3089. label_by_trial = AP_by_trial
  3090. elif distance_comparison == 'single_session':
  3091. label_by_trial = snum_by_trial
  3092. #Center trial factors
  3093. trial_factors, _, _ = pf.normalize_data(trial_factors, axis=0)
  3094. # trial_factors_clipped = np.zeros(trial_factors.shape)
  3095. #Get distance from each trial to cluster average
  3096. trial_type_unique = np.unique(label_by_trial)
  3097. trial_type_num = len(trial_type_unique)
  3098. num_trials = len(label_by_trial)
  3099. for i in range(1+distance_shuffles):
  3100. if i==0:
  3101. label_by_trial_current = np.copy(label_by_trial)
  3102. else:
  3103. shuffle_idxs = np.random.choice(range(len(label_by_trial)), size=len(label_by_trial), replace=False)
  3104. label_by_trial_current = label_by_trial[shuffle_idxs]
  3105. label_by_trial_current = np.random.choice(np.unique(label_by_trial), size=len(label_by_trial), replace=True)
  3106. avg_samples_per_class = int(np.floor(len(label_by_trial)/len(np.unique(label_by_trial))))
  3107. label_by_trial_new = np.zeros(len(label_by_trial), dtype=type(label_by_trial[0]))
  3108. for dtype_idx,dtype in enumerate(np.unique(label_by_trial)):
  3109. if dtype != np.unique(label_by_trial)[-1]:
  3110. label_by_trial_new[((dtype_idx)*avg_samples_per_class):((dtype_idx+1)*avg_samples_per_class)] = dtype
  3111. else:
  3112. label_by_trial_new[((dtype_idx)*avg_samples_per_class):] = dtype
  3113. shuffle_idxs = np.random.choice(range(len(label_by_trial_new)), size=len(label_by_trial_new), replace=False)
  3114. label_by_trial_current = label_by_trial_new[shuffle_idxs]
  3115. #Method 1: compare to centroid
  3116. distance_array = np.zeros((num_trials, trial_type_num))
  3117. class_centers = np.zeros((trial_type_num, TCA_dimensions))
  3118. for trial_type_idx, trial_type in enumerate(trial_type_unique):
  3119. trial_bool = label_by_trial_current == trial_type
  3120. trial_factor_type = trial_factors[trial_bool]
  3121. trial_type_center = np.average(trial_factor_type, axis=0)
  3122. class_centers[trial_type_idx] = trial_type_center
  3123. for trial in range(num_trials):
  3124. for trial_type_index, trial_type in enumerate(trial_type_unique):
  3125. tf = trial_factors[trial]
  3126. center = class_centers[trial_type_index]
  3127. d = np.linalg.norm(tf-center)
  3128. distance_array[trial, trial_type_index] = d
  3129. # #Take out outliers (IMPROVE!)
  3130. distances_ordered = np.sort(distance_array.ravel())
  3131. max_clip_distance = distances_ordered[int(distances_ordered.size * 0.95)] #Get 95th percentile
  3132. distance_array = np.clip(distance_array, 0, max_clip_distance)
  3133. # #Z-score
  3134. # if i == 0:
  3135. # dmean = np.mean(distance_array)
  3136. # dmean = np.mean(distance_array)
  3137. # distance_array = distance_array/dmean #+ 1
  3138. distance_array = 1 + (distance_array - np.average(distance_array))/np.std(distance_array)
  3139. #Get distances of each trial to each
  3140. for center_type_idx, center_type in enumerate(trial_type_unique):
  3141. for trial in range(num_trials):
  3142. trial_type = label_by_trial_current[trial]
  3143. d = distance_array[trial, center_type_idx]
  3144. if i == 0:
  3145. distance_dict_by_mouse_and_type[mnum, trial_type, center_type_idx].append(d)
  3146. else:
  3147. distance_dict_by_mouse_and_type_shuffle[mnum, trial_type, center_type_idx].append(d)
  3148. mtype_by_mouse = np.array([pparam.MOUSE_TYPE_LABEL_BY_MOUSE[mnum] for mnum in mouse_list])
  3149. mtype_unique = np.sort(np.unique(mtype_by_mouse))[::-1]
  3150. midxs_by_type = []
  3151. for mtype in mtype_unique:
  3152. midxs_by_type.append([midx for midx in range(len(mouse_list)) if mtype_by_mouse[midx] == mtype])
  3153. #Get within-between array
  3154. distance_dict_by_mouse_dtype_and_belonging = {(mnum, dtype, belonging_idx):[] for mnum in mouse_list
  3155. for dtype in range(distance_type_num)
  3156. for belonging_idx in range(2)}
  3157. distance_dict_by_mouse_dtype_and_belonging_shuffle = {(mnum, dtype, belonging_idx):[] for mnum in mouse_list
  3158. for dtype in range(distance_type_num)
  3159. for belonging_idx in range(2)}
  3160. #Get this information by mouse type
  3161. distance_dict_by_mtype_dtype_and_belonging = {(mtype, dtype, belonging_idx):[] for mtype in mtype_unique
  3162. for dtype in range(distance_type_num)
  3163. for belonging_idx in range(2)}
  3164. distance_dict_by_mtype_dtype_and_belonging_shuffle = {(mtype, dtype, belonging_idx):[] for mtype in mtype_unique
  3165. for dtype in range(distance_type_num)
  3166. for belonging_idx in range(2)}
  3167. #Difference between outside and inside
  3168. distance_diff_dict_by_mtype_dtype = {}
  3169. distance_diff_dict_by_mtype_dtype_shuffle = {}
  3170. for midx, mnum in enumerate(mouse_list):
  3171. mtype = mtype_by_mouse[midx]
  3172. for dtype1 in range(distance_type_num):
  3173. ddiff_array = np.zeros(len(distance_dict_by_mouse_and_type[mnum, dtype1, dtype1]))
  3174. ddiff_array_shuffle = np.zeros(len(distance_dict_by_mouse_and_type_shuffle[mnum, dtype1, dtype1]))
  3175. for dtype2 in range(distance_type_num):
  3176. belonging_idx = 1 - (dtype1 == dtype2) #0: within, 1:between
  3177. dlist = distance_dict_by_mouse_and_type[mnum, dtype1, dtype2]
  3178. distance_dict_by_mouse_dtype_and_belonging[mnum, dtype1, belonging_idx].extend(list(dlist))
  3179. distance_dict_by_mtype_dtype_and_belonging[mtype, dtype1, belonging_idx].extend(list(dlist))
  3180. dlist_shuffle = distance_dict_by_mouse_and_type_shuffle[mnum, dtype1, dtype2]
  3181. distance_dict_by_mouse_dtype_and_belonging_shuffle[mnum, dtype1, belonging_idx].extend(list(dlist_shuffle))
  3182. distance_dict_by_mtype_dtype_and_belonging_shuffle[mtype, dtype1, belonging_idx].extend(list(dlist_shuffle))
  3183. if belonging_idx != 0:
  3184. ddiff_array += (np.array(dlist) - np.array(distance_dict_by_mouse_and_type[mnum, dtype1, dtype1]))
  3185. ddiff_array_shuffle += (np.array(dlist_shuffle) - np.array(distance_dict_by_mouse_and_type_shuffle[mnum, dtype1, dtype1]))
  3186. ddiff_array /= (distance_type_num-1)
  3187. ddiff_array_shuffle /= (distance_type_num-1)
  3188. if (mtype, dtype1) not in distance_diff_dict_by_mtype_dtype.keys():
  3189. distance_diff_dict_by_mtype_dtype[mtype, dtype1] = (list(ddiff_array))
  3190. distance_diff_dict_by_mtype_dtype_shuffle[mtype, dtype1] = (list(ddiff_array_shuffle))
  3191. else:
  3192. distance_diff_dict_by_mtype_dtype[mtype, dtype1].extend(list(ddiff_array))
  3193. distance_diff_dict_by_mtype_dtype_shuffle[mtype, dtype1].extend(list(ddiff_array_shuffle))
  3194. #Dist diffs overall
  3195. distance_diff_dict_by_mtype = {}
  3196. distance_diff_dict_by_mtype_shuffle = {}
  3197. for mtype_idx, mtype in enumerate(mtype_unique):
  3198. distance_diff_dict_by_mtype[mtype] = []
  3199. distance_diff_dict_by_mtype_shuffle[mtype] = []
  3200. for dtype in range(distance_type_num):
  3201. distance_diff_dict_by_mtype[mtype].extend(distance_diff_dict_by_mtype_dtype[mtype, dtype])
  3202. distance_diff_dict_by_mtype_shuffle[mtype].extend(distance_diff_dict_by_mtype_dtype_shuffle[mtype, dtype])
  3203. #Get colormap array
  3204. distance_matrix_by_mouse_avg = np.zeros((len(mouse_list), distance_type_num, distance_type_num))
  3205. distance_matrix_by_mouse_std = np.zeros((len(mouse_list), distance_type_num, distance_type_num))
  3206. distance_matrix_by_mouse_avg_shuffle = np.zeros((len(mouse_list), distance_type_num, distance_type_num))
  3207. for midx, mnum in enumerate(mouse_list):
  3208. for dtype1 in range(distance_type_num):
  3209. for dtype2 in range(distance_type_num):
  3210. dlist = distance_dict_by_mouse_and_type[mnum, dtype1, dtype2]
  3211. distance_matrix_by_mouse_avg[midx, dtype1, dtype2] = np.average(dlist)
  3212. distance_matrix_by_mouse_std[midx, dtype1, dtype2] = np.std(dlist)/np.sqrt(len(dlist))
  3213. dlist = distance_dict_by_mouse_and_type_shuffle[mnum, dtype1, dtype2]
  3214. distance_matrix_by_mouse_avg_shuffle[midx, dtype1, dtype2] = np.average(dlist)
  3215. #Plot 1 - colormap by mouse
  3216. fig_cmap, axs_cmap = plt.subplots(2, 4, num=fig_num, figsize=(7,4)); fig_num += 1
  3217. fs = 12
  3218. vmin = np.min(distance_matrix_by_mouse_avg)
  3219. vmax = np.max(distance_matrix_by_mouse_avg)
  3220. for midx, mnum in enumerate(mouse_list):
  3221. ax = axs_cmap.ravel()[midx]
  3222. darray = distance_matrix_by_mouse_avg[midx]
  3223. ax.imshow(darray, cmap=pparam.DISTANCE_CMAP, vmin=vmin, vmax=vmax)
  3224. ax.set_xticks(range(distance_type_num), distance_type_names_short, fontsize=fs)
  3225. ax.set_yticks(range(distance_type_num), distance_type_names_short, fontsize=fs)
  3226. ax.set_title('M%d'%mnum, fontsize = fs+3, pad=0)
  3227. axs_cmap[0,0].set_ylabel('From', fontsize=fs+4)
  3228. axs_cmap[1,0].set_xlabel('To', fontsize=fs+4)
  3229. ### Add colorbar in the plot
  3230. # fig_cmap.subplots_adjust(right=0.90)
  3231. # cbar_ax = fig_cmap.add_axes([1.05, 0.17, 0.035, 0.7])
  3232. # norm = mpl.colors.Normalize(vmin=vmin, vmax=vmax)
  3233. # cbar = fig_cmap.colorbar(mpl.cm.ScalarMappable(norm=norm, cmap=pparam.DISTANCE_CMAP), cax=cbar_ax, orientation='vertical', fraction=0.01)
  3234. # cbar.ax.set_ylabel('Trial distance (norm.)', fontsize=fs+7, rotation=270, labelpad=25)
  3235. # cbar.ax.tick_params(axis='both', which='major', labelsize=fs+5)
  3236. fig_cmap.subplots_adjust(wspace=0, hspace=0)
  3237. fig_cmap.tight_layout()
  3238. fig_name = 'FigSI_distances_colormap_by_mouse_%s'%distance_comparison
  3239. save_figure(fig_cmap, fig_name)
  3240. #Plot 2 - colormap by mouse type
  3241. fig_cmap, axs_cmap = plt.subplots(1, 2, num=fig_num, figsize=(5,3)); fig_num += 1
  3242. fs = 16
  3243. for midx_list_idx, midx_list in enumerate(midxs_by_type):
  3244. ax = axs_cmap.ravel()[midx_list_idx]
  3245. darray = distance_matrix_by_mouse_avg[midx_list]
  3246. darray = np.average(darray, axis=0)
  3247. ax.imshow(darray, cmap=pparam.DISTANCE_CMAP, vmin=vmin, vmax=vmax)
  3248. ax.set_xticks(range(distance_type_num), distance_type_names_short, fontsize=fs)
  3249. ax.set_yticks(range(distance_type_num), distance_type_names_short, fontsize=fs)
  3250. ax.set_title(r"$\bf{%s}$"%mtype_unique[midx_list_idx], fontsize = fs+3, pad=10)
  3251. for dtype_idx in range(darray.shape[0]):
  3252. ax.add_patch(Rectangle((dtype_idx-0.5, dtype_idx-0.5), 1, 1, fill=False, edgecolor='black', lw=1))
  3253. axs_cmap[0].set_ylabel(r"$\bf{}$From", fontsize=fs+4)
  3254. axs_cmap[0].set_xlabel(r"$\bf{}$To", fontsize=fs+4)
  3255. fig_cmap.subplots_adjust(wspace=0, hspace=0)
  3256. fig_cmap.tight_layout()
  3257. fig_name = 'Fig4_distances_colormap_by_mtype_%s'%distance_comparison
  3258. save_figure(fig_cmap, fig_name)
  3259. ### Make colorbar separately
  3260. plt.figure(fig_num); fig_num += 1
  3261. fig_cbar = plt.gcf()
  3262. vmin_cbar, vmax_cbar = np.around([vmin, vmax], decimals=1)
  3263. cbar_ticks = np.linspace(vmin_cbar, vmax_cbar, num=4)
  3264. cbar_ticks = np.around(cbar_ticks, decimals=2)
  3265. vmin_cbar = cbar_ticks[0]; vmax_cbar = cbar_ticks[-1]
  3266. cbar = pf.add_distance_cbar(fig_cbar, pparam.DISTANCE_CMAP, vmin = vmin_cbar, vmax = vmax_cbar, fs=fs,
  3267. cbar_label = '',
  3268. cbar_kwargs = {'fraction':0.555, 'pad':0.04, 'aspect':15}
  3269. )
  3270. cbar.ax.set_yticks(cbar_ticks)
  3271. cbar.ax.tick_params(axis='y', labelsize=25)
  3272. cbar.ax.set_ylabel('Trial distance (norm.)', fontsize=fs+7, rotation=270, labelpad=25)
  3273. fig_name = 'Fig4_distances_colormap_colorbar_%s'%distance_comparison
  3274. save_figure(fig_cbar, fig_name)
  3275. #Plot 2.2 - Within distance time evolution
  3276. fs = 20
  3277. fig = plt.figure(fig_num); fig_num+=1
  3278. ax = plt.gca()
  3279. xx = range(len(np.unique(label_by_trial_current)))
  3280. delta_vals = {} #{(mtype, dtype): [delta of avg]}
  3281. delta_ypos_vals = {dtype:[] for dtype in trial_type_unique[:-1]} #{(dtype): [delta of avg]}
  3282. max_ypos_by_mtype = {mtype:0 for mtype in range(len(midxs_by_type))}
  3283. max_ypos = 0
  3284. for midx_list_idx, midx_list in enumerate(midxs_by_type):
  3285. mtype = mtype_unique[midx_list_idx]
  3286. color = pparam.MOUSE_TYPE_COLORS_BY_MTYPE[mtype]
  3287. avg_list = np.zeros(len(trial_type_unique))
  3288. std_list = np.zeros(len(trial_type_unique))
  3289. for trial_type_idx, trial_type in enumerate(trial_type_unique):
  3290. dds = distance_dict_by_mtype_dtype_and_belonging[mtype, trial_type_idx, 0]
  3291. avg_list[trial_type_idx] = np.average(dds)
  3292. std_list[trial_type_idx] = np.std(dds)/np.sqrt(len(dds))
  3293. ax.plot(xx, avg_list, 'o-', lw=3, color=color, label=mtype)
  3294. ax.fill_between(xx, avg_list-std_list, avg_list+std_list, color=color, alpha=0.5)
  3295. max_yval = np.max(np.array(avg_list+std_list))
  3296. max_ypos = np.maximum(max_ypos, max_yval)
  3297. max_ypos_by_mtype[midx_list_idx] = max_yval
  3298. for dtype_idx, dtype in enumerate(trial_type_unique[:-1]):
  3299. #Delta stat significances (compare the differences with probe)
  3300. delta_of_avg_list = []
  3301. for midx in midx_list:
  3302. mnum = mouse_list[midx]
  3303. d = distance_dict_by_mouse_dtype_and_belonging[mnum, dtype, 0]
  3304. dprobe = distance_dict_by_mouse_dtype_and_belonging[mnum, trial_type_unique[-1], 0]
  3305. delta_of_avg = np.average(d) - np.average(dprobe)
  3306. # deltas = [d1-d2 for d1 in d for d2 in dprobe]
  3307. # delta_of_avg = np.average(deltas)
  3308. delta_of_avg_list.append(delta_of_avg)
  3309. delta_vals[mtype,dtype] = delta_of_avg_list
  3310. if distance_comparison == 'session_type':
  3311. #Plot delta stat significances (compare the differences with probe)
  3312. for dtype_idx, dtype in enumerate(trial_type_unique[:-1]):
  3313. mtype = mtype_unique[midx_list_idx]
  3314. delta_list_id = delta_vals[MOUSE_TYPE_LABELS[0], dtype]
  3315. delta_list_dd = delta_vals[MOUSE_TYPE_LABELS[1], dtype]
  3316. # tstat, pval = scipy.stats.ttest_ind(delta_list_id, delta_list_dd, equal_var=True, permutations=None, alternative='greater')
  3317. tstat, pval = scipy.stats.mannwhitneyu(delta_list_id, delta_list_dd, use_continuity=False, alternative='greater')
  3318. # p00 = np.max(delta_ypos_vals[dtype])
  3319. p00 = max_ypos
  3320. p10 = xx[dtype_idx]; p11 = p10
  3321. d0 = 0; dp = 0; label_padding = 0
  3322. pf.draw_significance(ax, pval, p00, p10, p10, d0, dp, orientation='top', label_padding=label_padding, thresholds = [0.05], fs=fs+15)
  3323. print(delta_list_id, delta_list_dd)
  3324. print("pval", dtype, pval)
  3325. elif distance_comparison == 'airpuff':
  3326. #Just compare NoAP with AP for both mouse types
  3327. for midx_list_idx, midx_list in enumerate(midxs_by_type):
  3328. lists_of_values = []
  3329. for dtype_idx, dtype in enumerate(trial_type_unique):
  3330. vals = distance_dict_by_mtype_dtype_and_belonging[mtype, dtype_idx, 0]
  3331. lists_of_values.append(vals)
  3332. tstat, pval = scipy.stats.mannwhitneyu(lists_of_values[0], lists_of_values[1],
  3333. use_continuity=False, alternative='two-sided')
  3334. p00 = max_ypos_by_mtype[midx_list_idx]
  3335. p10 = xx[0]; p11 = p10
  3336. d0 = 0; dp = 0; label_padding = 0
  3337. pf.draw_significance(ax, pval, p00, p10, p10, d0, dp, orientation='top', label_padding=label_padding, thresholds = [0.05], fs=fs+15)
  3338. print(np.average(lists_of_values[0]), np.average(lists_of_values[1]))
  3339. print("pval", dtype, pval)
  3340. xlabels = [distance_type_names[i] for i in trial_type_unique]
  3341. ax.set_xticks(xx, xlabels, fontsize=fs)
  3342. ax.set_ylabel('Within class distances (norm.)', fontsize=fs-5)
  3343. ax.tick_params(axis='y', which='major', labelsize=fs)
  3344. ax.legend(fontsize=fs, frameon=False)
  3345. # ax.set_ylim([0.5, ax.get_ylim()[1]])
  3346. # ax.set_title("", fontsize=fs)
  3347. ax.spines[['right', 'top']].set_visible(False)
  3348. for axis in ['top','bottom','left','right']:
  3349. ax.spines[axis].set_linewidth(3)
  3350. fig.tight_layout()
  3351. fig_name = 'FigSI_within_distances_by_dtype_%s'%(distance_comparison)
  3352. save_figure(fig, fig_name)
  3353. #Plot 3: Shuffle colormap
  3354. fig_cmap, axs_cmap = plt.subplots(1, 2, num=fig_num, figsize=(5,3)); fig_num += 1
  3355. fs = 18
  3356. for midx_list_idx, midx_list in enumerate(midxs_by_type):
  3357. ax = axs_cmap.ravel()[midx_list_idx]
  3358. darray = distance_matrix_by_mouse_avg_shuffle[midx_list]
  3359. darray = np.average(darray, axis=0)
  3360. ax.imshow(darray, cmap=pparam.DISTANCE_CMAP, vmin=vmin, vmax=vmax)
  3361. ax.set_xticks(range(distance_type_num), distance_type_names_short, fontsize=fs)
  3362. ax.set_yticks(range(distance_type_num), distance_type_names_short, fontsize=fs)
  3363. ax.set_title(r"$\bf{%s}$"%mtype_unique[midx_list_idx], fontsize = fs+3, pad=10)
  3364. for dtype_idx in range(darray.shape[0]):
  3365. ax.add_patch(Rectangle((dtype_idx-0.5, dtype_idx-0.5), 1, 1, fill=False, edgecolor='black', lw=1))
  3366. axs_cmap[0].set_ylabel('From', fontsize=fs+4)
  3367. axs_cmap[0].set_xlabel('To', fontsize=fs+4)
  3368. fig_cmap.subplots_adjust(wspace=0, hspace=0)
  3369. fig_cmap.tight_layout()
  3370. fig_name = 'Fig4_distances_colormap_by_mtype_shuffle_%s'%distance_comparison
  3371. save_figure(fig_cmap, fig_name)
  3372. #Plot 4: Barplots, distance by type, includes shuffle
  3373. for midx_list_idx, midx_list in enumerate(midxs_by_type):
  3374. mtype = mtype_unique[midx_list_idx]
  3375. ## Within-without (by class and with shuffle)
  3376. fig = plt.figure(fig_num, figsize=(5,4)); fig_num+=1
  3377. ax = plt.gca()
  3378. all_bars_width = 0.75
  3379. num_of_bars = 4
  3380. barwidth = all_bars_width/(num_of_bars)
  3381. xpos_list = range(len(np.unique(label_by_trial_current)))
  3382. for trial_type_idx, trial_type in enumerate(trial_type_unique):
  3383. for belonging_idx, belonging in enumerate(pparam.DISTANCE_LABELS):
  3384. for analysis_type_idx, analysis_type in enumerate(['Normal', 'Shuffle']):
  3385. if analysis_type_idx == 0:
  3386. dlist = distance_dict_by_mtype_dtype_and_belonging[mtype_unique[midx_list_idx], trial_type_idx, belonging_idx]
  3387. color = distance_type_colors[trial_type_idx]
  3388. pltlabel = [None, belonging][trial_type_idx==0]
  3389. else:
  3390. dlist = distance_dict_by_mtype_dtype_and_belonging_shuffle[mtype_unique[midx_list_idx], trial_type_idx, belonging_idx]
  3391. color = pparam.SHUFFLE_DEFAULT_COLOR
  3392. pltlabel = [None, 'Shuffle'][trial_type_idx==0 and analysis_type_idx == 0]
  3393. xpos = xpos_list[trial_type_idx] - all_bars_width/2 + barwidth/2 + barwidth * belonging_idx + 2 * barwidth * analysis_type_idx
  3394. avg = np.average(dlist)
  3395. err = np.std(dlist)/np.sqrt(len(dlist))
  3396. # err = np.std(dlist)
  3397. alpha = [1, 0.5][belonging_idx == 1]
  3398. ax.bar(xpos, avg, width=barwidth, alpha=alpha, edgecolor=None, color=color, label=pltlabel)
  3399. ax.errorbar([xpos], avg, yerr=err, fmt='', markersize=35, markeredgewidth=5, elinewidth = 5, zorder=2,
  3400. color=color, alpha=alpha)
  3401. xlabels = [distance_type_names_short[i] for i in trial_type_unique]
  3402. ax.set_xticks(xpos_list, xlabels, fontsize=fs)
  3403. ax.set_ylabel('Trial factor distances', fontsize=fs)
  3404. ax.tick_params(axis='y', which='major', labelsize=fs)
  3405. ax.legend(fontsize=fs-3, frameon=False)
  3406. # ax.set_ylim([0.5, ax.get_ylim()[1]])
  3407. ax.set_title(r"$\bf{%s}$"%mtype_unique[midx_list_idx], fontsize=fs)
  3408. ax.spines[['right', 'top']].set_visible(False)
  3409. for axis in ['top','bottom','left','right']:
  3410. ax.spines[axis].set_linewidth(3)
  3411. fig.tight_layout()
  3412. fig_name = 'FigSI_distances_by_dtype_%s_%s'%(mtype, distance_comparison)
  3413. save_figure(fig, fig_name)
  3414. #Plot 5: Distance difference barplots by dtype
  3415. fs = 20
  3416. for midx_list_idx, midx_list in enumerate(midxs_by_type):
  3417. mtype = mtype_unique[midx_list_idx]
  3418. ## Within-without (by class and with shuffle)
  3419. fig = plt.figure(fig_num, figsize=(4

main_figures.py at commit 1221cc1, no license · at the source

Overview

Authors: Albert Miguel-López1, Negar Nikbahkt1, Carlos Wert-Carvajal1, Lena Johanna Gschossmann1, Martin Pofahl1, Heinz Beck1, Tatjana Tchumatchenko1
  1. Universität Bonn, Universitätsk-linikum Bonn, Institute for Experimental Epileptology and Cognition Research, Bonn 53127, Germany
Institutions: University of Bonn (Germany); University Hospital Bonn (Germany)
Dates: received 3 July 2025; accepted 28 March 2026; published online 12 June 2026; in print 16 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1073/pnas.2517639123 · PMID 42284325 · PMCID PMC13273363 · OpenAlex W4408279499
Open access: hybrid, a free copy (OpenAlex)
Status: code verified
Categories: mouse (organism), systems (subfield)
Methods: Statistics, Smoothing, state filtering, decompositions, Machine learning, Preprocessing, Single-unit activity, calcium imaging
Keywords: hippocampal circuit, computation, neural activity, spatial coding
MeSH: CA3 Region, Hippocampal*, Animals, Axons, Male, Mice (* major topic)
Topic: Constraint Satisfaction and Optimization (Computer Networks and Communications, Computer Science), according to OpenAlex
Funding: Deutsche Forschungsgemeinschaft (1089)
Citations: not cited yet (Europe PMC); 32 references in the paper
Notices: A correction to this paper has been published (42766758, from Europe PMC)

Abstract

Hippocampal circuits form cognitive maps that represent spatial position and integrate contextual information, including affective cues, into episodic memory representations. We investigated how spatial and affective information are combined in the population activity of CA3 axons by analyzing the activity of intermediate-to-dorsal and dorsal-to-dorsal axons in mice navigating a linear track before, during, and after exposure to an aversive air puff stimulus. Both axonal populations maintained a robust, time-invariant activity manifold that encoded spatial information independent of affective context. Deformations of this common manifold encoded the presence of the aversive stimulus without disrupting the spatial representation. Despite differences in spatial coding, both axonal populations encoded affective information with similar efficacy. This population-level encoding was distributed similarly across place and nonplace cells. Our findings demonstrate that hippocampal CA3 axons integrate spatial and affective information within a common representational geometry while maintaining the separability of each information type.

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

amiguello/aversive_analysis_2025

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 1221cc1dfd287a1784c695a5331068011ff46f1a, 4 March 2026
Languages: Python (6)
Size: 15 files, 6 scripts
Software Heritage: not archived
Found in: “Data, Materials, and Software Availability”
Holds: README, environment (poetry.lock, pyproject.toml)
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: Matplotlib (6 files), NumPy (6 files), scikit-learn (4 files), SciPy (4 files), h5py (2 files), pandas (1 file), statsmodels (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
7 files

The paper's code and data availability statement is in the Data section.

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;
  • 6 scripts, each with its path and the digest of its content;
  • 10 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

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

Data

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

Data, Materials, and Software Availability

All experimental data used in this study have been previously published (13) and their online source links are incorporated into our analysis code which is publicly available (https://github.com/amiguello/aversive_analysis_2025.git). DOI will be generated upon acceptance. All other data are included in the manuscript and/or SI Appendix (https://www.pnas.org/doi/10.1073/pnas.2517639123#supplementary-materials).

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

Versions

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

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 7 authors, 4 keywords, 5 MeSH terms, 1 funder, 29 references, 1 integrity notice.

Cite

This paper

Miguel-López, A., Nikbahkt, N., Wert-Carvajal, C., Gschossmann, L. J., Pofahl, M., Beck, H., & Tchumatchenko, T. (2026). Transformations of the spatial activity manifold convey aversive information in CA3. Proceedings of the National Academy of Sciences of the United States of America, 123(24), e2517639123. https://doi.org/10.1073/pnas.2517639123

BibTeX

@article{miguellopez2026transformations,
author = {Miguel-López, Albert and Nikbahkt, Negar and Wert-Carvajal, Carlos and Gschossmann, Lena Johanna and Pofahl, Martin and Beck, Heinz and Tchumatchenko, Tatjana},
title = {{Transformations of the spatial activity manifold convey aversive information in CA3}},
journal = {Proceedings of the National Academy of Sciences of the United States of America},
year = {2026},
month = jun,
volume = {123},
number = {24},
pages = {e2517639123},
publisher = {National Academy of Sciences},
issn = {0027-8424},
doi = {10.1073/pnas.2517639123},
url = {https://doi.org/10.1073/pnas.2517639123},
pmid = {42284325},
pmcid = {PMC13273363}
}

RIS

TY - JOUR
AU - Miguel-López, Albert
AU - Nikbahkt, Negar
AU - Wert-Carvajal, Carlos
AU - Gschossmann, Lena Johanna
AU - Pofahl, Martin
AU - Beck, Heinz
AU - Tchumatchenko, Tatjana
TI - Transformations of the spatial activity manifold convey aversive information in CA3
T2 - Proceedings of the National Academy of Sciences of the United States of America
J2 - Proc Natl Acad Sci U S A
PY - 2026
DA - 2026/06/12
VL - 123
IS - 24
SP - e2517639123
SN - 0027-8424
PB - National Academy of Sciences
DO - 10.1073/pnas.2517639123
UR - https://doi.org/10.1073/pnas.2517639123
LA - en
ER -

CSL-JSON

{
"id": "10.1073/pnas.2517639123",
"type": "article-journal",
"title": "Transformations of the spatial activity manifold convey aversive information in CA3",
"container-title": "Proceedings of the National Academy of Sciences of the United States of America",
"author": [
{
"family": "Miguel-López",
"given": "Albert"
},
{
"family": "Nikbahkt",
"given": "Negar"
},
{
"family": "Wert-Carvajal",
"given": "Carlos"
},
{
"family": "Gschossmann",
"given": "Lena Johanna"
},
{
"family": "Pofahl",
"given": "Martin"
},
{
"family": "Beck",
"given": "Heinz"
},
{
"family": "Tchumatchenko",
"given": "Tatjana"
}
],
"container-title-short": "Proc Natl Acad Sci U S A",
"volume": "123",
"issue": "24",
"page": "e2517639123",
"DOI": "10.1073/pnas.2517639123",
"PMID": "42284325",
"PMCID": "PMC13273363",
"ISSN": "0027-8424",
"publisher": "National Academy of Sciences",
"URL": "https://doi.org/10.1073/pnas.2517639123",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
12
]
]
}
}

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.1126/sciadv.aec4911 [code]
Dendritic shaft constrictions shape synaptic integration in neurons.
Journal: Science advances
In common: statsmodels, pandas, SciPy, 2 other tools, mouse, 2 authors
[2] doi:10.1371/journal.pcbi.1013488 [code]
Evaluating place cell detection methods in Rats and Humans: Implications for cross-species spatial coding.
Journal: PLoS computational biology
In common: statsmodels, scikit-learn, pandas, 3 other tools, 4 references
[3] doi:10.1371/journal.pbio.3003824 [code]
Flexible goal learning involves coordinated population activity in dCA1 and medial orbitofrontal cortex.
Journal: PLoS biology
In common: scikit-learn, pandas, SciPy, 2 other tools, systems, 4 references
[4] doi:10.1371/journal.pbio.3003755 [code]
Action information is integrated into entorhinal representations of conceptual space and is reflected in eye movements.
Journal: PLoS biology
In common: statsmodels, pandas, SciPy, 2 other tools, 4 references
[5] doi:10.1038/s41467-026-77240-6
Prefrontal-thalamic goal states organize spatially aligned hippocampal maps.
Journal: Nature communications
In common: systems, 6 references
[6] doi:10.1038/s41467-026-75455-1 [code]
Shared latent representations of speech production for cross-patient speech decoding.
Journal: Nature communications
In common: h5py, statsmodels, scikit-learn, 4 other tools, 2 references
[7] doi:10.1038/s41593-026-02232-0 [code]
Entorhinal cortex represents task-relevant remote locations independently of CA1.
Journal: Nature neuroscience
In common: h5py, statsmodels, scikit-learn, 4 other tools, systems, mouse, 1 reference
[8] doi:10.1038/s41467-026-76106-1 [code]
Facial expression discrimination emerges from partially overlapping neural subspaces of detection and identity.
Journal: Nature communications
In common: h5py, scikit-learn, pandas, 3 other tools, systems, 2 references
[9] doi:10.1016/j.isci.2026.117375 [code]
Motor priming is associated with widespread recruitment into neural ensembles and more rapid ensemble transitions.
Journal: iScience
In common: h5py, statsmodels, scikit-learn, 4 other tools, systems, 1 reference
[10] doi:10.1016/j.neuron.2026.07.016 [code]
Inferring brain-wide interactions using data-constrained recurrent neural network models.
Journal: Neuron
In common: Matplotlib, NumPy, systems, mouse, 4 references

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.