OSCR

Linear and categorical coding units in the mouse gustatory cortex drive population dynamics and behavior in taste decision-making.

Code ↔ Paper

15 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 15 matches
  1. [1] § Results › Population dynamics during mixture-based decision-making ↔ dPCA_master/python/dPCA/dPCA.py, lines 21–93 · score 0.76 · demixed principal component, dimensionality reduction techniques, stimulus component, dPCA, population activity, variance
  2. [2] § Methods › Behavioral data ↔ utilities/utils.py, lines 702–743 · score 0.76 · curve_fit, upper asymptote, lower asymptote, slope, optimize, psychometric
  3. [3] § Methods › Statistical tests ↔ analyze-rnns-manuscript.ipynb, lines 1795–1837 · score 0.69 · AnovaRM, post hoc, Bonferroni corrections, ANOVAs, Scipy
  4. [4] § Results › Model perturbations reveal behavioral significance of coding unit types ↔ analyze-rnns-manuscript.ipynb, lines 3077–3143 · score 0.66 · post hoc, Bonferroni corrected, stimulus projection, stimulus component, ablating, perception
  5. [5] § Methods › t-SNE ↔ dPCA_master/python/dPCA/dPCA.py, lines 438–471 · score 0.65 · dimensional space, low dimensional, scikit-learn, mapping, Python, matrix
  6. [6] § Methods › Electrophysiological data acquisition ↔ histology/BrowsingFunctions/allen_ccf_npx.m, lines 49–89 · score 0.65 · Neuropixels trajectory explorer, Allen CCF, brain, probes, atlas
  7. [7] § Results › Population dynamics during mixture-based decision-making ↔ utilities/utils.py, lines 68–173 · score 0.60 · spike trains, lateral lick, central lick, firing rate, selectivity, bins
  8. [8] § Results › Population dynamics during mixture-based decision-making ↔ experimental-data-analysis-manuscript.ipynb, lines 920–1021 · score 0.56 · spike trains, lateral lick, central lick, warped, behavioral, neuron
  9. [9] § Methods › Firing rate data ↔ experimental-data-analysis-manuscript.ipynb, lines 55–110 · score 0.56 · inter event, lateral lick, IEI, bins
  10. [10] § Methods › Electrophysiological data acquisition ↔ histology/BrowsingFunctions/allen_ccf_npx_4shank.m, lines 1–39 · score 0.55 · Neuropixels trajectory, Allen CCF, Bregma, probes, atlas
  11. [11] § Results › Model perturbations reveal behavioral significance of coding unit types ↔ analyze-rnns-manuscript.ipynb, lines 3077–3143 · score 0.54 · Bonferroni adjusted, post hoc, pseudo population, ablating, perception, components
  12. [12] § Methods › RNN: training ↔ utilities/RNN.py, lines 582–715 · score 0.53 · PyTorch, Adam, clipping, gradient, optimizer, network
  13. [13] § Methods › Firing rate data ↔ analyze-rnns-manuscript.ipynb, lines 1155–1203 · score 0.52 · inter event interval, PSTHs
  14. [14] § Methods › RNN: training ↔ train-rnns-manuscript.ipynb, lines 102–188 · score 0.52 · stimulus window, decision window, pre, loss, neural, network
  15. [15] § Methods › Behavioral data ↔ experimental-data-analysis-manuscript.ipynb, lines 181–217 · score 0.51 · curve fit, sucrose choice, optimize, psychometric, scipy, stimulus

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

Jupyter notebook · 3,875 lines · 153 KB · MIT · 4 matches

  1. # %% [markdown]
  2. # # Analyze RNNs
  3. # %%
  4. from utilities.utils import *
  5. from utilities.RNN import *
  6. import numpy as np
  7. import scipy
  8. import random
  9. import torch
  10. from sklearn.manifold import TSNE
  11. import pickle
  12. import pandas as pd
  13. from statsmodels.stats.anova import AnovaRM
  14. import matplotlib
  15. from matplotlib import pyplot as plt
  16. from matplotlib.colors import LinearSegmentedColormap
  17. from dPCA_master.python.dPCA import dPCA
  18. # %% [markdown]
  19. # ### Load experimental data and variables
  20. # %%
  21. # we'll want the experimental data for comparison ... load it now
  22. try:
  23. # if train-rnns-manuscript was run, this file will exist
  24. with open('data/model/experimental_input.pkl', 'rb') as file:
  25. experimental_input = pickle.load(file)
  26. file.close()
  27. all_data = experimental_input['all_data'] # a dictionary with a key for each included session
  28. included_sessions = list(all_data.keys())
  29. n_sessions = len(included_sessions)
  30. params = experimental_input['parameters']
  31. bin_size = params['bin_size']
  32. padding = params['padding']
  33. smooth_type = params['smooth_type']
  34. smooth_width = params['smooth_width']
  35. except:
  36. # otherwise, create it from scratch
  37. ## load raw data
  38. data = scipy.io.loadmat('data/experimental/data_all.mat')
  39. ## parameters for data prep
  40. bin_size = 0.050 # [s]
  41. padding = 1 # [s]
  42. min_n_neurons = 3
  43. smooth_type = 'gaussian'
  44. smooth_width = 11
  45. # filter sessions
  46. n_sessions_all = data['data_all'].shape[1]
  47. included_sessions = [i for i in range(n_sessions_all) if len(data['data_all'][0, i][-1]) >= min_n_neurons]
  48. n_sessions = len(included_sessions)
  49. # time warping
  50. all_data = {}
  51. for i_session in included_sessions:
  52. print('Preparing session', i_session + 1, '...')
  53. all_data[i_session] = get_warped_data(data, i_session, bin_size, padding, mean_correct_only=False)
  54. print('Done.')
  55. # smoothing
  56. for session in included_sessions:
  57. X, y_conc, y_choice, y_outcome, t = all_data[session]
  58. X_smooth = X.copy()
  59. for i in range(X_smooth.shape[0]):
  60. for j in range(X_smooth.shape[1]):
  61. X_smooth[i, j, :] = smooth(X_smooth[i, j, :], smooth_type, smooth_width)
  62. all_data[session] = (X_smooth, y_conc, y_choice, y_outcome, t)
  63. ## calculate additional variables
  64. unique_stims = np.unique(all_data[0][1]) # same stimuli for all sessions
  65. n_stims = len(unique_stims)
  66. t = all_data[0][4] # same time for all sessions
  67. n_bins = len(t)
  68. ## save model inputs
  69. with open('data/model/experimental_input.pkl', 'wb') as file:
  70. pickle.dump({
  71. 'parameters': {
  72. 'bin_size': bin_size,
  73. 'padding': padding,
  74. 'smooth_type': smooth_type,
  75. 'smooth_width': smooth_width
  76. },
  77. 'all_data': all_data
  78. }, file)
  79. file.close()
  80. # %% [markdown]
  81. # ### Load model data and variables
  82. # %%
  83. # load pre-trained models
  84. with open('data/model/all_models.pkl', 'rb') as file:
  85. all_models = pickle.load(file) # a dictionary with a key for each included session
  86. file.close()
  87. with open('data/model/hyperparameters.pkl', 'rb') as file:
  88. hyp = pickle.load(file) # a dictionary with a key for each included session
  89. file.close()
  90. # re-define important variables for the models (inputs, etc.)
  91. _, y_conc, _, _, t = all_data[included_sessions[0]]
  92. unique_stims = np.unique(y_conc)
  93. n_stims = len(unique_stims)
  94. colors = np.hstack([
  95. np.reshape(np.linspace(0, 1, n_stims), (-1, 1)),
  96. np.reshape(np.linspace(1, 0, n_stims), (-1, 1)),
  97. np.reshape(np.linspace(0, 1, n_stims), (-1, 1))
  98. ])
  99. n_bins = len(t)
  100. input_size = hyp['input_size']
  101. stim_start = hyp['stim_start']
  102. stim_end = hyp['stim_end']
  103. t_stim_window = (t >= stim_start) & (t <= stim_end)
  104. t_D = t[-1] + bin_size / 2 - padding # decision time point
  105. decision_pre_time = hyp['decision_pre_time'] # pre-decision time point window
  106. t_decision_window = (t > t_D - decision_pre_time) & (t <= t_D)
  107. t_target_window = (t < 0) | t_decision_window
  108. inputs = np.zeros((n_stims, n_bins, input_size))
  109. for i_stim, stim in enumerate(unique_stims):
  110. inputs[i_stim, :, :] = np.tile( np.array([stim / 100, 1 - stim / 100]), (n_bins, 1) )
  111. inputs[:, ~t_stim_window, :] = 0 # stimulus is only present within [stim_start, stim_end]
  112. inputs = torch.from_numpy(inputs).to(dtype=torch.float32) # convert to tensor
  113. window_width = 4
  114. window_step = 4
  115. fit_options = [
  116. {
  117. 'shape' : 'linear',
  118. 'y_min' : 0,
  119. 'y_max' : np.inf,
  120. 'min_y_range' : 0
  121. },
  122. {
  123. 'shape' : 'step',
  124. 'midpoints' : [40, 50, 60],
  125. 'min_y_range' : 0
  126. }
  127. ]
  128. n_windows = 1 + (n_bins - window_width) // window_step
  129. inds_L = [i * window_step for i in range(n_windows)]
  130. inds_R = [ind + window_width for ind in inds_L]
  131. t_downsample = [np.mean(t[ind_L:ind_R]) for ind_L, ind_R in zip(inds_L, inds_R)]
  132. # %% [markdown]
  133. # ### Find a noise level for each model that puts its accuracy near its corresponding animal's
  134. # %%
  135. try:
  136. with open('data/model/inputs_sigma.pkl', 'rb') as file:
  137. inputs_sigma = pickle.load(file) # a dictionary with a key for each included session
  138. file.close()
  139. except:
  140. accs_animal = []
  141. accs_model = []
  142. inputs_sigma = {}
  143. for session in included_sessions:
  144. print('Session {}'.format(session))
  145. acc_animal = 100 * np.mean(all_data[session][3])
  146. accs_animal.append(acc_animal)
  147. print(' Animal accuracy: {:.1f}%'.format(acc_animal))
  148. sigma = 0.4
  149. output_temp = model_simulation(
  150. all_models[session],
  151. inputs,
  152. unique_stims,
  153. t_decision_window,
  154. n_trials_per_stim=60,
  155. sigma=sigma,
  156. silenced=None,
  157. label_data=None,
  158. responses_over_time=False,
  159. fit_options=None,
  160. inds_L=None,
  161. inds_R=None,
  162. responses_beginning_end=False,
  163. window_beginning=None,
  164. window_end=None)
  165. acc_model = output_temp['accuracy']
  166. print(' sigma = {:.2f}, model accuracy: {:.1f}%'.format(sigma, acc_model))
  167. while np.abs(acc_model - acc_animal) > 5:
  168. if acc_model < acc_animal:
  169. sigma -= 0.05
  170. elif acc_model > acc_animal:
  171. sigma += 0.05
  172. output_temp = model_simulation(
  173. all_models[i_session],
  174. inputs,
  175. unique_stims,
  176. t_decision_window,
  177. n_trials_per_stim=60,
  178. sigma=sigma,
  179. silenced=None,
  180. label_data=None,
  181. responses_over_time=False,
  182. fit_options=None,
  183. inds_L=None,
  184. inds_R=None,
  185. responses_beginning_end=False,
  186. window_beginning=None,
  187. window_end=None)
  188. acc_model = output_temp['accuracy']
  189. print(' sigma = {:.2f}, model accuracy: {:.1f}%'.format(sigma, acc_model))
  190. accs_model.append(acc_model)
  191. inputs_sigma[session] = sigma
  192. with open('data/model/inputs_sigma.pkl', 'wb') as file:
  193. pickle.dump(inputs_sigma, file)
  194. file.close()
  195. # %%
  196. print('Model noise levels:')
  197. print(' Min: {:.4f}'.format(np.min(list(inputs_sigma.values()))))
  198. print(' Max: {:.4f}'.format(np.max(list(inputs_sigma.values()))))
  199. print(' Mean: {:.4f}'.format(np.mean(list(inputs_sigma.values()))))
  200. # %% [markdown]
  201. # ### Run 'control' (no ablations) simulations using these noise parameters
  202. # %%
  203. try:
  204. with open('data/model/output_tailored_noise.pkl', 'rb') as file:
  205. output_noise = pickle.load(file)
  206. file.close()
  207. except:
  208. output_noise = multi_model_simulation(
  209. list(all_models.values()),
  210. inputs,
  211. unique_stims,
  212. t_decision_window,
  213. n_trials_per_stim=20,
  214. sigma=list(inputs_sigma.values()),
  215. silenced=None,
  216. label_data=None,
  217. responses_over_time=True,
  218. fit_options=fit_options,
  219. inds_L=inds_L,
  220. inds_R=inds_R,
  221. responses_beginning_end=True,
  222. window_beginning=((t >= stim_start) & (t <= stim_start + 0.5)),
  223. window_end=((t >= t_D - 0.5) & (t <= t_D)))
  224. with open('data/model/output_tailored_noise.pkl', 'wb') as file:
  225. pickle.dump(output_noise, file)
  226. file.close()
  227. # %% [markdown]
  228. # ### Analyze the behavioral performance
  229. # %%
  230. fig, ax = plt.subplots(figsize=(4.5, 4));
  231. p_left_mat = np.empty([0, n_stims])
  232. p_left_mat_exp = np.empty([0, n_stims])
  233. for i_session, session in enumerate(included_sessions):
  234. p_left = output_noise['all_p_left'][i_session]
  235. p_left_mat = np.vstack([p_left_mat, p_left])
  236. p_left_exp = np.array([all_data[session][2].flatten()[all_data[session][1].flatten() == stim].mean()
  237. for stim in unique_stims])
  238. p_left_mat_exp = np.vstack([p_left_mat_exp, p_left_exp])
  239. ave_p_left = p_left_mat.mean(axis=0)
  240. std_p_left = p_left_mat.std(axis=0) / np.sqrt(p_left_mat.shape[0])
  241. psycho_res = fit_psychometric(unique_stims, ave_p_left)
  242. ax.plot(psycho_res['x'], psycho_res['y'], 'b');
  243. for i_stim, stim in enumerate(unique_stims):
  244. ax.plot([stim, stim], [ave_p_left[i_stim] - std_p_left[i_stim], ave_p_left[i_stim] + std_p_left[i_stim]], 'k');
  245. ax.plot(stim, ave_p_left[i_stim], '.k', markersize=10);
  246. ax.set_ylim([-0.05, 1.05]);
  247. ax.set_xlabel('Stimulus (% sucrose)');
  248. ax.set_ylabel('P(Sucrose choice)');
  249. accs = output_noise['all_accuracies']
  250. accs_exp = []
  251. for session in included_sessions:
  252. mask = np.array([s in unique_stims for s in all_data[session][1]])
  253. accs_exp.append(all_data[session][3][mask].mean())
  254. ax.set_title('Accuracy: {:.1f}% +/- {:.1f}%'.format(100 * np.mean(accs), 100 * np.std(accs) / np.sqrt(n_sessions)));
  255. # Figure 5B
  256. # plt.savefig('plots/model/model_accuracy.pdf');
  257. # %%
  258. # compare experiment and model
  259. Y = np.array([p_left_mat.mean(axis=0), p_left_mat_exp.mean(axis=0)])
  260. psycho_res_model = fit_psychometric(unique_stims, Y[0, :])
  261. print('Model psychometric slope: {:.4f}'.format(psycho_res_model['slope']))
  262. psycho_res_exp = fit_psychometric(unique_stims, Y[1, :])
  263. print('Experiment psychometric slope: {:.4f}\n'.format(psycho_res_exp['slope']))
  264. res = psychometric_comparison_test(unique_stims, Y)
  265. print('Psychometric comparison:')
  266. print(' F = {:.4f}'.format(res['F']))
  267. print(' p = {:.4f}\n'.format(res['p']))
  268. print('Model mean accuracy: {:.4f}%'.format(100 * np.mean(accs)))
  269. print('Animal mean accuracy: {:.4f}%'.format(100 * np.mean(accs_exp)))
  270. res = scipy.stats.ttest_ind(accs, accs_exp)
  271. print('Accuracy comparison via t test: p = {:.4f}'.format(res.pvalue))
  272. res = scipy.stats.mannwhitneyu(accs, accs_exp)
  273. print('Accuracy comparison via rank sum: p = {:.4f}'.format(res.pvalue))
  274. # %% [markdown]
  275. # ### Analyze goodness of fit to neural data
  276. # %%
  277. ys = []
  278. y_hats = []
  279. for session in included_sessions:
  280. X, y_conc, y_choice, y_outcome, _ = all_data[session]
  281. n_neurons = X.shape[1]
  282. y = np.zeros((n_stims, n_bins, n_neurons))
  283. for i_stim, stim in enumerate(unique_stims):
  284. trial_mask = (y_conc == stim) & (y_outcome == 1)
  285. for i_neuron in range(n_neurons):
  286. psth = X[trial_mask, i_neuron, :].mean(axis=0)
  287. y[i_stim, :, i_neuron] = psth
  288. net = all_models[session]
  289. # turn off internal noise too for this simulation
  290. net_no_noise = net.clone()
  291. net_no_noise.noise_std = 0
  292. output, traj = net_no_noise(inputs, initial_states=None, return_dynamics=True)
  293. y_hat = output[:, :, :n_neurons].detach().numpy()
  294. ys.append(y)
  295. y_hats.append(y_hat)
  296. sses_all = []
  297. for i_session in range(len(ys)):
  298. sses = []
  299. for i_neuron in range(ys[i_session].shape[2]):
  300. sse = 0
  301. for i_stim, stim in enumerate(unique_stims):
  302. y = ys[i_session][i_stim, :, i_neuron]
  303. y_hat = y_hats[i_session][i_stim, :, i_neuron]
  304. sse += ((y - y_hat) ** 2).sum()
  305. sse /= (n_stims * n_bins)
  306. sses.append(sse)
  307. sses_all.append(sses)
  308. # %%
  309. sses = []
  310. for sses_ in sses_all:
  311. sses += sses_
  312. print('Mean MSE: {:.4f}'.format(np.mean(sses)))
  313. print('Median MSE: {:.4f}'.format(np.median(sses)))
  314. print('Mean RMSE: {:.4f}'.format(np.mean(np.sqrt(sses))))
  315. print('Median RMSE: {:.4f}'.format(np.median(np.sqrt(sses))))
  316. # %%
  317. ## example constrained neuron (no noise)
  318. i_session = 2
  319. i_neuron = 9
  320. fig, axes = plt.subplots(1, 2, figsize=(11, 2.5));
  321. sse = 0
  322. for i_stim, stim in enumerate(unique_stims):
  323. y = ys[i_session][i_stim, :, i_neuron]
  324. y_hat = y_hats[i_session][i_stim, :, i_neuron]
  325. sse += ((y - y_hat) ** 2).sum()
  326. axes[0].plot(t, y, color=colors[i_stim, :]);
  327. axes[1].plot(t, y_hat, color=colors[i_stim, :]);
  328. sse = sse / n_stims / n_bins
  329. fig.legend([stim for stim in unique_stims]);
  330. for i_ax, ax in enumerate(axes):
  331. ax.set_xlim([t[0], t[-1]]);
  332. ax.set_xlabel('Warped time [s]');
  333. y_max = max([ax.get_ylim()[1] for ax in axes])
  334. ax.set_ylim([0, y_max]);
  335. ax.set_ylabel('Firing rate [Hz]');
  336. ax.axvline(0, color='k');
  337. ax.axvline(t_D, color='k');
  338. if (i_ax == 0):
  339. ax.title.set_text('Neuron {} PSTH'.format(i_neuron + 1));
  340. elif (i_ax == 1):
  341. ax.title.set_text('Model neuron {} PSTH'.format(i_neuron + 1));
  342. # Figure 5C
  343. # plt.savefig('plots/model/model_example_constrained_psth_5.pdf');
  344. print('SSE: {:.4f}'.format(sse))
  345. print('RSSE: {:.4f}'.format(np.sqrt(sse)))
  346. # %%
  347. ## example unconstrained neuron (no noise)
  348. i_session = 1
  349. i_neuron = 3
  350. fig, ax = plt.subplots(figsize=(11, 2.5));
  351. net = all_models[included_sessions[i_session]].clone()
  352. net.noise_std = 0
  353. n_neurons = net.observed_size
  354. _, traj = net(inputs, initial_states=None, return_dynamics=True)
  355. y_hats = net.non_linearity(traj[:, 1:, :].detach() + net.b.detach()).numpy()
  356. for i_stim, stim in enumerate(unique_stims):
  357. y_hat = y_hats[i_stim, :, n_neurons + i_neuron]
  358. ax.plot(t, y_hat, color=colors[i_stim, :]);
  359. fig.legend([stim for stim in unique_stims]);
  360. ax.set_xlim([t[0], t[-1]]);
  361. ax.set_xlabel('Warped time [s]');
  362. ax.set_ylim([0, ax.get_ylim()[1]]);
  363. ax.set_ylabel('Firing rate [Hz]');
  364. ax.axvline(0, color='k');
  365. ax.axvline(t_D, color='k');
  366. ax.title.set_text('Model neuron {} PSTH'.format(n_neurons + i_neuron));
  367. # Figure 5D
  368. # plt.savefig('plots/model/model_example_unconstrained_psth.pdf');
  369. # %% [markdown]
  370. # ### Responsivity and selectivity
  371. # %%
  372. try:
  373. with open('data/model/responsivity_and_selectivity.pkl', 'rb') as file:
  374. res = pickle.load(file) # a dictionary with a key for each i_session
  375. file.close()
  376. except:
  377. res = model_responsive_and_selective_tests(output_noise, list(range(len(included_sessions))), t,
  378. width=0.5, base_left=-0.5, correct_only=True)
  379. with open('data/model/responsivity_and_selectivity.pkl', 'wb') as file:
  380. pickle.dump(res, file)
  381. file.close()
  382. # %%
  383. alpha = 0.01
  384. is_taste_responsive_list, is_taste_selective_list = [], []
  385. is_delay_responsive_list, is_delay_selective_list = [], []
  386. taste_change_list = []
  387. delay_change_list = []
  388. psths_list = []
  389. for i_session in res:
  390. is_taste_responsive_i = []
  391. taste_change_i = []
  392. is_taste_selective_i = []
  393. is_delay_responsive_i = []
  394. delay_change_i = []
  395. is_delay_selective_i = []
  396. psths_i = []
  397. for i_neuron in res[i_session]:
  398. base = res[i_session][i_neuron]['fr_baseline'].mean()
  399. r = res[i_session][i_neuron]['fr_sampling'].mean()
  400. p = res[i_session][i_neuron]['p_responsive_sampling']
  401. is_taste_responsive_i.append(p < alpha)
  402. taste_change_i.append(r > base)
  403. r = res[i_session][i_neuron]['fr_delay'].mean()
  404. p = res[i_session][i_neuron]['p_responsive_delay']
  405. is_delay_responsive_i.append(p < alpha)
  406. delay_change_i.append(r > base)
  407. p = res[i_session][i_neuron]['p_selective_sampling']
  408. is_taste_selective_i.append(p < alpha)
  409. p = res[i_session][i_neuron]['p_selective_delay']
  410. is_delay_selective_i.append(p < alpha)
  411. psths_i.append(res[i_session][i_neuron]['psth'] - base)
  412. is_taste_responsive_list.append(np.array(is_taste_responsive_i))
  413. is_taste_selective_list.append(np.array(is_taste_selective_i))
  414. is_delay_responsive_list.append(np.array(is_delay_responsive_i))
  415. is_delay_selective_list.append(np.array(is_delay_selective_i))
  416. taste_change_list.append(np.array(taste_change_i))
  417. delay_change_list.append(np.array(delay_change_i))
  418. psths_list.append(np.array(psths_i))
  419. is_taste_responsive = np.hstack(is_taste_responsive_list)
  420. taste_change = np.hstack(taste_change_list)
  421. is_taste_selective = is_taste_responsive & np.hstack(is_taste_selective_list)
  422. is_delay_responsive = np.hstack(is_delay_responsive_list)
  423. delay_change = np.hstack(delay_change_list)
  424. is_delay_selective = is_delay_responsive & np.hstack(is_delay_selective_list)
  425. psths = np.vstack(psths_list)
  426. # %%
  427. print('{}/{} ({:.1f}%) neurons are taste-responsive'.format(
  428. is_taste_responsive.sum(), len(is_taste_responsive),
  429. 100 * is_taste_responsive.sum() / len(is_taste_responsive)
  430. ))
  431. print('{}/{} ({:.1f}%) neurons are taste-selective\n'.format(
  432. is_taste_selective.sum(), len(is_taste_selective),
  433. 100 * is_taste_selective.sum() / len(is_taste_selective)
  434. ))
  435. print('{}/{} ({:.1f}%) neurons are delay-responsive'.format(
  436. is_delay_responsive.sum(), len(is_delay_responsive),
  437. 100 * is_delay_responsive.sum() / len(is_delay_responsive)
  438. ))
  439. print('{}/{} ({:.1f}%) neurons are delay-selective\n'.format(
  440. is_delay_selective.sum(), len(is_delay_selective),
  441. 100 * is_delay_selective.sum() / len(is_delay_selective)
  442. ))
  443. print('{}/{} ({:.1f}%) neurons are responsive to both'.format(
  444. (is_delay_responsive & is_taste_responsive).sum(), len(is_delay_responsive),
  445. 100 * (is_delay_responsive & is_taste_responsive).sum() / len(is_delay_responsive)
  446. ))
  447. # %%
  448. print('Session | Total (Con / Unc) | Taste Responsive | Taste Selective | Delay Responsive | Delay Selective')
  449. print('-----------------------------------------------------------------------------------------------------')
  450. for i, sess in enumerate(res):
  451. n_total = len(is_taste_responsive_list[i])
  452. n_con = int(round(n_total / 5.88))
  453. if int(round(5.88 * n_con)) != n_total: raise ValueError('backwards rounding failed')
  454. print('{:7} | {:8}/{:8} | {:7}/{:8} | {:7}/{:7} | {:7}/{:8} | {:7}/{:7}'
  455. .format(sess, n_con, n_total - n_con,
  456. is_taste_responsive_list[i][:n_con].sum(), is_taste_responsive_list[i][n_con:].sum(),
  457. is_taste_selective_list[i][:n_con].sum(), is_taste_selective_list[i][n_con:].sum(),
  458. is_delay_responsive_list[i][:n_con].sum(), is_delay_responsive_list[i][n_con:].sum(),
  459. is_delay_selective_list[i][:n_con].sum(), is_delay_selective_list[i][n_con:].sum()))
  460. # %%
  461. fig, axes = plt.subplots(2, 3, figsize=(11, 6));
  462. t_align_T = t
  463. t_align_D = t - (t[-1] - 1 + 0.025)
  464. y_low, y_high = 0, 0
  465. mat = psths[is_taste_responsive & (taste_change == 1), :]
  466. mean_ = mat.mean(axis=0)
  467. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  468. patch = plt.Polygon(
  469. [[t_align_T[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_T[i], mean_[i] + std_[i]]
  470. for i in range(len(t))[::-1]],
  471. facecolor='b',
  472. edgecolor=None,
  473. alpha=0.3
  474. )
  475. axes[0, 0].add_patch(patch);
  476. axes[0, 0].plot(t_align_T, mean_, color='b', linewidth=2);
  477. axes[0, 0].axvline(0, color='k', linestyle=':');
  478. axes[0, 0].axhline(0, color='k', linestyle=':');
  479. axes[0, 0].set_title('Sampling responsive (increase)\n(N={})'.format(mat.shape[0]));
  480. axes[0, 0].set_xlim([-1, 1.5]);
  481. y_low, y_high = min(y_low, axes[0, 0].get_ylim()[0]), max(y_high, axes[0, 0].get_ylim()[1])
  482. mat = psths[is_taste_responsive & (taste_change == 0), :]
  483. mean_ = mat.mean(axis=0)
  484. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  485. patch = plt.Polygon(
  486. [[t_align_T[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_T[i], mean_[i] + std_[i]]
  487. for i in range(len(t))[::-1]],
  488. facecolor='b',
  489. edgecolor=None,
  490. alpha=0.3
  491. )
  492. axes[0, 1].add_patch(patch);
  493. axes[0, 1].plot(t_align_T, mean_, color='b', linewidth=2);
  494. axes[0, 1].axvline(0, color='k', linestyle=':');
  495. axes[0, 1].axhline(0, color='k', linestyle=':');
  496. axes[0, 1].set_title('Sampling responsive (decrease)\n(N={})'.format(mat.shape[0]));
  497. axes[0, 1].set_xlim([-1, 1.5]);
  498. y_low, y_high = min(y_low, axes[0, 1].get_ylim()[0]), max(y_high, axes[0, 1].get_ylim()[1])
  499. mat = psths[~is_taste_responsive, :]
  500. mean_ = mat.mean(axis=0)
  501. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  502. patch = plt.Polygon(
  503. [[t_align_T[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_T[i], mean_[i] + std_[i]]
  504. for i in range(len(t))[::-1]],
  505. facecolor='b',
  506. edgecolor=None,
  507. alpha=0.3
  508. )
  509. axes[0, 2].add_patch(patch);
  510. axes[0, 2].plot(t_align_T, mean_, color='b', linewidth=2);
  511. axes[0, 2].axvline(0, color='k', linestyle=':');
  512. axes[0, 2].axhline(0, color='k', linestyle=':');
  513. axes[0, 2].set_title('Sampling nonresponsive\n(N={})'.format(mat.shape[0]));
  514. axes[0, 2].set_xlim([-1, 1.5]);
  515. y_low, y_high = min(y_low, axes[0, 2].get_ylim()[0]), max(y_high, axes[0, 2].get_ylim()[1])
  516. for i in range(3): axes[0, i].set_ylim([y_low, y_high]);
  517. y_low, y_high = 0, 0
  518. mat = psths[is_delay_responsive & (delay_change == 1), :]
  519. mean_ = mat.mean(axis=0)
  520. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  521. patch = plt.Polygon(
  522. [[t_align_D[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_D[i], mean_[i] + std_[i]]
  523. for i in range(len(t))[::-1]],
  524. facecolor='b',
  525. edgecolor=None,
  526. alpha=0.3
  527. )
  528. axes[1, 0].add_patch(patch);
  529. axes[1, 0].plot(t_align_D, mean_, color='b', linewidth=2);
  530. axes[1, 0].axvline(0, color='k', linestyle=':');
  531. axes[1, 0].axhline(0, color='k', linestyle=':');
  532. axes[1, 0].set_title('Delay responsive (increase)\n(N={})'.format(mat.shape[0]));
  533. axes[1, 0].set_xlim([-1.5, 1]);
  534. y_low, y_high = min(y_low, axes[1, 0].get_ylim()[0]), max(y_high, axes[1, 0].get_ylim()[1])
  535. mat = psths[is_delay_responsive & (delay_change == 0), :]
  536. mean_ = mat.mean(axis=0)
  537. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  538. patch = plt.Polygon(
  539. [[t_align_D[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_D[i], mean_[i] + std_[i]]
  540. for i in range(len(t))[::-1]],
  541. facecolor='b',
  542. edgecolor=None,
  543. alpha=0.3
  544. )
  545. axes[1, 1].add_patch(patch);
  546. axes[1, 1].plot(t_align_D, mean_, color='b', linewidth=2);
  547. axes[1, 1].axvline(0, color='k', linestyle=':');
  548. axes[1, 1].axhline(0, color='k', linestyle=':');
  549. axes[1, 1].set_title('Delay responsive (decrease)\n(N={})'.format(mat.shape[0]));
  550. axes[1, 1].set_xlim([-1.5, 1]);
  551. y_low, y_high = min(y_low, axes[1, 1].get_ylim()[0]), max(y_high, axes[1, 1].get_ylim()[1])
  552. mat = psths[~is_delay_responsive, :]
  553. mean_ = mat.mean(axis=0)
  554. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  555. patch = plt.Polygon(
  556. [[t_align_D[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_D[i], mean_[i] + std_[i]]
  557. for i in range(len(t))[::-1]],
  558. facecolor='b',
  559. edgecolor=None,
  560. alpha=0.3
  561. )
  562. axes[1, 2].add_patch(patch);
  563. axes[1, 2].plot(t_align_D, mean_, color='b', linewidth=2);
  564. axes[1, 2].axvline(0, color='k', linestyle=':');
  565. axes[1, 2].axhline(0, color='k', linestyle=':');
  566. axes[1, 2].set_title('Delay nonresponsive\n(N={})'.format(mat.shape[0]));
  567. axes[1, 2].set_xlim([-1.5, 1]);
  568. y_low, y_high = min(y_low, axes[1, 2].get_ylim()[0]), max(y_high, axes[1, 2].get_ylim()[1])
  569. for i in range(3): axes[1, i].set_ylim([y_low, y_high]);
  570. # Figure 5 - Supplement 2 A-B
  571. # plt.savefig('plots/model/responsive_units.pdf');
  572. # %% [markdown]
  573. # ### Single unit response profiles
  574. # %%
  575. # constrained only
  576. labels_over_time = []
  577. for i_t in range(len(t_downsample)):
  578. labels = []
  579. for i_model in range(len(included_sessions)):
  580. labels_i = output_noise['all_labels'][i_model][i_t]
  581. n_total = len(labels_i)
  582. n_con = int(round(n_total / 5.88))
  583. if int(round(5.88 * n_con)) != n_total: raise ValueError('backwards rounding failed')
  584. labels += labels_i[:n_con]
  585. labels_over_time.append(labels)
  586. n_neurons_all = np.sum([all_models[i].observed_size for i in included_sessions])
  587. is_linear = np.full(n_neurons_all, False)
  588. is_perception = np.full(n_neurons_all, False)
  589. is_choice = np.full(n_neurons_all, False)
  590. for i_neuron in range(n_neurons_all):
  591. label_seq = [labels_over_time[i_t][i_neuron] for i_t in range(len(labels_over_time))]
  592. is_linear[i_neuron] = 'Linear' in label_seq
  593. is_perception[i_neuron] = 'Step (Perception)' in label_seq
  594. is_choice[i_neuron] = 'Step (Choice)' in label_seq
  595. N_coding = (is_linear | is_perception | is_choice).sum()
  596. den = n_neurons_all
  597. # %%
  598. fig, ax = plt.subplots(figsize=(11, 3));
  599. n_lin = np.array([sum([label == 'Linear' for label in labels]) for labels in labels_over_time])
  600. ax.plot(t_downsample, n_lin / den, '.r', markersize=15);
  601. p1, = ax.plot(t_downsample, n_lin / den, 'r', linewidth=2);
  602. n_perc = np.array([sum([label == 'Step (Perception)' for label in labels]) for labels in labels_over_time])
  603. ax.plot(t_downsample, n_perc / den, '.c', markersize=15);
  604. p2, = ax.plot(t_downsample, n_perc / den, 'c', linewidth=2);
  605. n_choice = np.array([sum([label == 'Step (Choice)' for label in labels]) for labels in labels_over_time])
  606. ax.plot(t_downsample, n_choice / den, '.b', markersize=15);
  607. p3, = ax.plot(t_downsample, n_choice / den, 'b', linewidth=2);
  608. ax.set_xlim([t[0], t[-1]]);
  609. ax.set_xlabel('Warped time [s]');
  610. ax.set_ylabel('Frac. all fits');
  611. ax.axvline(0, color='k', linestyle=':');
  612. ax.axvline(t_D, color='k', linestyle=':');
  613. ax.legend([p1, p2, p3], ['Linear', 'Step (Perception)', 'Step (Choice)'],
  614. bbox_to_anchor=(1.02, 1), loc='upper left', borderaxespad=0.);
  615. ax.set_title('Constrained model neuron response profiles over time');
  616. # Figure 6 - Supplement 1 C
  617. # plt.savefig('plots/review/model/model_response_profiles_over_time_constrained.pdf');
  618. print('Peak fraction of coding units (non-Other):', ((n_lin + n_perc + n_choice) / n_neurons_all).max())
  619. # %%
  620. print('Breakdown of constrained neurons ({}):\n'.format(n_neurons_all))
  621. print(' Linear AND Perception AND Choice: {} ({:.1f}%)'.format(
  622. (is_linear & is_perception & is_choice).sum(),
  623. 100 * (is_linear & is_perception & is_choice).sum() / n_neurons_all
  624. ))
  625. print(' Linear AND Perception AND NOT Choice: {} ({:.1f}%)'.format(
  626. (is_linear & is_perception & ~is_choice).sum(),
  627. 100 * (is_linear & is_perception & ~is_choice).sum() / n_neurons_all
  628. ))
  629. print(' Linear AND NOT Perception AND Choice: {} ({:.1f}%)'.format(
  630. (is_linear & ~is_perception & is_choice).sum(),
  631. 100 * (is_linear & ~is_perception & is_choice).sum() / n_neurons_all
  632. ))
  633. print(' Linear AND NOT Perception AND NOT Choice: {} ({:.1f}%)'.format(
  634. (is_linear & ~is_perception & ~is_choice).sum(),
  635. 100 * (is_linear & ~is_perception & ~is_choice).sum() / n_neurons_all
  636. ))
  637. print(' NOT Linear AND Perception AND Choice: {} ({:.1f}%)'.format(
  638. (~is_linear & is_perception & is_choice).sum(),
  639. 100 * (~is_linear & is_perception & is_choice).sum() / n_neurons_all
  640. ))
  641. print(' NOT Linear AND Perception AND NOT Choice: {} ({:.1f}%)'.format(
  642. (~is_linear & is_perception & ~is_choice).sum(),
  643. 100 * (~is_linear & is_perception & ~is_choice).sum() / n_neurons_all
  644. ))
  645. print(' NOT Linear AND NOT Perception AND Choice: {} ({:.1f}%)'.format(
  646. (~is_linear & ~is_perception & is_choice).sum(),
  647. 100 * (~is_linear & ~is_perception & is_choice).sum() / n_neurons_all
  648. ))
  649. print(' NOT Linear AND NOT Perception AND NOT Choice: {} ({:.1f}%)\n'.format(
  650. (~is_linear & ~is_perception & ~is_choice).sum(),
  651. 100 * (~is_linear & ~is_perception & ~is_choice).sum() / n_neurons_all
  652. ))
  653. print(' Linear: {} ({:.1f}%)'.format(
  654. is_linear.sum(),
  655. 100 * is_linear.sum() / n_neurons_all
  656. ))
  657. print(' Perception: {} ({:.1f}%)'.format(
  658. is_perception.sum(),
  659. 100 * is_perception.sum() / n_neurons_all
  660. ))
  661. print(' Choice: {} ({:.1f}%)\n'.format(
  662. is_choice.sum(),
  663. 100 * is_choice.sum() / n_neurons_all
  664. ))
  665. # %%
  666. # unconstrained only
  667. labels_over_time = []
  668. for i_t in range(len(t_downsample)):
  669. labels = []
  670. for i_model in range(len(included_sessions)):
  671. labels_i = output_noise['all_labels'][i_model][i_t]
  672. n_total = len(labels_i)
  673. n_con = int(round(n_total / 5.88))
  674. if int(round(5.88 * n_con)) != n_total: raise ValueError('backwards rounding failed')
  675. labels += labels_i[n_con:]
  676. labels_over_time.append(labels)
  677. n_neurons_all = np.sum([all_models[i].network_size - all_models[i].observed_size for i in included_sessions])
  678. is_linear = np.full(n_neurons_all, False)
  679. is_perception = np.full(n_neurons_all, False)
  680. is_choice = np.full(n_neurons_all, False)
  681. for i_neuron in range(n_neurons_all):
  682. label_seq = [labels_over_time[i_t][i_neuron] for i_t in range(len(labels_over_time))]
  683. is_linear[i_neuron] = 'Linear' in label_seq
  684. is_perception[i_neuron] = 'Step (Perception)' in label_seq
  685. is_choice[i_neuron] = 'Step (Choice)' in label_seq
  686. N_coding = (is_linear | is_perception | is_choice).sum()
  687. den = n_neurons_all
  688. # %%
  689. fig, ax = plt.subplots(figsize=(11, 3));
  690. n_lin = np.array([sum([label == 'Linear' for label in labels]) for labels in labels_over_time])
  691. ax.plot(t_downsample, n_lin / den, '.r', markersize=15);
  692. p1, = ax.plot(t_downsample, n_lin / den, 'r', linewidth=2);
  693. n_perc = np.array([sum([label == 'Step (Perception)' for label in labels]) for labels in labels_over_time])
  694. ax.plot(t_downsample, n_perc / den, '.c', markersize=15);
  695. p2, = ax.plot(t_downsample, n_perc / den, 'c', linewidth=2);
  696. n_choice = np.array([sum([label == 'Step (Choice)' for label in labels]) for labels in labels_over_time])
  697. ax.plot(t_downsample, n_choice / den, '.b', markersize=15);
  698. p3, = ax.plot(t_downsample, n_choice / den, 'b', linewidth=2);
  699. ax.set_xlim([t[0], t[-1]]);
  700. ax.set_xlabel('Warped time [s]');
  701. ax.set_ylabel('Frac. all fits');
  702. ax.axvline(0, color='k', linestyle=':');
  703. ax.axvline(t_D, color='k', linestyle=':');
  704. ax.legend([p1, p2, p3], ['Linear', 'Step (Perception)', 'Step (Choice)'],
  705. bbox_to_anchor=(1.02, 1), loc='upper left', borderaxespad=0.);
  706. ax.set_title('Unconstrained model neuron response profiles over time');
  707. # Figure 6 - Supplement 2 C
  708. # plt.savefig('plots/review/model/model_response_profiles_over_time_unconstrained.pdf');
  709. print('Peak fraction of coding units (non-Other):', ((n_lin + n_perc + n_choice) / n_neurons_all).max())
  710. # %%
  711. print('Breakdown of unconstrained neurons ({}):\n'.format(n_neurons_all))
  712. print(' Linear AND Perception AND Choice: {} ({:.1f}%)'.format(
  713. (is_linear & is_perception & is_choice).sum(),
  714. 100 * (is_linear & is_perception & is_choice).sum() / n_neurons_all
  715. ))
  716. print(' Linear AND Perception AND NOT Choice: {} ({:.1f}%)'.format(
  717. (is_linear & is_perception & ~is_choice).sum(),
  718. 100 * (is_linear & is_perception & ~is_choice).sum() / n_neurons_all
  719. ))
  720. print(' Linear AND NOT Perception AND Choice: {} ({:.1f}%)'.format(
  721. (is_linear & ~is_perception & is_choice).sum(),
  722. 100 * (is_linear & ~is_perception & is_choice).sum() / n_neurons_all
  723. ))
  724. print(' Linear AND NOT Perception AND NOT Choice: {} ({:.1f}%)'.format(
  725. (is_linear & ~is_perception & ~is_choice).sum(),
  726. 100 * (is_linear & ~is_perception & ~is_choice).sum() / n_neurons_all
  727. ))
  728. print(' NOT Linear AND Perception AND Choice: {} ({:.1f}%)'.format(
  729. (~is_linear & is_perception & is_choice).sum(),
  730. 100 * (~is_linear & is_perception & is_choice).sum() / n_neurons_all
  731. ))
  732. print(' NOT Linear AND Perception AND NOT Choice: {} ({:.1f}%)'.format(
  733. (~is_linear & is_perception & ~is_choice).sum(),
  734. 100 * (~is_linear & is_perception & ~is_choice).sum() / n_neurons_all
  735. ))
  736. print(' NOT Linear AND NOT Perception AND Choice: {} ({:.1f}%)'.format(
  737. (~is_linear & ~is_perception & is_choice).sum(),
  738. 100 * (~is_linear & ~is_perception & is_choice).sum() / n_neurons_all
  739. ))
  740. print(' NOT Linear AND NOT Perception AND NOT Choice: {} ({:.1f}%)\n'.format(
  741. (~is_linear & ~is_perception & ~is_choice).sum(),
  742. 100 * (~is_linear & ~is_perception & ~is_choice).sum() / n_neurons_all
  743. ))
  744. print(' Linear: {} ({:.1f}%)'.format(
  745. is_linear.sum(),
  746. 100 * is_linear.sum() / n_neurons_all
  747. ))
  748. print(' Perception: {} ({:.1f}%)'.format(
  749. is_perception.sum(),
  750. 100 * is_perception.sum() / n_neurons_all
  751. ))
  752. print(' Choice: {} ({:.1f}%)\n'.format(
  753. is_choice.sum(),
  754. 100 * is_choice.sum() / n_neurons_all
  755. ))
  756. # %%
  757. # all
  758. labels_over_time = []
  759. for i_t in range(len(t_downsample)):
  760. labels = []
  761. for i_model in range(len(included_sessions)):
  762. labels += output_noise['all_labels'][i_model][i_t]
  763. labels_over_time.append(labels)
  764. n_neurons_all = np.sum([all_models[i].network_size for i in included_sessions])
  765. is_linear = np.full(n_neurons_all, False)
  766. is_perception = np.full(n_neurons_all, False)
  767. is_choice = np.full(n_neurons_all, False)
  768. for i_neuron in range(n_neurons_all):
  769. label_seq = [labels_over_time[i_t][i_neuron] for i_t in range(len(labels_over_time))]
  770. is_linear[i_neuron] = 'Linear' in label_seq
  771. is_perception[i_neuron] = 'Step (Perception)' in label_seq
  772. is_choice[i_neuron] = 'Step (Choice)' in label_seq
  773. N_coding = (is_linear | is_perception | is_choice).sum()
  774. den = n_neurons_all
  775. # %%
  776. fig, ax = plt.subplots(figsize=(11, 3));
  777. n_lin = np.array([sum([label == 'Linear' for label in labels]) for labels in labels_over_time])
  778. ax.plot(t_downsample, n_lin / den, '.r', markersize=15);
  779. p1, = ax.plot(t_downsample, n_lin / den, 'r', linewidth=2);
  780. n_perc = np.array([sum([label == 'Step (Perception)' for label in labels]) for labels in labels_over_time])
  781. ax.plot(t_downsample, n_perc / den, '.c', markersize=15);
  782. p2, = ax.plot(t_downsample, n_perc / den, 'c', linewidth=2);
  783. n_choice = np.array([sum([label == 'Step (Choice)' for label in labels]) for labels in labels_over_time])
  784. ax.plot(t_downsample, n_choice / den, '.b', markersize=15);
  785. p3, = ax.plot(t_downsample, n_choice / den, 'b', linewidth=2);
  786. ax.set_xlim([t[0], t[-1]]);
  787. ax.set_xlabel('Warped time [s]');
  788. ax.set_ylabel('Frac. all fits');
  789. ax.axvline(0, color='k', linestyle=':');
  790. ax.axvline(t_D, color='k', linestyle=':');
  791. ax.legend([p1, p2, p3], ['Linear', 'Step (Perception)', 'Step (Choice)'],
  792. bbox_to_anchor=(1.02, 1), loc='upper left', borderaxespad=0.);
  793. ax.set_title('All model neuron response profiles over time');
  794. # Figure 6C
  795. # plt.savefig('plots/model/model_response_profiles_over_time.pdf');
  796. print('Peak fraction of coding units (non-Other):', ((n_lin + n_perc + n_choice) / n_neurons_all).max())
  797. # %%
  798. print('Breakdown of all neurons ({}):\n'.format(n_neurons_all))
  799. print(' Linear AND Perception AND Choice: {} ({:.1f}%)'.format(
  800. (is_linear & is_perception & is_choice).sum(),
  801. 100 * (is_linear & is_perception & is_choice).sum() / n_neurons_all
  802. ))
  803. print(' Linear AND Perception AND NOT Choice: {} ({:.1f}%)'.format(
  804. (is_linear & is_perception & ~is_choice).sum(),
  805. 100 * (is_linear & is_perception & ~is_choice).sum() / n_neurons_all
  806. ))
  807. print(' Linear AND NOT Perception AND Choice: {} ({:.1f}%)'.format(
  808. (is_linear & ~is_perception & is_choice).sum(),
  809. 100 * (is_linear & ~is_perception & is_choice).sum() / n_neurons_all
  810. ))
  811. print(' Linear AND NOT Perception AND NOT Choice: {} ({:.1f}%)'.format(
  812. (is_linear & ~is_perception & ~is_choice).sum(),
  813. 100 * (is_linear & ~is_perception & ~is_choice).sum() / n_neurons_all
  814. ))
  815. print(' NOT Linear AND Perception AND Choice: {} ({:.1f}%)'.format(
  816. (~is_linear & is_perception & is_choice).sum(),
  817. 100 * (~is_linear & is_perception & is_choice).sum() / n_neurons_all
  818. ))
  819. print(' NOT Linear AND Perception AND NOT Choice: {} ({:.1f}%)'.format(
  820. (~is_linear & is_perception & ~is_choice).sum(),
  821. 100 * (~is_linear & is_perception & ~is_choice).sum() / n_neurons_all
  822. ))
  823. print(' NOT Linear AND NOT Perception AND Choice: {} ({:.1f}%)'.format(
  824. (~is_linear & ~is_perception & is_choice).sum(),
  825. 100 * (~is_linear & ~is_perception & is_choice).sum() / n_neurons_all
  826. ))
  827. print(' NOT Linear AND NOT Perception AND NOT Choice: {} ({:.1f}%)\n'.format(
  828. (~is_linear & ~is_perception & ~is_choice).sum(),
  829. 100 * (~is_linear & ~is_perception & ~is_choice).sum() / n_neurons_all
  830. ))
  831. print(' Linear: {} ({:.1f}%)'.format(
  832. is_linear.sum(),
  833. 100 * is_linear.sum() / n_neurons_all
  834. ))
  835. print(' Perception: {} ({:.1f}%)'.format(
  836. is_perception.sum(),
  837. 100 * is_perception.sum() / n_neurons_all
  838. ))
  839. print(' Choice: {} ({:.1f}%)\n'.format(
  840. is_choice.sum(),
  841. 100 * is_choice.sum() / n_neurons_all
  842. ))
  843. # %%
  844. # session-by-session results
  845. print('Session | Total (Con/Unc) | Linear (Con/Unc) | Perception (Con/Unc) | Choice (Con/Unc) | Other (Con/Unc)')
  846. print('--------------------------------------------------------------------------------------------------------')
  847. for i_model in range(len(included_sessions)):
  848. n_total = len(output_noise['all_labels'][i_model][0])
  849. n_con = int(round(n_total / 5.88))
  850. if int(round(5.88 * n_con)) != n_total: raise ValueError('backwards rounding failed')
  851. n_lin_con, n_lin_unc = 0, 0
  852. n_perc_con, n_perc_unc = 0, 0
  853. n_choice_con, n_choice_unc = 0, 0
  854. n_other_con, n_other_unc = 0, 0
  855. for i_neuron in range(n_total):
  856. label_seq = labels_i = [output_noise['all_labels'][i_model][i_t][i_neuron] for i_t in range(len(t_downsample))]
  857. is_lin_ = 'Linear' in label_seq
  858. is_perc_ = 'Step (Perception)' in label_seq
  859. is_choice_ = 'Step (Choice)' in label_seq
  860. is_other_ = (not is_lin_) and (not is_perc_) and (not is_choice_)
  861. if is_lin_:
  862. if i_neuron < n_con:
  863. n_lin_con += 1
  864. else:
  865. n_lin_unc += 1
  866. if is_perc_:
  867. if i_neuron < n_con:
  868. n_perc_con += 1
  869. else:
  870. n_perc_unc += 1
  871. if is_choice_:
  872. if i_neuron < n_con:
  873. n_choice_con += 1
  874. else:
  875. n_choice_unc += 1
  876. if is_other_:
  877. if i_neuron < n_con:
  878. n_other_con += 1
  879. else:
  880. n_other_unc += 1
  881. print('{:7} {:5} / {:5} {:5} / {:5} {:5} / {:5} {:5} / {:5} {:5} / {:5}'
  882. .format(i_model, n_con, n_total - n_con,
  883. n_lin_con, n_lin_unc, n_perc_con, n_perc_unc, n_choice_con, n_choice_unc, n_other_con, n_other_unc))
  884. # %%
  885. is_other = (~is_linear & ~is_perception & ~is_choice)
  886. print('{}/{} ({:.1f}%) Other neurons are taste-responsive'.format(
  887. (is_other & is_taste_responsive).sum(), is_other.sum(),
  888. 100 * (is_other & is_taste_responsive).sum() / is_other.sum()
  889. ))
  890. print('{}/{} ({:.1f}%) Other neurons are taste-selective\n'.format(
  891. (is_other & is_taste_selective).sum(), is_other.sum(),
  892. 100 * (is_other & is_taste_selective).sum() / is_other.sum()
  893. ))
  894. print('{}/{} ({:.1f}%) Other neurons are delay-responsive'.format(
  895. (is_other & is_delay_responsive).sum(), is_other.sum(),
  896. 100 * (is_other & is_delay_responsive).sum() / is_other.sum()
  897. ))
  898. print('{}/{} ({:.1f}%) Other neurons are delay-selective\n'.format(
  899. (is_other & is_delay_selective).sum(), is_other.sum(),
  900. 100 * (is_other & is_delay_selective).sum() / is_other.sum()
  901. ))
  902. print('{}/{} ({:.1f}%) Other neurons are responsive to both'.format(
  903. (is_other & is_delay_responsive & is_taste_responsive).sum(), is_other.sum(),
  904. 100 * (is_other & is_delay_responsive & is_taste_responsive).sum() / is_other.sum()
  905. ))
  906. # %%
  907. fig, axes = plt.subplots(2, 3, figsize=(11, 6));
  908. t_align_T = t
  909. t_align_D = t - (t[-1] - 1 + 0.025)
  910. y_low, y_high = 0, 0
  911. mat = psths[is_other & is_taste_responsive & (taste_change == 1), :]
  912. mean_ = mat.mean(axis=0)
  913. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  914. patch = plt.Polygon(
  915. [[t_align_T[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_T[i], mean_[i] + std_[i]]
  916. for i in range(len(t))[::-1]],
  917. facecolor='b',
  918. edgecolor=None,
  919. alpha=0.3
  920. )
  921. axes[0, 0].add_patch(patch);
  922. axes[0, 0].plot(t_align_T, mean_, color='b', linewidth=2);
  923. axes[0, 0].axvline(0, color='k', linestyle=':');
  924. axes[0, 0].axhline(0, color='k', linestyle=':');
  925. axes[0, 0].set_title('Other sampling responsive (increase)\n(N={})'.format(mat.shape[0]));
  926. axes[0, 0].set_xlim([-1, 1.5]);
  927. y_low, y_high = min(y_low, axes[0, 0].get_ylim()[0]), max(y_high, axes[0, 0].get_ylim()[1])
  928. mat = psths[is_other & is_taste_responsive & (taste_change == 0), :]
  929. mean_ = mat.mean(axis=0)
  930. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  931. patch = plt.Polygon(
  932. [[t_align_T[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_T[i], mean_[i] + std_[i]]
  933. for i in range(len(t))[::-1]],
  934. facecolor='b',
  935. edgecolor=None,
  936. alpha=0.3
  937. )
  938. axes[0, 1].add_patch(patch);
  939. axes[0, 1].plot(t_align_T, mean_, color='b', linewidth=2);
  940. axes[0, 1].axvline(0, color='k', linestyle=':');
  941. axes[0, 1].axhline(0, color='k', linestyle=':');
  942. axes[0, 1].set_title('Other sampling responsive (decrease)\n(N={})'.format(mat.shape[0]));
  943. axes[0, 1].set_xlim([-1, 1.5]);
  944. y_low, y_high = min(y_low, axes[0, 1].get_ylim()[0]), max(y_high, axes[0, 1].get_ylim()[1])
  945. mat = psths[is_other & ~is_taste_responsive, :]
  946. mean_ = mat.mean(axis=0)
  947. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  948. patch = plt.Polygon(
  949. [[t_align_T[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_T[i], mean_[i] + std_[i]]
  950. for i in range(len(t))[::-1]],
  951. facecolor='b',
  952. edgecolor=None,
  953. alpha=0.3
  954. )
  955. axes[0, 2].add_patch(patch);
  956. axes[0, 2].plot(t_align_T, mean_, color='b', linewidth=2);
  957. axes[0, 2].axvline(0, color='k', linestyle=':');
  958. axes[0, 2].axhline(0, color='k', linestyle=':');
  959. axes[0, 2].set_title('Other sampling nonresponsive\n(N={})'.format(mat.shape[0]));
  960. axes[0, 2].set_xlim([-1, 1.5]);
  961. y_low, y_high = min(y_low, axes[0, 2].get_ylim()[0]), max(y_high, axes[0, 2].get_ylim()[1])
  962. for i in range(3): axes[0, i].set_ylim([y_low, y_high]);
  963. y_low, y_high = 0, 0
  964. mat = psths[is_other & is_delay_responsive & (delay_change == 1), :]
  965. mean_ = mat.mean(axis=0)
  966. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  967. patch = plt.Polygon(
  968. [[t_align_D[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_D[i], mean_[i] + std_[i]]
  969. for i in range(len(t))[::-1]],
  970. facecolor='b',
  971. edgecolor=None,
  972. alpha=0.3
  973. )
  974. axes[1, 0].add_patch(patch);
  975. axes[1, 0].plot(t_align_D, mean_, color='b', linewidth=2);
  976. axes[1, 0].axvline(0, color='k', linestyle=':');
  977. axes[1, 0].axhline(0, color='k', linestyle=':');
  978. axes[1, 0].set_title('Other delay responsive (increase)\n(N={})'.format(mat.shape[0]));
  979. axes[1, 0].set_xlim([-1.5, 1]);
  980. y_low, y_high = min(y_low, axes[1, 0].get_ylim()[0]), max(y_high, axes[1, 0].get_ylim()[1])
  981. mat = psths[is_other & is_delay_responsive & (delay_change == 0), :]
  982. mean_ = mat.mean(axis=0)
  983. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  984. patch = plt.Polygon(
  985. [[t_align_D[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_D[i], mean_[i] + std_[i]]
  986. for i in range(len(t))[::-1]],
  987. facecolor='b',
  988. edgecolor=None,
  989. alpha=0.3
  990. )
  991. axes[1, 1].add_patch(patch);
  992. axes[1, 1].plot(t_align_D, mean_, color='b', linewidth=2);
  993. axes[1, 1].axvline(0, color='k', linestyle=':');
  994. axes[1, 1].axhline(0, color='k', linestyle=':');
  995. axes[1, 1].set_title('Other delay responsive (decrease)\n(N={})'.format(mat.shape[0]));
  996. axes[1, 1].set_xlim([-1.5, 1]);
  997. y_low, y_high = min(y_low, axes[1, 1].get_ylim()[0]), max(y_high, axes[1, 1].get_ylim()[1])
  998. mat = psths[is_other & ~is_delay_responsive, :]
  999. mean_ = mat.mean(axis=0)
  1000. std_ = mat.std(axis=0) / np.sqrt(mat.shape[0])
  1001. patch = plt.Polygon(
  1002. [[t_align_D[i], mean_[i] - std_[i]] for i in range(len(t))] + [[t_align_D[i], mean_[i] + std_[i]]
  1003. for i in range(len(t))[::-1]],
  1004. facecolor='b',
  1005. edgecolor=None,
  1006. alpha=0.3
  1007. )
  1008. axes[1, 2].add_patch(patch);
  1009. axes[1, 2].plot(t_align_D, mean_, color='b', linewidth=2);
  1010. axes[1, 2].axvline(0, color='k', linestyle=':');
  1011. axes[1, 2].axhline(0, color='k', linestyle=':');
  1012. axes[1, 2].set_title('Other delay nonresponsive\n(N={})'.format(mat.shape[0]));
  1013. axes[1, 2].set_xlim([-1.5, 1]);
  1014. y_low, y_high = min(y_low, axes[1, 2].get_ylim()[0]), max(y_high, axes[1, 2].get_ylim()[1])
  1015. for i in range(3): axes[1, i].set_ylim([y_low, y_high]);
  1016. # Figure 5 Supplement 2 C-D
  1017. # plt.savefig('plots/model/responsive_other_units.pdf');
  1018. # %%
  1019. # example Other neuron that is task-responsive
  1020. ind = 0
  1021. inds = []
  1022. for i_session in res:
  1023. for i_neuron in res[i_session]:
  1024. if is_delay_responsive[ind] and is_taste_responsive[ind] and is_other[ind]:
  1025. inds.append((i_session, i_neuron))
  1026. ind += 1
  1027. i_session, i_neuron = inds[13]
  1028. fig, ax = plt.subplots(figsize=(11, 3));
  1029. psths = output_noise['all_psths_correct'][i_session][i_neuron]
  1030. for i_stim in range(psths.shape[0]):
  1031. ax.plot(t, psths[i_stim, :], color=colors[i_stim]);
  1032. ax.axvline(0, color='k', linestyle=':');
  1033. ax.axvline(t[-1] - 1 + 0.025, color='k', linestyle=':');
  1034. ax.axhline(res[i_session][i_neuron]['fr_baseline'].mean(), color='k', linestyle=':');
  1035. ax.axhline(res[i_session][i_neuron]['fr_delay'].mean(), color='b', linestyle=':');
  1036. ax.set_xlim([t[0], t[-1]]);
  1037. # Figure 5 - Supplement 3 B
  1038. # plt.savefig('plots/model/example_taste_selective_other_2.pdf');
  1039. # %%
  1040. # responsive neuron patterns
  1041. ind = 0
  1042. all_psths, all_other_psths = [], []
  1043. for i_session in res:
  1044. for i_neuron in res[i_session]:
  1045. if is_taste_responsive[ind] or is_delay_responsive[ind]:
  1046. # relative to baseline
  1047. psth = res[i_session][i_neuron]['psth'] - res[i_session][i_neuron]['fr_baseline'].mean()
  1048. # focus on inter-event interval
  1049. psth = psth[20:-20]
  1050. # normalize to max
  1051. psth = psth / np.max(np.abs(psth))
  1052. if is_other[ind]:
  1053. all_other_psths.append(psth)
  1054. else:
  1055. all_psths.append(psth)
  1056. ind += 1
  1057. def sort_heatmap(heatmap):
  1058. max_vals = np.array([heatmap[i, np.argmax(np.abs(heatmap[i, :]))] for i in range(heatmap.shape[0])])
  1059. tmax = np.array([np.argmax(heatmap[i, :]) for i in range(heatmap.shape[0])])
  1060. heatmap = heatmap[np.argsort(tmax)[::-1], :]
  1061. return heatmap
  1062. cmap = LinearSegmentedColormap.from_list('my_cmap', [(0, [0, 0, 1]), (0.5, [0, 0, 0]), (1, [1, 0, 0])])
  1063. fig, ax = plt.subplots(figsize=(11, 4));
  1064. im = ax.imshow((sort_heatmap(np.array(all_psths)) + 1) / 2, cmap=cmap, aspect='auto');
  1065. ax.set_title('Responsive and Linear-/Perception-/Choice-coding (N={})'.format(len(all_psths)));
  1066. plt.colorbar(im, ax=ax);
  1067. # Figure 5 - Supplement 3 A
  1068. # plt.savefig('plots/model/responsive_coding_heatmap.pdf');
  1069. fig, ax = plt.subplots(figsize=(11, 4));
  1070. im = ax.imshow((sort_heatmap(np.array(all_other_psths)) + 1) / 2, cmap=cmap, aspect='auto');
  1071. ax.set_title('Responsive and Other (N={})'.format(len(all_other_psths)));
  1072. plt.colorbar(im, ax=ax);
  1073. # Figure 5 - Supplement 3 A
  1074. # plt.savefig('plots/model/responsive_other_heatmap.pdf');
  1075. # # for time scale
  1076. # fig, ax = plt.subplots(figsize=(11, 3));
  1077. # ax.axvline(0, color='k');
  1078. # ax.axvline(t[-1] - 1 + 0.025, color='k');
  1079. # # plt.savefig('plots/model/responsive_other_heatmap_scale.pdf');
  1080. # %% [markdown]
  1081. # ### Ablation experiments
  1082. # %%
  1083. # define 'beginning' and 'end' windows
  1084. inds_beginning = np.where((t >= 0) & (t < 1.2))[0]
  1085. inds_beginning_downsample = np.where((np.array(t_downsample) >= 0) & (np.array(t_downsample) < 1.2))[0]
  1086. inds_end = np.where((t >= t_D - 1.2) & (t < t_D))[0]
  1087. inds_end_downsample = np.where((np.array(t_downsample) >= t_D - 1.2) & (np.array(t_downsample) < t_D))[0]
  1088. # create 'label_data' for ablation experiments
  1089. label_data = []
  1090. for i_session, session in enumerate(included_sessions):
  1091. labels_over_time = output_noise['all_labels'][i_session]
  1092. ind_lin = []
  1093. ind_lin_con = []
  1094. ind_lin_unc = []
  1095. ind_lin_beginning = []
  1096. ind_lin_end = []
  1097. ind_perception = []
  1098. ind_perception_con = []
  1099. ind_perception_unc = []
  1100. ind_perception_beginning = []
  1101. ind_perception_end = []
  1102. ind_choice = []
  1103. ind_choice_con = []
  1104. ind_choice_unc = []
  1105. ind_choice_beginning = []
  1106. ind_choice_end = []
  1107. ind_other = []
  1108. ind_other_con = []
  1109. ind_other_unc = []
  1110. n_neurons_constrained = all_models[session].observed_size
  1111. n_neurons = all_models[session].network_size
  1112. for i_neuron in range(n_neurons):
  1113. label_over_time = [labels_over_time[i_t][i_neuron] for i_t in range(len(labels_over_time))]
  1114. is_linear = 'Linear' in label_over_time
  1115. is_perception = 'Step (Perception)' in label_over_time
  1116. is_choice = 'Step (Choice)' in label_over_time
  1117. is_constrained = i_neuron < n_neurons_constrained
  1118. if is_linear:
  1119. ind_lin.append(i_neuron)
  1120. if is_constrained:
  1121. ind_lin_con.append(i_neuron)
  1122. else:
  1123. ind_lin_unc.append(i_neuron)
  1124. if is_perception:
  1125. ind_perception.append(i_neuron)
  1126. if is_constrained:
  1127. ind_perception_con.append(i_neuron)
  1128. else:
  1129. ind_perception_unc.append(i_neuron)
  1130. if is_choice:
  1131. ind_choice.append(i_neuron)
  1132. if is_constrained:
  1133. ind_choice_con.append(i_neuron)
  1134. else:
  1135. ind_choice_unc.append(i_neuron)
  1136. if all([not is_linear, not is_perception, not is_choice]):
  1137. ind_other.append(i_neuron)
  1138. if is_constrained:
  1139. ind_other_con.append(i_neuron)
  1140. else:
  1141. ind_other_unc.append(i_neuron)
  1142. label_over_time_beginning = [label_over_time[i] for i in inds_beginning_downsample]
  1143. label_over_time_end = [label_over_time[i] for i in inds_end_downsample]
  1144. if 'Linear' in label_over_time_beginning: ind_lin_beginning.append(i_neuron)
  1145. if 'Linear' in label_over_time_end: ind_lin_end.append(i_neuron)
  1146. if 'Step (Perception)' in label_over_time_beginning: ind_perception_beginning.append(i_neuron)
  1147. if 'Step (Perception)' in label_over_time_end: ind_perception_end.append(i_neuron)
  1148. if 'Step (Choice)' in label_over_time_beginning: ind_choice_beginning.append(i_neuron)
  1149. if 'Step (Choice)' in label_over_time_end: ind_choice_end.append(i_neuron)
  1150. label_data.append({
  1151. 'ind_lin': ind_lin.copy(),
  1152. 'ind_lin_con': ind_lin_con.copy(),
  1153. 'ind_lin_unc': ind_lin_unc.copy(),
  1154. 'ind_lin_beginning': ind_lin_beginning.copy(),
  1155. 'ind_lin_end': ind_lin_end.copy(),
  1156. 'ind_perception': ind_perception.copy(),
  1157. 'ind_perception_con': ind_perception_con.copy(),
  1158. 'ind_perception_unc': ind_perception_unc.copy(),
  1159. 'ind_perception_beginning': ind_perception_beginning.copy(),
  1160. 'ind_perception_end': ind_perception_end.copy(),
  1161. 'ind_choice': ind_choice.copy(),
  1162. 'ind_choice_con': ind_choice_con.copy(),
  1163. 'ind_choice_unc': ind_choice_unc.copy(),
  1164. 'ind_choice_beginning': ind_choice_beginning.copy(),
  1165. 'ind_choice_end': ind_choice_end.copy(),
  1166. 'ind_other': ind_other.copy(),
  1167. 'ind_other_con': ind_other_con.copy(),
  1168. 'ind_other_unc': ind_other_unc.copy(),
  1169. })
  1170. # %%
  1171. ## temporally-restricted silencing was not built into the original MyRNN class
  1172. ## 'upgrade' models to new subclass
  1173. for key in all_models:
  1174. all_models[key].__class__ = MyRNNWithDynamicSilencing
  1175. # %% [markdown]
  1176. # ### Ablate 'Other'
  1177. # %%
  1178. try:
  1179. with open('data/model/output_tailored_noise_ablate_other.pkl', 'rb') as file:
  1180. output_noise_ablate_other = pickle.load(file)
  1181. file.close()
  1182. except:
  1183. output_noise_ablate_other = multi_model_simulation(
  1184. list(all_models.values()),
  1185. inputs,
  1186. unique_stims,
  1187. t_decision_window,
  1188. n_trials_per_stim=20,
  1189. sigma=list(inputs_sigma.values()),
  1190. silenced='ind_other',
  1191. label_data=label_data,
  1192. responses_over_time=False,
  1193. fit_options=None,
  1194. inds_L=None,
  1195. inds_R=None,
  1196. responses_beginning_end=False,
  1197. window_beginning=None,
  1198. window_end=None)
  1199. with open('data/model/output_tailored_noise_ablate_other.pkl', 'wb') as file:
  1200. pickle.dump(output_noise_ablate_other, file)
  1201. file.close()
  1202. # %%
  1203. try:
  1204. with open('data/model/output_tailored_noise_ablate_other_con.pkl', 'rb') as file:
  1205. output_noise_ablate_other_con = pickle.load(file)
  1206. file.close()
  1207. except:
  1208. output_noise_ablate_other_con = multi_model_simulation(
  1209. list(all_models.values()),
  1210. inputs,
  1211. unique_stims,
  1212. t_decision_window,
  1213. n_trials_per_stim=20,
  1214. sigma=list(inputs_sigma.values()),
  1215. silenced='ind_other_con',
  1216. label_data=label_data,
  1217. responses_over_time=False,
  1218. fit_options=None,
  1219. inds_L=None,
  1220. inds_R=None,
  1221. responses_beginning_end=False,
  1222. window_beginning=None,
  1223. window_end=None)
  1224. with open('data/model/output_tailored_noise_ablate_other_con.pkl', 'wb') as file:
  1225. pickle.dump(output_noise_ablate_other_con, file)
  1226. file.close()
  1227. # %%
  1228. try:
  1229. with open('data/model/output_tailored_noise_ablate_other_unc.pkl', 'rb') as file:
  1230. output_noise_ablate_other_unc = pickle.load(file)
  1231. file.close()
  1232. except:
  1233. output_noise_ablate_other_unc = multi_model_simulation(
  1234. list(all_models.values()),
  1235. inputs,
  1236. unique_stims,
  1237. t_decision_window,
  1238. n_trials_per_stim=20,
  1239. sigma=list(inputs_sigma.values()),
  1240. silenced='ind_other_unc',
  1241. label_data=label_data,
  1242. responses_over_time=False,
  1243. fit_options=None,
  1244. inds_L=None,
  1245. inds_R=None,
  1246. responses_beginning_end=False,
  1247. window_beginning=None,
  1248. window_end=None)
  1249. with open('data/model/output_tailored_noise_ablate_other_unc.pkl', 'wb') as file:
  1250. pickle.dump(output_noise_ablate_other_unc, file)
  1251. file.close()
  1252. # %% [markdown]
  1253. # ### Ablate 'Linear'
  1254. # %%
  1255. try:
  1256. with open('data/model/output_tailored_noise_ablate_linear.pkl', 'rb') as file:
  1257. output_noise_ablate_linear = pickle.load(file)
  1258. file.close()
  1259. except:
  1260. output_noise_ablate_linear = multi_model_simulation(
  1261. list(all_models.values()),
  1262. inputs,
  1263. unique_stims,
  1264. t_decision_window,
  1265. n_trials_per_stim=20,
  1266. sigma=list(inputs_sigma.values()),
  1267. silenced='ind_lin',
  1268. label_data=label_data,
  1269. responses_over_time=False,
  1270. fit_options=None,
  1271. inds_L=None,
  1272. inds_R=None,
  1273. responses_beginning_end=False,
  1274. window_beginning=None,
  1275. window_end=None)
  1276. with open('data/model/output_tailored_noise_ablate_linear.pkl', 'wb') as file:
  1277. pickle.dump(output_noise_ablate_linear, file)
  1278. file.close()
  1279. # %%
  1280. try:
  1281. with open('data/model/output_tailored_noise_ablate_linear_con.pkl', 'rb') as file:
  1282. output_noise_ablate_linear_con = pickle.load(file)
  1283. file.close()
  1284. except:
  1285. output_noise_ablate_linear_con = multi_model_simulation(
  1286. list(all_models.values()),
  1287. inputs,
  1288. unique_stims,
  1289. t_decision_window,
  1290. n_trials_per_stim=20,
  1291. sigma=list(inputs_sigma.values()),
  1292. silenced='ind_lin_con',
  1293. label_data=label_data,
  1294. responses_over_time=False,
  1295. fit_options=None,
  1296. inds_L=None,
  1297. inds_R=None,
  1298. responses_beginning_end=False,
  1299. window_beginning=None,
  1300. window_end=None)
  1301. with open('data/model/output_tailored_noise_ablate_linear_con.pkl', 'wb') as file:
  1302. pickle.dump(output_noise_ablate_linear_con, file)
  1303. file.close()
  1304. # %%
  1305. try:
  1306. with open('data/model/output_tailored_noise_ablate_linear_unc.pkl', 'rb') as file:
  1307. output_noise_ablate_linear_unc = pickle.load(file)
  1308. file.close()
  1309. except:
  1310. output_noise_ablate_linear_unc = multi_model_simulation(
  1311. list(all_models.values()),
  1312. inputs,
  1313. unique_stims,
  1314. t_decision_window,
  1315. n_trials_per_stim=20,
  1316. sigma=list(inputs_sigma.values()),
  1317. silenced='ind_lin_unc',
  1318. label_data=label_data,
  1319. responses_over_time=False,
  1320. fit_options=None,
  1321. inds_L=None,
  1322. inds_R=None,
  1323. responses_beginning_end=False,
  1324. window_beginning=None,
  1325. window_end=None)
  1326. with open('data/model/output_tailored_noise_ablate_linear_unc.pkl', 'wb') as file:
  1327. pickle.dump(output_noise_ablate_linear_unc, file)
  1328. file.close()
  1329. # %%
  1330. try:
  1331. with open('data/model/output_tailored_noise_ablate_linear_beginning.pkl', 'rb') as file:
  1332. output_noise_ablate_linear_beginning = pickle.load(file)
  1333. file.close()
  1334. except:
  1335. output_noise_ablate_linear_beginning = multi_model_simulation_with_dynamic_silencing(
  1336. list(all_models.values()),
  1337. inputs,
  1338. unique_stims,
  1339. t_decision_window,
  1340. n_trials_per_stim=20,
  1341. sigma=list(inputs_sigma.values()),
  1342. silenced='ind_lin_beginning',
  1343. silence_time=inds_beginning,
  1344. label_data=label_data,
  1345. responses_over_time=False,
  1346. fit_options=None,
  1347. inds_L=None,
  1348. inds_R=None,
  1349. responses_beginning_end=False,
  1350. window_beginning=None,
  1351. window_end=None)
  1352. with open('data/model/output_tailored_noise_ablate_linear_beginning.pkl', 'wb') as file:
  1353. pickle.dump(output_noise_ablate_linear_beginning, file)
  1354. file.close()
  1355. # %%
  1356. try:
  1357. with open('data/model/output_tailored_noise_ablate_linear_end.pkl', 'rb') as file:
  1358. output_noise_ablate_linear_end = pickle.load(file)
  1359. file.close()
  1360. except:
  1361. output_noise_ablate_linear_end = multi_model_simulation_with_dynamic_silencing(
  1362. list(all_models.values()),
  1363. inputs,
  1364. unique_stims,
  1365. t_decision_window,
  1366. n_trials_per_stim=20,
  1367. sigma=list(inputs_sigma.values()),
  1368. silenced='ind_lin_end',
  1369. silence_time=inds_end,
  1370. label_data=label_data,
  1371. responses_over_time=False,
  1372. fit_options=None,
  1373. inds_L=None,
  1374. inds_R=None,
  1375. responses_beginning_end=False,
  1376. window_beginning=None,
  1377. window_end=None)
  1378. with open('data/model/output_tailored_noise_ablate_linear_end.pkl', 'wb') as file:
  1379. pickle.dump(output_noise_ablate_linear_end, file)
  1380. file.close()
  1381. # %% [markdown]
  1382. # ### Ablate 'Step-Perception'
  1383. # %%
  1384. try:
  1385. with open('data/model/output_tailored_noise_ablate_perception.pkl', 'rb') as file:
  1386. output_noise_ablate_perception = pickle.load(file)
  1387. file.close()
  1388. except:
  1389. output_noise_ablate_perception = multi_model_simulation(
  1390. list(all_models.values()),
  1391. inputs,
  1392. unique_stims,
  1393. t_decision_window,
  1394. n_trials_per_stim=20,
  1395. sigma=list(inputs_sigma.values()),
  1396. silenced='ind_perception',
  1397. label_data=label_data,
  1398. responses_over_time=False,
  1399. fit_options=None,
  1400. inds_L=None,
  1401. inds_R=None,
  1402. responses_beginning_end=False,
  1403. window_beginning=None,
  1404. window_end=None)
  1405. with open('data/model/output_tailored_noise_ablate_perception.pkl', 'wb') as file:
  1406. pickle.dump(output_noise_ablate_perception, file)
  1407. file.close()
  1408. # %%
  1409. try:
  1410. with open('data/model/output_tailored_noise_ablate_perception_con.pkl', 'rb') as file:
  1411. output_noise_ablate_perception_con = pickle.load(file)
  1412. file.close()
  1413. except:
  1414. output_noise_ablate_perception_con = multi_model_simulation(
  1415. list(all_models.values()),
  1416. inputs,
  1417. unique_stims,
  1418. t_decision_window,
  1419. n_trials_per_stim=20,
  1420. sigma=list(inputs_sigma.values()),
  1421. silenced='ind_perception_con',
  1422. label_data=label_data,
  1423. responses_over_time=False,
  1424. fit_options=None,
  1425. inds_L=None,
  1426. inds_R=None,
  1427. responses_beginning_end=False,
  1428. window_beginning=None,
  1429. window_end=None)
  1430. with open('data/model/output_tailored_noise_ablate_perception_con.pkl', 'wb') as file:
  1431. pickle.dump(output_noise_ablate_perception_con, file)
  1432. file.close()
  1433. # %%
  1434. try:
  1435. with open('data/model/output_tailored_noise_ablate_perception_unc.pkl', 'rb') as file:
  1436. output_noise_ablate_perception_unc = pickle.load(file)
  1437. file.close()
  1438. except:
  1439. output_noise_ablate_perception_unc = multi_model_simulation(
  1440. list(all_models.values()),
  1441. inputs,
  1442. unique_stims,
  1443. t_decision_window,
  1444. n_trials_per_stim=20,
  1445. sigma=list(inputs_sigma.values()),
  1446. silenced='ind_perception_unc',
  1447. label_data=label_data,
  1448. responses_over_time=False,
  1449. fit_options=None,
  1450. inds_L=None,
  1451. inds_R=None,
  1452. responses_beginning_end=False,
  1453. window_beginning=None,
  1454. window_end=None)
  1455. with open('data/model/output_tailored_noise_ablate_perception_unc.pkl', 'wb') as file:
  1456. pickle.dump(output_noise_ablate_perception_unc, file)
  1457. file.close()
  1458. # %%
  1459. try:
  1460. with open('data/model/output_tailored_noise_ablate_perception_beginning.pkl', 'rb') as file:
  1461. output_noise_ablate_perception_beginning = pickle.load(file)
  1462. file.close()
  1463. except:
  1464. output_noise_ablate_perception_beginning = multi_model_simulation_with_dynamic_silencing(
  1465. list(all_models.values()),
  1466. inputs,
  1467. unique_stims,
  1468. t_decision_window,
  1469. n_trials_per_stim=20,
  1470. sigma=list(inputs_sigma.values()),
  1471. silenced='ind_perception_beginning',
  1472. silence_time=inds_beginning,
  1473. label_data=label_data,
  1474. responses_over_time=False,
  1475. fit_options=None,
  1476. inds_L=None,
  1477. inds_R=None,
  1478. responses_beginning_end=False,
  1479. window_beginning=None,
  1480. window_end=None)
  1481. with open('data/model/output_tailored_noise_ablate_perception_beginning.pkl', 'wb') as file:
  1482. pickle.dump(output_noise_ablate_perception_beginning, file)
  1483. file.close()
  1484. # %%
  1485. try:
  1486. with open('data/model/output_tailored_noise_ablate_perception_end.pkl', 'rb') as file:
  1487. output_noise_ablate_perception_end = pickle.load(file)
  1488. file.close()
  1489. except:
  1490. output_noise_ablate_perception_end = multi_model_simulation_with_dynamic_silencing(
  1491. list(all_models.values()),
  1492. inputs,
  1493. unique_stims,
  1494. t_decision_window,
  1495. n_trials_per_stim=20,
  1496. sigma=list(inputs_sigma.values()),
  1497. silenced='ind_perception_end',
  1498. silence_time=inds_end,
  1499. label_data=label_data,
  1500. responses_over_time=False,
  1501. fit_options=None,
  1502. inds_L=None,
  1503. inds_R=None,
  1504. responses_beginning_end=False,
  1505. window_beginning=None,
  1506. window_end=None)
  1507. with open('data/model/output_tailored_noise_ablate_perception_end.pkl', 'wb') as file:
  1508. pickle.dump(output_noise_ablate_perception_end, file)
  1509. file.close()
  1510. # %% [markdown]
  1511. # ### Ablate 'Step-Choice'
  1512. # %%
  1513. try:
  1514. with open('data/model/output_tailored_noise_ablate_choice.pkl', 'rb') as file:
  1515. output_noise_ablate_choice = pickle.load(file)
  1516. file.close()
  1517. except:
  1518. output_noise_ablate_choice = multi_model_simulation(
  1519. list(all_models.values()),
  1520. inputs,
  1521. unique_stims,
  1522. t_decision_window,
  1523. n_trials_per_stim=20,
  1524. sigma=list(inputs_sigma.values()),
  1525. silenced='ind_choice',
  1526. label_data=label_data,
  1527. responses_over_time=False,
  1528. fit_options=None,
  1529. inds_L=None,
  1530. inds_R=None,
  1531. responses_beginning_end=False,
  1532. window_beginning=None,
  1533. window_end=None)
  1534. with open('data/model/output_tailored_noise_ablate_choice.pkl', 'wb') as file:
  1535. pickle.dump(output_noise_ablate_choice, file)
  1536. file.close()
  1537. # %%
  1538. try:
  1539. with open('data/model/output_tailored_noise_ablate_choice_con.pkl', 'rb') as file:
  1540. output_noise_ablate_choice_con = pickle.load(file)
  1541. file.close()
  1542. except:
  1543. output_noise_ablate_choice_con = multi_model_simulation(
  1544. list(all_models.values()),
  1545. inputs,
  1546. unique_stims,
  1547. t_decision_window,
  1548. n_trials_per_stim=20,
  1549. sigma=list(inputs_sigma.values()),
  1550. silenced='ind_choice_con',
  1551. label_data=label_data,
  1552. responses_over_time=False,
  1553. fit_options=None,
  1554. inds_L=None,
  1555. inds_R=None,
  1556. responses_beginning_end=False,
  1557. window_beginning=None,
  1558. window_end=None)
  1559. with open('data/model/output_tailored_noise_ablate_choice_con.pkl', 'wb') as file:
  1560. pickle.dump(output_noise_ablate_choice_con, file)
  1561. file.close()
  1562. # %%
  1563. try:
  1564. with open('data/model/output_tailored_noise_ablate_choice_unc.pkl', 'rb') as file:
  1565. output_noise_ablate_choice_unc = pickle.load(file)
  1566. file.close()
  1567. except:
  1568. output_noise_ablate_choice_unc = multi_model_simulation(
  1569. list(all_models.values()),
  1570. inputs,
  1571. unique_stims,
  1572. t_decision_window,
  1573. n_trials_per_stim=20,
  1574. sigma=list(inputs_sigma.values()),
  1575. silenced='ind_choice_unc',
  1576. label_data=label_data,
  1577. responses_over_time=False,
  1578. fit_options=None,
  1579. inds_L=None,
  1580. inds_R=None,
  1581. responses_beginning_end=False,
  1582. window_beginning=None,
  1583. window_end=None)
  1584. with open('data/model/output_tailored_noise_ablate_choice_unc.pkl', 'wb') as file:
  1585. pickle.dump(output_noise_ablate_choice_unc, file)
  1586. file.close()
  1587. # %%
  1588. try:
  1589. with open('data/model/output_tailored_noise_ablate_choice_beginning.pkl', 'rb') as file:
  1590. output_noise_ablate_choice_beginning = pickle.load(file)
  1591. file.close()
  1592. except:
  1593. output_noise_ablate_choice_beginning = multi_model_simulation_with_dynamic_silencing(
  1594. list(all_models.values()),
  1595. inputs,
  1596. unique_stims,
  1597. t_decision_window,
  1598. n_trials_per_stim=20,
  1599. sigma=list(inputs_sigma.values()),
  1600. silenced='ind_choice_beginning',
  1601. silence_time=inds_beginning,
  1602. label_data=label_data,
  1603. responses_over_time=False,
  1604. fit_options=None,
  1605. inds_L=None,
  1606. inds_R=None,
  1607. responses_beginning_end=False,
  1608. window_beginning=None,
  1609. window_end=None)
  1610. with open('data/model/output_tailored_noise_ablate_choice_beginning.pkl', 'wb') as file:
  1611. pickle.dump(output_noise_ablate_choice_beginning, file)
  1612. file.close()
  1613. # %%
  1614. try:
  1615. with open('data/model/output_tailored_noise_ablate_choice_end.pkl', 'rb') as file:
  1616. output_noise_ablate_choice_end = pickle.load(file)
  1617. file.close()
  1618. except:
  1619. output_noise_ablate_choice_end = multi_model_simulation_with_dynamic_silencing(
  1620. list(all_models.values()),
  1621. inputs,
  1622. unique_stims,
  1623. t_decision_window,
  1624. n_trials_per_stim=20,
  1625. sigma=list(inputs_sigma.values()),
  1626. silenced='ind_choice_end',
  1627. silence_time=inds_end,
  1628. label_data=label_data,
  1629. responses_over_time=False,
  1630. fit_options=None,
  1631. inds_L=None,
  1632. inds_R=None,
  1633. responses_beginning_end=False,
  1634. window_beginning=None,
  1635. window_end=None)
  1636. with open('data/model/output_tailored_noise_ablate_choice_end.pkl', 'wb') as file:
  1637. pickle.dump(output_noise_ablate_choice_end, file)
  1638. file.close()
  1639. # %% [markdown]
  1640. # ### Stats on accuracies
  1641. # %%
  1642. # 1-way repeated-measures ANOVA
  1643. data = [
  1644. (output_noise['all_accuracies'], 'control'),
  1645. (output_noise_ablate_linear['all_accuracies'], 'linear'),
  1646. (output_noise_ablate_perception['all_accuracies'], 'perception'),
  1647. (output_noise_ablate_choice['all_accuracies'], 'choice'),
  1648. (output_noise_ablate_other['all_accuracies'], 'other')
  1649. ]
  1650. n_subjects = n_sessions
  1651. n_conditions = len(data)
  1652. subject = [i for i in range(n_subjects)] * n_conditions
  1653. condition, measurement = [], []
  1654. for i in range(n_conditions):
  1655. condition += ['{}'.format(i)] * n_subjects
  1656. measurement += data[i][0]
  1657. df = pd.DataFrame({
  1658. 'subject': subject,
  1659. 'condition': condition,
  1660. 'measurement': measurement
  1661. })
  1662. # Perform repeated measures ANOVA
  1663. rm_anova = AnovaRM(df, 'measurement', 'subject', within=['condition'])
  1664. results = rm_anova.fit()
  1665. # Print the results
  1666. display(results.anova_table)
  1667. # post-hoc tests
  1668. correction = scipy.special.comb(n_conditions, 2) # bonferroni correction for number of paired tests
  1669. print('\nPost-hoc tests (Bonferroni-adjusted p-values):\n')
  1670. for i in range(n_conditions - 1):
  1671. for j in range(i + 1, n_conditions):
  1672. x1, label1 = data[i]
  1673. x2, label2 = data[j]
  1674. print(' {} vs. {}: p = {:.4e}'.format(label1, label2, correction * scipy.stats.ttest_rel(x1, x2)[1]))
  1675. # %%
  1676. # plot accuracies
  1677. data = [
  1678. ([_ for _ in output_noise['all_accuracies']], 'control'),
  1679. ([_ for _ in output_noise_ablate_linear['all_accuracies']], 'linear'),
  1680. ([_ for _ in output_noise_ablate_perception['all_accuracies']], 'perc'),
  1681. ([_ for _ in output_noise_ablate_choice['all_accuracies']], 'choice'),
  1682. ([_ for _ in output_noise_ablate_other['all_accuracies']], 'other'),
  1683. ]
  1684. fig, ax = plt.subplots(figsize=(12, 3));
  1685. ax.bar([data[i][1] for i in range(len(data))], [np.mean(data[i][0]) for i in range(len(data))]);
  1686. for i in range(len(data)):
  1687. ax.plot(i + 0.025 * np.random.normal(size=len(data[i][0])), data[i][0], '.k');
  1688. # Figure 7C
  1689. # plt.savefig('plots/model/accuracies.pdf');
  1690. # %%
  1691. # plot psychometrics
  1692. def plot_psychometric(all_p_left, color, style):
  1693. p_left_mat = np.empty([0, n_stims])
  1694. for p_left in all_p_left:
  1695. p_left_mat = np.vstack([p_left_mat, p_left])
  1696. ave_p_left = p_left_mat.mean(axis=0)
  1697. std_p_left = p_left_mat.std(axis=0) / np.sqrt(p_left_mat.shape[0])
  1698. psycho_res = fit_psychometric(unique_stims, ave_p_left)
  1699. ax.plot(psycho_res['x'], psycho_res['y'], color=color, linestyle=style);
  1700. for i_stim, stim in enumerate(unique_stims):
  1701. ax.plot([stim, stim], [ave_p_left[i_stim] - std_p_left[i_stim], ave_p_left[i_stim] + std_p_left[i_stim]],
  1702. color=color, linestyle=style);
  1703. ax.plot(stim, ave_p_left[i_stim], '.', color=color, markersize=10);
  1704. fig, ax = plt.subplots(figsize=(4.5, 4));
  1705. plot_psychometric(output_noise['all_p_left'], 'k', '-')
  1706. plot_psychometric(output_noise_ablate_linear['all_p_left'], 'r', '-')
  1707. plot_psychometric(output_noise_ablate_perception['all_p_left'], 'c', '-')
  1708. plot_psychometric(output_noise_ablate_choice['all_p_left'], 'b', '-')
  1709. plot_psychometric(output_noise_ablate_other['all_p_left'], 'g', '-')
  1710. ax.set_ylim([-0.05, 1.05]);
  1711. ax.set_xlabel('Stimulus (% sucrose)');
  1712. ax.set_ylabel('P(Sucrose choice)');
  1713. # Figure 7F
  1714. # plt.savefig('plots/model/ablation_psychometrics.pdf');
  1715. # %%
  1716. # stats on psychometrics (extra-sum-of-squares F tests)
  1717. p_left_control = np.array(output_noise['all_p_left']).mean(axis=0)
  1718. p_left_other = np.array(output_noise_ablate_other['all_p_left']).mean(axis=0)
  1719. p_left_linear = np.array(output_noise_ablate_linear['all_p_left']).mean(axis=0)
  1720. p_left_perception = np.array(output_noise_ablate_perception['all_p_left']).mean(axis=0)
  1721. p_left_choice = np.array(output_noise_ablate_choice['all_p_left']).mean(axis=0)
  1722. correction = scipy.special.comb(n_conditions, 2)
  1723. # control vs. ablated other
  1724. Y = [p_left_control, p_left_other]
  1725. res = psychometric_comparison_test(unique_stims, Y)
  1726. print('Control vs. ablated other, F = {:.4f}, p = {:.4f}, corrected p = {:.4f}'.format(
  1727. res['F'], res['p'], correction * res['p']
  1728. ))
  1729. # control vs. ablated linear
  1730. Y = [p_left_control, p_left_linear]
  1731. res = psychometric_comparison_test(unique_stims, Y)
  1732. print('Control vs. ablated linear, F = {:.4f}, p = {:.4f}, corrected p = {:.4f}'.format(
  1733. res['F'], res['p'], correction * res['p']
  1734. ))
  1735. # control vs. ablated perception
  1736. Y = [p_left_control, p_left_perception]
  1737. res = psychometric_comparison_test(unique_stims, Y)
  1738. print('Control vs. ablated perception, F = {:.4f}, p = {:.4f}, corrected p = {:.4f}'.format(
  1739. res['F'], res['p'], correction * res['p']
  1740. ))
  1741. # control vs. ablated choice
  1742. Y = [p_left_control, p_left_choice]
  1743. res = psychometric_comparison_test(unique_stims, Y)
  1744. print('Control vs. ablated choice, F = {:.4f}, p = {:.4f}, corrected p = {:.4f}'.format(
  1745. res['F'], res['p'], correction * res['p']
  1746. ))
  1747. # %%
  1748. '''
  1749. 2-way within-subjects ANOVA on behavioral performance
  1750. factor 1: coding type (4 levels)
  1751. factor 2: constrained vs unconstrained (2 levels)
  1752. There is only 1 control group b/c factor 2 does not apply when ablating nothing
  1753. Remove it as a level of factor 1 to avoid an empty cell in the 2-way ANOVA
  1754. To compare groups to control, separately use Dunnett's test on all 9 groups
  1755. '''
  1756. n_subjects = n_sessions
  1757. coding_type = ['linear', 'perception', 'choice', 'other']
  1758. constraint = ['constrained', 'unconstrained']
  1759. col_subject = [i for i in included_sessions] * (len(coding_type) * len(constraint))
  1760. col_coding_type, col_constraint, col_measurement = [], [], []
  1761. for coding_type_i in coding_type:
  1762. for constraint_i in constraint:
  1763. col_coding_type += [coding_type_i] * n_subjects
  1764. col_constraint += [constraint_i] * n_subjects
  1765. if coding_type_i == 'linear':
  1766. if constraint_i == 'constrained':
  1767. col_measurement += list(output_noise_ablate_linear_con['all_accuracies'])
  1768. else:
  1769. col_measurement += list(output_noise_ablate_linear_unc['all_accuracies'])
  1770. elif coding_type_i == 'perception':
  1771. if constraint_i == 'constrained':
  1772. col_measurement += list(output_noise_ablate_perception_con['all_accuracies'])
  1773. else:
  1774. col_measurement += list(output_noise_ablate_perception_unc['all_accuracies'])
  1775. if coding_type_i == 'choice':
  1776. if constraint_i == 'constrained':
  1777. col_measurement += list(output_noise_ablate_choice_con['all_accuracies'])
  1778. else:
  1779. col_measurement += list(output_noise_ablate_choice_unc['all_accuracies'])
  1780. if coding_type_i == 'other':
  1781. if constraint_i == 'constrained':
  1782. col_measurement += list(output_noise_ablate_other_con['all_accuracies'])
  1783. else:
  1784. col_measurement += list(output_noise_ablate_other_unc['all_accuracies'])
  1785. df = pd.DataFrame({
  1786. 'id': col_subject,
  1787. 'iv1': col_coding_type,
  1788. 'iv2': col_constraint,
  1789. 'dv': col_measurement
  1790. })
  1791. # Perform repeated measures ANOVA
  1792. print('2-way ANOVA')
  1793. print('Factor 1: coding_type')
  1794. print('Factor 2: constraint\n')
  1795. res_anova = rmAnova2Way(df)
  1796. # %%
  1797. # post-hoc tests
  1798. all_groups = [
  1799. ('linear', 'constrained'),
  1800. ('linear', 'unconstrained'),
  1801. ('perception', 'constrained'),
  1802. ('perception', 'unconstrained'),
  1803. ('choice', 'constrained'),
  1804. ('choice', 'unconstrained'),
  1805. ('other', 'constrained'),
  1806. ('other', 'unconstrained')
  1807. ]
  1808. n_groups = len(all_groups)
  1809. correction = scipy.special.comb(n_groups, 2) # bonferroni correction for number of paired tests
  1810. p_mat = np.full([n_groups, n_groups], np.nan)
  1811. print('\nPost-hoc tests (Bonferroni-adjusted p-values):\n')
  1812. for i in range(n_groups - 1):
  1813. for j in range(i + 1, n_groups):
  1814. coding_type_i, constraint_i = all_groups[i]
  1815. x_i = np.array([_ for _ in df[(df['iv1'] == coding_type_i) & (df['iv2'] == constraint_i)]['dv']])
  1816. label_i = coding_type_i + '_' + constraint_i
  1817. coding_type_j, constraint_j = all_groups[j]
  1818. x_j = np.array([_ for _ in df[(df['iv1'] == coding_type_j) & (df['iv2'] == constraint_j)]['dv']])
  1819. label_j = coding_type_j + '_' + constraint_j
  1820. p_adj = correction * scipy.stats.ttest_rel(x_i, x_j)[1]
  1821. print(' {} vs. {}: p = {:.4e}'.format(label_i, label_j, p_adj))
  1822. p_mat[i, j] = (p_adj < 0.01).astype(float)
  1823. print('')
  1824. print(p_mat)
  1825. # %%
  1826. ## dunnett test
  1827. data = [
  1828. ([_ for _ in output_noise_ablate_linear_con['all_accuracies']], 'linear_con'),
  1829. ([_ for _ in output_noise_ablate_linear_unc['all_accuracies']], 'linear_unc'),
  1830. ([_ for _ in output_noise_ablate_perception_con['all_accuracies']], 'perc_con'),
  1831. ([_ for _ in output_noise_ablate_perception_unc['all_accuracies']], 'perc_unc'),
  1832. ([_ for _ in output_noise_ablate_choice_con['all_accuracies']], 'choice_con'),
  1833. ([_ for _ in output_noise_ablate_choice_unc['all_accuracies']], 'choice_unc'),
  1834. ([_ for _ in output_noise_ablate_other_con['all_accuracies']], 'other_con'),
  1835. ([_ for _ in output_noise_ablate_other_unc['all_accuracies']], 'other_unc')
  1836. ]
  1837. samples = [np.array(data_[0]) for data_ in data]
  1838. labels = [data_[1] for data_ in data]
  1839. control = np.array([_ for _ in output_noise['all_accuracies']])
  1840. res = scipy.stats.dunnett(*samples, control=control)
  1841. print('Dunnett test:\n')
  1842. for i in range(len(data)):
  1843. print('Control vs {}: p = {:.4e}'.format(labels[i], res.pvalue[i]))
  1844. # %%
  1845. # plot accuracies
  1846. data = [
  1847. ([_ for _ in output_noise['all_accuracies']], 'control'),
  1848. ([_ for _ in output_noise_ablate_linear_con['all_accuracies']], 'linear_con'),
  1849. ([_ for _ in output_noise_ablate_linear_unc['all_accuracies']], 'linear_unc'),
  1850. ([_ for _ in output_noise_ablate_perception_con['all_accuracies']], 'perc_con'),
  1851. ([_ for _ in output_noise_ablate_perception_unc['all_accuracies']], 'perc_unc'),
  1852. ([_ for _ in output_noise_ablate_choice_con['all_accuracies']], 'choice_con'),
  1853. ([_ for _ in output_noise_ablate_choice_unc['all_accuracies']], 'choice_unc'),
  1854. ([_ for _ in output_noise_ablate_other_con['all_accuracies']], 'other_con'),
  1855. ([_ for _ in output_noise_ablate_other_unc['all_accuracies']], 'other_unc')
  1856. ]
  1857. fig, ax = plt.subplots(figsize=(12, 3));
  1858. ax.bar([data[i][1] for i in range(len(data))], [np.mean(data[i][0]) for i in range(len(data))]);
  1859. for i in range(len(data)):
  1860. ax.plot(i + 0.025 * np.random.normal(size=len(data[i][0])), data[i][0], '.k');
  1861. # Figure 7 - Supplement 1 C
  1862. # plt.savefig('plots/model/accuracies_con_vs_unc.pdf');
  1863. # %%
  1864. '''
  1865. Repeat above for beginning vs end
  1866. '''
  1867. n_subjects = n_sessions
  1868. coding_type = ['linear', 'perception', 'choice']
  1869. window = ['beginning', 'end']
  1870. col_subject = [i for i in included_sessions] * (len(coding_type) * len(window))
  1871. col_coding_type, col_window, col_measurement = [], [], []
  1872. for coding_type_i in coding_type:
  1873. for window_i in window:
  1874. col_coding_type += [coding_type_i] * n_subjects
  1875. col_window += [window_i] * n_subjects
  1876. if coding_type_i == 'linear':
  1877. if window_i == 'beginning':
  1878. col_measurement += list(output_noise_ablate_linear_beginning['all_accuracies'])
  1879. else:
  1880. col_measurement += list(output_noise_ablate_linear_end['all_accuracies'])
  1881. elif coding_type_i == 'perception':
  1882. if window_i == 'beginning':
  1883. col_measurement += list(output_noise_ablate_perception_beginning['all_accuracies'])
  1884. else:
  1885. col_measurement += list(output_noise_ablate_perception_end['all_accuracies'])
  1886. if coding_type_i == 'choice':
  1887. if window_i == 'beginning':
  1888. col_measurement += list(output_noise_ablate_choice_beginning['all_accuracies'])
  1889. else:
  1890. col_measurement += list(output_noise_ablate_choice_end['all_accuracies'])
  1891. df = pd.DataFrame({
  1892. 'id': col_subject,
  1893. 'iv1': col_coding_type,
  1894. 'iv2': col_window,
  1895. 'dv': col_measurement
  1896. })
  1897. # Perform repeated measures ANOVA
  1898. print('2-way ANOVA')
  1899. print('Factor 1: coding_type')
  1900. print('Factor 2: window\n')
  1901. res_anova = rmAnova2Way(df)
  1902. # %%
  1903. # post-hoc tests
  1904. all_groups = [
  1905. ('linear', 'beginning'),
  1906. ('linear', 'end'),
  1907. ('perception', 'beginning'),
  1908. ('perception', 'end'),
  1909. ('choice', 'beginning'),
  1910. ('choice', 'end')
  1911. ]
  1912. n_groups = len(all_groups)
  1913. correction = scipy.special.comb(n_groups, 2) # bonferroni correction for number of paired tests
  1914. p_mat = np.full([n_groups, n_groups], np.nan)
  1915. print('\nPost-hoc tests (Bonferroni-adjusted p-values):\n')
  1916. for i in range(n_groups - 1):
  1917. for j in range(i + 1, n_groups):
  1918. coding_type_i, window_i = all_groups[i]
  1919. x_i = np.array([_ for _ in df[(df['iv1'] == coding_type_i) & (df['iv2'] == window_i)]['dv']])
  1920. label_i = coding_type_i + '_' + window_i
  1921. coding_type_j, window_j = all_groups[j]
  1922. x_j = np.array([_ for _ in df[(df['iv1'] == coding_type_j) & (df['iv2'] == window_j)]['dv']])
  1923. label_j = coding_type_j + '_' + window_j
  1924. p_adj = correction * scipy.stats.ttest_rel(x_i, x_j)[1]
  1925. print(' {} vs. {}: p = {:.4e}'.format(label_i, label_j, p_adj))
  1926. p_mat[i, j] = (p_adj < 0.01).astype(float)
  1927. print('')
  1928. print(p_mat)
  1929. # %%
  1930. ## dunnett test
  1931. data = [
  1932. ([_ for _ in output_noise_ablate_linear_beginning['all_accuracies']], 'linear_beg'),
  1933. ([_ for _ in output_noise_ablate_linear_end['all_accuracies']], 'linear_end'),
  1934. ([_ for _ in output_noise_ablate_perception_beginning['all_accuracies']], 'perc_beg'),
  1935. ([_ for _ in output_noise_ablate_perception_end['all_accuracies']], 'perc_end'),
  1936. ([_ for _ in output_noise_ablate_choice_beginning['all_accuracies']], 'choice_beg'),
  1937. ([_ for _ in output_noise_ablate_choice_end['all_accuracies']], 'choice_end')
  1938. ]
  1939. samples = [np.array(data_[0]) for data_ in data]
  1940. labels = [data_[1] for data_ in data]
  1941. control = np.array([_ for _ in output_noise['all_accuracies']])
  1942. res = scipy.stats.dunnett(*samples, control=control)
  1943. print('Dunnett test:\n')
  1944. for i in range(len(data)):
  1945. print('Control vs {}: p = {:.4e}'.format(labels[i], res.pvalue[i]))
  1946. # %%
  1947. data = [
  1948. ([_ for _ in output_noise['all_accuracies']], 'control'),
  1949. ([_ for _ in output_noise_ablate_linear_beginning['all_accuracies']], 'linear_beg'),
  1950. ([_ for _ in output_noise_ablate_linear_end['all_accuracies']], 'linear_end'),
  1951. ([_ for _ in output_noise_ablate_perception_beginning['all_accuracies']], 'perc_beg'),
  1952. ([_ for _ in output_noise_ablate_perception_end['all_accuracies']], 'perc_end'),
  1953. ([_ for _ in output_noise_ablate_choice_beginning['all_accuracies']], 'choice_beg'),
  1954. ([_ for _ in output_noise_ablate_choice_end['all_accuracies']], 'choice_end')
  1955. ]
  1956. fig, ax = plt.subplots(figsize=(12, 3));
  1957. ax.bar([data[i][1] for i in range(len(data))], [np.mean(data[i][0]) for i in range(len(data))]);
  1958. for i in range(len(data)):
  1959. ax.plot(i + 0.025 * np.random.normal(size=len(data[i][0])), data[i][0], '.k');
  1960. # Figure 7 - Supplement 2 C
  1961. # plt.savefig('plots/review/model/accuracies_beg_vs_end.pdf');
  1962. # %%
  1963. # numbers of each
  1964. for i_session, session_key in enumerate(included_sessions):
  1965. print('Session {}'.format(session_key))
  1966. n_con = all_models[session_key].observed_size
  1967. n_tot = all_models[session_key].network_size
  1968. for coding_type in ['Linear', 'Step (Perception)', 'Step (Choice)', 'Other']:
  1969. n_c, n_u = 0, 0
  1970. for i_neuron in range(n_tot):
  1971. label_seq = [
  1972. output_noise['all_labels'][i_session][i_t][i_neuron]
  1973. for i_t in range(len(output_noise['all_labels'][i_session]))
  1974. ]
  1975. if coding_type != 'Other':
  1976. if coding_type in label_seq:
  1977. if i_neuron < n_con:
  1978. n_c += 1
  1979. else:
  1980. n_u += 1
  1981. else:
  1982. if all([label == 'Other' for label in label_seq]):
  1983. if i_neuron < n_con:
  1984. n_c += 1
  1985. else:
  1986. n_u += 1
  1987. print(' {}'.format(coding_type))
  1988. print(' Constrained: {}'.format(n_c))
  1989. print(' Unconstrained: {}'.format(n_u))
  1990. print('---- TOTALS ----')
  1991. print('Number of constrained linear: {}'
  1992. .format(sum([len(label_data[i]['ind_lin_con']) for i in range(len(label_data))])))
  1993. print('Number of unconstrained linear: {}'
  1994. .format(sum([len(label_data[i]['ind_lin_unc']) for i in range(len(label_data))])))
  1995. print('Number of constrained perception: {}'
  1996. .format(sum([len(label_data[i]['ind_perception_con']) for i in range(len(label_data))])))
  1997. print('Number of unconstrained perception: {}'
  1998. .format(sum([len(label_data[i]['ind_perception_unc']) for i in range(len(label_data))])))
  1999. print('Number of constrained choice: {}'
  2000. .format(sum([len(label_data[i]['ind_choice_con']) for i in range(len(label_data))])))
  2001. print('Number of unconstrained choice: {}'
  2002. .format(sum([len(label_data[i]['ind_choice_unc']) for i in range(len(label_data))])))
  2003. print('Number of constrained other: {}'
  2004. .format(sum([len(label_data[i]['ind_other_con']) for i in range(len(label_data))])))
  2005. print('Number of unconstrained other: {}'
  2006. .format(sum([len(label_data[i]['ind_other_unc']) for i in range(len(label_data))])))
  2007. # %% [markdown]
  2008. # ### --- dPCA ---
  2009. # %% [markdown]
  2010. # #### Original (complete ablation of each coding category)
  2011. # %%
  2012. filter_models = True # filter out models that do not have at least 2 trials per choice across all ablation conditions
  2013. # this is done to improve the linear prediction of missing trial types
  2014. outputs = {
  2015. 'baseline': output_noise,
  2016. 'other': output_noise_ablate_other,
  2017. 'linear': output_noise_ablate_linear,
  2018. 'perception': output_noise_ablate_perception,
  2019. 'choice': output_noise_ablate_choice
  2020. }
  2021. conditions = outputs.keys()
  2022. # filtering criteria
  2023. good_model_inds = []
  2024. for i_model in range(len(included_sessions)):
  2025. is_good = True
  2026. if filter_models:
  2027. for condition in conditions:
  2028. y_choice = outputs[condition]['y_lefts'][i_model]
  2029. if (y_choice.sum() <= 1) or (y_choice.sum() >= len(y_choice) - 1):
  2030. is_good = False
  2031. break
  2032. if is_good:
  2033. good_model_inds.append(i_model)
  2034. if filter_models:
  2035. print('{}/{} models passed.'.format(len(good_model_inds), len(included_sessions)))
  2036. else:
  2037. print('All models included.')
  2038. file_name = 'filtered' if filter_models else 'all'
  2039. # %%
  2040. try:
  2041. with open('data/model/dPCA_{}_sessions_results.pkl'.format(file_name), 'rb') as file:
  2042. dPCA_results = pickle.load(file)
  2043. file.close()
  2044. except:
  2045. dPCA_results = {}
  2046. for condition in conditions:
  2047. # get maximum number of trials for any parameter combination across sessions
  2048. max_num_trials = 0
  2049. n_neurons_pseudo = 0
  2050. flagged_models = []
  2051. for i_model in good_model_inds:
  2052. X = outputs[condition]['Xs'][i_model]
  2053. y_stim = outputs[condition]['y_stims'][i_model]
  2054. y_choice = outputs[condition]['y_lefts'][i_model]
  2055. n_neurons_pseudo += X.shape[1]
  2056. for stim in unique_stims:
  2057. for choice in [0, 1]:
  2058. n = ((y_stim == stim) & (y_choice == choice)).sum()
  2059. if (n == 0) and (i_model not in flagged_models):
  2060. flagged_models.append(i_model)
  2061. max_num_trials = max(max_num_trials, n)
  2062. # assemble pseudo-population trial-by-trial tensor
  2063. # X_pseudo: (trials x neurons x stimuli x decisions x time)
  2064. X_pseudo = np.full([max_num_trials, n_neurons_pseudo, len(unique_stims), 2, len(t)], np.nan)
  2065. n_neurons_cumul = 0
  2066. for i_model in good_model_inds:
  2067. X = outputs[condition]['Xs'][i_model]
  2068. y_stim = outputs[condition]['y_stims'][i_model]
  2069. y_choice = outputs[condition]['y_lefts'][i_model]
  2070. n_neurons = X.shape[1]
  2071. ###----------------------------------------------------------------------------------------------------
  2072. if i_model in flagged_models:
  2073. ### issue: some data are missing for this session
  2074. ### workaround: fit a simple linear model to existing data and use it to predict the missing parts
  2075. predictors = np.vstack([
  2076. y_stim.reshape([1, -1]),
  2077. y_choice.reshape([1, -1]),
  2078. (y_stim * y_choice).reshape([1, -1]),
  2079. np.ones([1, X.shape[0]])
  2080. ])
  2081. betas = [[scipy.linalg.pinv(predictors @ predictors.T) @ predictors @ X[:, i_n, i_t].reshape([-1, 1])
  2082. for i_t in range(len(t))] for i_n in range(n_neurons)]
  2083. ###-----------------------------------------------------------------------------------------------------
  2084. for i_stim, stim in enumerate(unique_stims):
  2085. for i_choice, choice in enumerate([0, 1]):
  2086. trial_mask = (y_stim == stim) & (y_choice == choice)
  2087. n_trials = trial_mask.sum()
  2088. if n_trials == 0:
  2089. ### use the linear model's prediction ------------------------------------------------------
  2090. for i_n in range(n_neurons):
  2091. for i_t in range(len(t)):
  2092. r = max((betas[i_n][i_t].flatten() * np.array([stim, choice, stim * choice, 1])).sum(), 0)
  2093. X_pseudo[:2, n_neurons_cumul + i_n, i_stim, i_choice, i_t] = np.array([r, r])
  2094. ### ----------------------------------------------------------------------------------------
  2095. elif n_trials == 1:
  2096. X_pseudo[:2, n_neurons_cumul:(n_neurons_cumul + n_neurons), i_stim, i_choice, :] = \
  2097. np.concatenate([
  2098. X[trial_mask, :, :].reshape([1, n_neurons, len(t)]),
  2099. X[trial_mask, :, :].reshape([1, n_neurons, len(t)])],
  2100. axis=0)
  2101. else:
  2102. X_pseudo[:n_trials, n_neurons_cumul:(n_neurons_cumul + n_neurons), i_stim, i_choice, :] = \
  2103. X[trial_mask, :, :]
  2104. n_neurons_cumul += n_neurons
  2105. # get the pseudo-population trial-averaged (PSTH) tensor
  2106. X_pseudo_psth = np.nanmean(X_pseudo, axis=0)
  2107. # do the dPCA
  2108. dpca = dPCA.dPCA(labels='sdt', regularizer='auto',
  2109. join={'st': ['s', 'st'], 'dt': ['d', 'dt'], 'sdt': ['sd', 'sdt']})
  2110. dpca.protect = ['t']
  2111. Z = dpca.fit_transform(X_pseudo_psth, X_pseudo)
  2112. dPCA_results[condition] = {'X_pseudo': X_pseudo, 'X_pseudo_psth': X_pseudo_psth, 'dpca': dpca, 'Z': Z}
  2113. with open('data/model/dPCA_{}_sessions_results.pkl'.format(file_name), 'wb') as file:
  2114. pickle.dump(dPCA_results, file)
  2115. file.close()
  2116. # %%
  2117. fig, axes = plt.subplots(len(conditions), 1, figsize=(11, 3 * len(conditions)));
  2118. y_limits = []
  2119. mean_stim_projections = []
  2120. for i_condition, condition in enumerate(conditions):
  2121. temp = []
  2122. for i_stim, stim in enumerate(unique_stims):
  2123. for i_choice, choice in enumerate([0, 1]):
  2124. style = '-' if (((stim < 50) and (choice == 0)) or ((stim > 50) and (choice == 1))) else ':'
  2125. v = dPCA_results['baseline']['dpca'].D['st'][:, 0]
  2126. X_tensor = dPCA_results[condition]['X_pseudo_psth']
  2127. X_tilde = X_tensor.reshape([X_tensor.shape[0], -1])
  2128. X_tilde -= X_tilde.mean(axis=1).reshape([-1, 1])
  2129. X_tensor = X_tilde.reshape(X_tensor.shape)
  2130. X = X_tensor[:, i_stim, i_choice, :]
  2131. y = [(v * X[:, k]).sum() for k in range(len(t))]
  2132. temp.append(np.mean(np.abs(y)))
  2133. axes[i_condition].plot(t, y, color=colors[i_stim, :], linestyle=style);
  2134. axes[i_condition].axvline(0, color='k', linestyle=':');
  2135. axes[i_condition].axvline(t_D, color='k', linestyle=':');
  2136. axes[i_condition].set_title('Stimulus coding ({})'.format(condition));
  2137. y_limits.append(axes[i_condition].get_ylim())
  2138. mean_stim_projections.append(temp)
  2139. global_limit = np.max([np.max(np.abs(lims)) for lims in y_limits])
  2140. for ax in axes:
  2141. ax.set_xlim([t[0], t[-1]]);
  2142. ax.set_ylim([-global_limit, global_limit]);
  2143. # Figure 6A, 7A (left column)
  2144. # plt.savefig('plots/model/dPCA_stimulus_coding.pdf');
  2145. # %%
  2146. fig, axes = plt.subplots(len(conditions), 1, figsize=(11, 3 * len(conditions)));
  2147. y_limits = []
  2148. mean_choice_projections = []
  2149. for i_condition, condition in enumerate(conditions):
  2150. temp = []
  2151. for i_stim, stim in enumerate(unique_stims):
  2152. for i_choice, choice in enumerate([0, 1]):
  2153. style = '-' if (((stim < 50) and (choice == 0)) or ((stim > 50) and (choice == 1))) else ':'
  2154. v = dPCA_results['baseline']['dpca'].D['dt'][:, 0]
  2155. X_tensor = dPCA_results[condition]['X_pseudo_psth']
  2156. X_tilde = X_tensor.reshape([X_tensor.shape[0], -1])
  2157. X_tilde -= X_tilde.mean(axis=1).reshape([-1, 1])
  2158. X_tensor = X_tilde.reshape(X_tensor.shape)
  2159. X = X_tensor[:, i_stim, i_choice, :]
  2160. y = [(v * X[:, k]).sum() for k in range(len(t))]
  2161. temp.append(np.mean(np.abs(y)))
  2162. axes[i_condition].plot(t, y, color=colors[i_stim, :], linestyle=style);
  2163. axes[i_condition].axvline(0, color='k', linestyle=':');
  2164. axes[i_condition].axvline(t_D, color='k', linestyle=':');
  2165. axes[i_condition].set_title('Choice coding ({})'.format(condition));
  2166. y_limits.append(axes[i_condition].get_ylim())
  2167. mean_choice_projections.append(temp)
  2168. global_limit = np.max([np.max(np.abs(lims)) for lims in y_limits])
  2169. for ax in axes:
  2170. ax.set_xlim([t[0], t[-1]]);
  2171. ax.set_ylim([-global_limit, global_limit]);
  2172. # Figure 6B, 7A (right column)
  2173. # plt.savefig('plots/model/dPCA_choice_coding.pdf');
  2174. # %%
  2175. # calculate component overlaps
  2176. stim_ov_mat = np.zeros([len(conditions), len(conditions)])
  2177. choice_ov_mat = np.zeros([len(conditions), len(conditions)])
  2178. for i_condition, condition_i in enumerate(conditions):
  2179. for j_condition, condition_j in enumerate(conditions):
  2180. v_stim_i = dPCA_results[condition_i]['dpca'].P['st'][:, 0]
  2181. v_stim_j = dPCA_results[condition_j]['dpca'].P['st'][:, 0]
  2182. stim_ov_mat[i_condition, j_condition] = overlap(v_stim_i, v_stim_j)
  2183. v_choice_i = dPCA_results[condition_i]['dpca'].P['dt'][:, 0]
  2184. v_choice_j = dPCA_results[condition_j]['dpca'].P['dt'][:, 0]
  2185. choice_ov_mat[i_condition, j_condition] = overlap(v_choice_i, v_choice_j)
  2186. # %%
  2187. # plot stimulus component overlaps
  2188. fig, ax = plt.subplots(figsize=(5, 5));
  2189. cmap = LinearSegmentedColormap.from_list('my_cmap', [(0, [1, 1, 1]), (1, [0, 0, 0])])
  2190. ax.matshow(stim_ov_mat, cmap=cmap);
  2191. ax.set_xticks(range(len(conditions)));
  2192. ax.set_xticklabels(conditions);
  2193. ax.set_yticks(range(len(conditions)));
  2194. ax.set_yticklabels(conditions);
  2195. ax.set_title('Stimulus coding direction');
  2196. plt.colorbar(plt.cm.ScalarMappable(cmap=cmap), ax=ax);
  2197. # Figure 7B (top)
  2198. # plt.savefig('plots/model/dpca_stimulus_coding_direction_overlaps.pdf');
  2199. # %%
  2200. print('Stimulus component overlaps:\n')
  2201. for i_condition, condition_i in enumerate(conditions):
  2202. for j_condition, condition_j in enumerate(conditions):
  2203. v_stim_i = dPCA_results[condition_i]['dpca'].P['st'][:, 0]
  2204. v_stim_j = dPCA_results[condition_j]['dpca'].P['st'][:, 0]
  2205. ov = overlap(v_stim_i, v_stim_j)
  2206. print(' {} vs. {}: {:.4f}'.format(condition_i, condition_j, ov))
  2207. # %%
  2208. # stats on stimulus component projections
  2209. data_stim = [(mean_stim_projections[i_condition], condition) for i_condition, condition in enumerate(conditions)]
  2210. data_stim[1], data_stim[2], data_stim[3], data_stim[4] = (
  2211. data_stim[2], data_stim[3], data_stim[4], data_stim[1]) # swap order for plot
  2212. n_subjects = 16 # 'subjects' here are trial conditions (8 stims * 2 choices)
  2213. n_conditions = len(data_stim)
  2214. subject = [i for i in range(n_subjects)] * n_conditions
  2215. condition, measurement = [], []
  2216. for i in range(n_conditions):
  2217. condition += ['{}'.format(i)] * n_subjects
  2218. measurement += data_stim[i][0]
  2219. df = pd.DataFrame({
  2220. 'subject': subject,
  2221. 'condition': condition,
  2222. 'measurement': measurement
  2223. })
  2224. # Perform repeated measures ANOVA
  2225. rm_anova = AnovaRM(df, 'measurement', 'subject', within=['condition'])
  2226. results = rm_anova.fit()
  2227. # Print the results
  2228. display(results.anova_table)
  2229. # post-hoc tests
  2230. correction = scipy.special.comb(n_conditions, 2) # bonferroni correction for number of paired tests
  2231. print('\nPost-hoc tests (Bonferroni-adjusted p-values):\n')
  2232. for i in range(n_conditions - 1):
  2233. for j in range(i + 1, n_conditions):
  2234. x1, label1 = data_stim[i]
  2235. x2, label2 = data_stim[j]
  2236. print(' {} vs. {}: p = {:.4e}'.format(label1, label2, correction * scipy.stats.ttest_rel(x1, x2)[1]))
  2237. # %%
  2238. # plot choice component overlaps
  2239. fig, ax = plt.subplots(figsize=(5, 5));
  2240. cmap = LinearSegmentedColormap.from_list('my_cmap', [(0, [1, 1, 1]), (1, [0, 0, 0])])
  2241. ax.matshow(choice_ov_mat, cmap=cmap);
  2242. ax.set_xticks(range(len(conditions)));
  2243. ax.set_xticklabels(conditions);
  2244. ax.set_yticks(range(len(conditions)));
  2245. ax.set_yticklabels(conditions);
  2246. ax.set_title('Choice coding direction');
  2247. plt.colorbar(plt.cm.ScalarMappable(cmap=cmap), ax=ax);
  2248. # Figure 7B (bottom)
  2249. # plt.savefig('plots/model/dpca_choice_coding_direction_overlaps.pdf');
  2250. # %%
  2251. print('\nChoice component overlaps:\n')
  2252. for i_condition, condition_i in enumerate(conditions):
  2253. for j_condition, condition_j in enumerate(conditions):
  2254. v_choice_i = dPCA_results[condition_i]['dpca'].P['dt'][:, 0]
  2255. v_choice_j = dPCA_results[condition_j]['dpca'].P['dt'][:, 0]
  2256. ov = overlap(v_choice_i, v_choice_j)
  2257. print(' {} vs. {}: {:.4f}'.format(condition_i, condition_j, ov))
  2258. # %%
  2259. # stats on choice component projections
  2260. data_choice = [(mean_choice_projections[i_condition], condition) for i_condition, condition in enumerate(conditions)]
  2261. data_choice[1], data_choice[2], data_choice[3], data_choice[4] = (
  2262. data_choice[2], data_choice[3], data_choice[4], data_choice[1]) # swap order for plot
  2263. n_subjects = 16 # 'subjects' here are trial conditions (8 stims * 2 choices)
  2264. n_conditions = len(data_choice)
  2265. subject = [i for i in range(n_subjects)] * n_conditions
  2266. condition, measurement = [], []
  2267. for i in range(n_conditions):
  2268. condition += ['{}'.format(i)] * n_subjects
  2269. measurement += data_choice[i][0]
  2270. df = pd.DataFrame({
  2271. 'subject': subject,
  2272. 'condition': condition,
  2273. 'measurement': measurement
  2274. })
  2275. # Perform repeated measures ANOVA
  2276. rm_anova = AnovaRM(df, 'measurement', 'subject', within=['condition'])
  2277. results = rm_anova.fit()
  2278. # Print the results
  2279. display(results.anova_table)
  2280. # post-hoc tests
  2281. correction = scipy.special.comb(n_conditions, 2) # bonferroni correction for number of paired tests
  2282. print('\nPost-hoc tests (Bonferroni-adjusted p-values):\n')
  2283. for i in range(n_conditions - 1):
  2284. for j in range(i + 1, n_conditions):
  2285. x1, label1 = data_choice[i]
  2286. x2, label2 = data_choice[j]
  2287. print(' {} vs. {}: p = {:.4e}'.format(label1, label2, correction * scipy.stats.ttest_rel(x1, x2)[1]))
  2288. # %% [markdown]
  2289. # ### --- dPCA ---
  2290. # %% [markdown]
  2291. # #### dPCA for constrained units only, no ablations
  2292. # %%
  2293. # filter in the original way so F6S1A-B can be directly compared to 6A-B
  2294. filter_models = True # filter out models that do not have at least 2 trials per choice across all ablation conditions
  2295. # this is done to improve the linear prediction of missing trial types
  2296. outputs = {
  2297. 'baseline': output_noise,
  2298. 'other': output_noise_ablate_other,
  2299. 'linear': output_noise_ablate_linear,
  2300. 'perception': output_noise_ablate_perception,
  2301. 'choice': output_noise_ablate_choice
  2302. }
  2303. conditions = outputs.keys()
  2304. # filtering criteria
  2305. good_model_inds = []
  2306. for i_model in range(len(included_sessions)):
  2307. is_good = True
  2308. if filter_models:
  2309. for condition in conditions:
  2310. y_choice = outputs[condition]['y_lefts'][i_model]
  2311. if (y_choice.sum() <= 1) or (y_choice.sum() >= len(y_choice) - 1):
  2312. is_good = False
  2313. break
  2314. if is_good:
  2315. good_model_inds.append(i_model)
  2316. if filter_models:
  2317. print('{}/{} models passed.'.format(len(good_model_inds), len(included_sessions)))
  2318. else:
  2319. print('All models included.')
  2320. file_name = 'filtered' if filter_models else 'all'
  2321. # %%
  2322. try:
  2323. with open('data/model/dPCA_{}_constrained_results.pkl'.format(file_name), 'rb') as file:
  2324. dPCA_results = pickle.load(file)
  2325. file.close()
  2326. except:
  2327. dPCA_results = {}
  2328. # get maximum number of trials for any parameter combination across sessions
  2329. max_num_trials = 0
  2330. n_neurons_pseudo = 0
  2331. flagged_models = []
  2332. for i_model in good_model_inds:
  2333. X = outputs['baseline']['Xs'][i_model]
  2334. y_stim = outputs['baseline']['y_stims'][i_model]
  2335. y_choice = outputs['baseline']['y_lefts'][i_model]
  2336. n_neurons_pseudo += int(round(X.shape[1] / 5.88))
  2337. for stim in unique_stims:
  2338. for choice in [0, 1]:
  2339. n = ((y_stim == stim) & (y_choice == choice)).sum()
  2340. if (n == 0) and (i_model not in flagged_models):
  2341. flagged_models.append(i_model)
  2342. max_num_trials = max(max_num_trials, n)
  2343. # assemble pseudo-population trial-by-trial tensor
  2344. # X_pseudo: (trials x neurons x stimuli x decisions x time)
  2345. X_pseudo = np.full([max_num_trials, n_neurons_pseudo, len(unique_stims), 2, len(t)], np.nan)
  2346. n_neurons_cumul = 0
  2347. for i_model in good_model_inds:
  2348. X = outputs['baseline']['Xs'][i_model]
  2349. X = X[:, :int(round(X.shape[1] / 5.88)), :]
  2350. y_stim = outputs['baseline']['y_stims'][i_model]
  2351. y_choice = outputs['baseline']['y_lefts'][i_model]
  2352. n_neurons = X.shape[1]
  2353. ###----------------------------------------------------------------------------------------------------
  2354. if i_model in flagged_models:
  2355. ### issue: some data are missing for this session
  2356. ### workaround: fit a simple linear model to existing data and use it to predict the missing parts
  2357. predictors = np.vstack([
  2358. y_stim.reshape([1, -1]),
  2359. y_choice.reshape([1, -1]),
  2360. (y_stim * y_choice).reshape([1, -1]),
  2361. np.ones([1, X.shape[0]])
  2362. ])
  2363. betas = [[scipy.linalg.pinv(predictors @ predictors.T) @ predictors @ X[:, i_n, i_t].reshape([-1, 1])
  2364. for i_t in range(len(t))] for i_n in range(n_neurons)]
  2365. ###-----------------------------------------------------------------------------------------------------
  2366. for i_stim, stim in enumerate(unique_stims):
  2367. for i_choice, choice in enumerate([0, 1]):
  2368. trial_mask = (y_stim == stim) & (y_choice == choice)
  2369. n_trials = trial_mask.sum()
  2370. if n_trials == 0:
  2371. ### use the linear model's prediction ------------------------------------------------------
  2372. for i_n in range(n_neurons):
  2373. for i_t in range(len(t)):
  2374. r = max((betas[i_n][i_t].flatten() * np.array([stim, choice, stim * choice, 1])).sum(), 0)
  2375. X_pseudo[:2, n_neurons_cumul + i_n, i_stim, i_choice, i_t] = np.array([r, r])
  2376. ### ----------------------------------------------------------------------------------------
  2377. elif n_trials == 1:
  2378. X_pseudo[:2, n_neurons_cumul:(n_neurons_cumul + n_neurons), i_stim, i_choice, :] = \
  2379. np.concatenate([
  2380. X[trial_mask, :, :].reshape([1, n_neurons, len(t)]),
  2381. X[trial_mask, :, :].reshape([1, n_neurons, len(t)])],
  2382. axis=0)
  2383. else:
  2384. X_pseudo[:n_trials, n_neurons_cumul:(n_neurons_cumul + n_neurons), i_stim, i_choice, :] = \
  2385. X[trial_mask, :, :]
  2386. n_neurons_cumul += n_neurons
  2387. # get the pseudo-population trial-averaged (PSTH) tensor
  2388. X_pseudo_psth = np.nanmean(X_pseudo, axis=0)
  2389. # do the dPCA
  2390. dpca = dPCA.dPCA(labels='sdt', regularizer='auto',
  2391. join={'st': ['s', 'st'], 'dt': ['d', 'dt'], 'sdt': ['sd', 'sdt']})
  2392. dpca.protect = ['t']
  2393. Z = dpca.fit_transform(X_pseudo_psth, X_pseudo)
  2394. dPCA_results['baseline'] = {'X_pseudo': X_pseudo, 'X_pseudo_psth': X_pseudo_psth, 'dpca': dpca, 'Z': Z}
  2395. with open('data/model/dPCA_{}_constrained_results.pkl'.format(file_name), 'wb') as file:
  2396. pickle.dump(dPCA_results, file)
  2397. file.close()
  2398. # %%
  2399. fig, axes = plt.subplots(figsize=(11, 3));
  2400. for i_stim, stim in enumerate(unique_stims):
  2401. for i_choice, choice in enumerate([0, 1]):
  2402. style = '-' if (((stim < 50) and (choice == 0)) or ((stim > 50) and (choice == 1))) else ':'
  2403. v = dPCA_results['baseline']['dpca'].D['st'][:, 0]
  2404. X_tensor = dPCA_results['baseline']['X_pseudo_psth']
  2405. X_tilde = X_tensor.reshape([X_tensor.shape[0], -1])
  2406. X_tilde -= X_tilde.mean(axis=1).reshape([-1, 1])
  2407. X_tensor = X_tilde.reshape(X_tensor.shape)
  2408. X = X_tensor[:, i_stim, i_choice, :]
  2409. y = [(v * X[:, k]).sum() for k in range(len(t))]
  2410. axes.plot(t, y, color=colors[i_stim, :], linestyle=style);
  2411. axes.axvline(0, color='k', linestyle=':');
  2412. axes.axvline(t_D, color='k', linestyle=':');
  2413. axes.set_title('Stimulus coding');
  2414. axes.set_xlim([t[0], t[-1]]);
  2415. # Figure 6 - Supplement 1 A
  2416. # plt.savefig('plots/review/model/dPCA_stimulus_coding_constrained.pdf');
  2417. # %%
  2418. fig, axes = plt.subplots(figsize=(11, 3));
  2419. for i_stim, stim in enumerate(unique_stims):
  2420. for i_choice, choice in enumerate([0, 1]):
  2421. style = '-' if (((stim < 50) and (choice == 0)) or ((stim > 50) and (choice == 1))) else ':'
  2422. v = dPCA_results['baseline']['dpca'].D['dt'][:, 0]
  2423. X_tensor = dPCA_results['baseline']['X_pseudo_psth']
  2424. X_tilde = X_tensor.reshape([X_tensor.shape[0], -1])
  2425. X_tilde -= X_tilde.mean(axis=1).reshape([-1, 1])
  2426. X_tensor = X_tilde.reshape(X_tensor.shape)
  2427. X = X_tensor[:, i_stim, i_choice, :]
  2428. y = [(v * X[:, k]).sum() for k in range(len(t))]
  2429. axes.plot(t, y, color=colors[i_stim, :], linestyle=style);
  2430. axes.axvline(0, color='k', linestyle=':');
  2431. axes.axvline(t_D, color='k', linestyle=':');
  2432. axes.set_title('Choice coding');
  2433. axes.set_xlim([t[0], t[-1]]);
  2434. y_lim = np.max(np.abs(axes.get_ylim())); axes.set_ylim([-y_lim, y_lim]);
  2435. # Figure 6 - Supplement 1 B
  2436. # plt.savefig('plots/review/model/dPCA_choice_coding_constrained.pdf');
  2437. # %% [markdown]
  2438. # ### --- dPCA ---
  2439. # %% [markdown]
  2440. # #### dPCA for unconstrained units only, no ablations
  2441. # %%
  2442. # filter in the original way so F6S2A-B can be directly compared to 6A-B
  2443. filter_models = True # filter out models that do not have at least 2 trials per choice across all ablation conditions
  2444. # this is done to improve the linear prediction of missing trial types
  2445. outputs = {
  2446. 'baseline': output_noise,
  2447. 'other': output_noise_ablate_other,
  2448. 'linear': output_noise_ablate_linear,
  2449. 'perception': output_noise_ablate_perception,
  2450. 'choice': output_noise_ablate_choice
  2451. }
  2452. conditions = outputs.keys()
  2453. # filtering criteria
  2454. good_model_inds = []
  2455. for i_model in range(len(included_sessions)):
  2456. is_good = True
  2457. if filter_models:
  2458. for condition in conditions:
  2459. y_choice = outputs[condition]['y_lefts'][i_model]
  2460. if (y_choice.sum() <= 1) or (y_choice.sum() >= len(y_choice) - 1):
  2461. is_good = False
  2462. break
  2463. if is_good:
  2464. good_model_inds.append(i_model)
  2465. if filter_models:
  2466. print('{}/{} models passed.'.format(len(good_model_inds), len(included_sessions)))
  2467. else:
  2468. print('All models included.')
  2469. file_name = 'filtered' if filter_models else 'all'
  2470. # %%
  2471. try:
  2472. with open('data/model/dPCA_{}_unconstrained_results.pkl'.format(file_name), 'rb') as file:
  2473. dPCA_results = pickle.load(file)
  2474. file.close()
  2475. except:
  2476. dPCA_results = {}
  2477. # get maximum number of trials for any parameter combination across sessions
  2478. max_num_trials = 0
  2479. n_neurons_pseudo = 0
  2480. flagged_models = []
  2481. for i_model in good_model_inds:
  2482. X = outputs['baseline']['Xs'][i_model]
  2483. y_stim = outputs['baseline']['y_stims'][i_model]
  2484. y_choice = outputs['baseline']['y_lefts'][i_model]
  2485. n_neurons_pseudo += (X.shape[1] - int(round(X.shape[1] / 5.88)))
  2486. for stim in unique_stims:
  2487. for choice in [0, 1]:
  2488. n = ((y_stim == stim) & (y_choice == choice)).sum()
  2489. if (n == 0) and (i_model not in flagged_models):
  2490. flagged_models.append(i_model)
  2491. max_num_trials = max(max_num_trials, n)
  2492. # assemble pseudo-population trial-by-trial tensor
  2493. # X_pseudo: (trials x neurons x stimuli x decisions x time)
  2494. X_pseudo = np.full([max_num_trials, n_neurons_pseudo, len(unique_stims), 2, len(t)], np.nan)
  2495. n_neurons_cumul = 0
  2496. for i_model in good_model_inds:
  2497. X = outputs['baseline']['Xs'][i_model]
  2498. X = X[:, int(round(X.shape[1] / 5.88)):, :]
  2499. y_stim = outputs['baseline']['y_stims'][i_model]
  2500. y_choice = outputs['baseline']['y_lefts'][i_model]
  2501. n_neurons = X.shape[1]
  2502. ###----------------------------------------------------------------------------------------------------
  2503. if i_model in flagged_models:
  2504. ### issue: some data are missing for this session
  2505. ### workaround: fit a simple linear model to existing data and use it to predict the missing parts
  2506. predictors = np.vstack([
  2507. y_stim.reshape([1, -1]),
  2508. y_choice.reshape([1, -1]),
  2509. (y_stim * y_choice).reshape([1, -1]),
  2510. np.ones([1, X.shape[0]])
  2511. ])
  2512. betas = [[scipy.linalg.pinv(predictors @ predictors.T) @ predictors @ X[:, i_n, i_t].reshape([-1, 1])
  2513. for i_t in range(len(t))] for i_n in range(n_neurons)]
  2514. ###-----------------------------------------------------------------------------------------------------
  2515. for i_stim, stim in enumerate(unique_stims):
  2516. for i_choice, choice in enumerate([0, 1]):
  2517. trial_mask = (y_stim == stim) & (y_choice == choice)
  2518. n_trials = trial_mask.sum()
  2519. if n_trials == 0:
  2520. ### use the linear model's prediction ------------------------------------------------------
  2521. for i_n in range(n_neurons):
  2522. for i_t in range(len(t)):
  2523. r = max((betas[i_n][i_t].flatten() * np.array([stim, choice, stim * choice, 1])).sum(), 0)
  2524. X_pseudo[:2, n_neurons_cumul + i_n, i_stim, i_choice, i_t] = np.array([r, r])
  2525. ### ----------------------------------------------------------------------------------------
  2526. elif n_trials == 1:
  2527. X_pseudo[:2, n_neurons_cumul:(n_neurons_cumul + n_neurons), i_stim, i_choice, :] = \
  2528. np.concatenate([
  2529. X[trial_mask, :, :].reshape([1, n_neurons, len(t)]),
  2530. X[trial_mask, :, :].reshape([1, n_neurons, len(t)])],
  2531. axis=0)
  2532. else:
  2533. X_pseudo[:n_trials, n_neurons_cumul:(n_neurons_cumul + n_neurons), i_stim, i_choice, :] = \
  2534. X[trial_mask, :, :]
  2535. n_neurons_cumul += n_neurons
  2536. # get the pseudo-population trial-averaged (PSTH) tensor
  2537. X_pseudo_psth = np.nanmean(X_pseudo, axis=0)
  2538. # do the dPCA
  2539. dpca = dPCA.dPCA(labels='sdt', regularizer='auto',
  2540. join={'st': ['s', 'st'], 'dt': ['d', 'dt'], 'sdt': ['sd', 'sdt']})
  2541. dpca.protect = ['t']
  2542. Z = dpca.fit_transform(X_pseudo_psth, X_pseudo)
  2543. dPCA_results['baseline'] = {'X_pseudo': X_pseudo, 'X_pseudo_psth': X_pseudo_psth, 'dpca': dpca, 'Z': Z}
  2544. with open('data/model/dPCA_{}_unconstrained_results.pkl'.format(file_name), 'wb') as file:
  2545. pickle.dump(dPCA_results, file)
  2546. file.close()
  2547. # %%
  2548. fig, axes = plt.subplots(figsize=(11, 3));
  2549. for i_stim, stim in enumerate(unique_stims):
  2550. for i_choice, choice in enumerate([0, 1]):
  2551. style = '-' if (((stim < 50) and (choice == 0)) or ((stim > 50) and (choice == 1))) else ':'
  2552. v = dPCA_results['baseline']['dpca'].D['st'][:, 0]
  2553. X_tensor = dPCA_results['baseline']['X_pseudo_psth']
  2554. X_tilde = X_tensor.reshape([X_tensor.shape[0], -1])
  2555. X_tilde -= X_tilde.mean(axis=1).reshape([-1, 1])
  2556. X_tensor = X_tilde.reshape(X_tensor.shape)
  2557. X = X_tensor[:, i_stim, i_choice, :]
  2558. y = [(v * X[:, k]).sum() for k in range(len(t))]
  2559. axes.plot(t, y, color=colors[i_stim, :], linestyle=style);
  2560. axes.axvline(0, color='k', linestyle=':');
  2561. axes.axvline(t_D, color='k', linestyle=':');
  2562. axes.set_title('Stimulus coding');
  2563. axes.set_xlim([t[0], t[-1]]);
  2564. # Figure 6 - Supplement 2 A
  2565. # plt.savefig('plots/review/model/dPCA_stimulus_coding_unconstrained.pdf');
  2566. # %%
  2567. fig, axes = plt.subplots(figsize=(11, 3));
  2568. for i_stim, stim in enumerate(unique_stims):
  2569. for i_choice, choice in enumerate([0, 1]):
  2570. style = '-' if (((stim < 50) and (choice == 0)) or ((stim > 50) and (choice == 1))) else ':'
  2571. v = dPCA_results['baseline']['dpca'].D['dt'][:, 0]
  2572. X_tensor = dPCA_results['baseline']['X_pseudo_psth']
  2573. X_tilde = X_tensor.reshape([X_tensor.shape[0], -1])
  2574. X_tilde -= X_tilde.mean(axis=1).reshape([-1, 1])
  2575. X_tensor = X_tilde.reshape(X_tensor.shape)
  2576. X = X_tensor[:, i_stim, i_choice, :]
  2577. y = [(v * X[:, k]).sum() for k in range(len(t))]
  2578. axes.plot(t, y, color=colors[i_stim, :], linestyle=style);
  2579. axes.axvline(0, color='k', linestyle=':');
  2580. axes.axvline(t_D, color='k', linestyle=':');
  2581. axes.set_title('Choice coding');
  2582. axes.set_xlim([t[0], t[-1]]);
  2583. y_lim = np.max(np.abs(axes.get_ylim())); axes.set_ylim([-y_lim, y_lim]);
  2584. # Figure 6 - Supplement 2 B
  2585. # plt.savefig('plots/review/model/dPCA_choice_coding_unconstrained.pdf');
  2586. # %% [markdown]
  2587. # ### --- dPCA ---
  2588. # %% [markdown]
  2589. # #### dPCA for constrained vs unconstrained with ablations
  2590. # %%
  2591. filter_models = True # filter out models that do not have at least 2 trials per choice across all ablation conditions
  2592. # this is done to improve the linear prediction of missing trial types
  2593. outputs = {
  2594. 'baseline': output_noise,
  2595. 'other': output_noise_ablate_other,
  2596. 'other_con': output_noise_ablate_other_con,
  2597. 'other_unc': output_noise_ablate_other_unc,
  2598. 'linear': output_noise_ablate_linear,
  2599. 'linear_con': output_noise_ablate_linear_con,
  2600. 'linear_unc': output_noise_ablate_linear_unc,
  2601. 'perception': output_noise_ablate_perception,
  2602. 'perception_con': output_noise_ablate_perception_con,
  2603. 'perception_unc': output_noise_ablate_perception_unc,
  2604. 'choice': output_noise_ablate_choice,
  2605. 'choice_con': output_noise_ablate_choice_con,
  2606. 'choice_unc': output_noise_ablate_choice_unc
  2607. }
  2608. conditions = outputs.keys()
  2609. # filtering criteria
  2610. good_model_inds = []
  2611. for i_model in range(len(included_sessions)):
  2612. is_good = True
  2613. if filter_models:
  2614. for condition in conditions:
  2615. y_choice = outputs[condition]['y_lefts'][i_model]
  2616. if (y_choice.sum() <= 1) or (y_choice.sum() >= len(y_choice) - 1):
  2617. is_good = False
  2618. break
  2619. if is_good:
  2620. good_model_inds.append(i_model)
  2621. if filter_models:
  2622. print('{}/{} models passed.'.format(len(good_model_inds), len(included_sessions)))
  2623. else:
  2624. print('All models included.')
  2625. file_name = 'filtered' if filter_models else 'all'
  2626. # %%
  2627. try:
  2628. with open('data/model/dPCA_{}_sessions_constrainedAndUnconstrained_results.pkl'.format(file_name), 'rb') as file:
  2629. dPCA_results = pickle.load(file)
  2630. file.close()
  2631. except:
  2632. dPCA_results = {}
  2633. for condition in conditions:
  2634. # get maximum number of trials for any parameter combination across sessions
  2635. max_num_trials = 0
  2636. n_neurons_pseudo = 0
  2637. flagged_models = []
  2638. for i_model in good_model_inds:
  2639. X = outputs[condition]['Xs'][i_model]
  2640. y_stim = outputs[condition]['y_stims'][i_model]
  2641. y_choice = outputs[condition]['y_lefts'][i_model]
  2642. n_neurons_pseudo += X.shape[1]
  2643. for stim in unique_stims:
  2644. for choice in [0, 1]:
  2645. n = ((y_stim == stim) & (y_choice == choice)).sum()
  2646. if (n == 0) and (i_model not in flagged_models):
  2647. flagged_models.append(i_model)
  2648. max_num_trials = max(max_num_trials, n)
  2649. # assemble pseudo-population trial-by-trial tensor
  2650. # X_pseudo: (trials x neurons x stimuli x decisions x time)
  2651. X_pseudo = np.full([max_num_trials, n_neurons_pseudo, len(unique_stims), 2, len(t)], np.nan)
  2652. n_neurons_cumul = 0
  2653. for i_model in good_model_inds:
  2654. X = outputs[condition]['Xs'][i_model]
  2655. y_stim = outputs[condition]['y_stims'][i_model]
  2656. y_choice = outputs[condition]['y_lefts'][i_model]
  2657. n_neurons = X.shape[1]
  2658. ###----------------------------------------------------------------------------------------------------
  2659. if i_model in flagged_models:
  2660. ### issue: some data are missing for this session
  2661. ### workaround: fit a simple linear model to existing data and use it to predict the missing parts
  2662. predictors = np.vstack([
  2663. y_stim.reshape([1, -1]),
  2664. y_choice.reshape([1, -1]),
  2665. (y_stim * y_choice).reshape([1, -1]),
  2666. np.ones([1, X.shape[0]])
  2667. ])
  2668. betas = [[scipy.linalg.pinv(predictors @ predictors.T) @ predictors @ X[:, i_n, i_t].reshape([-1, 1])
  2669. for i_t in range(len(t))] for i_n in range(n_neurons)]
  2670. ###-----------------------------------------------------------------------------------------------------
  2671. for i_stim, stim in enumerate(unique_stims):
  2672. for i_choice, choice in enumerate([0, 1]):
  2673. trial_mask = (y_stim == stim) & (y_choice == choice)
  2674. n_trials = trial_mask.sum()
  2675. if n_trials == 0:
  2676. ### use the linear model's prediction ------------------------------------------------------
  2677. for i_n in range(n_neurons):
  2678. for i_t in range(len(t)):
  2679. r = max((betas[i_n][i_t].flatten() * np.array([stim, choice, stim * choice, 1])).sum(), 0)
  2680. X_pseudo[:2, n_neurons_cumul + i_n, i_stim, i_choice, i_t] = np.array([r, r])
  2681. ### ----------------------------------------------------------------------------------------
  2682. elif n_trials == 1:
  2683. X_pseudo[:2, n_neurons_cumul:(n_neurons_cumul + n_neurons), i_stim, i_choice, :] = \
  2684. np.concatenate([
  2685. X[trial_mask, :, :].reshape([1, n_neurons, len(t)]),
  2686. X[trial_mask, :, :].reshape([1, n_neurons, len(t)])],
  2687. axis=0)
  2688. else:
  2689. X_pseudo[:n_trials, n_neurons_cumul:(n_neurons_cumul + n_neurons), i_stim, i_choice, :] = \
  2690. X[trial_mask, :, :]
  2691. n_neurons_cumul += n_neurons
  2692. # get the pseudo-population trial-averaged (PSTH) tensor
  2693. X_pseudo_psth = np.nanmean(X_pseudo, axis=0)
  2694. # do the dPCA
  2695. dpca = dPCA.dPCA(labels='sdt', regularizer='auto',
  2696. join={'st': ['s', 'st'], 'dt': ['d', 'dt'], 'sdt': ['sd', 'sdt']})
  2697. dpca.protect = ['t']
  2698. Z = dpca.fit_transform(X_pseudo_psth, X_pseudo)
  2699. dPCA_results[condition] = {'X_pseudo': X_pseudo, 'X_pseudo_psth': X_pseudo_psth, 'dpca': dpca, 'Z': Z}
  2700. with open('data/model/dPCA_{}_sessions_constrainedAndUnconstrained_results.pkl'.format(file_name), 'wb') as file:
  2701. pickle.dump(dPCA_results, file)
  2702. file.close()
  2703. # %%
  2704. fig, axes = plt.subplots(len(conditions), 1, figsize=(11, 3 * len(conditions)));
  2705. y_limits = []
  2706. mean_stim_projections_con_vs_unc = {}
  2707. for i_condition, condition in enumerate(conditions):
  2708. temp = []
  2709. for i_stim, stim in enumerate(unique_stims):
  2710. for i_choice, choice in enumerate([0, 1]):
  2711. style = '-' if (((stim < 50) and (choice == 0)) or ((stim > 50) and (choice == 1))) else ':'
  2712. v = dPCA_results['baseline']['dpca'].D['st'][:, 0]
  2713. X_tensor = dPCA_results[condition]['X_pseudo_psth']
  2714. X_tilde = X_tensor.reshape([X_tensor.shape[0], -1])
  2715. X_tilde -= X_tilde.mean(axis=1).reshape([-1, 1])
  2716. X_tensor = X_tilde.reshape(X_tensor.shape)
  2717. X = X_tensor[:, i_stim, i_choice, :]
  2718. y = [(v * X[:, k]).sum() for k in range(len(t))]
  2719. temp.append(np.mean(np.abs(y)))
  2720. axes[i_condition].plot(t, y, color=colors[i_stim, :], linestyle=style);
  2721. axes[i_condition].axvline(0, color='k', linestyle=':');
  2722. axes[i_condition].axvline(t_D, color='k', linestyle=':');
  2723. axes[i_condition].set_title('Stimulus coding ({})'.format(condition));
  2724. y_limits.append(axes[i_condition].get_ylim())
  2725. mean_stim_projections_con_vs_unc[condition] = temp
  2726. global_limit = np.max([np.max(np.abs(lims)) for lims in y_limits])
  2727. for ax in axes:
  2728. ax.set_xlim([t[0], t[-1]]);
  2729. ax.set_ylim([-global_limit, global_limit]);
  2730. # Figure 7 - Supplement 1 A
  2731. # plt.savefig('plots/review/model/dPCA_con_vs_unc_stimulus_coding.pdf');
  2732. # %%
  2733. fig, axes = plt.subplots(len(conditions), 1, figsize=(11, 3 * len(conditions)));
  2734. y_limits = []
  2735. mean_choice_projections_con_vs_unc = {}
  2736. for i_condition, condition in enumerate(conditions):
  2737. temp = []
  2738. for i_stim, stim in enumerate(unique_stims):
  2739. for i_choice, choice in enumerate([0, 1]):
  2740. style = '-' if (((stim < 50) and (choice == 0)) or ((stim > 50) and (choice == 1))) else ':'
  2741. v = dPCA_results['baseline']['dpca'].D['dt'][:, 0]
  2742. X_tensor = dPCA_results[condition]['X_pseudo_psth']
  2743. X_tilde = X_tensor.reshape([X_tensor.shape[0], -1])
  2744. X_tilde -= X_tilde.mean(axis=1).reshape([-1, 1])
  2745. X_tensor = X_tilde.reshape(X_tensor.shape)
  2746. X = X_tensor[:, i_stim, i_choice, :]
  2747. y = [(v * X[:, k]).sum() for k in range(len(t))]
  2748. temp.append(np.mean(np.abs(y)))
  2749. axes[i_condition].plot(t, y, color=colors[i_stim, :], linestyle=style);
  2750. axes[i_condition].axvline(0, color='k', linestyle=':');
  2751. axes[i_condition].axvline(t_D, color='k', linestyle=':');
  2752. axes[i_condition].set_title('Choice coding ({})'.format(condition));
  2753. y_limits.append(axes[i_condition].get_ylim())
  2754. mean_choice_projections_con_vs_unc[condition] = temp
  2755. global_limit = np.max([np.max(np.abs(lims)) for lims in y_limits])
  2756. for ax in axes:
  2757. ax.set_xlim([t[0], t[-1]]);
  2758. ax.set_ylim([-global_limit, global_limit]);
  2759. # Figure 7 - Supplement 1 B
  2760. # plt.savefig('plots/review/model/dPCA_con_vs_unc_choice_coding.pdf');
  2761. # %%
  2762. '''
  2763. 2-way within-subjects ANOVA on stimulus component projections
  2764. factor 1: coding type (4 levels)
  2765. factor 2: constrained vs unconstrained (2 levels)
  2766. There is only 1 control group b/c factor 2 does not apply when ablating nothing
  2767. Remove it as a level of factor 1 to avoid an empty cell in the 2-way ANOVA
  2768. To compare groups to control, separately use Dunnett's test on all 9 groups
  2769. This is pseudo-population data, so the subjects are no longer individual models
  2770. Instead, they are the 16 trial types
  2771. '''
  2772. n_subjects = 16
  2773. coding_type = ['linear', 'perception', 'choice', 'other']
  2774. constraint = ['con', 'unc']
  2775. col_subject = list(range(16)) * (len(coding_type) * len(constraint))
  2776. col_coding_type, col_constraint, col_measurement = [], [], []
  2777. for coding_type_i in coding_type:
  2778. for constraint_i in constraint:
  2779. col_coding_type += [coding_type_i] * n_subjects
  2780. col_constraint += [constraint_i] * n_subjects
  2781. col_measurement += list(mean_stim_projections_con_vs_unc[coding_type_i + '_' + constraint_i])
  2782. df = pd.DataFrame({
  2783. 'id': col_subject,
  2784. 'iv1': col_coding_type,
  2785. 'iv2': col_constraint,
  2786. 'dv': col_measurement
  2787. })
  2788. # Perform repeated measures ANOVA
  2789. res_anova = rmAnova2Way(df)
  2790. # post-hoc tests
  2791. all_groups = [
  2792. ('linear', 'con'),
  2793. ('linear', 'unc'),
  2794. ('perception', 'con'),
  2795. ('perception', 'unc'),
  2796. ('choice', 'con'),
  2797. ('choice', 'unc'),
  2798. ('other', 'con'),
  2799. ('other', 'unc')
  2800. ]
  2801. n_groups = len(all_groups)
  2802. correction = scipy.special.comb(n_groups, 2) # bonferroni correction for number of paired tests
  2803. p_mat = np.full([n_groups, n_groups], np.nan)
  2804. print('\nPost-hoc tests (Bonferroni-adjusted p-values):\n')
  2805. for i in range(n_groups - 1):
  2806. for j in range(i + 1, n_groups):
  2807. coding_type_i, constraint_i = all_groups[i]
  2808. x_i = np.array([_ for _ in df[(df['iv1'] == coding_type_i) & (df['iv2'] == constraint_i)]['dv']])
  2809. label_i = coding_type_i + '_' + constraint_i
  2810. coding_type_j, constraint_j = all_groups[j]
  2811. x_j = np.array([_ for _ in df[(df['iv1'] == coding_type_j) & (df['iv2'] == constraint_j)]['dv']])
  2812. label_j = coding_type_j + '_' + constraint_j
  2813. p_adj = correction * scipy.stats.ttest_rel(x_i, x_j)[1]
  2814. print(' {} vs. {}: p = {:.4e}'.format(label_i, label_j, p_adj))
  2815. p_mat[i, j] = (p_adj < 0.01).astype(float)
  2816. print('')
  2817. print(p_mat)
  2818. # %%
  2819. ## dunnett test
  2820. data = [
  2821. (mean_stim_projections_con_vs_unc[key], key)
  2822. for key in [
  2823. 'linear_con', 'linear_unc', 'perception_con', 'perception_unc', 'choice_con', 'choice_unc', 'other_con', 'other_unc'
  2824. ]
  2825. ]
  2826. samples = [np.array(data_[0]) for data_ in data]
  2827. labels = [data_[1] for data_ in data]
  2828. control = np.array(mean_stim_projections_con_vs_unc['baseline'])
  2829. res = scipy.stats.dunnett(*samples, control=control)
  2830. print('Dunnett test:\n')
  2831. for i in range(len(data)):
  2832. print('Control vs {}: p = {:.4e}'.format(labels[i], res.pvalue[i]))
  2833. # %%
  2834. '''
  2835. 2-way within-subjects ANOVA on choice component projections
  2836. factor 1: coding type (4 levels)
  2837. factor 2: constrained vs unconstrained (2 levels)
  2838. There is only 1 control group b/c factor 2 does not apply when ablating nothing
  2839. Remove it as a level of factor 1 to avoid an empty cell in the 2-way ANOVA
  2840. To compare groups to control, separately use Dunnett's test on all 9 groups
  2841. This is pseudo-population data, so the subjects are no longer individual models
  2842. Instead, they are the 16 trial types
  2843. '''
  2844. n_subjects = 16
  2845. coding_type = ['linear', 'perception', 'choice', 'other']
  2846. constraint = ['con', 'unc']
  2847. col_subject = list(range(16)) * (len(coding_type) * len(constraint))
  2848. col_coding_type, col_constraint, col_measurement = [], [], []
  2849. for coding_type_i in coding_type:
  2850. for constraint_i in constraint:
  2851. col_coding_type += [coding_type_i] * n_subjects
  2852. col_constraint += [constraint_i] * n_subjects
  2853. col_measurement += list(mean_choice_projections_con_vs_unc[coding_type_i + '_' + constraint_i])
  2854. df = pd.DataFrame({
  2855. 'id': col_subject,
  2856. 'iv1': col_coding_type,
  2857. 'iv2': col_constraint,
  2858. 'dv': col_measurement
  2859. })
  2860. # Perform repeated measures ANOVA
  2861. res_anova = rmAnova2Way(df)
  2862. # post-hoc tests
  2863. all_groups = [
  2864. ('linear', 'con'),
  2865. ('linear', 'unc'),
  2866. ('perception', 'con'),
  2867. ('perception', 'unc'),
  2868. ('choice', 'con'),
  2869. ('choice', 'unc'),
  2870. ('other', 'con'),
  2871. ('other', 'unc')
  2872. ]
  2873. n_groups = len(all_groups)
  2874. correction = scipy.special.comb(n_groups, 2) # bonferroni correction for number of paired tests
  2875. p_mat = np.full([n_groups, n_groups], np.nan)
  2876. print('\nPost-hoc tests (Bonferroni-adjusted p-values):\n')
  2877. for i in range(n_groups - 1):
  2878. for j in range(i + 1, n_groups):
  2879. coding_type_i, constraint_i = all_groups[i]
  2880. x_i = np.array([_ for _ in df[(df['iv1'] == coding_type_i) & (df['iv2'] == constraint_i)]['dv']])
  2881. label_i = coding_type_i + '_' + constraint_i
  2882. coding_type_j, constraint_j = all_groups[j]
  2883. x_j = np.array([_ for _ in df[(df['iv1'] == coding_type_j) & (df['iv2'] == constraint_j)]['dv']])
  2884. label_j = coding_type_j + '_' + constraint_j
  2885. p_adj = correction * scipy.stats.ttest_rel(x_i, x_j)[1]
  2886. print(' {} vs. {}: p = {:.4e}'.format(label_i, label_j, p_adj))
  2887. p_mat[i, j] = (p_adj < 0.01).astype(float)
  2888. print('')
  2889. print(p_mat)
  2890. # %%
  2891. ## dunnett test
  2892. data = [
  2893. (mean_choice_projections_con_vs_unc[key], key)
  2894. for key in [
  2895. 'linear_con', 'linear_unc', 'perception_con', 'perception_unc', 'choice_con', 'choice_unc', 'other_con', 'other_unc'
  2896. ]
  2897. ]
  2898. samples = [np.array(data_[0]) for data_ in data]
  2899. labels = [data_[1] for data_ in data]
  2900. control = np.array(mean_choice_projections_con_vs_unc['baseline'])
  2901. res = scipy.stats.dunnett(*samples, control=control)
  2902. print('Dunnett test:\n')
  2903. for i in range(len(data)):
  2904. print('Control vs {}: p = {:.4e}'.format(labels[i], res.pvalue[i]))
  2905. # %% [markdown]
  2906. # ### --- dPCA ---
  2907. # %% [markdown]
  2908. # #### Temporally-restricted ablations
  2909. # %%
  2910. filter_models = True # filter out models that do not have at least 2 trials per choice across all ablation conditions
  2911. # this is done to improve the linear prediction of missing trial types
  2912. outputs = {
  2913. 'baseline': output_noise,
  2914. 'linear': output_noise_ablate_linear,
  2915. 'perception': output_noise_ablate_perception,
  2916. 'choice': output_noise_ablate_choice,
  2917. 'other': output_noise_ablate_other,
  2918. 'linear_beginning': output_noise_ablate_linear_beginning,
  2919. 'linear_end': output_noise_ablate_linear_end,
  2920. 'perception_beginning': output_noise_ablate_perception_beginning,
  2921. 'perception_end': output_noise_ablate_perception_end,
  2922. 'choice_beginning': output_noise_ablate_choice_beginning,
  2923. 'choice_end': output_noise_ablate_choice_end
  2924. }
  2925. conditions = outputs.keys()
  2926. # filtering criteria
  2927. good_model_inds = []
  2928. for i_model in range(len(included_sessions)):
  2929. is_good = True
  2930. if filter_models:
  2931. for condition in conditions:
  2932. y_choice = outputs[condition]['y_lefts'][i_model]
  2933. if (y_choice.sum() <= 1) or (y_choice.sum() >= len(y_choice) - 1):
  2934. is_good = False
  2935. break
  2936. if is_good:
  2937. good_model_inds.append(i_model)
  2938. if filter_models:
  2939. print('{}/{} models passed.'.format(len(good_model_inds), len(included_sessions)))
  2940. else:
  2941. print('All models included.')
  2942. file_name = 'filtered' if filter_models else 'all'
  2943. # %%
  2944. try:
  2945. with open('data/model/dPCA_{}_sessions_beginningAndEnd_results.pkl'.format(file_name), 'rb') as file:
  2946. dPCA_results = pickle.load(file)
  2947. file.close()
  2948. except:
  2949. dPCA_results = {}
  2950. for condition in conditions:
  2951. # get maximum number of trials for any parameter combination across sessions
  2952. max_num_trials = 0
  2953. n_neurons_pseudo = 0
  2954. flagged_models = []
  2955. for i_model in good_model_inds:
  2956. X = outputs[condition]['Xs'][i_model]
  2957. y_stim = outputs[condition]['y_stims'][i_model]
  2958. y_choice = outputs[condition]['y_lefts'][i_model]
  2959. n_neurons_pseudo += X.shape[1]
  2960. for stim in unique_stims:
  2961. for choice in [0, 1]:
  2962. n = ((y_stim == stim) & (y_choice == choice)).sum()
  2963. if (n == 0) and (i_model not in flagged_models):
  2964. flagged_models.append(i_model)
  2965. max_num_trials = max(max_num_trials, n)
  2966. # assemble pseudo-population trial-by-trial tensor
  2967. # X_pseudo: (trials x neurons x stimuli x decisions x time)
  2968. X_pseudo = np.full([max_num_trials, n_neurons_pseudo, len(unique_stims), 2, len(t)], np.nan)
  2969. n_neurons_cumul = 0
  2970. for i_model in good_model_inds:
  2971. X = outputs[condition]['Xs'][i_model]
  2972. y_stim = outputs[condition]['y_stims'][i_model]
  2973. y_choice = outputs[condition]['y_lefts'][i_model]
  2974. n_neurons = X.shape[1]
  2975. ###----------------------------------------------------------------------------------------------------
  2976. if i_model in flagged_models:
  2977. ### issue: some data are missing for this session
  2978. ### workaround: fit a simple linear model to existing data and use it to predict the missing parts
  2979. predictors = np.vstack([
  2980. y_stim.reshape([1, -1]),
  2981. y_choice.reshape([1, -1]),
  2982. (y_stim * y_choice).reshape([1, -1]),
  2983. np.ones([1, X.shape[0]])
  2984. ])
  2985. betas = [[scipy.linalg.pinv(predictors @ predictors.T) @ predictors @ X[:, i_n, i_t].reshape([-1, 1])
  2986. for i_t in range(len(t))] for i_n in range(n_neurons)]
  2987. ###-----------------------------------------------------------------------------------------------------
  2988. for i_stim, stim in enumerate(unique_stims):
  2989. for i_choice, choice in enumerate([0, 1]):
  2990. trial_mask = (y_stim == stim) & (y_choice == choice)
  2991. n_trials = trial_mask.sum()
  2992. if n_trials == 0:
  2993. ### use the linear model's prediction ------------------------------------------------------
  2994. for i_n in range(n_neurons):
  2995. for i_t in range(len(t)):
  2996. r = max((betas[i_n][i_t].flatten() * np.array([stim, choice, stim * choice, 1])).sum(), 0)
  2997. X_pseudo[:2, n_neurons_cumul + i_n, i_stim, i_choice, i_t] = np.array([r, r])
  2998. ### ----------------------------------------------------------------------------------------
  2999. elif n_trials == 1:
  3000. X_pseudo[:2, n_neurons_cumul:(n_neurons_cumul + n_neurons), i_stim, i_choice, :] = \
  3001. np.concatenate([
  3002. X[trial_mask, :, :].reshape([1, n_neurons, len(t)]),
  3003. X[trial_mask, :, :].reshape([1, n_neurons, len(t)])],
  3004. axis=0)
  3005. else:
  3006. X_pseudo[:n_trials, n_neurons_cumul:(n_neurons_cumul + n_neurons), i_stim, i_choice, :] = \
  3007. X[trial_mask, :, :]
  3008. n_neurons_cumul += n_neurons
  3009. # get the pseudo-population trial-averaged (PSTH) tensor
  3010. X_pseudo_psth = np.nanmean(X_pseudo, axis=0)
  3011. # do the dPCA
  3012. dpca = dPCA.dPCA(labels='sdt', regularizer='auto',
  3013. join={'st': ['s', 'st'], 'dt': ['d', 'dt'], 'sdt': ['sd', 'sdt']})
  3014. dpca.protect = ['t']
  3015. Z = dpca.fit_transform(X_pseudo_psth, X_pseudo)
  3016. dPCA_results[condition] = {'X_pseudo': X_pseudo, 'X_pseudo_psth': X_pseudo_psth, 'dpca': dpca, 'Z': Z}
  3017. with open('data/model/dPCA_{}_sessions_beginningAndEnd_results.pkl'.format(file_name), 'wb') as file:
  3018. pickle.dump(dPCA_results, file)
  3019. file.close()
  3020. # %%
  3021. fig, axes = plt.subplots(len(conditions), 1, figsize=(11, 3 * len(conditions)));
  3022. y_limits = []
  3023. mean_stim_projections_beg_vs_end = {}
  3024. for i_condition, condition in enumerate(conditions):
  3025. temp = []
  3026. for i_stim, stim in enumerate(unique_stims):
  3027. for i_choice, choice in enumerate([0, 1]):
  3028. style = '-' if (((stim < 50) and (choice == 0)) or ((stim > 50) and (choice == 1))) else ':'
  3029. v = dPCA_results['baseline']['dpca'].D['st'][:, 0]
  3030. X_tensor = dPCA_results[condition]['X_pseudo_psth']
  3031. X_tilde = X_tensor.reshape([X_tensor.shape[0], -1])
  3032. X_tilde -= X_tilde.mean(axis=1).reshape([-1, 1])
  3033. X_tensor = X_tilde.reshape(X_tensor.shape)
  3034. X = X_tensor[:, i_stim, i_choice, :]
  3035. y = [(v * X[:, k]).sum() for k in range(len(t))]
  3036. temp.append(np.mean(np.abs(y)))
  3037. axes[i_condition].plot(t, y, color=colors[i_stim, :], linestyle=style);
  3038. axes[i_condition].axvline(0, color='k', linestyle=':');
  3039. axes[i_condition].axvline(t_D, color='k', linestyle=':');
  3040. axes[i_condition].set_title('Stimulus coding ({})'.format(condition));
  3041. y_limits.append(axes[i_condition].get_ylim())
  3042. mean_stim_projections_beg_vs_end[condition] = temp
  3043. global_limit = np.max([np.max(np.abs(lims)) for lims in y_limits])
  3044. for ax in axes:
  3045. ax.set_xlim([t[0], t[-1]]);
  3046. ax.set_ylim([-global_limit, global_limit]);
  3047. # Figure 7 - Supplement 2 A
  3048. # plt.savefig('plots/review/model/dPCA_beg_vs_end_stimulus_coding.pdf');
  3049. # %%
  3050. fig, axes = plt.subplots(len(conditions), 1, figsize=(11, 3 * len(conditions)));
  3051. y_limits = []
  3052. mean_choice_projections_beg_vs_end = {}
  3053. for i_condition, condition in enumerate(conditions):
  3054. temp = []
  3055. for i_stim, stim in enumerate(unique_stims):
  3056. for i_choice, choice in enumerate([0, 1]):
  3057. style = '-' if (((stim < 50) and (choice == 0)) or ((stim > 50) and (choice == 1))) else ':'
  3058. v = dPCA_results['baseline']['dpca'].D['dt'][:, 0]
  3059. X_tensor = dPCA_results[condition]['X_pseudo_psth']
  3060. X_tilde = X_tensor.reshape([X_tensor.shape[0], -1])
  3061. X_tilde -= X_tilde.mean(axis=1).reshape([-1, 1])
  3062. X_tensor = X_tilde.reshape(X_tensor.shape)
  3063. X = X_tensor[:, i_stim, i_choice, :]
  3064. y = [(v * X[:, k]).sum() for k in range(len(t))]
  3065. temp.append(np.mean(np.abs(y)))
  3066. axes[i_condition].plot(t, y, color=colors[i_stim, :], linestyle=style);
  3067. axes[i_condition].axvline(0, color='k', linestyle=':');
  3068. axes[i_condition].axvline(t_D, color='k', linestyle=':');
  3069. axes[i_condition].set_title('Choice coding ({})'.format(condition));
  3070. y_limits.append(axes[i_condition].get_ylim())
  3071. mean_choice_projections_beg_vs_end[condition] = temp
  3072. global_limit = np.max([np.max(np.abs(lims)) for lims in y_limits])
  3073. for ax in axes:
  3074. ax.set_xlim([t[0], t[-1]]);
  3075. ax.set_ylim([-global_limit, global_limit]);
  3076. # Figure 7 - Supplement 2 B
  3077. # plt.savefig('plots/review/model/dPCA_beg_vs_end_choice_coding.pdf');
  3078. # %%
  3079. '''
  3080. 2-way within-subjects ANOVA on stimulus component projections
  3081. factor 1: coding type (3 levels)
  3082. factor 2: beginning vs end (2 levels)
  3083. There is only 1 control group b/c factor 2 does not apply when ablating nothing
  3084. Remove it as a level of factor 1 to avoid an empty cell in the 2-way ANOVA
  3085. To compare groups to control, separately use Dunnett's test on all 6 groups
  3086. This is pseudo-population data, so the subjects are no longer individual models
  3087. Instead, they are the 16 trial types
  3088. '''
  3089. n_subjects = 16
  3090. coding_type = ['linear', 'perception', 'choice']
  3091. window = ['beginning', 'end']
  3092. col_subject = list(range(n_subjects)) * (len(coding_type) * len(window))
  3093. col_coding_type, col_window, col_measurement = [], [], []
  3094. for coding_type_i in coding_type:
  3095. for window_i in window:
  3096. col_coding_type += [coding_type_i] * n_subjects
  3097. col_window += [window_i] * n_subjects
  3098. col_measurement += list(mean_stim_projections_beg_vs_end[coding_type_i + '_' + window_i])
  3099. df = pd.DataFrame({
  3100. 'id': col_subject,
  3101. 'iv1': col_coding_type,
  3102. 'iv2': col_window,
  3103. 'dv': col_measurement
  3104. })
  3105. # Perform repeated measures ANOVA
  3106. res_anova = rmAnova2Way(df)
  3107. # post-hoc tests
  3108. all_groups = [
  3109. ('linear', 'beginning'),
  3110. ('linear', 'end'),
  3111. ('perception', 'beginning'),
  3112. ('perception', 'end'),
  3113. ('choice', 'beginning'),
  3114. ('choice', 'end')
  3115. ]
  3116. n_groups = len(all_groups)
  3117. correction = scipy.special.comb(n_groups, 2) # bonferroni correction for number of paired tests
  3118. p_mat = np.full([n_groups, n_groups], np.nan)
  3119. print('\nPost-hoc tests (Bonferroni-adjusted p-values):\n')
  3120. for i in range(n_groups - 1):
  3121. for j in range(i + 1, n_groups):
  3122. coding_type_i, window_i = all_groups[i]
  3123. x_i = np.array([_ for _ in df[(df['iv1'] == coding_type_i) & (df['iv2'] == window_i)]['dv']])
  3124. label_i = coding_type_i + '_' + window_i
  3125. coding_type_j, window_j = all_groups[j]
  3126. x_j = np.array([_ for _ in df[(df['iv1'] == coding_type_j) & (df['iv2'] == window_j)]['dv']])
  3127. label_j = coding_type_j + '_' + window_j
  3128. p_adj = correction * scipy.stats.ttest_rel(x_i, x_j)[1]
  3129. print(' {} vs. {}: p = {:.4e}'.format(label_i, label_j, p_adj))
  3130. p_mat[i, j] = (p_adj < 0.01).astype(float)
  3131. print('')
  3132. print(p_mat)
  3133. # %%
  3134. ## dunnett test
  3135. data = [
  3136. (mean_stim_projections_beg_vs_end[key], key)
  3137. for key in [
  3138. 'linear_beginning', 'linear_end', 'perception_beginning', 'perception_end', 'choice_beginning', 'choice_end'
  3139. ]
  3140. ]
  3141. samples = [np.array(data_[0]) for data_ in data]
  3142. labels = [data_[1] for data_ in data]
  3143. control = np.array(mean_stim_projections_beg_vs_end['baseline'])
  3144. res = scipy.stats.dunnett(*samples, control=control)
  3145. print('Dunnett test:\n')
  3146. for i in range(len(data)):
  3147. print('Control vs {}: p = {:.4e}'.format(labels[i], res.pvalue[i]))
  3148. # %%
  3149. '''
  3150. 2-way within-subjects ANOVA on choice component projections
  3151. factor 1: coding type (3 levels)
  3152. factor 2: beginning vs end (2 levels)
  3153. There is only 1 control group b/c factor 2 does not apply when ablating nothing
  3154. Remove it as a level of factor 1 to avoid an empty cell in the 2-way ANOVA
  3155. To compare groups to control, separately use Dunnett's test on all 6 groups
  3156. This is pseudo-population data, so the subjects are no longer individual models
  3157. Instead, they are the 16 trial types
  3158. '''
  3159. n_subjects = 16
  3160. coding_type = ['linear', 'perception', 'choice']
  3161. window = ['beginning', 'end']
  3162. col_subject = list(range(n_subjects)) * (len(coding_type) * len(window))
  3163. col_coding_type, col_window, col_measurement = [], [], []
  3164. for coding_type_i in coding_type:
  3165. for window_i in window:
  3166. col_coding_type += [coding_type_i] * n_subjects
  3167. col_window += [window_i] * n_subjects
  3168. col_measurement += list(mean_choice_projections_beg_vs_end[coding_type_i + '_' + window_i])
  3169. df = pd.DataFrame({
  3170. 'id': col_subject,
  3171. 'iv1': col_coding_type,
  3172. 'iv2': col_window,
  3173. 'dv': col_measurement
  3174. })
  3175. # Perform repeated measures ANOVA
  3176. res_anova = rmAnova2Way(df)
  3177. # post-hoc tests
  3178. all_groups = [
  3179. ('linear', 'beginning'),
  3180. ('linear', 'end'),
  3181. ('perception', 'beginning'),
  3182. ('perception', 'end'),
  3183. ('choice', 'beginning'),
  3184. ('choice', 'end')
  3185. ]
  3186. n_groups = len(all_groups)
  3187. correction = scipy.special.comb(n_groups, 2) # bonferroni correction for number of paired tests
  3188. p_mat = np.full([n_groups, n_groups], np.nan)
  3189. print('\nPost-hoc tests (Bonferroni-adjusted p-values):\n')
  3190. for i in range(n_groups - 1):
  3191. for j in range(i + 1, n_groups):
  3192. coding_type_i, window_i = all_groups[i]
  3193. x_i = np.array([_ for _ in df[(df['iv1'] == coding_type_i) & (df['iv2'] == window_i)]['dv']])
  3194. label_i = coding_type_i + '_' + window_i
  3195. coding_type_j, window_j = all_groups[j]
  3196. x_j = np.array([_ for _ in df[(df['iv1'] == coding_type_j) & (df['iv2'] == window_j)]['dv']])
  3197. label_j = coding_type_j + '_' + window_j
  3198. p_adj = correction * scipy.stats.ttest_rel(x_i, x_j)[1]
  3199. print(' {} vs. {}: p = {:.4e}'.format(label_i, label_j, p_adj))
  3200. p_mat[i, j] = (p_adj < 0.01).astype(float)
  3201. print('')
  3202. print(p_mat)
  3203. # %%
  3204. ## dunnett test
  3205. data = [
  3206. (mean_choice_projections_beg_vs_end[key], key)
  3207. for key in [
  3208. 'linear_beginning', 'linear_end', 'perception_beginning', 'perception_end', 'choice_beginning', 'choice_end'
  3209. ]
  3210. ]
  3211. samples = [np.array(data_[0]) for data_ in data]
  3212. labels = [data_[1] for data_ in data]
  3213. control = np.array(mean_choice_projections_beg_vs_end['baseline'])
  3214. res = scipy.stats.dunnett(*samples, control=control)
  3215. print('Dunnett test:\n')
  3216. for i in range(len(data)):
  3217. print('Control vs {}: p = {:.4e}'.format(labels[i], res.pvalue[i]))
  3218. # %% [markdown]
  3219. # ### Additional examples of RNN units and target experimental PSTHs
  3220. # %%
  3221. smooth_type = 'gaussian'
  3222. smooth_width = 11
  3223. window_width = 4
  3224. window_step = 4
  3225. fit_options = [
  3226. {
  3227. 'shape' : 'linear',
  3228. 'y_min' : 0,
  3229. 'y_max' : np.inf,
  3230. 'min_y_range' : 0
  3231. },
  3232. {
  3233. 'shape' : 'step',
  3234. 'midpoints' : [40, 50, 60],
  3235. 'min_y_range' : 0
  3236. }
  3237. ]
  3238. n_windows = 1 + (len(t) - window_width) // window_step
  3239. inds_L = [i * window_step for i in range(n_windows)]
  3240. inds_R = [ind + window_width for ind in inds_L]
  3241. t_downsample = [np.mean(t[ind_L:ind_R]) for ind_L, ind_R in zip(inds_L, inds_R)]
  3242. exp_labels_over_time = []
  3243. for ind_L, ind_R in zip(inds_L, inds_R):
  3244. labels = []
  3245. for sess in included_sessions:
  3246. X, y_conc, y_choice, y_outcome, t = all_data[sess]
  3247. # smooth individual trials
  3248. X_smooth = X.copy()
  3249. for i in range(X_smooth.shape[0]):
  3250. for j in range(X_smooth.shape[1]):
  3251. X_smooth[i, j, :] = smooth(X_smooth[i, j, :], smooth_type, smooth_width)
  3252. for i_neuron in range(X_smooth.shape[1]):
  3253. # form response profiles
  3254. response_profile_correct = np.full(len(unique_stims), np.nan)
  3255. response_profile_error = np.full(len(unique_stims), np.nan)
  3256. for i_stim, stim in enumerate(unique_stims):
  3257. trial_mask_correct = (y_conc == stim) & (y_outcome == 1)
  3258. if trial_mask_correct.sum() > 0:
  3259. r = X_smooth[trial_mask_correct, i_neuron, :].mean(axis=0)[ind_L:ind_R].mean()
  3260. response_profile_correct[i_stim] = r
  3261. trial_mask_error = (y_conc == stim) & (y_outcome == 0)
  3262. if trial_mask_error.sum() > 0:
  3263. r = X_smooth[trial_mask_error, i_neuron, :].mean(axis=0)[ind_L:ind_R].mean()
  3264. response_profile_error[i_stim] = r
  3265. # fit shape templates
  3266. fits = fit_shape_templates(unique_stims, response_profile_correct, fit_options)
  3267. # classify
  3268. classification = classify_fits(fits)
  3269. # parse label
  3270. label = parse_label(classification, response_profile_correct, response_profile_error)
  3271. labels.append(label)
  3272. exp_labels_over_time.append(labels)
  3273. # %%
  3274. exp_data = {}
  3275. pseudo_idx_neuron = -1
  3276. for sess in included_sessions:
  3277. exp_data[sess] = {}
  3278. X, y_conc, _, y_outcome, _ = all_data[sess]
  3279. # smooth individual trials
  3280. X_smooth = X.copy()
  3281. for i in range(X_smooth.shape[0]):
  3282. for j in range(X_smooth.shape[1]):
  3283. X_smooth[i, j, :] = smooth(X_smooth[i, j, :], smooth_type, smooth_width)
  3284. for i_neuron in range(X_smooth.shape[1]):
  3285. pseudo_idx_neuron += 1
  3286. psth_correct = np.full([len(unique_stims), len(t)], np.nan)
  3287. psth_error = np.full([len(unique_stims), len(t)], np.nan)
  3288. for i_stim, stim in enumerate(unique_stims):
  3289. trial_mask_correct = (y_conc == stim) & (y_outcome == 1)
  3290. if trial_mask_correct.sum() > 0:
  3291. psth_correct[i_stim, :] = X_smooth[trial_mask_correct, i_neuron, :].mean(axis=0)
  3292. trial_mask_error = (y_conc == stim) & (y_outcome == 0)
  3293. if trial_mask_error.sum() > 0:
  3294. psth_error[i_stim, :] = X_smooth[trial_mask_error, i_neuron, :].mean(axis=0)
  3295. label_seq = [exp_labels_over_time[i_t][pseudo_idx_neuron] for i_t in range(len(t_downsample))]
  3296. is_linear = 'Linear' in label_seq
  3297. t_linear = None
  3298. if is_linear:
  3299. inds = [int(_) for _ in np.where([label == 'Linear' for label in label_seq])[0]]
  3300. t_linear = [t_downsample[ind] for ind in inds]
  3301. is_perception = 'Step (Perception)' in label_seq
  3302. t_perception = None
  3303. if is_perception:
  3304. inds = [int(_) for _ in np.where(['Step' in label for label in label_seq])[0]]
  3305. t_perception = [t_downsample[ind] for ind in inds]
  3306. is_choice = 'Step (Choice)' in label_seq
  3307. t_choice = None
  3308. if is_choice:
  3309. inds = [int(_) for _ in np.where(['Step' in label for label in label_seq])[0]]
  3310. t_choice = [t_downsample[ind] for ind in inds]
  3311. is_other = (not is_linear) and (not is_perception) and (not is_choice)
  3312. exp_data[sess][i_neuron] = {
  3313. 'psth_correct': psth_correct.copy(),
  3314. 'psth_error': psth_error.copy(),
  3315. 'is_linear': is_linear,
  3316. 'is_perception': is_perception,
  3317. 'is_choice': is_choice,
  3318. 'is_other': is_other,
  3319. 't_linear': t_linear.copy() if t_linear is not None else None,
  3320. 't_perception': t_perception.copy() if t_perception is not None else None,
  3321. 't_choice': t_choice.copy() if t_choice is not None else None
  3322. }
  3323. # %%
  3324. mdl_data = {}
  3325. for i_model, sess in enumerate(included_sessions):
  3326. mdl_data[sess] = {}
  3327. net = all_models[sess]
  3328. net_no_noise = net.clone()
  3329. net_no_noise.noise_std = 0
  3330. output = net_no_noise(inputs).detach().numpy()
  3331. n_total = len(output_noise['all_labels'][i_model][0])
  3332. n_con = int(round(n_total / 5.88))
  3333. if int(round(5.88 * n_con)) != n_total: raise ValueError('backwards rounding failed')
  3334. for i_neuron in range(n_con):
  3335. label_seq = [
  3336. output_noise['all_labels'][i_model][i_t][i_neuron]
  3337. for i_t in range(len(t_downsample))
  3338. ]
  3339. is_linear = 'Linear' in label_seq
  3340. t_linear = None
  3341. if is_linear:
  3342. inds = [int(_) for _ in np.where([label == 'Linear' for label in label_seq])[0]]
  3343. t_linear = [t_downsample[ind] for ind in inds]
  3344. is_perception = 'Step (Perception)' in label_seq
  3345. t_perception = None
  3346. if is_perception:
  3347. inds = [int(_) for _ in np.where(['Step' in label for label in label_seq])[0]]
  3348. t_perception = [t_downsample[ind] for ind in inds]
  3349. is_choice = 'Step (Choice)' in label_seq
  3350. t_choice = None
  3351. if is_choice:
  3352. inds = [int(_) for _ in np.where(['Step' in label for label in label_seq])[0]]
  3353. t_choice = [t_downsample[ind] for ind in inds]
  3354. is_other = (not is_linear) and (not is_perception) and (not is_choice)
  3355. psth_correct = output_noise['all_psths_correct'][i_model][i_neuron]
  3356. psth_error = output_noise['all_psths_error'][i_model][i_neuron]
  3357. psth_noiseless = output[:, :, i_neuron]
  3358. mdl_data[sess][i_neuron] = {
  3359. 'psth_correct': psth_correct.copy(),
  3360. 'psth_error': psth_error.copy(),
  3361. 'psth_noiseless': psth_noiseless.copy(),
  3362. 'is_linear': is_linear,
  3363. 'is_perception': is_perception,
  3364. 'is_choice': is_choice,
  3365. 'is_other': is_other,
  3366. 't_linear': t_linear.copy() if t_linear is not None else None,
  3367. 't_perception': t_perception.copy() if t_perception is not None else None,
  3368. 't_choice': t_choice.copy() if t_choice is not None else None
  3369. }
  3370. # %%
  3371. candidate_linear, candidate_perception, candidate_choice = [], [], []
  3372. for sess in exp_data:
  3373. for i_neuron in exp_data[sess]:
  3374. if exp_data[sess][i_neuron]['is_linear'] and mdl_data[sess][i_neuron]['is_linear']:
  3375. candidate_linear.append((sess, i_neuron))
  3376. if exp_data[sess][i_neuron]['is_perception'] and mdl_data[sess][i_neuron]['is_perception']:
  3377. candidate_perception.append((sess, i_neuron))
  3378. if exp_data[sess][i_neuron]['is_choice'] and mdl_data[sess][i_neuron]['is_choice']:
  3379. candidate_choice.append((sess, i_neuron))
  3380. print('{} candidate linear units'.format(len(candidate_linear)))
  3381. print('{} candidate perception units'.format(len(candidate_perception)))
  3382. print('{} candidate choice units'.format(len(candidate_choice)))
  3383. # %%
  3384. ## example constrained neuron
  3385. # (i_session, i_neuron) = candidate_linear[1], 'linear'
  3386. # (i_session, i_neuron), sub = candidate_perception[6], 'perception'
  3387. (i_session, i_neuron), sub = candidate_choice[5], 'choice'
  3388. fig, axes = plt.subplots(2, 3, figsize=(11.5, 3.5));
  3389. for i_stim, stim in enumerate(unique_stims):
  3390. y = exp_data[i_session][i_neuron]['psth_correct'][i_stim, :]
  3391. axes[0, 0].plot(t, y, color=colors[i_stim, :]);
  3392. y = exp_data[i_session][i_neuron]['psth_error'][i_stim, :]
  3393. axes[1, 0].plot(t, y, ':', color=colors[i_stim, :]);
  3394. y = mdl_data[i_session][i_neuron]['psth_correct'][i_stim, :]
  3395. axes[0, 1].plot(t, y, color=colors[i_stim, :]);
  3396. y = mdl_data[i_session][i_neuron]['psth_error'][i_stim, :]
  3397. axes[1, 1].plot(t, y, ':', color=colors[i_stim, :]);
  3398. y = mdl_data[i_session][i_neuron]['psth_noiseless'][i_stim, :]
  3399. axes[0, 2].plot(t, y, color=colors[i_stim, :]);
  3400. for i_row, ax_row in enumerate(axes):
  3401. for i_col, ax in enumerate(ax_row):
  3402. ax.set_xlim([t[0], t[-1]]);
  3403. ax.set_xlabel('Warped time [s]');
  3404. y_max = max([ax_.get_ylim()[1] for ax_ in ax_row])
  3405. ax.set_ylim([0, y_max]);
  3406. ax.set_ylabel('Firing rate [Hz]');
  3407. ax.axvline(0, color='k');
  3408. ax.axvline(t_D, color='k');
  3409. if (i_row == 0) and (i_col == 0):
  3410. ax.title.set_text('Exp Neuron PSTH');
  3411. elif (i_row == 0) and (i_col == 1):
  3412. ax.title.set_text('Model Unit w/ Noise');
  3413. elif (i_row == 0) and (i_col == 2):
  3414. ax.title.set_text('Model Unit w/o Noise');
  3415. for point in exp_data[i_session][i_neuron]['t_{}'.format(sub)]:
  3416. axes[0, 0].plot(point, axes[0, 0].get_ylim()[1], '.b');
  3417. for point in mdl_data[i_session][i_neuron]['t_{}'.format(sub)]:
  3418. axes[0, 1].plot(point, axes[0, 1].get_ylim()[1], '.b');
  3419. # Figure 5 - Supplement 1 A
  3420. # plt.savefig('plots/review/model/ex_unit_{}.pdf'.format(sub))
  3421. # %%
  3422. frs = []
  3423. for sess in included_sessions:
  3424. ## get experimental data
  3425. X, y_conc, _, y_outcome, _ = all_data[sess]
  3426. heatmap_exp = np.full((X.shape[1], X.shape[2]), np.nan)
  3427. X_smooth = X.copy() # smooth individual trials
  3428. for i in range(X_smooth.shape[0]):
  3429. for j in range(X_smooth.shape[1]):
  3430. X_smooth[i, j, :] = smooth(X_smooth[i, j, :], smooth_type, smooth_width)
  3431. for i_neuron in range(X_smooth.shape[1]):
  3432. psth = np.full([len(unique_stims), len(t)], np.nan)
  3433. for i_stim, stim in enumerate(unique_stims):
  3434. trial_mask = (y_conc == stim) & (y_outcome == 1)
  3435. if trial_mask.sum() > 0:
  3436. psth[i_stim, :] = X_smooth[trial_mask, i_neuron, :].mean(axis=0).copy()
  3437. heatmap_exp[i_neuron, :] = np.nanmean(psth, axis=0).copy()
  3438. ## get model data
  3439. net = all_models[sess]
  3440. net_no_noise = net.clone()
  3441. net_no_noise.noise_std = 0
  3442. output = net_no_noise(inputs).detach().numpy()[:, :, :net_no_noise.observed_size]
  3443. heatmap_mdl = np.swapaxes(output, 1, 2).mean(axis=0)
  3444. global_min = 0 # np.min([np.min(heatmap_exp.flatten()), np.min(heatmap_mdl.flatten())])
  3445. global_max = np.max([np.max(heatmap_exp.flatten()), np.max(heatmap_mdl.flatten())])
  3446. shared_norm = matplotlib.colors.Normalize(vmin=global_min, vmax=global_max)
  3447. fig, axes = plt.subplots(1, 2, figsize=(11.5, X_smooth.shape[1] / 626 * 6));
  3448. im = axes[0].imshow(heatmap_exp, norm=shared_norm, cmap='hot', aspect='auto');
  3449. axes[0].set_title('Experiment');
  3450. plt.colorbar(im, ax=axes[0]);
  3451. im = axes[1].imshow(heatmap_mdl, norm=shared_norm, cmap='hot', aspect='auto');
  3452. axes[1].set_title('Model');
  3453. plt.colorbar(im, ax=axes[1]);
  3454. for ax in axes:
  3455. ax.axvline(19.5, color='w', linestyle=':', linewidth=1);
  3456. ax.axvline(len(t) - 20.5, color='w', linestyle=':', linewidth=1);
  3457. ax.set_xlabel('Time');
  3458. ax.set_ylabel('Neurons');
  3459. frs.append(global_max)
  3460. # Figure 5 - Supplement 1 B
  3461. # plt.savefig('plots/review/model/all_session_fr_session_{}.pdf'.format(sess));
  3462. print('Range of max firing rates: {} - {}'.format(np.min(frs), np.max(frs)))
  3463. # %%

analyze-rnns-manuscript.ipynb at commit 04daac5, under MIT · at the source

Overview

Authors: Liam Lang1,2,3, Camelia Yuejiao Zheng1,2,3,4, Jennifer M Blackwell1,3, Giancarlo La Camera1,2,3, Alfredo Fontanini1,2,3,4
  1. Department of Neurobiology and Behavior, Stony Brook University Stony Brook United States
  2. Graduate Program in Neuroscience, Stony Brook University Stony Brook United States
  3. Center for Neural Circuit Dynamics, Stony Brook University Stony Brook United States
  4. Medical Scientist Training Program, Stony Brook University Stony Brook United States
Institutions: Stony Brook University (United States)
Journal: eLife, volume 14, article RP109313
Dates: published online 16 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.7554/elife.109313 · PMID 42299851 · PMCID PMC13271740 · OpenAlex W7114908353
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: computational modeling (no new data) (modality), mouse (organism), cognitive (subfield)
Methods: Connectivity, Statistics, Smoothing, state filtering, decompositions, Machine learning, Single-unit activity, calcium imaging, Physiology & signal measures
Keywords: decision-making, computational modeling, gustatory cortex, recurrent neural network, Mouse
MeSH: Cerebral Cortex*, Decision Making*, Neurons*, Taste*, Taste Perception*, Animals, Male, Mice, Models, Neurological, Recurrent Neural Networks (* major topic)
Journal subjects: Neuroscience
Topic: Biochemical Analysis and Sensing Techniques (Nutrition and Dietetics, Nursing), according to OpenAlex
Citations: not cited yet (Europe PMC); 65 references in the paper

Abstract

Cortical circuits produce time-varying patterns of population and single-neuron activity that play a fundamental role in perceptual and behavioral processes. However, the functional contributions of individual neuron activity to population dynamics and behavior remain unclear. Here, we addressed this issue focusing on the mouse gustatory cortex (GC) and using a taste mixture-based decision-making task, high-density electrophysiology, and computational modeling. GC population dynamics represented stimuli linearly during taste sampling, and choices categorically before decisions. Single neurons were classified by their linear and categorical activity patterns, revealing sub-populations encoding sensory, perceptual, and decisional variables. To test their functional role, we built a recurrent neural network model of GC. Model perturbations showed linear and categorical neurons were essential for driving normal population dynamics and behavioral performance, whereas many units with other activity patterns could be silenced without consequence. These results have implications that extend beyond GC and demonstrate the role of linear and categorical coding neurons in cortical dynamics and behavior during perceptual decision-making.

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

llang6/linear-categorical

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 04daac524da9485f777e48c7a19495932f2197fa, 15 June 2026
Languages: MATLAB (89), Python (6), Jupyter (4)
Size: 139 files, 99 scripts
Software Heritage: archived
Found in: “Data availability”
Holds: README, license file, environment (dPCA_master/python/requirements.txt, dPCA_master/python/setup.py), 4 notebooks
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: Image Processing Toolbox (13 files), NumPy (9 files), Statistics and Machine Learning Toolbox (7 files), SciPy (6 files), scikit-learn (4 files), PyTorch (3 files), Matplotlib (2 files), Numba (2 files), pandas (2 files), statsmodels (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
101 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;
  • 99 scripts, each with its path and the digest of its content;
  • 15 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

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

Data

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

Data availability

Experimental dataset, modeling dataset, and code for all analyses available at: https://github.com/llang6/linear-categorical (copy archived at Lang, 2026).

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, pages, dates, 5 authors, 5 keywords, 10 MeSH terms, 3 funders, 55 references.

Cite

This paper

Lang, L., Zheng, C. Y., Blackwell, J. M., La Camera, G., & Fontanini, A. (2026). Linear and categorical coding units in the mouse gustatory cortex drive population dynamics and behavior in taste decision-making. eLife, 14, RP109313. https://doi.org/10.7554/elife.109313

BibTeX

@article{lang2026linear,
author = {Lang, Liam and Zheng, Camelia Yuejiao and Blackwell, Jennifer M and La Camera, Giancarlo and Fontanini, Alfredo},
title = {{Linear and categorical coding units in the mouse gustatory cortex drive population dynamics and behavior in taste decision-making}},
journal = {eLife},
year = {2026},
month = jun,
volume = {14},
pages = {RP109313},
publisher = {eLife Sciences Publications, Ltd},
issn = {2050-084X},
doi = {10.7554/elife.109313},
url = {https://doi.org/10.7554/elife.109313},
pmid = {42299851},
pmcid = {PMC13271740}
}

RIS

TY - JOUR
AU - Lang, Liam
AU - Zheng, Camelia Yuejiao
AU - Blackwell, Jennifer M
AU - La Camera, Giancarlo
AU - Fontanini, Alfredo
TI - Linear and categorical coding units in the mouse gustatory cortex drive population dynamics and behavior in taste decision-making
T2 - eLife
J2 - eLife
PY - 2026
DA - 2026/06/16
VL - 14
SP - RP109313
SN - 2050-084X
PB - eLife Sciences Publications, Ltd
DO - 10.7554/elife.109313
UR - https://doi.org/10.7554/elife.109313
LA - en
ER -

CSL-JSON

{
"id": "10.7554/elife.109313",
"type": "article-journal",
"title": "Linear and categorical coding units in the mouse gustatory cortex drive population dynamics and behavior in taste decision-making",
"container-title": "eLife",
"author": [
{
"family": "Lang",
"given": "Liam"
},
{
"family": "Zheng",
"given": "Camelia Yuejiao"
},
{
"family": "Blackwell",
"given": "Jennifer M"
},
{
"family": "La Camera",
"given": "Giancarlo"
},
{
"family": "Fontanini",
"given": "Alfredo"
}
],
"container-title-short": "eLife",
"volume": "14",
"page": "RP109313",
"DOI": "10.7554/elife.109313",
"PMID": "42299851",
"PMCID": "PMC13271740",
"ISSN": "2050-084X",
"publisher": "eLife Sciences Publications, Ltd",
"URL": "https://doi.org/10.7554/elife.109313",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
16
]
]
}
}

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.1016/j.neuron.2026.07.016 [code]
Inferring brain-wide interactions using data-constrained recurrent neural network models.
Journal: Neuron
In common: Image Processing Toolbox, Statistics and Machine Learning Toolbox, Matplotlib, 1 other tool, computational modeling (no new data), mouse, 5 references
[2] doi:10.1038/s41593-026-02232-0 [code]
Entorhinal cortex represents task-relevant remote locations independently of CA1.
Journal: Nature neuroscience
In common: Numba, Image Processing Toolbox, statsmodels, 7 other tools, mouse, 1 reference
[3] doi:10.1038/s41467-026-76581-6 [code]
Thalamocortical bursts encode reward contingencies and drive associative learning.
Journal: Nature communications
In common: Image Processing Toolbox, Statistics and Machine Learning Toolbox, scikit-learn, 4 other tools, mouse, 3 references
[4] doi:10.1038/s41467-026-71725-0 [code]
Interactions across hemispheres in prefrontal cortex reflect global cognitive processing.
Journal: Nature communications
In common: Image Processing Toolbox, Statistics and Machine Learning Toolbox, scikit-learn, 4 other tools, cognitive, 3 references
[5] doi:10.1038/s41467-026-75347-4 [code]
Sleep reveals dynamics integrating and segregating movement and stimulus representations in V1.
Journal: Nature communications
In common: Numba, Image Processing Toolbox, PyTorch, 6 other tools, mouse, 1 reference
[6] doi:10.1038/s41467-026-71151-2 [code]
Common and distinct neural correlates of social interaction processing and theory of mind in narratives.
Journal: Nature communications
In common: Numba, Image Processing Toolbox, statsmodels, 6 other tools, cognitive
[7] doi:10.7554/elife.109717 [code]
Retrosplenial cortex enables context-dependent goal-directed sensorimotor transformation.
Journal: eLife
In common: Numba, statsmodels, Statistics and Machine Learning Toolbox, 5 other tools, mouse, 1 reference
[8] doi:10.1038/s41592-026-03076-z [code]
Neuropixels Opto: combining high-resolution electrophysiology and optogenetics.
Journal: Nature methods
In common: PyTorch, pandas, SciPy, 2 other tools, mouse, 4 references
[9] doi:10.7554/elife.111876 [code]
Distinct sensorimotor encoding in tuft dendrites and somata associated with action, correction, and learning.
Journal: eLife
In common: statsmodels, PyTorch, scikit-learn, 4 other tools, mouse, 2 references
[10] doi:10.1016/j.patter.2026.101590 [code]
Density-based longitudinal neuron tracking in high-density electrophysiological recordings.
Journal: Patterns (New York, N.Y.)
In common: statsmodels, PyTorch, Statistics and Machine Learning Toolbox, 5 other tools, 2 references

Contribute

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

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

Request its removal

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

Discussion, reproductions, activity

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

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

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