OSCR

Protocol for analyzing slow cortical dynamics in mouse neuronal recordings.

Code ↔ Paper

30 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 30 matches
  1. [1] § Step-by-step method details › Trial-to-trial variability analysis ↔ slow_dynamics_isi_correlation.ipynb, lines 36–62 · score 0.78 · 0–1.2 s, Responsive cells, stimulus onset, stimulus triggered, threshold, post
  2. [2] § Step-by-step method details › Trial-to-trial variability analysis ↔ functions/sd_utils.py, lines 380–427 · score 0.71 · stimulus response window, peak exceeds, Responsive cells, threshold, 2 s
  3. [3] § Step-by-step method details › Intrinsic timescale analysis ↔ functions/sd_utils.py, lines 782–809 · score 0.65 · tau_cell, tau_net, network intrinsic timescales, sd
  4. [4] § Step-by-step method details › Intrinsic timescale analysis ↔ functions/sd_utils.py, lines 782–809 · score 0.65 · tau_cell, tau_net, network intrinsic timescales, sd
  5. [5] § Step-by-step method details › Intrinsic timescale analysis ↔ slow_dynamics_net_int.ipynb, lines 55–66 · score 0.64 · sd.get_network_intrinsic_timescales, tau_cell, tau_net
  6. [6] § Step-by-step method details › Intrinsic timescale analysis ↔ slow_dynamics_net_int.ipynb, lines 55–66 · score 0.64 · sd.get_network_intrinsic_timescales, tau_cell, tau_net
  7. [7] § Step-by-step method details › Trial-to-trial variability analysis ↔ slow_dynamics_isi_correlation.ipynb, lines 36–62 · score 0.64 · plot_t_tuning, trial_frames_tuning, sd.get_frames
  8. [8] § Step-by-step method details › Trial-to-trial variability analysis ↔ slow_dynamics_isi_correlation.py, lines 41–65 · score 0.64 · plot_t_tuning, trial_frames_tuning, sd.get_frames
  9. [9] § Step-by-step method details › Intrinsic timescale analysis ↔ functions/sd_utils.py, lines 505–511 · score 0.64 · Euclidean distance, population vector, baseline, Firing rates, network, frame
  10. [10] § Step-by-step method details › Trial-to-trial variability analysis ↔ slow_dynamics_isi_correlation_comb.ipynb, lines 167–190 · score 0.63 · mouse_idx2, trial_idx, n_freq, echo
  11. [11] § Step-by-step method details › Trial-to-trial variability analysis ↔ slow_dynamics_isi_correlation.ipynb, lines 114–132 · score 0.63 · mouse_idx2, trial_idx, n_freq, echo
  12. [12] § Step-by-step method details › Trial-to-trial variability analysis ↔ slow_dynamics_isi_correlation.ipynb, lines 114–132 · score 0.62 · mouse_list, mouse_idx2, mouse_tag
  13. [13] § Step-by-step method details › Trial-to-trial variability analysis ↔ slow_dynamics_isi_correlation.py, lines 125–139 · score 0.62 · mouse_list, mouse_idx2, mouse_tag
  14. [14] § Step-by-step method details › Intrinsic timescale analysis ↔ slow_dynamics_net_int.ipynb, lines 34–53 · score 0.61 · num_initial_trials_skip, flatten_runs, untrained, firing rate, oddball, network
  15. [15] § Step-by-step method details › Intrinsic timescale analysis ↔ slow_dynamics_net_int.py, lines 39–58 · score 0.61 · num_initial_trials_skip, flatten_runs, untrained, firing rate, oddball, network
  16. [16] § Step-by-step method details › Data loading and preprocessing ↔ slow_dynamics_decoder_comb_data_rnn.ipynb, lines 86–101 · score 0.60 · trial_frames_rnn, plot_t_rnn, sd.get_frames
  17. [17] § Step-by-step method details › Data loading and preprocessing ↔ slow_dynamics_decoder_comb.py, lines 91–96 · score 0.60 · trial_frames_rnn, plot_t_rnn, sd.get_frames
  18. [18] § Step-by-step method details › Data loading and preprocessing ↔ slow_dynamics_decoder.ipynb, lines 1–14 · score 0.59 · import functions, pipeline_dir, slow dynamics
  19. [19] § Step-by-step method details › Data loading and preprocessing ↔ slow_dynamics_decoder_comb_data_rnn.ipynb, lines 1–12 · score 0.59 · import functions, pipeline_dir, slow dynamics
  20. [20] § Step-by-step method details › Trial-to-trial variability analysis ↔ slow_dynamics_isi_correlation.ipynb, lines 134–143 · score 0.58 · n_freq, mouse_tag, title_tag
  21. [21] § Step-by-step method details › Trial-to-trial variability analysis ↔ slow_dynamics_isi_correlation.py, lines 143–143 · score 0.58 · n_freq, mouse_tag, title_tag
  22. [22] § Step-by-step method details › Trial-to-trial variability analysis ↔ slow_dynamics_isi_correlation.py, lines 125–139 · score 0.58 · mouse_idx2, trial_idx, stim_trig_resp
  23. [23] § Step-by-step method details › Intrinsic timescale analysis ↔ slow_dynamics_net_int.ipynb, lines 82–110 · score 0.57 · Freq neuron, sd.stat_compare, CaIm, Intrinsic, timescale
  24. [24] § Step-by-step method details › Intrinsic timescale analysis ↔ slow_dynamics_net_int.ipynb, lines 82–110 · score 0.57 · Ob neuron, sd.stat_compare, CaIm, Intrinsic, timescale
  25. [25] § Step-by-step method details › Data loading and preprocessing ↔ functions/sd_utils.py, lines 57–183 · score 0.56 · CaImAn, derivative, rectified, oasis, firing rate, smoothdfdt
  26. [26] § Step-by-step method details › Intrinsic timescale analysis ↔ functions/sd_utils.py, lines 513–563 · score 0.55 · zero lag, fit, autocorrelation, trace, Intrinsic, timescale
  27. [27] § Step-by-step method details › Binwise decoder ↔ functions/sd_decoder.py, lines 90–141 · score 0.55 · temporal generalization, decoder trained, cross, Binwise
  28. [28] § Step-by-step method details › Binwise decoder ↔ slow_dynamics_decoder.ipynb, lines 71–92 · score 0.51 · run_binwise_dec, cross validation, diagonal, decoder
  29. [29] § Step-by-step method details › Binwise decoder ↔ slow_dynamics_decoder.py, lines 67–88 · score 0.51 · run_binwise_dec, cross validation, diagonal, decoder
  30. [30] § Before you begin › Data preparation ↔ slow_dynamics_decoder_comb_data_rnn.ipynb, lines 36–55 · score 0.50 · neuronal activity, RNNs trained, untrained, oddball, network, 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

Python · 1,678 lines · 69 KB · no license · 6 matches

  1. # -*- coding: utf-8 -*-
  2. """
  3. Created on Wed Nov 26 13:32:23 2025
  4. @author: ys2605
  5. """
  6. import os
  7. import h5py
  8. import numpy as np
  9. from scipy.stats import norm, wilcoxon, mannwhitneyu, ttest_ind, ttest_rel, f_oneway
  10. from scipy.spatial.distance import pdist, squareform, cdist
  11. from scipy.signal import correlate #, correlation_lags
  12. import matplotlib.pyplot as plt
  13. from matplotlib.ticker import FuncFormatter, NullFormatter
  14. from datetime import datetime
  15. from sd_decoder import plot_diag_binwise_dec, plot_full_binwise_dec
  16. #%%
  17. def get_fnames_from_dir(data_dir, ext_list = [], tags = None, f_list=None):
  18. # list files in data_dir, keeping those with an ext_list extension AND containing every tag; f_list overrides the listing
  19. if f_list is None:
  20. if not os.path.isdir(data_dir):
  21. raise FileNotFoundError('Data directory not found: %s' % data_dir)
  22. f_list = os.listdir(data_dir)
  23. f_list2 = []
  24. for fil1 in f_list:
  25. if len(ext_list):
  26. for ext1 in ext_list:
  27. if fil1.endswith(ext1):
  28. f_list2.append(fil1)
  29. else:
  30. f_list2.append(fil1)
  31. if tags is not None:
  32. if type(tags) is str:
  33. tags = [tags]
  34. f_list_out = []
  35. for fil1 in f_list2:
  36. has_tag = True
  37. for tag in tags:
  38. if tag not in fil1:
  39. has_tag = False
  40. if has_tag:
  41. f_list_out.append(fil1)
  42. else:
  43. f_list_out = f_list2
  44. return f_list_out
  45. #%%
  46. def load_caim_data_mat(data_dir, ext_list = [], tags = None, num_files=None, data_tag = 'results_cnmf_sort.mat', proc_tag = 'processed_data.mat', deconvolution='oasis', smooth_std_duration=0.1, norm_first=False):
  47. # load matched caiman .mat sessions from data_dir into a list of dataset dicts (one per session);
  48. # main fields: firing_rates (neurons x time), trial_types, stim_times, volume_period, isi.
  49. # deconvolution methods are either oasis (caiman default) or smoothdftd - smoothed, rectified first derivative
  50. # norm_first: False (default) = smooth then peak-normalize (original); True = peak-normalize raw S then smooth (MATLAB order)
  51. # smooth_std_duration in sec
  52. dir_files = get_fnames_from_dir(data_dir)
  53. flist = get_fnames_from_dir(data_dir, ext_list=ext_list, tags=tags, f_list=dir_files)
  54. if len(flist) == 0:
  55. print('Warning: no files found in %s matching ext %s and tags %s' % (data_dir, ext_list, tags))
  56. return []
  57. if num_files is None:
  58. num_files = len(flist)
  59. else:
  60. num_files = np.min([len(flist), num_files])
  61. data_out = []
  62. for n_fl in range(num_files):
  63. fname_core = flist[n_fl]
  64. if data_tag in fname_core:
  65. fname_core = fname_core.removesuffix(data_tag)
  66. if proc_tag in fname_core:
  67. fname_core = fname_core.removesuffix(proc_tag)
  68. flist_data = get_fnames_from_dir(data_dir, ext_list = ['.mat'], tags = [fname_core, data_tag], f_list=dir_files)
  69. flist_proc = get_fnames_from_dir(data_dir, ext_list = ['.mat'], tags = [fname_core, proc_tag], f_list=dir_files)
  70. do_load = False
  71. if len(flist_proc):
  72. if len(flist_data):
  73. do_load = True
  74. else:
  75. print(fname_core + " data file with " + data_tag + " tag not found, skipping")
  76. else:
  77. print(fname_core + " proc file with " + proc_tag + " tag not found, skipping")
  78. if do_load:
  79. data_slice = {'flist_data': flist_data,
  80. 'flist_proc': flist_proc}
  81. f_proc = h5py.File(data_dir + '/' + flist_proc[0], 'r')
  82. vid_cuts_trace = f_proc[f_proc['data']['file_cuts_params'][0][0]]['vid_cuts_trace'][()].flatten().astype(bool)
  83. trial_types = f_proc['data']['trial_types'][()].flatten().astype(int)
  84. stim_times = f_proc[f_proc['data']['stim_times_frame'][0][0]][()].flatten().astype(int)
  85. if 'volume_period' in f_proc['data']['frame_data'].keys():
  86. data_slice['volume_period'] = f_proc['data']['frame_data']['volume_period'][()].flatten()[0]
  87. if data_slice.get('volume_period', 0) <= 0:
  88. print('Warning: %s missing or invalid volume_period in _processed_data; using default 33.3 ms (~30 Hz).' % fname_core)
  89. data_slice['volume_period'] = 33.3 # ms, default ~30 Hz
  90. if 'isi' in f_proc['data']['stim_params'].keys():
  91. data_slice['isi'] = f_proc['data']['stim_params']['isi'][()].flatten()[0]
  92. if 'MMN_orientations' in f_proc['data'].keys():
  93. data_slice['MMN_ori'] = f_proc['data']['MMN_orientations'][()].flatten().astype(int)
  94. if 'MMN_freq' in f_proc['data']['stim_params'].keys():
  95. data_slice['MMN_ori'] = f_proc['data']['stim_params']['MMN_freq'][()].flatten().astype(int)
  96. f_proc.close()
  97. firing_rates_all = []
  98. dset_idx_all = []
  99. for n_fl in range(len(flist_data)):
  100. fname_data = flist_data[n_fl]
  101. f = h5py.File(os.path.join(data_dir, fname_data), 'r')
  102. d_est = f['est']
  103. d_proc = f['proc']
  104. comp_acc = d_proc['comp_accepted'][()].flatten().astype(bool)
  105. if deconvolution == 'oasis':
  106. firing_rates_cut = d_est['S'][()][:,comp_acc].T
  107. elif deconvolution == 'smoothdfdt':
  108. C = d_est['C'][()]
  109. YrA = d_est['YrA'][()]
  110. ca_traces_cut = (C + YrA)[:,comp_acc].T
  111. firing_rates_cut = smooth_dfdt(ca_traces_cut, sigma_frames=1000/data_slice['volume_period']*0.1, do_smooth=True)
  112. sigma_fr = 1000/data_slice['volume_period']*smooth_std_duration
  113. if norm_first:
  114. # MATLAB order: peak-normalize the raw deconvolved S first, then smooth
  115. peak_rate = np.max(firing_rates_cut, axis=1)[:,None]
  116. peak_rate[peak_rate == 0] = 1
  117. firing_rates_cutn = gauss_smooth(firing_rates_cut/peak_rate, sigma_frames=sigma_fr)
  118. else:
  119. # original order: smooth, then peak-normalize the smoothed trace
  120. firing_rates_cut = gauss_smooth(firing_rates_cut, sigma_frames=sigma_fr)
  121. peak_rate = np.max(firing_rates_cut, axis=1)[:,None]
  122. peak_rate[peak_rate == 0] = 1
  123. firing_rates_cutn = firing_rates_cut/peak_rate
  124. firing_rates = np.zeros((firing_rates_cutn.shape[0], vid_cuts_trace.shape[0]))
  125. firing_rates[:, vid_cuts_trace] = firing_rates_cutn
  126. firing_rates_all.append(firing_rates)
  127. dset_idx_all.append(np.ones(firing_rates.shape[0], dtype=int)*(n_fl))
  128. f.close()
  129. firing_rates = np.vstack(firing_rates_all)
  130. dset_idx = np.hstack(dset_idx_all)
  131. data_slice['fname_core'] = fname_core
  132. data_slice['firing_rates'] = firing_rates
  133. data_slice['trial_types'] = trial_types
  134. data_slice['stim_times'] = stim_times
  135. data_slice['vid_cuts_trace'] = vid_cuts_trace
  136. data_slice['files_loaded'] = do_load
  137. data_slice['dset_idx'] = dset_idx
  138. data_out.append(data_slice)
  139. if len(data_out):
  140. mouse_ids = list(np.unique(get_mouse_id(data_out)))
  141. print('Loaded %d datasets from %d mice: %s' % (len(data_out), len(mouse_ids), mouse_ids))
  142. return data_out
  143. def load_caim_data_mat2(*args, **kwargs):
  144. # same as load_caim_data_mat but with the MATLAB preprocessing order:
  145. # peak-normalize the raw deconvolved S first, then smooth (norm_first=True)
  146. kwargs['norm_first'] = True
  147. return load_caim_data_mat(*args, **kwargs)
  148. def h5_load_group(group, keys=None):
  149. # load the given keys (default: all) of an open h5py group into a plain dict
  150. if keys is None:
  151. keys = group.keys()
  152. data = {}
  153. for key1 in keys:
  154. data[key1] = group[key1][()]
  155. return data
  156. def get_values(data_out, key):
  157. # collect data_slice[key] across the datasets in the list (skips slices missing key)
  158. values = []
  159. for data_slice in data_out:
  160. if key in data_slice.keys():
  161. values.append(data_slice[key])
  162. return values
  163. def get_mouse_id(data_echo):
  164. # mouse id per dataset = text before the first underscore of fname_core (e.g. 'M4372')
  165. fnames = get_values(data_echo, 'fname_core')
  166. mouse_id = []
  167. for fname in fnames:
  168. mouse_id.append(fname.split('_')[0])
  169. return mouse_id
  170. def get_example_mouse(data_echo, example_mouse='M4372'):
  171. # return (per-dataset mouse-id array, example_mouse); error if example_mouse is not among the loaded datasets
  172. mouse_list = np.array(get_mouse_id(data_echo))
  173. mouse_uq = np.unique(mouse_list)
  174. if example_mouse not in mouse_uq:
  175. raise ValueError('Example mouse %s not found among loaded datasets. The trial-to-trial '
  176. 'similarity demo requires all six variable-ISI (echo) datasets: '
  177. 'M226, M4264, M4265, M4266, M4371, M4372. Loaded mice: %s'
  178. % (example_mouse, list(mouse_uq)))
  179. return mouse_list, example_mouse
  180. def smooth_dfdt(data, do_smooth=True, sigma_frames=1, rectify=True, normalize=True):
  181. # smoothdfdt firing-rate estimate: per-neuron smoothed first derivative, optionally rectified and peak-normalized
  182. num_cells, num_frames = data.shape
  183. firing_rates = np.zeros((num_cells, num_frames));
  184. if sigma_frames == 0:
  185. do_smooth=False
  186. if do_smooth:
  187. s_fr = np.ceil(sigma_frames).astype(int)
  188. x = np.linspace(-3*s_fr, 3*s_fr, s_fr*6+1)
  189. gauss_kernel = np.exp(-x**2 / (2 * sigma_frames**2))
  190. for n_cell in range(num_cells):
  191. temp_data = np.diff(data[n_cell,:], prepend=0)
  192. if do_smooth:
  193. temp_data = np.convolve(temp_data, gauss_kernel, mode='same')
  194. if rectify:
  195. temp_data = np.maximum(temp_data, 0)
  196. if normalize:
  197. temp_data = temp_data - np.mean(temp_data)
  198. temp_data = temp_data/np.max(temp_data)
  199. firing_rates[n_cell,:] = temp_data;
  200. return firing_rates
  201. def gauss_smooth(firing_rates, sigma_frames=1):
  202. # gaussian smoothing of each neuron's trace (sigma in frames; 0 = no smoothing)
  203. if sigma_frames:
  204. num_cells, num_frames = firing_rates.shape
  205. s_fr = np.ceil(sigma_frames).astype(int)
  206. x = np.linspace(-3*s_fr, 3*s_fr, s_fr*6+1)
  207. gauss_kernel = np.exp(-x**2 / (2 * sigma_frames**2))
  208. firing_rates_sm = np.zeros((num_cells, num_frames));
  209. for n_cell in range(num_cells):
  210. firing_rates_sm[n_cell,:] = np.convolve(firing_rates[n_cell,:], gauss_kernel, mode='same')
  211. else:
  212. firing_rates_sm = firing_rates
  213. return firing_rates_sm
  214. def normalize(rates):
  215. # assumes neurons x time
  216. rates_n = rates - np.min(rates, axis=1)[:, None]
  217. # ignoring cells that are always zero
  218. max_rates = np.max(rates_n, axis=1)
  219. has_max = max_rates > 0
  220. rates_n = rates_n[has_max,:] / max_rates[has_max][:,None]
  221. return rates_n
  222. def get_frames(trial_win = [-0.05, .95], frame_rate = 30):
  223. # anchor at 0
  224. frame_start = np.ceil(trial_win[0] * frame_rate)
  225. frame_end = np.ceil(trial_win[1] * frame_rate)
  226. trial_frames = [int(frame_start), int(frame_end)]
  227. plot_t = np.round(np.arange(frame_start/frame_rate, frame_end/frame_rate, 1/frame_rate), decimals=4)
  228. return trial_frames, plot_t
  229. def get_stim_trig_resp(firing_rates, stim_times, trial_frames = [-29, 85]):
  230. # input: cells x time
  231. num_cells, T = firing_rates.shape
  232. num_trials = len(stim_times)
  233. win_size = trial_frames[1] - trial_frames[0]
  234. stim_trig_resp = np.zeros((num_cells, win_size, num_trials))
  235. for n_tr in range(num_trials):
  236. cur_frame = round(stim_times[n_tr]-1) # correct for matlab to python
  237. raw_start = cur_frame + trial_frames[0]
  238. raw_end = cur_frame + trial_frames[1]
  239. # clamp the window to the recording bounds; trials near the edges stay zero-padded
  240. src_start = np.max([raw_start, 0])
  241. src_end = np.min([raw_end, T])
  242. if src_end > src_start:
  243. dst_start = src_start - raw_start
  244. stim_trig_resp[:, dst_start:dst_start + (src_end - src_start), n_tr] = firing_rates[:, src_start:src_end]
  245. return stim_trig_resp
  246. def save_fig(fig, path='/', name_tag=''):
  247. # save the figure as svg + png (1200 dpi), named from the first-axis title + date + name_tag
  248. plt.rcParams['svg.fonttype'] = 'none'
  249. name1 = fig.axes[0].title.get_text()
  250. now1 = datetime.now()
  251. date_tag = '%d_%d_%d_%dh_%dm' % (now1.year, now1.month, now1.day, now1.hour, now1.minute)
  252. fig.savefig('%s/%s_%s%s.svg' % (path, name1, date_tag, name_tag))
  253. fig.savefig('%s/%s_%s%s.png' % (path, name1, date_tag, name_tag), dpi=1200)
  254. def get_trial_peak(trial_ave, peak_size=3):
  255. # per-neuron peak response: mean over a peak_size-frame window centred on each neuron's argmax; returns (peak_vals, peak_locs)
  256. num_cells, num_bins = trial_ave.shape
  257. pad_left = np.floor((peak_size-1)/2).astype(int)
  258. pad_right = np.ceil((peak_size-1)/2).astype(int)
  259. peak_locs = np.argmax(trial_ave,axis=1)
  260. peak_start = peak_locs - pad_left
  261. peak_end = peak_locs + pad_right + 1
  262. idx_fix_start = peak_start < 0
  263. if sum(idx_fix_start):
  264. peak_start_to_fix = peak_start[idx_fix_start]
  265. peak_start[idx_fix_start] = peak_start[idx_fix_start] - peak_start_to_fix
  266. peak_end[idx_fix_start] = peak_end[idx_fix_start] - peak_start_to_fix
  267. idx_fix_end = peak_end > num_bins
  268. if sum(idx_fix_end):
  269. peak_end_to_fix = peak_end[idx_fix_end]
  270. peak_end[idx_fix_end] = peak_end[idx_fix_end] - peak_end_to_fix + num_bins
  271. peak_start[idx_fix_end] = peak_start[idx_fix_end] - peak_end_to_fix + num_bins
  272. peak_vals = np.zeros(num_cells)
  273. for n_cell in range(num_cells):
  274. peak_vals[n_cell] = np.mean(trial_ave[n_cell,peak_start[n_cell]:peak_end[n_cell]])
  275. return peak_vals, peak_locs
  276. def compute_tuning(stim_trig_resp, trial_types, trials_analyze, plot_t, num_samp=2000, z_thresh = 3, sig_resp_win = [0, 1.5], seed=None):
  277. # find stimulus-responsive cells per trial type by comparing each cell's peak response to a shuffled null
  278. # (z_thresh over num_samp shuffles, within sig_resp_win); returns a (cells x trial types) responsive-cell mask
  279. tt_use_idx = np.sum(trial_types == trials_analyze[:,None],axis=0).astype(bool)
  280. trial_types_use = trial_types[tt_use_idx]
  281. stim_trig_resp_use = stim_trig_resp[:,:,tt_use_idx]
  282. num_cells, _, num_trials = stim_trig_resp_use.shape
  283. num_tt = len(trials_analyze)
  284. # get data
  285. peak_vals = np.full((num_cells, num_tt), np.nan)
  286. peak_locs = np.full((num_cells, num_tt), np.nan)
  287. for n_tt in range(num_tt):
  288. tt1_idx = trial_types_use == trials_analyze[n_tt]
  289. if np.any(tt1_idx):
  290. trial_ave1 = np.mean(stim_trig_resp_use[:,:,tt1_idx], axis=2)
  291. peak_vals[:,n_tt], peak_locs[:,n_tt] = get_trial_peak(trial_ave1, peak_size=3)
  292. # null distribution of peaks from random trial resamples (with replacement)
  293. trials_per_stim = np.zeros(num_tt, dtype=int)
  294. for n_tt in range(num_tt):
  295. trials_per_stim[n_tt] = np.sum(trial_types_use == trials_analyze[n_tt])
  296. trials_per_stim_ave = np.round(np.mean(trials_per_stim)).astype(int)
  297. samp_peak_vals = np.full((num_cells, num_samp), np.nan)
  298. samp_peak_locs = np.full((num_cells, num_samp), np.nan)
  299. rng = seed if isinstance(seed, np.random.Generator) else np.random.default_rng(seed)
  300. for n_cell in range(num_cells):
  301. random_integers = rng.integers(low=0, high=num_trials, size=(trials_per_stim_ave, num_samp))
  302. samp_trial_ave = np.mean(stim_trig_resp_use[n_cell,:,random_integers], axis=0)
  303. samp_peak_vals[n_cell,:], samp_peak_locs[n_cell,:] = get_trial_peak(samp_trial_ave, peak_size=3)
  304. idx1 = ~np.isnan(peak_locs[0,:])
  305. peak_locs_t = np.full((num_cells, num_tt), np.nan)
  306. peak_locs_t[:,idx1] = plot_t[peak_locs[:,idx1].astype(int)]
  307. peak_in_resp_win = np.logical_and(peak_locs_t >= sig_resp_win[0], peak_locs_t <= sig_resp_win[1])
  308. # responsive if the observed peak exceeds the z-threshold percentile of the null and peaks within the response window
  309. peak_prcntle = norm.cdf(z_thresh)*100
  310. prc_thresh = np.percentile(samp_peak_vals, peak_prcntle, axis=1)
  311. resp_cells_peak = np.zeros((num_cells, num_tt), dtype=bool)
  312. resp_cells_peak[:,idx1] = np.logical_and(peak_vals[:,idx1] > prc_thresh[:,None], peak_in_resp_win[:,idx1])
  313. return resp_cells_peak
  314. def compute_correlation(stim_trig_resp, trial_types, trials_analyze, resp_cells=None, min_resp_cells=5, subtract_mean=False, add_noise_sigma=1e-5, metric='correlation', cell_select='resp_marg', drop_zero_cells=True, seed=None):
  315. # cell_select: 'resp_marg' (union of responsive cells; = MATLAB 'Resp marg'),
  316. # 'resp_split' (per-frequency responsive cells; = MATLAB 'Resp split'), or 'all'
  317. # drop_zero_cells: True drops cells with non-positive mean across trials (Python-only; MATLAB keeps all)
  318. rng = seed if isinstance(seed, np.random.Generator) else np.random.default_rng(seed)
  319. corr_vals = np.full((len(trials_analyze)), np.nan)
  320. if subtract_mean:
  321. stim_trig_resp = stim_trig_resp - np.mean(stim_trig_resp)
  322. if add_noise_sigma: # add some uncorrelated noise to add stability
  323. stim_trig_resp = stim_trig_resp + rng.normal(0, add_noise_sigma, size=stim_trig_resp.shape)
  324. if resp_cells is not None:
  325. resp_marg = np.sum(resp_cells, axis=1).astype(bool)
  326. for n_tn in range(len(trials_analyze)):
  327. if resp_cells is not None:
  328. trig = sum(resp_cells[:,n_tn]) > min_resp_cells
  329. else:
  330. trig = 1
  331. if trig:
  332. tn1 = trials_analyze[n_tn]
  333. tr_idx = trial_types == tn1
  334. stim_trig_resp2 = stim_trig_resp[:,:,tr_idx]
  335. # cell selection mode
  336. if resp_cells is None or cell_select == 'all':
  337. cell_mask = np.ones(stim_trig_resp2.shape[0], dtype=bool)
  338. elif cell_select == 'resp_split':
  339. cell_mask = resp_cells[:,n_tn].astype(bool) # per-frequency responsive cells
  340. else: # 'resp_marg'
  341. cell_mask = resp_marg # union of responsive cells
  342. stim_trig_resp3 = stim_trig_resp2[cell_mask,:,:]
  343. stim_trig_resp4 = np.mean(stim_trig_resp3, axis=1)
  344. if drop_zero_cells:
  345. act_idx = np.mean(stim_trig_resp4, axis=1) > 0
  346. stim_trig_resp5 = stim_trig_resp4[act_idx,:]
  347. else:
  348. stim_trig_resp5 = stim_trig_resp4
  349. distances = squareform(pdist(stim_trig_resp5.T, metric=metric)) # cosine, correlation
  350. # trial-to-trial similarity = 1 - distance; average over the unique trial pairs (lower triangle)
  351. SI = 1 - distances
  352. SI2 = np.tril(SI, k=-1)
  353. SI2_vals = SI2[SI2.astype(bool)]
  354. corr_vals[n_tn] = np.mean(SI2_vals) if len(SI2_vals) else np.nan
  355. # if 0:
  356. # if tn1==4:
  357. # plt.figure()
  358. # plt.imshow(SI)
  359. # plt.title('isi = ' + str(data_echo[n_fl]['isi']))
  360. return corr_vals
  361. def compute_correlation_mat(stim_trig_resp, subtract_mean=False, add_noise_sigma=1e-5, metric='correlation', seed=None):
  362. # trial-by-trial similarity matrix (1 - pairwise distance) of each trial's mean stimulus response
  363. rng = seed if isinstance(seed, np.random.Generator) else np.random.default_rng(seed)
  364. if subtract_mean:
  365. stim_trig_resp = stim_trig_resp - np.mean(stim_trig_resp)
  366. if add_noise_sigma: # add some uncorrelated noise to add stability
  367. stim_trig_resp = stim_trig_resp + rng.normal(0, add_noise_sigma, size=stim_trig_resp.shape)
  368. stim_trig_resp2 = np.mean(stim_trig_resp, axis=1)
  369. distances = squareform(pdist(stim_trig_resp2.T, metric=metric)) # cosine, correlation
  370. SI = 1 - distances
  371. return SI
  372. #%%
  373. def get_network_distance(firing_rates):
  374. # euclidean distance of the population vector from its time-averaged baseline at each frame: d(t) = ||r(t) - mean_t r||
  375. base_pop_vec = np.mean(firing_rates, axis=1)
  376. tr_dist = cdist(np.reshape(base_pop_vec, (1,len(base_pop_vec))), firing_rates.T, 'euclidean')[0]
  377. return tr_dist
  378. def get_trace_tau(trace, sm_bin = 0):
  379. # intrinsic timescale of a trace = first autocorrelation lag (in frames) where it drops below 0.5; returns (tau, autocorr)
  380. #sm_bin = 10#round(1/params['dt'])*50;
  381. #trial_len = out_temp_all.shape[1]
  382. # z-score the trace so the zero-lag autocorrelation equals 1 (half-max crossing = 0.5)
  383. tracen = trace - np.mean(trace)
  384. trace_std = np.std(tracen)
  385. if trace_std == 0:
  386. return np.nan, np.full(len(trace), np.nan)
  387. tracen = tracen/trace_std
  388. corr1 = correlate(tracen, tracen)/len(tracen)
  389. #lags = correlation_lags(len(tracen), len(tracen))
  390. if sm_bin:
  391. kernel = np.ones(sm_bin)/sm_bin
  392. corr1_sm = np.convolve(corr1, kernel, mode='same')
  393. corr1_smn = corr1_sm - np.mean(corr1_sm)
  394. corr1_smn = corr1_smn/np.max(corr1_smn)
  395. else:
  396. corr1_smn = corr1
  397. corr1_smn2 = corr1_smn[len(trace)-1:] # positive lags only (zero lag is at index len-1)
  398. # plt.figure(); plt.plot(corr1)
  399. below = np.where(corr1_smn2 < 0.5)[0]
  400. tau_corr = below[0] if len(below) else np.nan
  401. # x = np.arange(corr_len)+1
  402. # y = corr1[num_trials2*num_run:num_trials2*num_run+corr_len]
  403. # yn = y - np.min(y)+0.01
  404. # yn = yn/np.max(yn)
  405. # fit = np.polyfit(x, np.log(yn), 1)
  406. # y_fit = np.exp(x*fit[0]+fit[1])
  407. # tau_corr = np.log(1/2)/fit[0]*params['dt']
  408. # x = np.random.rand(1000)
  409. # corrx = correlate(x, x)
  410. # plt.figure(); plt.plot(corrx)
  411. return tau_corr, corr1_smn2
  412. #%%
  413. def load_rnn_test(data_dir, fname_data, fname_params, max_net_load = 999, limit_network_types=[], flatten_runs = False, max_trial_types = 10, max_trials = 500, cut_zero_trials = False, num_initial_trials_skip = 0, seed=None):
  414. rng = seed if isinstance(seed, np.random.Generator) else np.random.default_rng(seed)
  415. # input is the spectrogram inputs
  416. # target is the index when oddball trial happens
  417. # loaded rates shape is (time, run, neurons) - inside converted to
  418. # output rates are converted to (runs, neurons, time)
  419. # flatten runs makes output rates 2D (neurons, time)
  420. # max trial types and max trials only apply during flattened runs
  421. # returns a list of dicts (one per network), matching the calcium-data format;
  422. # main field 'firing_rates' is (neurons, time) if flatten_runs else (runs, neurons, time)
  423. data_path = os.path.join(data_dir, fname_data)
  424. params_path = os.path.join(data_dir, fname_params)
  425. for p in (data_path, params_path):
  426. if not os.path.isfile(p):
  427. raise FileNotFoundError('RNN test file not found: %s' % p)
  428. if '_ob_data' in fname_data:
  429. print('Note: the oddball (_ob_data) RNN file is large (tens of GB) and is loaded fully into memory '
  430. 'before subsetting, so max_net_load does not reduce the peak memory. Ensure sufficient RAM, '
  431. 'or use the smaller control file (_cont_data).')
  432. test_data_load = np.load(data_path, allow_pickle=True).item()
  433. data_key = list(test_data_load.keys())[0]
  434. test_data_load_all = test_data_load[data_key]
  435. # for n1 in range(len(test_data_load_all)):
  436. # for n2 in range(len(test_data_load_all[n1])):
  437. # del test_data_load_all[n1][n2]['input']
  438. # del test_data_load_all[n1][n2]['target']
  439. # del test_data_load_all[n1][n2]['output']
  440. # test_data_load_all[0][0].keys()
  441. # test_data_load_all[0][0]['rates'].shape
  442. # np.save(data_dir + 'RNN_test_data_2024_5_24_9h_42m_ob_data2.npy', test_data_load)
  443. deets_load = np.load(params_path, allow_pickle=True).item()
  444. params_all = deets_load['params_all']
  445. params_test = deets_load['params_test']
  446. # deets_load['ob_data'].keys()
  447. # del deets_load['ob_data']['input_oddball']
  448. # del deets_load['ob_data']['target_oddball_freq']
  449. # del deets_load['cont_data']['input_control']
  450. # del deets_load['cont_data']['target_control']
  451. # np.save(data_dir + 'RNN_test_data_2024_5_24_9h_42m2_params.npy', deets_load)
  452. data_all = []
  453. if len(limit_network_types) == 0:
  454. limit_network_types = list(np.unique(deets_load['rnn_leg']))
  455. use_net_type_lim = True
  456. if 'ob' in data_key:
  457. is_ob = True
  458. else:
  459. is_ob = False
  460. # oddball tained, cont trained, and untrained
  461. for n_net in range(len(test_data_load_all)):
  462. num_rnn = np.min([len(test_data_load_all[n_net]), max_net_load])
  463. if use_net_type_lim:
  464. if deets_load['rnn_leg'][n_net] in limit_network_types:
  465. for n_rnn in range(num_rnn):
  466. net1 = test_data_load_all[n_net][n_rnn]
  467. T, num_runs, num_neurons = net1['rates'].shape
  468. # get stim times
  469. stim_times_runs = []
  470. trial_types_runs = []
  471. if is_ob:
  472. trial_types_red_dd_runs = []
  473. for n_run in range(num_runs):
  474. # deets_load['ob_data'].keys()
  475. # deets_load['ob_data']['trials_oddball_freq']
  476. if 1:
  477. stim_on_trace = 1 - deets_load['ob_data']['target_oddball_ctx3'][:,n_run,0]
  478. else:
  479. stim_trace = np.max(net1['input'][:,n_run], axis=1)
  480. stim_trace_n = stim_trace - np.percentile(stim_trace, 20)
  481. stim_trace_n = stim_trace_n / np.max(stim_trace_n)
  482. stim_on_trace = (stim_trace_n > 0.5).astype(int)
  483. stim_onset_trace = np.diff(stim_on_trace, prepend=[0]) > 0
  484. stim_times = np.where(stim_onset_trace)[0]
  485. stim_times_runs.append(stim_times)
  486. if is_ob:
  487. ob_deets = deets_load['ob_data']
  488. trial_idx = ob_deets['trials_oddball_ctx3'][:,n_run] > 0
  489. trials_types_red_dd = ob_deets['trials_oddball_ctx3'][:,n_run][trial_idx] - 1
  490. trials_types = ob_deets['trials_oddball_freq'][:, n_run][trial_idx]
  491. red_dd_seq = ob_deets['red_dd_seq']
  492. trial_types_red_dd_runs.append(trials_types_red_dd)
  493. else:
  494. cont_deets = deets_load['cont_data']
  495. trial_idx = cont_deets['trials_control_freq'][:,n_run] > 0
  496. trials_types = cont_deets['trials_control_freq'][:, n_run][trial_idx]
  497. trial_types_runs.append(trials_types)
  498. stim_times_runs2 = np.vstack(stim_times_runs)
  499. trial_types_runs2 = np.vstack(trial_types_runs)
  500. if is_ob:
  501. trial_types_red_dd_runs2 = np.vstack(trial_types_red_dd_runs)
  502. if cut_zero_trials or num_initial_trials_skip > 0:
  503. trial_len = round((params_test['stim_duration'] + params_test['isi_duration']) / params_test['dt'])
  504. if T % trial_len != 0:
  505. raise ValueError('RNN trace length T=%d is not divisible by trial_len=%d; cannot reshape into whole trials (check stim_duration/isi_duration/dt).' % (T, trial_len))
  506. rates3d = net1['rates']
  507. rates4d = np.reshape(rates3d, (round(T/trial_len), trial_len, num_runs, num_neurons), order='C')
  508. if cut_zero_trials:
  509. num_skip = params_test['num_prepend_zeros'] + num_initial_trials_skip
  510. else:
  511. num_skip = num_initial_trials_skip
  512. if num_skip >= rates4d.shape[0]:
  513. raise ValueError('num_initial_trials_skip (+ prepended zeros) = %d exceeds the %d available trials; reduce num_initial_trials_skip.' % (num_skip, rates4d.shape[0]))
  514. rates4d_cut = rates4d[num_skip:,:,:]
  515. rates2 = np.reshape(rates4d_cut, ((round(T/trial_len) - num_skip) * trial_len, num_runs, num_neurons), order='C')
  516. stim_times_runs2 = stim_times_runs2[:,num_skip:]
  517. trial_types_runs2 = trial_types_runs2[:,num_skip:]
  518. else:
  519. rates2 = net1['rates']
  520. T, num_runs, num_neurons = rates2.shape
  521. if flatten_runs:
  522. stim_times_runs_flat = (stim_times_runs2 + np.arange(num_runs)[:,None] * T).flatten()
  523. trial_types_runs_flat = trial_types_runs2.flatten()
  524. if is_ob:
  525. trial_types_red_dd_runs2 = trial_types_red_dd_runs2.flatten()
  526. rates = np.reshape(rates2, (T * num_runs, num_neurons), order='F').T
  527. trials_uq = np.unique(trial_types_runs_flat)
  528. num_tt = len(trials_uq)
  529. num_stim_all = len(stim_times_runs_flat)
  530. if max_trial_types < num_tt:
  531. step = num_tt/10
  532. trial_sel = np.arange(step/2, num_tt, step=step, dtype=int)
  533. sel_trial_idx = np.sum(trial_types_runs_flat[:,None] == trials_uq[trial_sel][None,:], axis=1).astype(bool)
  534. sel_trials = np.where(sel_trial_idx)[0]
  535. else:
  536. sel_trials = np.arange(num_stim_all)
  537. num_stim_all2 = len(sel_trials)
  538. # limit trials to max number
  539. if max_trials < num_stim_all:
  540. trial_idx = rng.choice(sel_trials, size=np.min([max_trials, num_stim_all2]), replace=False)
  541. trial_idx.sort()
  542. else:
  543. trial_idx = sel_trials
  544. stim_times_runs3 = stim_times_runs_flat[trial_idx]
  545. trial_types_runs3 = trial_types_runs_flat[trial_idx]
  546. if is_ob:
  547. trial_types_red_dd_runs3 = trial_types_red_dd_runs2[trial_idx]
  548. else:
  549. stim_times_runs3 = stim_times_runs2
  550. trial_types_runs3 = trial_types_runs2
  551. if is_ob:
  552. trial_types_red_dd_runs3 = trial_types_red_dd_runs2
  553. rates = np.transpose(rates2, (1, 2, 0))
  554. data_slice = {}
  555. data_slice['training'] = deets_load['rnn_leg'][n_net]
  556. data_slice['test_data_key'] = data_key
  557. data_slice['firing_rates'] = rates
  558. data_slice['stim_times'] = stim_times_runs3
  559. data_slice['trial_types'] = trial_types_runs3
  560. if is_ob:
  561. data_slice['trial_types_red_dd'] = trial_types_red_dd_runs3
  562. data_slice['red_dd_seq'] = red_dd_seq
  563. data_slice['num_runs'] = num_runs
  564. data_slice['volume_period'] = params_test['dt']*1000 # for frame rate equivalent
  565. data_slice['params_train'] = params_all[n_net][n_rnn]
  566. data_slice['params_test'] = params_test
  567. data_all.append(data_slice)
  568. del net1
  569. if len(data_all):
  570. trainings = [d['training'] for d in data_all]
  571. uq, cnt = np.unique(trainings, return_counts=True)
  572. print('Loaded %d RNN networks: %s' % (len(data_all), dict(zip(list(uq), [int(c) for c in cnt]))))
  573. else:
  574. print('Warning: no RNN networks loaded (file %s, limit_network_types=%s)' % (fname_data, limit_network_types))
  575. return data_all
  576. #%%
  577. def get_network_intrinsic_timescales(firing_rates, frame_rate):
  578. # per recording (or RNN run): network tau from the baseline-distance autocorrelation and per-neuron tau; both in seconds
  579. if len(firing_rates.shape) == 2:
  580. rates3d = firing_rates[None,:,:]
  581. else:
  582. rates3d = firing_rates
  583. num_runs, num_neurons, _ = rates3d.shape
  584. tau_net = np.zeros(num_runs)
  585. tau_cell = np.full(shape=(num_runs, num_neurons), fill_value=np.nan)
  586. for n_run in range(num_runs):
  587. tr_dist = get_network_distance(rates3d[n_run,:,:])
  588. tau_net1, _ = get_trace_tau(tr_dist, sm_bin = 0)
  589. tau_net[n_run] = tau_net1/frame_rate
  590. for n_nr in range(num_neurons):
  591. neur = rates3d[n_run,n_nr,:]
  592. if np.sum(neur) > 0.1:
  593. tau_neur1, _ = get_trace_tau(neur, sm_bin = 0)
  594. tau_cell[n_run, n_nr] = tau_neur1/frame_rate
  595. return tau_net, tau_cell
  596. #%%
  597. def plot_fig_raster(firing_rates_ob, firing_rates_rnn, frame_rate=None, frame_rate_rnn=None):
  598. # ---- plot example raster ----
  599. fig, ax = plt.subplots(1,2, figsize=(12,5), layout='constrained')
  600. fig.set_constrained_layout_pads(w_pad=0.2, h_pad=0.3, hspace=0, wspace=0)
  601. fig.text(0.015, .89, 'A', fontsize=18)
  602. fig.text(0.515, .89, 'B', fontsize=18)
  603. num_cells, num_t = firing_rates_ob.shape
  604. if frame_rate is not None:
  605. x_end = num_t/frame_rate
  606. x_lab = 'Time (sec)'
  607. else:
  608. x_end = num_t
  609. x_lab = 'Frames'
  610. ax[0].imshow(firing_rates_ob,
  611. aspect='auto',
  612. cmap='gist_yarg',
  613. vmin=0,
  614. vmax=.5,
  615. extent=[0, x_end, 1, num_cells],
  616. interpolation='none')
  617. ax[0].set_title('CaIm control data')
  618. ax[0].set_ylabel('Neurons')
  619. ax[0].set_xlabel(x_lab)
  620. num_cells_rnn, num_t_rnn = firing_rates_rnn.shape
  621. if frame_rate_rnn is not None:
  622. x_end = num_t_rnn/frame_rate_rnn
  623. x_lab = 'Time (sec)'
  624. else:
  625. x_end = num_t_rnn
  626. x_lab = 'Frames'
  627. ax[1].imshow(normalize(firing_rates_rnn),
  628. aspect='auto',
  629. cmap='gist_yarg',
  630. vmin=0,
  631. vmax=1,
  632. extent=[0, x_end, 1, num_cells_rnn],
  633. interpolation='none')
  634. ax[1].set_title('RNN control test data')
  635. ax[1].set_ylabel('Neurons')
  636. ax[1].set_xlabel(x_lab)
  637. return fig
  638. def plot_fig_diag_decoder(dec_data_all, dec_data_all_rnn, training_type, plot_t=None, plot_t_rnn=None, add_sig=True):
  639. # ---- Plotting diagonal decoder results ----
  640. # add_sig=True overlays binwise data-vs-shuffle significance (see plot_diag_binwise_dec)
  641. fig_diag, ax = plt.subplots(1,2, figsize=(12,5), layout='constrained')
  642. fig_diag.set_constrained_layout_pads(w_pad=0.2, h_pad=0.1, hspace=0, wspace=0)
  643. fig_diag.text(0.015, .87, 'A', fontsize=18)
  644. fig_diag.text(0.515, .87, 'B', fontsize=18)
  645. fig_diag.suptitle('Binwise decoder')
  646. plot_diag_binwise_dec(
  647. dec_data_all,
  648. plot_t=plot_t,
  649. plot_legend=('Data', 'Shuff'),
  650. plot_start=-1, # plot window start
  651. plot_end=3, # plot window end
  652. axis = ax[0],
  653. title_tag='CaIm data',
  654. add_sig=add_sig)
  655. plot_diag_binwise_dec(
  656. np.array(dec_data_all_rnn)[training_type == 'ob trained'],
  657. plot_t=plot_t_rnn,
  658. plot_start=-1,
  659. plot_end=8,
  660. axis=ax[1],
  661. title_tag='RNN test data',
  662. colors = ['limegreen', 'darkgreen'],
  663. add_sig=add_sig)
  664. plot_diag_binwise_dec(
  665. np.array(dec_data_all_rnn)[training_type == 'freq trained'],
  666. plot_t=plot_t_rnn,
  667. plot_start=-1,
  668. plot_end=8,
  669. axis=ax[1],
  670. colors = ['orange', 'saddlebrown'],
  671. add_sig=add_sig)
  672. # legend from the mean-trace lines only (exclude the '.-' binwise-significance lines)
  673. main_lines = [l for l in ax[1].get_lines() if l.get_marker() in ('', 'None', None)]
  674. ax[1].legend(main_lines, ['Oddball trained', 'Oddball trained shuff', 'Freq trained', 'Freq trained shuff'])
  675. return fig_diag
  676. def plot_fig_full_decoder_caim(dec_data_full, plot_t, clim=[0, 0.5]):
  677. # ---- Plotting full decoder results CaIm data ----
  678. fig_full, ax = plt.subplots(1, 3, gridspec_kw={'width_ratios': [20, 20, 1]}, figsize=(12, 5.7), layout='constrained')
  679. fig_full.set_constrained_layout_pads(w_pad=0.01, h_pad=0.1, hspace=0, wspace=0)
  680. fig_full.text(0.01, .88, 'A', fontsize=16)
  681. fig_full.text(0.47, .88, 'B', fontsize=16)
  682. fig_full.suptitle('CaIm data full space decoder', fontsize=14)
  683. plot_full_binwise_dec(
  684. dec_data_full,
  685. plot_t=plot_t,
  686. plot_legend=('Data', 'Shuff'),
  687. plot_start=-1,
  688. plot_end=2,
  689. clim=clim,
  690. axis=ax,
  691. title_tag='')
  692. return fig_full
  693. def plot_fig_full_decoder_RNN(dec_data_full_ob, dec_data_full_freq, plot_t):
  694. # ---- Plotting full decoder results RNN data ----
  695. fig_full, ax = plt.subplots(2, 3, gridspec_kw={'width_ratios': [20, 20, 1]}, figsize=(12,11), layout='constrained')
  696. fig_full.set_constrained_layout_pads(w_pad=0.05, h_pad=0.1, hspace=0, wspace=0)
  697. fig_full.text(0.005, .94, 'A', fontsize=16)
  698. fig_full.text(0.465, .94, 'B', fontsize=16)
  699. fig_full.text(0.005, .46, 'C', fontsize=16)
  700. fig_full.text(0.465, .46, 'D', fontsize=16)
  701. fig_full.suptitle('RNN neurons full space decoder', fontsize=14)
  702. plot_full_binwise_dec(
  703. dec_data_full_ob,
  704. plot_t=plot_t,
  705. plot_legend=('data', 'shuff'),
  706. plot_start=-1,
  707. plot_end=10,
  708. clim=[0, 1],
  709. axis=ax[0,:],
  710. title_tag='RNN Oddball trained')
  711. plot_full_binwise_dec(
  712. dec_data_full_freq,
  713. plot_t=plot_t,
  714. plot_legend=('data', 'shuff'),
  715. plot_start=-1,
  716. plot_end=10,
  717. clim=[0, 1],
  718. axis=ax[1,:],
  719. title_tag='RNN Control freq trained')
  720. return fig_full
  721. def plot_fig_full_decoder_comb(dec_data_full, dec_data_full_ob, dec_data_full_freq, plot_t, plot_t_rnn):
  722. # ---- plotting full decoder results ----
  723. fig_full, ax = plt.subplots(3, 3, gridspec_kw={'width_ratios': [10, 10, 1]}, figsize=(15, 12))
  724. fig_full.text(0.09, .88, 'A', fontsize=16)
  725. fig_full.text(0.51, .88, 'B', fontsize=16)
  726. plot_full_binwise_dec(
  727. dec_data_full,
  728. plot_t=plot_t,
  729. plot_legend=('Data', 'Shuff'),
  730. plot_start=-1,
  731. plot_end=2,
  732. clim=[0, 0.5],
  733. axis=ax[0,:],
  734. title_tag='CaIm data')
  735. plot_full_binwise_dec(
  736. dec_data_full_ob,
  737. plot_t=plot_t_rnn,
  738. plot_legend=('Oddball trained', 'Shuff'),
  739. plot_start=-1,
  740. plot_end=10,
  741. clim=[0, 1],
  742. axis=ax[1,:],
  743. title_tag='RNN Ob trained')
  744. plot_full_binwise_dec(
  745. dec_data_full_freq,
  746. plot_t=plot_t_rnn,
  747. plot_legend=('Freq trained', 'Shuff'),
  748. plot_start=-1,
  749. plot_end=10,
  750. clim=[0, 1],
  751. axis=ax[2,:],
  752. title_tag='RNN Freq trained')
  753. return fig_full
  754. #%%
  755. def plot_fig_isi_corr_trials(corr_vals, isi_list, colormap='jet', metric_tag = None):
  756. # trial-to-trial correlation vs ISI: panel A = mean over frequencies, panel B = per-frequency lines; returns (fig, ax, groups)
  757. fig, ax = plt.subplots(1,2, figsize=(12,5))
  758. fig.text(0.07, .88, 'A', fontsize=16)
  759. fig.text(0.51, .88, 'B', fontsize=16)
  760. idx_uq = np.unique(isi_list)
  761. num_trials = corr_vals.shape[1]
  762. col1 = plt.colormaps[colormap](np.linspace(0, 1, num_trials))
  763. corr_tn_all = np.zeros((num_trials, len(idx_uq)))
  764. for n_tn in range(num_trials):
  765. corr_tn = np.full(len(idx_uq), np.nan)
  766. for n_isi in range(len(idx_uq)):
  767. idx1 = (idx_uq[n_isi] == np.array(isi_list)).flatten()
  768. if np.sum(~np.isnan(corr_vals[idx1,n_tn])):
  769. corr_tn[n_isi] = np.nanmean(corr_vals[idx1,n_tn])
  770. corr_tn_all[n_tn, n_isi] = np.nanmean(corr_vals[idx1,n_tn])
  771. if np.sum(~np.isnan(corr_tn)):
  772. ax[1].plot(idx_uq, corr_tn, '-o', color=col1[n_tn])
  773. if metric_tag is not None:
  774. ax[1].set_ylabel(metric_tag)
  775. else:
  776. ax[1].set_ylabel('Correlation')
  777. ax[1].set_xlabel('ISI duration (sec)')
  778. freqs = np.logspace(np.log10(2), np.log10(76.9), num_trials) # kHz, log-spaced 2 -> 76.9 (x1.5/step)
  779. ax[1].legend(['%g' % round(f, 1) for f in freqs], loc='upper right', title='Freq (kHz)')
  780. ax[1].set_title('Individual freqs.')
  781. mean1 = np.nanmean(corr_tn_all, axis=0)
  782. sem1 = np.nanstd(corr_tn_all, axis=0)/np.sqrt(np.sum(~np.isnan(corr_tn_all), axis=0) - 1)
  783. ax[0].errorbar(idx_uq, mean1, sem1, fmt='-o', color='k', capsize=4)
  784. ax[0].set_ylabel('Correlation')
  785. ax[0].set_xlabel('ISI duration (sec)')
  786. ax[0].set_title('Average over freqs.')
  787. # per-ISI groups for stat_compare; each entry is (data_array, x-position = ISI value).
  788. # here we include every individual (dataset x frequency) correlation value at each ISI
  789. # (not the per-dataset mean), so all the individual points enter the test. Pass ax[0], e.g.
  790. # sd.stat_compare(ax[0], groups, 'ISI 0.5', 'ISI 1', test='mannwhitney', alternative='two-sided')
  791. groups = {}
  792. for isi in idx_uq:
  793. sel = (np.array(isi_list) == isi).flatten()
  794. d = corr_vals[sel, :].flatten() # all datasets x frequencies at this ISI
  795. groups['ISI %g' % isi] = (d[np.isfinite(d)], float(isi))
  796. return fig, ax, groups
  797. def plot_fig_isi_corr_trials2(corr_vals, isi_list, colormap='jet', metric_tag=None):
  798. # nicer version: panel A = individual (dataset x freq) points + mean line + shaded SEM;
  799. # panel B = per-frequency lines. Returns (fig, ax, groups) like plot_fig_isi_corr_trials.
  800. fig, ax = plt.subplots(1, 2, figsize=(12, 5))
  801. fig.text(0.07, .88, 'A', fontsize=16)
  802. fig.text(0.51, .88, 'B', fontsize=16)
  803. idx_uq = np.unique(isi_list)
  804. num_trials = corr_vals.shape[1]
  805. isi_arr = np.array(isi_list).flatten()
  806. rng = np.random.default_rng(0) # reproducible jitter
  807. # ---- panel B: per-frequency correlation vs ISI ----
  808. col1 = plt.colormaps[colormap](np.linspace(0, 1, num_trials))
  809. corr_tn_all = np.full((num_trials, len(idx_uq)), np.nan)
  810. for n_tn in range(num_trials):
  811. for n_isi in range(len(idx_uq)):
  812. vals = corr_vals[idx_uq[n_isi] == isi_arr, n_tn]
  813. if np.sum(~np.isnan(vals)):
  814. corr_tn_all[n_tn, n_isi] = np.nanmean(vals)
  815. ax[1].plot(idx_uq, corr_tn_all[n_tn, :], '-o', color=col1[n_tn], markersize=4, linewidth=1.2)
  816. ax[1].set_xlabel('ISI duration (sec)')
  817. ax[1].set_ylabel(metric_tag if metric_tag is not None else 'Correlation')
  818. ax[1].set_title('Individual freqs.')
  819. ax[1].legend([str(n + 1) for n in range(num_trials)], loc='upper right', fontsize=8, ncol=2, frameon=False)
  820. # ---- panel A: individual points + mean line + shaded SEM ----
  821. means = np.full(len(idx_uq), np.nan)
  822. sems = np.full(len(idx_uq), np.nan)
  823. for n_isi in range(len(idx_uq)):
  824. pts = corr_vals[idx_uq[n_isi] == isi_arr, :].flatten()
  825. pts = pts[np.isfinite(pts)]
  826. if len(pts) == 0:
  827. continue
  828. jit = (rng.random(len(pts)) - 0.5) * 0.08
  829. ax[0].plot(np.full(len(pts), idx_uq[n_isi]) + jit, pts, '.', color='0.7',
  830. markersize=4, alpha=0.5, zorder=1)
  831. means[n_isi] = np.mean(pts)
  832. sems[n_isi] = np.std(pts) / np.sqrt(len(pts) - 1) if len(pts) > 1 else np.nan
  833. ax[0].fill_between(idx_uq, means - sems, means + sems, color='steelblue', alpha=0.3, zorder=2)
  834. ax[0].plot(idx_uq, means, '-o', color='steelblue', markersize=6, linewidth=2, zorder=3)
  835. ax[0].set_xlabel('ISI duration (sec)')
  836. ax[0].set_ylabel(metric_tag if metric_tag is not None else 'Correlation')
  837. ax[0].set_title('Average over freqs.')
  838. for a in ax:
  839. a.spines['top'].set_visible(False)
  840. a.spines['right'].set_visible(False)
  841. # ---- groups (all individual dataset x freq points per ISI) for stat_compare ----
  842. groups = {}
  843. for isi in idx_uq:
  844. d = corr_vals[isi_arr == isi, :].flatten()
  845. groups['ISI %g' % isi] = (d[np.isfinite(d)], float(isi))
  846. return fig, ax, groups
  847. def plot_fig_SI_mat(SI_list, isi_uq, title_tag = ''):
  848. # plot the per-ISI trial-by-trial similarity matrices side by side with a shared colorbar
  849. fig, ax = plt.subplots(1,len(isi_uq)+1, gridspec_kw={'width_ratios': list(np.ones(len(isi_uq))*10) + [1]}, figsize=(12, 2.8))
  850. for n_isi in range(len(SI_list)):
  851. ax1 = ax.flatten()[n_isi]
  852. im = ax1.imshow(SI_list[n_isi], vmin=0, vmax=.7)
  853. ax1.set_title('isi %s sec' % (isi_uq[n_isi]))
  854. if title_tag is not None:
  855. fig.suptitle(title_tag)
  856. ax.flatten()[0].set_ylabel('trials')
  857. ax.flatten()[0].set_xlabel('trials')
  858. fig.colorbar(im, cax=ax.flatten()[-1])
  859. ax.flatten()[-1].set_ylabel('cosine similarity')
  860. return fig
  861. def plot_fig_tau_networks_comb3(tau_ob_net_all, tau_rnn_net_all, training_type, tau_ob_cell_all, tau_rnn_cell_all, do_log=True):
  862. # legacy two-panel network/neuron tau figure (superseded by plot_fig_tau_networks_comb)
  863. fig, ax, = plt.subplots(1, 2, sharey=True, figsize=(12,5))
  864. fig.text(0.075, .88, 'A', fontsize=16)
  865. fig.text(0.495, .88, 'B', fontsize=16)
  866. data_all = [np.array(tau_ob_net_all).flatten()] + tau_rnn_net_all
  867. labels_all = np.array(['Caim data'] + list(training_type))
  868. plot_int_violin2(data_all,
  869. net_labels = labels_all,
  870. axis=ax[0],
  871. points=1000,
  872. mean_std=True,
  873. showmeans=True,
  874. showmedians=False,
  875. quantile = [0.05, 0.95],
  876. colors=['blue', 'green', 'orange', 'gray'],
  877. do_log=True)
  878. tau_rnn_cell_all2 = []
  879. for n_net in range(len(tau_rnn_cell_all)):
  880. tau_rnn_cell_all2.append(np.nanmean(tau_rnn_cell_all[n_net], axis=0))
  881. data_cell_all = [np.hstack(tau_ob_cell_all).flatten()] + tau_rnn_cell_all2
  882. plot_int_violin2(data_cell_all,
  883. net_labels = labels_all,
  884. axis=ax[1],
  885. points=1000,
  886. mean_std=True,
  887. showmeans=True,
  888. showmedians=False,
  889. quantile = [0.05, 0.95],
  890. colors=['blue', 'green', 'orange', 'gray'],
  891. do_log=True)
  892. ax[1].yaxis.set_tick_params(labelleft=True)
  893. ax[0].set_title('Network tau')
  894. ax[1].set_title('Neuron tau')
  895. return fig
  896. def get_tau_groups(tau_ob_net_all, tau_rnn_net_all, training_type, tau_ob_cell_all, tau_rnn_cell_all, gap=1.0):
  897. # build the 8 tau groups (network block + neuron block) shared by the plot and the stats.
  898. # returns:
  899. # groups: dict {label: (data_array, x_position)} e.g. 'CaIm net', 'CaIm neuron', ...
  900. # meta: dict with 'labels', 'x_labels', 'positions', 'colors', 'n_net', 'data'
  901. tt = np.array(training_type)
  902. uq_types = list(dict.fromkeys(tt.tolist())) # RNN types in order of appearance
  903. short = {'ob trained': 'Ob', 'freq trained': 'Freq', 'untrained': 'Untr'}
  904. sh = lambda l: short.get(l, l)
  905. # network tau, one group per dataset type
  906. net_groups = [np.array(tau_ob_net_all).flatten()]
  907. net_names = ['CaIm']
  908. for u in uq_types:
  909. net_groups.append(np.hstack([np.asarray(tau_rnn_net_all[i]).flatten()
  910. for i in range(len(tau_rnn_net_all)) if tt[i] == u]))
  911. net_names.append(sh(u))
  912. # neuron tau (mean over runs per network, then pooled per type)
  913. tau_rnn_cell_mean = [np.nanmean(tau_rnn_cell_all[n], axis=0) for n in range(len(tau_rnn_cell_all))]
  914. neuron_groups = [np.hstack(tau_ob_cell_all).flatten()]
  915. for u in uq_types:
  916. neuron_groups.append(np.hstack([np.asarray(tau_rnn_cell_mean[i]).flatten()
  917. for i in range(len(tau_rnn_cell_mean)) if tt[i] == u]))
  918. data_all = [g[np.isfinite(g)] for g in (net_groups + neuron_groups)]
  919. labels = [n + ' net' for n in net_names] + [n + ' neuron' for n in net_names]
  920. x_labels = net_names + net_names
  921. base_colors = ['blue', 'green', 'orange', 'gray']
  922. colors = base_colors[:len(net_names)] + base_colors[:len(net_names)]
  923. n_net = len(net_names)
  924. positions = list(range(n_net)) + [p + n_net + gap for p in range(n_net)] # e.g. [0,1,2,3, 5,6,7,8]
  925. groups = {labels[i]: (data_all[i], positions[i]) for i in range(len(labels))}
  926. meta = {'labels': labels, 'x_labels': x_labels, 'positions': positions,
  927. 'colors': colors, 'n_net': n_net, 'data': data_all}
  928. return groups, meta
  929. def plot_fig_tau_networks_comb2(tau_ob_net_all, tau_rnn_net_all, training_type, tau_ob_cell_all, tau_rnn_cell_all, do_log=True):
  930. # all 8 groups (network block + neuron block) on one shared-y axis.
  931. # returns (fig, ax, groups); pass ax + groups to stat_compare() to add significance brackets.
  932. groups, meta = get_tau_groups(tau_ob_net_all, tau_rnn_net_all, training_type, tau_ob_cell_all, tau_rnn_cell_all)
  933. data_all = meta['data']
  934. positions = meta['positions']
  935. colors_all = meta['colors']
  936. x_labels = meta['x_labels']
  937. n_net = meta['n_net']
  938. num = len(data_all)
  939. fig, ax1 = plt.subplots(1, 1, figsize=(12, 5))
  940. parts = ax1.violinplot(data_all, positions=positions, showmeans=False, showextrema=False,
  941. quantiles=[[0.05, 0.95]] * num, points=1000)
  942. if 'cquantiles' in parts:
  943. parts['cquantiles'].set_color('k')
  944. for i in range(num):
  945. parts['bodies'][i].set_facecolor(colors_all[i])
  946. parts['bodies'][i].set_edgecolor(colors_all[i])
  947. for i in range(num):
  948. y = data_all[i]
  949. ax1.plot(positions[i], np.mean(y), '_', color='black', mew=2, markersize=30)
  950. ax1.errorbar(positions[i], np.mean(y), np.std(y), fmt='o', color='black', markersize=4, linewidth=2, capsize=8)
  951. ax1.set_xticks(positions)
  952. ax1.set_xticklabels(x_labels, rotation=0)
  953. if do_log:
  954. ax1.set_yscale('log')
  955. ax1.set_ylabel('Tau (sec)')
  956. ax1.set_title('Network and neuron intrinsic timescales')
  957. # group underlines + labels beneath (x in data coords, y in axes fraction)
  958. trans = ax1.get_xaxis_transform()
  959. net_pos = positions[:n_net]
  960. neuron_pos = positions[n_net:]
  961. for xs, name, col in [(net_pos, 'Networks', 'black'), (neuron_pos, 'Neurons', 'red')]:
  962. ax1.plot([xs[0] - 0.45, xs[-1] + 0.45], [-0.13, -0.13], transform=trans,
  963. color=col, lw=4, clip_on=False)
  964. ax1.text(np.mean(xs), -0.19, name, transform=trans, ha='center', va='top',
  965. fontsize=13, fontweight='bold', color=col)
  966. fig.subplots_adjust(bottom=0.26)
  967. return fig, ax1, groups
  968. def plot_fig_tau_networks_comb(tau_ob_net_all, tau_rnn_net_all, training_type, tau_ob_cell_all, tau_rnn_cell_all, do_log=True):
  969. # cleaner version of comb2: violin body + a slim inner boxplot (median + IQR + whiskers),
  970. # instead of the mean/SD dash + 5/95 quantile lines. Returns (fig, ax, groups) for stat_compare().
  971. groups, meta = get_tau_groups(tau_ob_net_all, tau_rnn_net_all, training_type, tau_ob_cell_all, tau_rnn_cell_all)
  972. data_all = meta['data']
  973. positions = meta['positions']
  974. colors_all = meta['colors']
  975. x_labels = meta['x_labels']
  976. n_net = meta['n_net']
  977. num = len(data_all)
  978. fig, ax1 = plt.subplots(1, 1, figsize=(12, 5))
  979. # violin bodies only (no extrema / quantile lines)
  980. parts = ax1.violinplot(data_all, positions=positions, showmeans=False, showextrema=False,
  981. showmedians=False, points=1000)
  982. for i in range(num):
  983. parts['bodies'][i].set_facecolor(colors_all[i])
  984. parts['bodies'][i].set_edgecolor(colors_all[i])
  985. parts['bodies'][i].set_alpha(0.55)
  986. # slim inner boxplot: median + IQR + whiskers, no outliers
  987. ax1.boxplot(data_all, positions=positions, widths=0.12, showfliers=False, patch_artist=True,
  988. medianprops=dict(color='black', linewidth=1.5),
  989. boxprops=dict(facecolor='white', edgecolor='black', linewidth=1.0),
  990. whiskerprops=dict(color='black', linewidth=1.0),
  991. capprops=dict(color='black', linewidth=1.0))
  992. ax1.set_xticks(positions)
  993. ax1.set_xticklabels(x_labels, rotation=0)
  994. if do_log:
  995. ax1.set_yscale('log')
  996. # plain-number tick labels (0.1, 1, 10, ...) instead of 10^0 exponential form
  997. ax1.yaxis.set_major_formatter(FuncFormatter(lambda v, _: ('%g' % v)))
  998. ax1.yaxis.set_minor_formatter(NullFormatter())
  999. ax1.set_ylabel('Tau (sec)')
  1000. ax1.set_title('Network and neuron intrinsic timescales')
  1001. # group underlines + labels beneath (x in data coords, y in axes fraction)
  1002. trans = ax1.get_xaxis_transform()
  1003. net_pos = positions[:n_net]
  1004. neuron_pos = positions[n_net:]
  1005. for xs, name, col in [(net_pos, 'Networks', 'black'), (neuron_pos, 'Neurons', 'red')]:
  1006. ax1.plot([xs[0] - 0.45, xs[-1] + 0.45], [-0.13, -0.13], transform=trans,
  1007. color=col, lw=4, clip_on=False)
  1008. ax1.text(np.mean(xs), -0.19, name, transform=trans, ha='center', va='top',
  1009. fontsize=13, fontweight='bold', color=col)
  1010. fig.subplots_adjust(bottom=0.26)
  1011. return fig, ax1, groups
  1012. def p_to_star(p):
  1013. # p-value to significance stars: *** <0.001, ** <0.01, * <0.05, else n.s.
  1014. if p < 0.001: return '***'
  1015. if p < 0.01: return '**'
  1016. if p < 0.05: return '*'
  1017. return 'n.s.'
  1018. def draw_sig_bracket(ax, x1, x2, p, star=True, color='black', y=None):
  1019. # draw a significance bracket + label between x1 and x2 in the space above the data,
  1020. # extending the y-axis to fit. Stacked calls space evenly (pixel gaps) on log or linear.
  1021. sig = p_to_star(p)
  1022. def _yoff(y_data, pts):
  1023. x0 = ax.get_xlim()[0]
  1024. yd = ax.transData.transform((x0, y_data))[1] + pts * ax.figure.dpi / 72.0
  1025. return ax.transData.inverted().transform((x0, yd))[1]
  1026. step_pts = 15 # spacing between stacked brackets
  1027. margin_pts = 24 # space above the topmost bracket/star and the plot edge
  1028. last = getattr(ax, '_sig_last_y', None)
  1029. if y is not None:
  1030. y_line = y
  1031. elif last is None:
  1032. y_line = _yoff(ax.get_ylim()[1], step_pts)
  1033. else:
  1034. y_line = _yoff(last, step_pts)
  1035. ax._sig_last_y = y_line
  1036. y_tick = _yoff(y_line, -5)
  1037. ax.set_ylim(top=_yoff(y_line, margin_pts))
  1038. ax.plot([x1, x1, x2, x2], [y_tick, y_line, y_line, y_tick], color=color, lw=1.5, clip_on=False)
  1039. # '*' renders high in its box (va='top'); text like 'n.s.'/'p=..' sits low (va='bottom')
  1040. label = sig if star else ('p=%.3g' % p)
  1041. if set(label) <= set('*'):
  1042. va, dy = 'top', 7
  1043. else:
  1044. va, dy = 'bottom', 3
  1045. ax.annotate(label, xy=((x1 + x2) / 2.0, y_line), xytext=(0, dy), textcoords='offset points',
  1046. ha='center', va=va, color=color, fontsize=12, annotation_clip=False)
  1047. def stat_compare(ax, groups, name1, name2, test='mannwhitney', alternative='two-sided',
  1048. stat_file=None, y=None, star=True, color='black'):
  1049. # compare two groups from get_tau_groups(); print result, append to a stats file,
  1050. # and draw a significance bracket on the same panel (ax) between the two groups.
  1051. # test: 'mannwhitney' (unpaired) or 'wilcoxon' (paired; needs equal-length data)
  1052. # alternative: 'two-sided', 'greater', or 'less' (name1 vs name2)
  1053. d1, x1 = groups[name1]
  1054. d2, x2 = groups[name2]
  1055. d1 = np.asarray(d1, dtype=float); d1 = d1[np.isfinite(d1)]
  1056. d2 = np.asarray(d2, dtype=float); d2 = d2[np.isfinite(d2)]
  1057. if test in ('wilcoxon', 'paired'):
  1058. if len(d1) != len(d2):
  1059. raise ValueError('wilcoxon needs paired equal-length data (got %d and %d); '
  1060. 'use test="mannwhitney" or pass paired per-dataset values' % (len(d1), len(d2)))
  1061. stat, p = wilcoxon(d1, d2, alternative=alternative)
  1062. n_str = 'n=%d pairs' % len(d1)
  1063. elif test in ('mannwhitney', 'mannwhitneyu', 'mwu'):
  1064. stat, p = mannwhitneyu(d1, d2, alternative=alternative)
  1065. n_str = 'n1=%d; n2=%d' % (len(d1), len(d2))
  1066. elif test in ('ttest', 'ttest_ind', 't'):
  1067. stat, p = ttest_ind(d1, d2, equal_var=False, alternative=alternative) # Welch's t-test
  1068. n_str = 'n1=%d; n2=%d' % (len(d1), len(d2))
  1069. elif test in ('ttest_rel', 'ttest_paired', 'paired_t'):
  1070. if len(d1) != len(d2):
  1071. raise ValueError('paired t-test needs equal-length data (got %d and %d)' % (len(d1), len(d2)))
  1072. stat, p = ttest_rel(d1, d2, alternative=alternative)
  1073. n_str = 'n=%d pairs' % len(d1)
  1074. else:
  1075. raise ValueError('unknown test: %s (use "mannwhitney", "wilcoxon", "ttest", or "ttest_rel")' % test)
  1076. sig = p_to_star(p)
  1077. print('%s vs %s | %s (%s) | %s | stat=%.4g | p=%.4g | %s'
  1078. % (name1, name2, test, alternative, n_str, stat, p, sig))
  1079. if stat_file is not None:
  1080. new = not os.path.isfile(stat_file)
  1081. with open(stat_file, 'a') as f:
  1082. if new:
  1083. f.write('group1,group2,test,alternative,n,statistic,p_value,significance\n')
  1084. f.write('%s,%s,%s,%s,%s,%.6g,%.6g,%s\n'
  1085. % (name1, name2, test, alternative, n_str, stat, p, sig))
  1086. # draw the significance bracket on the panel
  1087. draw_sig_bracket(ax, x1, x2, p, star=star, color=color, y=y)
  1088. return {'name1': name1, 'name2': name2, 'test': test, 'alternative': alternative,
  1089. 'statistic': stat, 'p': p, 'significance': sig}
  1090. def _bh_fdr(p):
  1091. # Benjamini-Hochberg FDR-adjusted p-values (matches MATLAB f_FDR_correction.m)
  1092. p = np.asarray(p, dtype=float)
  1093. n = p.size
  1094. order = np.argsort(p) # ascending
  1095. ranked = p[order] * n / (np.arange(n) + 1) # p * n / rank
  1096. ranked = np.minimum.accumulate(ranked[::-1])[::-1] # monotone from the top
  1097. ranked = np.minimum(ranked, 1.0)
  1098. out = np.empty(n)
  1099. out[order] = ranked
  1100. return out
  1101. def isi_stats(groups, ax=None, draw='consecutive', color='black', posthoc='tukey'):
  1102. # Stats for the ISI correlation figure: one-way ANOVA across ALL ISI groups + post-hoc.
  1103. # groups: dict {label: (data, x_position)} from plot_fig_isi_corr_trials.
  1104. # posthoc: 'tukey' (Tukey HSD) OR 'fisher_fdr' = MATLAB f_dv_plot_anova1: Fisher LSD using the
  1105. # pooled ANOVA error MSE, t with df=n1+n2-1, two-tailed, then Benjamini-Hochberg FDR.
  1106. # draw: None, 'consecutive' (0.5-1,1-2,2-4), 'first' (vs first group), or 'all'.
  1107. labels = list(groups)
  1108. data = [np.asarray(groups[k][0], dtype=float) for k in labels]
  1109. data = [d[np.isfinite(d)] for d in data]
  1110. ncat = len(data)
  1111. F, p_anova = f_oneway(*data)
  1112. print('One-way ANOVA across %d groups: F=%.4g, p=%.4g' % (ncat, F, p_anova))
  1113. if posthoc in ('fisher_fdr', 'fisher', 'matlab'):
  1114. from scipy.stats import t as _tdist
  1115. ns = np.array([len(d) for d in data])
  1116. means = np.array([d.mean() for d in data])
  1117. N = int(ns.sum())
  1118. MSE = float(sum(((d - d.mean()) ** 2).sum() for d in data)) / (N - ncat) # pooled ANOVA error
  1119. pairs = [(i, j) for i in range(ncat) for j in range(i + 1, ncat)]
  1120. praw = np.array([2 * _tdist.sf(abs((means[i] - means[j]) / np.sqrt(MSE / ns[i] + MSE / ns[j])),
  1121. ns[i] + ns[j] - 1) for (i, j) in pairs])
  1122. padj = _bh_fdr(praw)
  1123. pmat = np.ones((ncat, ncat))
  1124. for k, (i, j) in enumerate(pairs):
  1125. pmat[i, j] = pmat[j, i] = padj[k]
  1126. method = 'Fisher LSD (pooled MSE) + Benjamini-Hochberg FDR'
  1127. else:
  1128. try:
  1129. from scipy.stats import tukey_hsd
  1130. pmat = np.asarray(tukey_hsd(*data).pvalue)
  1131. method = 'Tukey HSD'
  1132. except Exception:
  1133. m = ncat * (ncat - 1) // 2
  1134. pmat = np.ones((ncat, ncat))
  1135. for i in range(ncat):
  1136. for j in range(i + 1, ncat):
  1137. _, pp = ttest_ind(data[i], data[j], equal_var=False)
  1138. pmat[i, j] = pmat[j, i] = min(pp * m, 1.0)
  1139. method = 'pairwise Welch t-test + Bonferroni'
  1140. print('Post-hoc (%s):' % method)
  1141. for i in range(len(labels)):
  1142. for j in range(i + 1, len(labels)):
  1143. print(' %s vs %s: p=%.4g %s' % (labels[i], labels[j], pmat[i, j], p_to_star(pmat[i, j])))
  1144. if ax is not None and draw:
  1145. if draw == 'consecutive':
  1146. pairs = [(i, i + 1) for i in range(len(labels) - 1)]
  1147. elif draw == 'first':
  1148. pairs = [(0, j) for j in range(1, len(labels))]
  1149. elif draw == 'all':
  1150. pairs = [(i, j) for i in range(len(labels)) for j in range(i + 1, len(labels))]
  1151. else:
  1152. pairs = []
  1153. for (i, j) in pairs:
  1154. draw_sig_bracket(ax, groups[labels[i]][1], groups[labels[j]][1], pmat[i, j], color=color)
  1155. return {'labels': labels, 'F': F, 'p_anova': p_anova, 'pvalue': pmat, 'method': method}
  1156. def plot_fig_tau_networks(tau_ob_net_all, tau_rnn_net_all, training_type):
  1157. # two-panel network-tau figure (CaIm | RNN by training type) drawn as violins
  1158. fig, ax, = plt.subplots(1, 2, sharey=True, gridspec_kw={'width_ratios': [1, 3]}, figsize=(6,5))
  1159. fig.text(0.02, .89, 'A', fontsize=16)
  1160. fig.text(0.32, .89, 'B', fontsize=16)
  1161. fig.suptitle('Network Tau')
  1162. plot_int_violin2(tau_ob_net_all,
  1163. net_labels = None,
  1164. title_tag = 'CaIm',
  1165. axis=ax[0],
  1166. points=1000,
  1167. mean_std=True,
  1168. showmeans=True,
  1169. showmedians=False,
  1170. quantile = [0.05, 0.95],
  1171. colors=['blue', 'magenta', 'green'],
  1172. do_log=True)
  1173. plot_int_violin2(tau_rnn_net_all,
  1174. net_labels = training_type,
  1175. title_tag = 'RNN',
  1176. axis=ax[1],
  1177. points=1000,
  1178. mean_std=True,
  1179. showmeans=True,
  1180. showmedians=False,
  1181. quantile = [0.05, 0.95],
  1182. colors=['blue', 'magenta', 'green'],
  1183. do_log=True)
  1184. ax[1].set_ylabel(None)
  1185. return fig
  1186. def plot_fig_tau_networks2(tau_ob_net_all, tau_rnn_net_all, training_type):
  1187. # single-panel network-tau violins (CaIm data + each RNN training type on one axis)
  1188. fig, ax, = plt.subplots(1, 1, figsize=(6,5))
  1189. fig.text(0.02, .89, 'A', fontsize=16)
  1190. fig.text(0.32, .89, 'B', fontsize=16)
  1191. fig.suptitle('Network Tau')
  1192. data_all = [np.array(tau_ob_net_all).flatten()] + tau_rnn_net_all
  1193. labels_all = np.array(['Caim data'] + list(training_type))
  1194. plot_int_violin2(data_all,
  1195. net_labels = labels_all,
  1196. axis=ax,
  1197. points=1000,
  1198. mean_std=True,
  1199. showmeans=True,
  1200. showmedians=False,
  1201. quantile = [0.05, 0.95],
  1202. colors=['blue', 'green', 'orange', 'gray'],
  1203. do_log=True)
  1204. return fig
  1205. #%%
  1206. def plot_cat_data2(y_data_in, rnn_leg, title_tag = '', do_log=False):
  1207. # scatter of each category's values with a mean +/- std marker
  1208. num_cat = len(y_data_in)
  1209. plt.figure()
  1210. ax1 = plt.subplot(111)
  1211. ax1.bar(rnn_leg, np.zeros(num_cat))
  1212. for n_net in range(num_cat):
  1213. y_data = np.concatenate(y_data_in[n_net])
  1214. x_data = ((np.random.rand(len(y_data)))-0.5)/5+n_net
  1215. ax1.plot(x_data, y_data, '.', color='gray')
  1216. ax1.plot(n_net, np.mean(y_data), '_', color='black', mew=2, markersize=40)
  1217. ax1.errorbar(n_net, np.mean(y_data), np.std(y_data), fmt='o', color='black', mew=2, markersize=5, linewidth=2, capsize=10)
  1218. ax1.set_title(title_tag)
  1219. if do_log:
  1220. ax1.set_yscale('log')
  1221. def plot_cat_data_violin(y_data_in, rnn_leg, title_tag = '', points=100, mean_std=True, showmeans=False, showmedians=False, quantile = [], colors=[], do_log=False):
  1222. # violin plot of categorical groups, optional mean +/- std overlay and log y-axis
  1223. num_cat = len(y_data_in)
  1224. plt.figure()
  1225. ax1 = plt.subplot(111)
  1226. ax1.bar(rnn_leg, np.zeros(num_cat))
  1227. parts = ax1.violinplot(y_data_in, positions=range(num_cat), showmeans=showmeans, showextrema=False, showmedians=showmedians, quantiles=[quantile, quantile, quantile], points=points)
  1228. for key in ['cmeans', 'cmedians', 'cquantiles']:
  1229. if key in parts:
  1230. parts[key].set_color('k')
  1231. if len(colors):
  1232. for n_net in range(num_cat):
  1233. pc = parts['bodies'][n_net]
  1234. pc.set_facecolor(colors[n_net])
  1235. pc.set_edgecolor(colors[n_net])
  1236. if mean_std:
  1237. for n_net in range(num_cat):
  1238. y_data = y_data_in[n_net]
  1239. ax1.plot(n_net, np.mean(y_data), '_', color='black', mew=2, markersize=40)
  1240. ax1.errorbar(n_net, np.mean(y_data), np.std(y_data), fmt='o', color='black', mew=2, markersize=5, linewidth=2, capsize=10)
  1241. ax1.set_title(title_tag)
  1242. if do_log:
  1243. ax1.set_yscale('log')
  1244. def plot_cat_data_bar(y_data_in, rnn_leg, title_tag = '', do_sem=True, colors=[]):
  1245. # bar plot (mean +/- sem or std) of categorical groups
  1246. num_cat = len(y_data_in)
  1247. plt.figure()
  1248. plt.bar(rnn_leg, np.zeros(num_cat))
  1249. for n_net in range(num_cat):
  1250. y_data = y_data_in[n_net]
  1251. if do_sem:
  1252. stds = np.std(y_data)/np.sqrt(len(y_data)-1)
  1253. else:
  1254. stds = np.std(y_data)
  1255. if len(colors):
  1256. plt.bar(n_net, np.mean(y_data), color=colors[n_net], alpha=0.5, edgecolor=colors[n_net])
  1257. else:
  1258. plt.bar(n_net, np.mean(y_data))
  1259. plt.plot(n_net, np.mean(y_data), '_', color='black', mew=2, markersize=40)
  1260. plt.errorbar(n_net, np.mean(y_data), stds, fmt='o', color='black', mew=2, markersize=5, linewidth=2, capsize=10)
  1261. plt.title(title_tag)
  1262. def plot_int_violin2(tau_net_list, net_labels = None, data_lab = ['CaIm data'], title_tag = '', axis=None, points=100, mean_std=True, showmeans=False, showmedians=False, quantile = [0.05, 0.95], colors=['blue', 'green', 'orange', 'gray'], do_log=False):
  1263. # violin plot of tau distributions grouped by net_labels, with optional mean +/- std overlay and log y-axis
  1264. num_net = len(tau_net_list)
  1265. net_type_list = []
  1266. net_idx = np.zeros(num_net, dtype=int)
  1267. rnn_leg = []
  1268. quantiles_all = []
  1269. if net_labels is not None:
  1270. _, idx1 = np.unique(net_labels, return_index=True)
  1271. idx1.sort()
  1272. net_uq = net_labels[idx1]
  1273. for n_net in range(len(net_uq)):
  1274. net_idx[net_labels == net_uq[n_net]] = n_net
  1275. temp_bin = []
  1276. for n_net2 in range(num_net):
  1277. if net_uq[n_net] == net_labels[n_net2]:
  1278. temp_bin.append(tau_net_list[n_net2].flatten())
  1279. net_type_list.append(np.hstack(temp_bin).flatten())
  1280. rnn_leg.append(net_uq[n_net].capitalize())
  1281. quantiles_all.append(quantile)
  1282. else:
  1283. net_uq = data_lab
  1284. net_type_list.append(np.array(tau_net_list).flatten())
  1285. quantiles_all.append(quantile)
  1286. rnn_leg = data_lab
  1287. num_net_types = len(net_uq)
  1288. if axis is None:
  1289. plt.figure()
  1290. ax1 = plt.subplot(111)
  1291. else:
  1292. ax1 = axis
  1293. ax1.bar(rnn_leg, np.zeros(num_net_types))
  1294. parts = ax1.violinplot(net_type_list, positions=range(num_net_types), showmeans=showmeans, showextrema=False, showmedians=showmedians, quantiles=quantiles_all, points=points)
  1295. for key in ['cmeans', 'cmedians', 'cquantiles']:
  1296. if key in parts:
  1297. parts[key].set_color('k')
  1298. if len(colors):
  1299. for n_net in range(num_net_types):
  1300. pc = parts['bodies'][n_net]
  1301. pc.set_facecolor(colors[n_net])
  1302. pc.set_edgecolor(colors[n_net])
  1303. if mean_std:
  1304. for n_net in range(num_net_types):
  1305. y_data = net_type_list[n_net]
  1306. ax1.plot(n_net, np.mean(y_data), '_', color='black', mew=2, markersize=40)
  1307. ax1.errorbar(n_net, np.mean(y_data), np.std(y_data), fmt='o', color='black', mew=2, markersize=5, linewidth=2, capsize=10)
  1308. ax1.set_title(title_tag)
  1309. if do_log:
  1310. ax1.set_yscale('log')
  1311. ax1.set_ylabel('Tau (sec)')
  1312. else:
  1313. ax1.set_ylabel('Tau (sec)')

sd_utils.py at commit 53e1dfd, no license · at the source

Overview

Authors: Yuriy Shymkiv1, Rafael Yuste1
ORCID iDs: Yuriy Shymkiv
  1. Neurotechnology Center, Department of Biological Sciences, Columbia University, New York, NY, USA
Institutions: Columbia University (United States)
Journal: STAR protocols, volume 7, issue 3, article 104804
Dates: published online 27 August 2026; in print August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1016/j.xpro.2026.104804 · PMID 42658683 · PMCID PMC13544393 · OpenAlex W7204494996
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: mouse (organism)
Methods: Statistics, Smoothing, state filtering, decompositions, Machine learning, Preprocessing, Evoked potentials, Connectivity, Single-unit activity, calcium imaging
Keywords: Neuroscience, Cognitive Neuroscience, Computer sciences
Topic: Neural dynamics and brain function (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Citations: not cited yet (Europe PMC); 31 references in the paper
Research resources: Scikit-learn (v1.7.1) RRID:SCR_002577, SciPy (v1.16.3) RRID:SCR_008058, Python 3 (v3.12.12) RRID:SCR_008394, Matplotlib (v3.10.7) RRID:SCR_008624, NumPy (v2.3.5) RRID:SCR_008633, Jupyter Notebook (v7.5.0) RRID:SCR_018315, Miniconda RRID:SCR_018317, h5py (v3.15.1) RRID:SCR_024812

Abstract

Temporal information processing is critical for brain function, supporting neural computations such as novelty detection, adaptation, and temporal normalization. Its disruption is implicated in schizophrenia. We present a protocol for analyzing ongoing neuronal network activity using binwise decoding, trial-to-trial variability analysis, and estimation of network-intrinsic timescales (INTs). We apply these techniques to identify slow dynamics that encode the memory of recent stimuli in neuronal populations in the mouse auditory cortex and in artificial neural networks trained on a novelty-detection task.

For complete details on the use and execution of this protocol, please refer to Shymkiv et al.1

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

Repositories

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

shymkivy/slow_dynamics_protocol

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 53e1dfd5b9a9b053e0e8ad5c13f815a5861b789b, 10 July 2026
Languages: Jupyter (11), Python (7)
Size: 22 files, 18 scripts
Software Heritage: not archived
Found in: “Data and code availability”
Holds: README, 6 notebooks
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (18 files), Matplotlib (12 files), SciPy (3 files), h5py (1 file), scikit-learn (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
19 files

doi:10.5061/dryad.xsj3tx9q6

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: the link answers
Software Heritage: not checked
Found in: “Data and code availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)

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:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 18 scripts, each with its path and the digest of its content;
  • 30 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 and code availability

• The calcium imaging and RNN datasets analyzed in this protocol are publicly available at Dryad: https://doi.org/10.5061/dryad.xsj3tx9q6. • All original analysis code is publicly available on GitHub: https://github.com/shymkivy/slow_dynamics_protocol. • Any additional information required to reanalyze the data reported in this paper is available from the lead contact upon request.

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

Materials availability

This study did not generate new unique reagents or materials. All software and datasets used are listed in the key resources table.

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

Versions

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

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 2 authors, 3 keywords, 3 funders, 28 references, 8 RRIDs.

Cite

This paper

Shymkiv, Y., & Yuste, R. (2026). Protocol for analyzing slow cortical dynamics in mouse neuronal recordings. STAR protocols, 7(3), 104804. https://doi.org/10.1016/j.xpro.2026.104804

BibTeX

@article{shymkiv2026protocol,
author = {Shymkiv, Yuriy and Yuste, Rafael},
title = {{Protocol for analyzing slow cortical dynamics in mouse neuronal recordings}},
journal = {STAR protocols},
year = {2026},
month = aug,
volume = {7},
number = {3},
pages = {104804},
publisher = {Elsevier},
issn = {2666-1667},
doi = {10.1016/j.xpro.2026.104804},
url = {https://doi.org/10.1016/j.xpro.2026.104804},
pmid = {42658683},
pmcid = {PMC13544393}
}

RIS

TY - JOUR
AU - Shymkiv, Yuriy
AU - Yuste, Rafael
TI - Protocol for analyzing slow cortical dynamics in mouse neuronal recordings
T2 - STAR protocols
J2 - STAR Protoc
PY - 2026
DA - 2026/08/27
VL - 7
IS - 3
SP - 104804
SN - 2666-1667
PB - Elsevier
DO - 10.1016/j.xpro.2026.104804
UR - https://doi.org/10.1016/j.xpro.2026.104804
LA - en
ER -

CSL-JSON

{
"id": "10.1016/j.xpro.2026.104804",
"type": "article-journal",
"title": "Protocol for analyzing slow cortical dynamics in mouse neuronal recordings",
"container-title": "STAR protocols",
"author": [
{
"family": "Shymkiv",
"given": "Yuriy"
},
{
"family": "Yuste",
"given": "Rafael"
}
],
"container-title-short": "STAR Protoc",
"volume": "7",
"issue": "3",
"page": "104804",
"DOI": "10.1016/j.xpro.2026.104804",
"PMID": "42658683",
"PMCID": "PMC13544393",
"ISSN": "2666-1667",
"publisher": "Elsevier",
"URL": "https://doi.org/10.1016/j.xpro.2026.104804",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
27
]
]
}
}

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

Similar papers

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

[1] doi:10.1126/sciadv.aed6417 [code]
Intrinsic timing, not temporal prediction, underlies ramping dynamics in visual and parietal cortex during passive behavior.
Journal: Science advances
In common: h5py, scikit-learn, SciPy, 2 other tools, mouse, 3 references
[2] doi:10.1371/journal.pcbi.1014162 [code]
Exploring neural manifolds across a wide range of intrinsic dimensions.
Journal: PLoS computational biology
In common: h5py, scikit-learn, SciPy, 2 other tools, 3 references
[3] doi:10.1016/j.crmeth.2026.101421 [code]
EthoPy provides an accessible platform for reproducible behavioral neuroscience.
Journal: Cell reports methods
In common: h5py, scikit-learn, SciPy, 2 other tools, mouse, 3 references
[4] doi:10.1038/s41467-026-74816-0 [code]
Hierarchical optimization predicts plasticity in the macaque inferior temporal cortex following object training.
Journal: Nature communications
In common: h5py, scikit-learn, SciPy, 2 other tools, 3 references
[5] doi:10.3390/biomimetics11080569 [code]
Pretraining of Embodied Recurrent Networks Bridges the Gap Between Artificial and Cortical Neural Activities.
Journal: Biomimetics (Basel, Switzerland)
In common: scikit-learn, SciPy, Matplotlib, 1 other tool, 4 references
[6] doi:10.1038/s41593-026-02333-w [code]
Learning shapes neural geometry in the primate prefrontal cortex.
Journal: Nature neuroscience
In common: scikit-learn, SciPy, Matplotlib, 1 other tool, 3 references
[7] doi:10.1038/s41467-026-76106-1 [code]
Facial expression discrimination emerges from partially overlapping neural subspaces of detection and identity.
Journal: Nature communications
In common: h5py, scikit-learn, SciPy, 2 other tools, 2 references
[8] doi:10.1038/s41467-026-75924-7 [code]
Data-driven reduced modeling of neural dynamics.
Journal: Nature communications
In common: h5py, scikit-learn, SciPy, 2 other tools, 2 references
[9] doi:10.1016/j.isci.2026.117492 [code]
Neural subspace reorganization reflects value-based decision-making.
Journal: iScience
In common: scikit-learn, SciPy, Matplotlib, 1 other tool, 3 references
[10] doi:10.1038/s41467-026-74347-8 [code]
Compositionality of social gaze in the prefrontal-amygdala circuits.
Journal: Nature communications
In common: scikit-learn, SciPy, Matplotlib, 1 other tool, 3 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.