OSCR

Shallow recurrent decoders for neural and behavioural dynamics.

Code ↔ Paper

5 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 5 matches
  1. [1] § Model systems › Mouse › SHRED reconstruction of neural response to visual stimuli ↔ mice/shred_mice_expB_i_flash.ipynb, lines 560–629 · score 0.59 · SCig, MB, POL, VISp, APN, POST
  2. [2] § Model systems › Mouse › SHRED reconstruction of neural response to visual stimuli ↔ mice/shred_mice_expB_i_dg.ipynb, lines 589–658 · score 0.59 · SCig, MB, POL, VISp, APN, POST
  3. [3] § Model systems › Mouse › SHRED reconstruction of local field potential from pupil area response to a visual stimulus ↔ mice/shred_mice_expB_iv.ipynb, lines 125–150 · score 0.56 · cubic spline, pupil area, stimuli
  4. [4] § Model systems › Control comparison ↔ mice/shred_mice_expB_i_dg.ipynb, lines 756–826 · score 0.55 · power spectral density, linear regression, PSD, neurons, SHRED
  5. [5] § Model systems › Control comparison ↔ mice/shred_mice_expB_i_dg.ipynb, lines 756–826 · score 0.52 · Power spectral densities, linear regression, ground truth, PSDs, neurons, SHRED

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 · 1,010 lines · 35 KB · no license · 3 matches

  1. # %% [markdown]
  2. # # Applying SHRED to the Mice Dataset
  3. # %% [markdown]
  4. # Dataset obtained from Allen Brain Atlas https://allensdk.readthedocs.io/en/latest/visual_coding_neuropixels.html
  5. # %% [markdown]
  6. # SHRED (SHallow REcurrent Decoder) models are a network architecture that merges a recurrent layer (LSTM) with a shallow decoder network (SDN) to reconstruct high-dimensional spatio-temporal fields from a trajectory of sensor measurements of the field. More formally, the SHRED architecture can be written as
  7. # $$ \mathcal {H} \left( \{ y_i \} _{i=t-k}^t \right) = \mathcal {F} \left( \mathcal {G} \left( \{ y_i \} _{i=t-k}^t \right) ; W_{RN}) ; W_{SD} \right)$$
  8. # where $\mathcal F$ is a feed forward network parameterized by weights $W_{SD}$, $\mathcal G$ is a LSTM network parameterized by weights $W_{RN}$, and $\{ y_i \} _{i=t-k}^t$ is a trajectory of sensor measurements of a high-dimensional spatio-temporal field $\{ x_i \} _{i=t-k}^t$.
  9. # %% [markdown]
  10. # Goal: Reconstruct LFP potential of different brain regions from only one session 756029989 and 180 degree orientation drifting gratings
  11. #
  12. # LFP band:
  13. # - ~2.5 kHz original sample rate
  14. # - 1000 Hz analog lo-pass
  15. # - 625 Hz digital lo-pass
  16. # - NWB includes every 2nd sample and every 4th channel
  17. # %%
  18. # Importing all dependencies
  19. import numpy as np
  20. import torch
  21. import subprocess
  22. import os
  23. import torch
  24. import matplotlib.pyplot as plt
  25. from sklearn.preprocessing import MinMaxScaler
  26. import random
  27. import mne
  28. # %%
  29. import pandas as pd
  30. from scipy.ndimage import gaussian_filter
  31. from pathlib import Path
  32. import json
  33. from IPython.display import display
  34. from PIL import Image
  35. from allensdk.brain_observatory.ecephys.visualization import plot_mean_waveforms, plot_spike_counts, raster_plot
  36. from allensdk.brain_observatory.visualization import plot_running_speed
  37. # tell pandas to show all columns when we display a DataFrame
  38. pd.set_option("display.max_columns", None)
  39. # %%
  40. #Set up data cache
  41. from allensdk.brain_observatory.ecephys.ecephys_project_cache import EcephysProjectCache
  42. HOME_DIR = "/Users/amyrude/Downloads/Kutz_Research/SHRED_neuro/neuralSHRED/"
  43. data_directory = os.path.join(HOME_DIR, "mice/data") #where data will be stored
  44. manifest_path = os.path.join(data_directory, 'manifest.json')
  45. cache = EcephysProjectCache.from_warehouse(manifest = manifest_path)
  46. # %%
  47. sessions = cache.get_session_table()
  48. print('Total number of sessions: ' + str(len(sessions)))
  49. sessions.head()
  50. # Can also filter which mice to select
  51. filtered_sessions = sessions[(sessions.index == 756029989)]
  52. filtered_sessions.head()
  53. # %%
  54. #Loading dataset for specific session
  55. session_id = 756029989 # for example
  56. session = cache.get_session_data(session_id)
  57. print([attr_or_method for attr_or_method in dir(session) if attr_or_method[0] != '_'])
  58. # %% [markdown]
  59. # **Importing local field potential (LFP) data**
  60. # %%
  61. # list the probes recorded from in this session
  62. session.probes.head()
  63. # %%
  64. {session.probes.loc[probe_id].description :
  65. list(session.channels[session.channels.probe_id == probe_id].ecephys_structure_acronym.unique())
  66. for probe_id in session.probes.index.values}
  67. # %%
  68. probe_id = session.probes.index.values[2] ### Select 0 to probe A, 2 for probe C
  69. lfp = session.get_lfp(probe_id)
  70. # %%
  71. print(lfp)
  72. # %%
  73. # now use a utility to associate intervals of /rows with structures
  74. ## For probe C, mismatch between channels in lfp["channel"] and session.channels.index
  75. lfp_channels = lfp["channel"].values.tolist()
  76. valid_channels = [ch for ch in lfp_channels if ch in session.channels.index]
  77. structure_acronyms, intervals = session.channel_structure_intervals(valid_channels)
  78. interval_midpoints = [aa + (bb - aa) / 2 for aa, bb in zip(intervals[:-1], intervals[1:])]
  79. print(structure_acronyms)
  80. print(intervals)
  81. print(interval_midpoints)
  82. # %% [markdown]
  83. # Aligning the LFP data to a particular stimulus
  84. # %%
  85. stim_table = session.get_stimulus_table('drifting_gratings')
  86. stim_table[stim_table['stimulus_condition_id'] == 246].head(10)
  87. # %%
  88. # presentation_table = session.stimulus_presentations[session.stimulus_presentations.stimulus_name == 'drifting_gratings']
  89. presentation_table = session.stimulus_presentations[session.stimulus_presentations.stimulus_condition_id == 246]
  90. presentation_times = presentation_table.start_time.values
  91. presentation_ids = presentation_table.index.values
  92. # %%
  93. print(presentation_table)
  94. # %%
  95. # %%
  96. ### Probe A and C
  97. dt = 1/500
  98. sr = 1/dt
  99. trial_window = np.arange(0, 2, dt)
  100. time_selection = np.concatenate([trial_window + t for t in presentation_times])
  101. inds = pd.MultiIndex.from_product((presentation_ids, trial_window),
  102. names=('presentation_id', 'time_from_presentation_onset'))
  103. ds = lfp.sel(time = time_selection, method='nearest').to_dataset(name = 'aligned_lfp')
  104. ds = ds.assign(time=inds).unstack('time')
  105. aligned_lfp = ds['aligned_lfp']
  106. # %%
  107. print(aligned_lfp.shape)
  108. # %%
  109. fig, ax = plt.subplots()
  110. im = ax.imshow(aligned_lfp.mean(dim='presentation_id'), aspect='auto', origin='lower', vmin=-1e-4, vmax=1e-4)
  111. ax.set_yticks(intervals)
  112. ax.set_yticks(interval_midpoints, minor=True)
  113. ax.set_yticklabels(structure_acronyms, minor=True)
  114. plt.tick_params("y", which="major", labelleft=False, length=40)
  115. num_time_labels = 8
  116. time_label_indices = np.around(np.linspace(1, len(trial_window), num_time_labels)).astype(int) - 1
  117. # time_labels = [ f"{val:1.3}" for val in lfp["time"].values[window][time_label_indices]]
  118. # ax.set_xticks(time_label_indices + 0.5)
  119. # ax.set_xticklabels(time_labels)
  120. ax.set_xlabel("Sample Number", fontsize=20)
  121. plt.colorbar(im, fraction=0.036, pad=0.04)
  122. plt.show()
  123. ## just one
  124. fig, ax = plt.subplots()
  125. im = ax.imshow(aligned_lfp[:,0,:], aspect='auto', origin='lower', vmin=-0.5e-3, vmax=0.5e-3)
  126. ax.set_yticks(intervals)
  127. ax.set_yticks(interval_midpoints, minor=True)
  128. ax.set_yticklabels(structure_acronyms, minor=True)
  129. plt.tick_params("y", which="major", labelleft=False, length=40)
  130. num_time_labels = 8
  131. time_label_indices = np.around(np.linspace(1, len(trial_window), num_time_labels)).astype(int) - 1
  132. # time_labels = [ f"{val:1.3}" for val in lfp["time"].values[window][time_label_indices]]
  133. # ax.set_xticks(time_label_indices + 0.5)
  134. # ax.set_xticklabels(time_labels)
  135. ax.set_xlabel("Sample Number", fontsize=20)
  136. plt.colorbar(im, fraction=0.036, pad=0.04)
  137. plt.show()
  138. # %%
  139. ### save data as np arrays
  140. sens_max = 76 #74
  141. probe = 'c'
  142. lfp_np = aligned_lfp.data
  143. data = np.hstack((lfp_np[:sens_max,0,:], lfp_np[:sens_max,1,:], lfp_np[:sens_max,2,:]))
  144. print(data.shape)
  145. n_t = data.shape[1]
  146. n_s = data.shape[0]
  147. ## Plot the numpy array
  148. fig, ax = plt.subplots(figsize = (5,6))
  149. im = ax.imshow(data, aspect='auto', origin='lower', vmin=-0.4e-3, vmax=0.4e-3, cmap = 'Reds')
  150. ax.set_yticks(intervals)
  151. ax.set_yticks(interval_midpoints, minor=True)
  152. ax.set_yticklabels(structure_acronyms, minor=True)
  153. plt.tick_params("y", which="major", labelleft=False, length=40)
  154. for y in [0,1000, 2000]:
  155. plt.axvline(y, color='white', linestyle='--', linewidth = '1.25')
  156. ax.set_xlim((-15,data.shape[1]))
  157. ax.set_ylim((0,sens_max))
  158. num_time_labels = 8
  159. time_label_indices = np.around(np.linspace(1, len(trial_window), num_time_labels)).astype(int) - 1
  160. ax.set_xlabel("Sample Number", fontsize=20)
  161. plt.colorbar(im, fraction=0.036, pad=0.04)
  162. plt.show()
  163. # %% [markdown]
  164. # **Applying SHRED on Dataset**
  165. # %%
  166. #Importing packages
  167. os.environ["CUDA_VISIBLE_DEVICES"]="0"
  168. SHRED_DIR = os.path.join(HOME_DIR, "neuralSHRED/sindy-shred-main")
  169. import sys
  170. sys.path.append(SHRED_DIR)
  171. from processdata import load_data
  172. from processdata import TimeSeriesDataset
  173. import sindy
  174. import sindy_shred
  175. os.environ["CUDA_VISIBLE_DEVICES"]="0"
  176. # %%
  177. load_X = data.T
  178. latent_dim = 32
  179. poly_order = 1
  180. num_neurons = 3
  181. lags = 100
  182. test_val_size = 1000 - lags
  183. ## Ignore below
  184. include_sine = False
  185. library_dim = sindy.library_size(latent_dim, poly_order, include_sine, True)
  186. train_indices = np.arange(0, n_t - lags - test_val_size)
  187. mask = np.ones(n_t - lags)
  188. mask[train_indices] = 0
  189. valid_test_indices = np.arange(0, n_t - lags)[np.where(mask!=0)[0]]
  190. valid_indices = valid_test_indices[int(test_val_size/2):test_val_size]
  191. test_indices = valid_test_indices[:int(test_val_size/2)]
  192. # %%
  193. ## Plot the lag
  194. fig, ax = plt.subplots(figsize = (7,3))
  195. im = ax.imshow(data[:, :lags], aspect='auto', origin='lower', vmin=-0.4e-3, vmax=0.4e-3, cmap = 'Reds')
  196. ax.set_yticks(intervals)
  197. ax.set_yticks(interval_midpoints, minor=True)
  198. ax.set_yticklabels(structure_acronyms, minor=True)
  199. plt.tick_params("y", which="major", labelleft=False, length=40)
  200. # for y in [0,1000, 2000]:
  201. # plt.axvline(y, color='white', linestyle='--', linewidth = '1.25')
  202. # ax.set_xlim((-15,data.shape[1]))
  203. ax.set_ylim((0,sens_max))
  204. num_time_labels = 8
  205. time_label_indices = np.around(np.linspace(1, len(trial_window), num_time_labels)).astype(int) - 1
  206. ax.set_xlabel("Sample Number", fontsize=20)
  207. plt.colorbar(im, fraction=0.036, pad=0.04)
  208. plt.show()
  209. # %%
  210. # Randomly select the sensors from the CA1 region for probe A (35 - 51) and SUB for probe C (31 - 54)
  211. print(intervals)
  212. print(structure_acronyms)
  213. for k in range(4,5):
  214. neuron_locations = np.random.choice(np.arange(31, 54), size=num_neurons, replace=False)
  215. file_path = os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/neuron_loc_{probe}_{k}.npy")
  216. np.save(file_path, neuron_locations)
  217. sc = MinMaxScaler()
  218. sc = sc.fit(load_X[train_indices])
  219. transformed_X = sc.transform(load_X)
  220. ### Generate input sequences to a SHRED model
  221. all_data_in = np.zeros((n_t - lags, lags, num_neurons))
  222. for i in range(len(all_data_in)):
  223. all_data_in[i] = transformed_X[i:i+lags, neuron_locations]
  224. ### Generate training validation and test datasets both for reconstruction of states and forecasting sensors
  225. device = 'cuda' if torch.cuda.is_available() else 'cpu'
  226. train_data_in = torch.tensor(all_data_in[train_indices], dtype=torch.float32).to(device)
  227. valid_data_in = torch.tensor(all_data_in[valid_indices], dtype=torch.float32).to(device)
  228. test_data_in = torch.tensor(all_data_in[test_indices], dtype=torch.float32).to(device)
  229. ### -1 to have output be at the same time as final sensor measurements
  230. train_data_out = torch.tensor(transformed_X[train_indices + lags - 1], dtype=torch.float32).to(device)
  231. valid_data_out = torch.tensor(transformed_X[valid_indices + lags - 1], dtype=torch.float32).to(device)
  232. test_data_out = torch.tensor(transformed_X[test_indices + lags - 1], dtype=torch.float32).to(device)
  233. train_dataset = TimeSeriesDataset(train_data_in, train_data_out)
  234. valid_dataset = TimeSeriesDataset(valid_data_in, valid_data_out)
  235. test_dataset = TimeSeriesDataset(test_data_in, test_data_out)
  236. shred = sindy_shred.SINDy_SHRED(num_neurons, n_s, hidden_size=latent_dim, hidden_layers=2, l1=350, l2=400, dropout=0.1,
  237. library_dim=library_dim, poly_order=poly_order,
  238. include_sine=include_sine, dt=dt, layer_norm=False, sindy = False).to(device)
  239. validation_errors = sindy_shred.fit(shred, train_dataset, valid_dataset, batch_size=128, num_epochs=200,
  240. lr=1e-3, verbose=True, threshold=0.25, patience=5, sindy_regularization= 0.0,
  241. optimizer="AdamW", thres_epoch=100)
  242. test_recons = sc.inverse_transform(shred(test_dataset.X).detach().cpu().numpy())
  243. test_ground_truth = sc.inverse_transform(test_dataset.Y.detach().cpu().numpy())
  244. train_recons = sc.inverse_transform(shred(train_dataset.X).detach().cpu().numpy())
  245. train_ground_truth = sc.inverse_transform(train_dataset.Y.detach().cpu().numpy())
  246. mask = np.ones(train_ground_truth.shape[1], dtype=bool)
  247. mask[neuron_locations] = False
  248. train_ex = train_recons.T[mask]
  249. train_gt_ex = train_ground_truth.T[mask]
  250. mask = np.ones(test_ground_truth.shape[1], dtype=bool)
  251. mask[neuron_locations] = False
  252. test_ex = test_recons.T[mask]
  253. test_gt_ex = test_ground_truth.T[mask]
  254. mse_train = np.linalg.norm(train_ex - train_gt_ex) / np.linalg.norm(train_gt_ex)
  255. mse_test = np.linalg.norm(test_ex - test_gt_ex) / np.linalg.norm(test_gt_ex)
  256. print('mse test', mse_test)
  257. print('mse train', mse_train)
  258. file_path = os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/train_recons_{probe}_{k}.npy")
  259. np.save(file_path, train_recons)
  260. file_path = os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/test_recons_{probe}_{k}.npy")
  261. np.save(file_path, test_recons)
  262. file_path = os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/train_gt_{probe}_{k}.npy")
  263. np.save(file_path, train_ground_truth)
  264. file_path = os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/test_gt_{probe}_{k}.npy")
  265. np.save(file_path, test_ground_truth)
  266. file_path = os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/mse_{probe}_{k}.npy")
  267. np.save(file_path, [mse_train, mse_test])
  268. # %% [markdown]
  269. # **Plotting the Data**
  270. # %%
  271. ## Probe A
  272. structure_acronyms = ['APN', 'DG', 'CA1', 'VISam', 'nan']
  273. intervals = [0, 27, 35, 51, 74 ,87]
  274. interval_midpoints = [13.5, 31.0, 43.0, 62.5, 80.5]
  275. ### Load the data
  276. sens_max = 74
  277. probe = 'a'
  278. trial = 2
  279. train_recon = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/train_recons_a_{trial}.npy"))
  280. test_recon = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/test_recons_a_{trial}.npy"))
  281. data_recon = np.vstack((train_recon, test_recon))
  282. train_gt = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/train_gt_a_{trial}.npy"))
  283. test_gt = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/test_gt_a_{trial}.npy"))
  284. data_gt= np.vstack((train_gt, test_gt))
  285. neuron_loc = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/neuron_loc_a_{trial}.npy"))
  286. mse = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/mse_a_{trial}.npy"))
  287. ## Plot the Reconstruction
  288. fig, ax = plt.subplots(figsize = (6.5,3))
  289. im = ax.imshow(data_recon.T, aspect='auto', origin='lower', vmin=-0.4e-3, vmax=0.4e-3, cmap = 'YlOrRd')
  290. ax.set_yticks(intervals)
  291. ax.set_yticks(interval_midpoints, minor=True)
  292. ax.set_yticklabels(structure_acronyms, minor=True)
  293. for tick in ax.yaxis.get_minor_ticks():
  294. tick.label1.set_fontsize(12)
  295. plt.tick_params("y", which="major", labelleft=False, length=40)
  296. for y in [0,1000, 2000]:
  297. plt.axvline(y, color='white', linestyle='--', linewidth = '1.25')
  298. for k in range(len(neuron_loc)):
  299. ax.plot(-5, neuron_loc[k], 'o', color='red', zorder = 10, clip_on=False, markersize = 5)
  300. ax.set_clip_on(False)
  301. ax.set_xlim((-15,data_recon.shape[0]))
  302. ax.set_ylim((0,sens_max))
  303. plt.tick_params(axis='both', which='major', labelsize=13)
  304. num_time_labels = 8
  305. time_label_indices = np.around(np.linspace(1, len(trial_window), num_time_labels)).astype(int) - 1
  306. ax.set_xlabel("Sample Number", fontsize=15)
  307. cbar = plt.colorbar(im, fraction=0.036, pad=0.04)
  308. cbar.set_label("V", fontsize = 12)
  309. cbar.ax.tick_params(labelsize=11)
  310. plt.tight_layout()
  311. plt.savefig(os.path.join(HOME_DIR, "mice/data_output/figs/lfp_dg_a_recon.png"), transparent = True, dpi = 400)
  312. plt.show()
  313. ## Plot the Ground Truth
  314. fig, ax = plt.subplots(figsize = (6.5,3))
  315. im = ax.imshow(data_gt.T, aspect='auto', origin='lower', vmin=-0.4e-3, vmax=0.4e-3, cmap = 'YlOrRd')
  316. ax.set_yticks(intervals)
  317. ax.set_yticks(interval_midpoints, minor=True)
  318. ax.set_yticklabels(structure_acronyms, minor=True)
  319. for tick in ax.yaxis.get_minor_ticks():
  320. tick.label1.set_fontsize(12)
  321. plt.tick_params("y", which="major", labelleft=False, length=40)
  322. for y in [0,1000, 2000]:
  323. plt.axvline(y, color='white', linestyle='--', linewidth = '1.25')
  324. for k in range(len(neuron_loc)):
  325. ax.plot(-5, neuron_loc[k], 'o', color='red', zorder = 10, clip_on=False, markersize = 5)
  326. ax.set_clip_on(False)
  327. ax.set_xlim((-15,data_gt.shape[0]))
  328. ax.set_ylim((0,sens_max))
  329. plt.tick_params(axis='both', which='major', labelsize=13)
  330. num_time_labels = 8
  331. time_label_indices = np.around(np.linspace(1, len(trial_window), num_time_labels)).astype(int) - 1
  332. ax.set_xlabel("Sample Number", fontsize=15)
  333. cbar = plt.colorbar(im, fraction=0.036, pad=0.04)
  334. cbar.set_label("V", fontsize = 12)
  335. cbar.ax.tick_params(labelsize=11)
  336. plt.tight_layout()
  337. plt.savefig(os.path.join(HOME_DIR, "mice/data_output/figs/lfp_dg_a_gt.png"), transparent = True, dpi = 400)
  338. plt.show()
  339. print('mse', mse)
  340. # %%
  341. structure_acronyms = ['APN', 'DG', 'CA1', 'VISam', 'nan']
  342. intervals = [0, 27, 35, 51, 74 ,87]
  343. APN_train_gt = train_gt[0:8,:]
  344. DG_train_gt = train_gt[27:35,:]
  345. CA1_train_gt = train_gt[35:43,:]
  346. VISam_train_gt = train_gt[51:59,:]
  347. APN_test_gt = test_gt[0:8,:]
  348. DG_test_gt = test_gt[27:35,:]
  349. CA1_test_gt = test_gt[35:43,:]
  350. VISam_test_gt = test_gt[51:59,:]
  351. APN_train_recon = train_recon[0:8,:]
  352. DG_train_recon = train_recon[27:35,:]
  353. CA1_train_recon = train_recon[35:43,:]
  354. VISam_train_recon = train_recon[51:59,:]
  355. APN_test_recon = test_recon[0:8,:]
  356. DG_test_recon = test_recon[27:35,:]
  357. CA1_test_recon = test_recon[35:43,:]
  358. VISam_test_recon = test_recon[51:59,:]
  359. ### CA1 to CA1
  360. mse_ca1_train = np.linalg.norm(CA1_train_recon - CA1_train_gt) / np.linalg.norm(CA1_train_gt)
  361. mse_ca1_test = np.linalg.norm(CA1_test_recon - CA1_test_gt) / np.linalg.norm(CA1_test_gt)
  362. ### CA1 to VISam
  363. mse_visam_train = np.linalg.norm(VISam_train_recon - VISam_train_gt) / np.linalg.norm(VISam_train_gt)
  364. mse_visam_test = np.linalg.norm(VISam_test_recon - VISam_test_gt) / np.linalg.norm(VISam_test_gt)
  365. ### CA1 to DG
  366. mse_dg_train = np.linalg.norm(DG_train_recon - DG_train_gt) / np.linalg.norm(DG_train_gt)
  367. mse_dg_test = np.linalg.norm(DG_test_recon - DG_test_gt) / np.linalg.norm(DG_test_gt)
  368. ### CA1 to APN
  369. mse_apn_train = np.linalg.norm(APN_train_recon - APN_train_gt) / np.linalg.norm(APN_train_gt)
  370. mse_apn_test = np.linalg.norm(APN_test_recon - APN_test_gt) / np.linalg.norm(APN_test_gt)
  371. print('mse visam train', mse_visam_train, mse_visam_test)
  372. print('mse dg train', mse_dg_train, mse_dg_test)
  373. print('mse apn train', mse_apn_train, mse_apn_test)
  374. print('mse ca1 train', mse_ca1_train, mse_ca1_test)
  375. # %%
  376. ## Probe C
  377. structure_acronyms = ['POL', '', '', 'SCig', '', 'SUB', 'VISp', 'nan']
  378. intervals = [ 0, 5 , 9 ,12 ,29, 31 ,54, 76 ,81]
  379. interval_midpoints = [2.5, 7.0, 10.5, 20.5, 30.0, 42.5, 65.0, 78.5]
  380. ### Load the data
  381. sens_max = 76
  382. probe = 'c'
  383. trial = 3 ## trial 4- 64 latent dimension
  384. train_recon = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/train_recons_c_{trial}.npy"))
  385. test_recon = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/test_recons_c_{trial}.npy"))
  386. data_recon = np.vstack((train_recon, test_recon))
  387. train_gt = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/train_gt_c_{trial}.npy"))
  388. test_gt = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/test_gt_c_{trial}.npy"))
  389. data_gt= np.vstack((train_gt, test_gt))
  390. neuron_loc = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/neuron_loc_c_{trial}.npy"))
  391. mse = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/mse_c_{trial}.npy"))
  392. ## Plot the Reconstruction
  393. fig, ax = plt.subplots(figsize = (6.5,3))
  394. im = ax.imshow(data_recon.T, aspect='auto', origin='lower', vmin=-0.4e-3, vmax=0.4e-3, cmap = 'GnBu')
  395. ax.set_yticks(intervals)
  396. ax.set_yticks(interval_midpoints, minor=True)
  397. ax.set_yticklabels(structure_acronyms, minor=True)
  398. for tick in ax.yaxis.get_minor_ticks():
  399. tick.label1.set_fontsize(12)
  400. plt.tick_params("y", which="major", labelleft=False, length=40)
  401. for y in [0,1000, 2000]:
  402. plt.axvline(y, color='white', linestyle='--', linewidth = '1.25')
  403. for k in range(len(neuron_loc)):
  404. ax.plot(-5, neuron_loc[k], 'o', color='red', zorder = 10, clip_on=False, markersize = 5)
  405. ax.set_clip_on(False)
  406. ax.set_xlim((-15,data_recon.shape[0]))
  407. ax.set_ylim((0,sens_max))
  408. plt.tick_params(axis='both', which='major', labelsize=13)
  409. num_time_labels = 8
  410. time_label_indices = np.around(np.linspace(1, len(trial_window), num_time_labels)).astype(int) - 1
  411. ax.set_xlabel("Sample Number", fontsize=15)
  412. cbar = plt.colorbar(im, fraction=0.036, pad=0.04)
  413. cbar.ax.tick_params(labelsize=11)
  414. cbar.set_label("V", fontsize = 12)
  415. plt.tight_layout()
  416. plt.savefig(os.path.join(HOME_DIR, "mice/data_output/figs/lfp_dg_c_recon.png"), transparent = True, dpi = 400)
  417. plt.show()
  418. ## Plot the Ground Truth
  419. fig, ax = plt.subplots(figsize = (6.5,3))
  420. im = ax.imshow(data_gt.T, aspect='auto', origin='lower', vmin=-0.4e-3, vmax=0.4e-3, cmap = 'GnBu')
  421. ax.set_yticks(intervals)
  422. ax.set_yticks(interval_midpoints, minor=True)
  423. ax.set_yticklabels(structure_acronyms, minor=True)
  424. for tick in ax.yaxis.get_minor_ticks():
  425. tick.label1.set_fontsize(12)
  426. plt.tick_params("y", which="major", labelleft=False, length=40)
  427. for y in [0,1000, 2000]:
  428. plt.axvline(y, color='white', linestyle='--', linewidth = '1.25')
  429. for k in range(len(neuron_loc)):
  430. ax.plot(-5, neuron_loc[k], 'o', color='red', zorder = 10, clip_on=False, markersize = 5)
  431. ax.set_clip_on(False)
  432. ax.set_xlim((-15,data_gt.shape[0]))
  433. ax.set_ylim((0,sens_max))
  434. plt.tick_params(axis='both', which='major', labelsize=13)
  435. num_time_labels = 8
  436. time_label_indices = np.around(np.linspace(1, len(trial_window), num_time_labels)).astype(int) - 1
  437. ax.set_xlabel("Sample Number", fontsize=15)
  438. cbar = plt.colorbar(im, fraction=0.036, pad=0.04)
  439. cbar.ax.tick_params(labelsize=11)
  440. cbar.set_label("V", fontsize = 12)
  441. plt.tight_layout()
  442. plt.savefig(os.path.join(HOME_DIR, "mice/data_output/figs/lfp_dg_c_gt.png"), transparent = True, dpi = 400)
  443. plt.show()
  444. print('mse', mse)
  445. # %%
  446. structure_acronyms = ['POL', 'APN', 'MB', 'SCig', 'POST', 'SUB', 'VISp', 'nan']
  447. intervals = [ 0, 5 , 9 ,12 ,29, 31 ,54, 76 ,81]
  448. POL_train_gt = train_gt[0:2,:]
  449. APN_train_gt = train_gt[5:7,:]
  450. MB_train_gt = train_gt[9:11,:]
  451. SCig_train_gt = train_gt[12:14,:]
  452. POST_train_gt = train_gt[29:31,:]
  453. SUB_train_gt = train_gt[31:33,:]
  454. VISp_train_gt = train_gt[54:57,:]
  455. POL_train_recon = train_recon[0:2,:]
  456. APN_train_recon = train_recon[5:7,:]
  457. MB_train_recon = train_recon[9:11,:]
  458. SCig_train_recon = train_recon[12:14,:]
  459. POST_train_recon = train_recon[29:31,:]
  460. SUB_train_recon = train_recon[31:33,:]
  461. VISp_train_recon = train_recon[54:57,:]
  462. POL_test_gt = test_gt[0:2,:]
  463. APN_test_gt = test_gt[5:7,:]
  464. MB_test_gt = test_gt[9:11,:]
  465. SCig_test_gt = test_gt[12:14,:]
  466. POST_test_gt = test_gt[29:31,:]
  467. SUB_test_gt = test_gt[31:33,:]
  468. VISp_test_gt = test_gt[54:57,:]
  469. POL_test_recon = test_recon[0:2,:]
  470. APN_test_recon = test_recon[5:7,:]
  471. MB_test_recon = test_recon[9:11,:]
  472. SCig_test_recon = test_recon[12:14,:]
  473. POST_test_recon = test_recon[29:31,:]
  474. SUB_test_recon = test_recon[31:33,:]
  475. VISp_test_recon = test_recon[54:57,:]
  476. ### APN
  477. mse_apn_train = np.linalg.norm(APN_train_recon - APN_train_gt) / np.linalg.norm(APN_train_gt)
  478. mse_apn_test = np.linalg.norm(APN_test_recon - APN_test_gt) / np.linalg.norm(APN_test_gt)
  479. print('apn', mse_apn_train, mse_apn_test)
  480. ### VISp
  481. mse_visp_train = np.linalg.norm(VISp_train_recon - VISp_train_gt) / np.linalg.norm(VISp_train_gt)
  482. mse_visp_test = np.linalg.norm(VISp_test_recon - VISp_test_gt) / np.linalg.norm(VISp_test_gt)
  483. print('visp', mse_visp_train, mse_visp_test)
  484. ### MB
  485. mse_mb_train = np.linalg.norm(MB_train_recon - MB_train_gt) / np.linalg.norm(MB_train_gt)
  486. mse_mb_test = np.linalg.norm(MB_test_recon - MB_test_gt) / np.linalg.norm(MB_test_gt)
  487. print('mb', mse_mb_train, mse_mb_test)
  488. ### SCig
  489. mse_scig_train = np.linalg.norm(SCig_train_recon - SCig_train_gt) / np.linalg.norm(SCig_train_gt)
  490. mse_scig_test = np.linalg.norm(SCig_test_recon - SCig_test_gt) / np.linalg.norm(SCig_test_gt)
  491. print('scig', mse_scig_train, mse_scig_test)
  492. ### POL
  493. mse_pol_train = np.linalg.norm(POL_train_recon - POL_train_gt) / np.linalg.norm(POL_train_gt)
  494. mse_pol_test = np.linalg.norm(POL_test_recon - POL_test_gt) / np.linalg.norm(POL_test_gt)
  495. print('pol', mse_pol_train, mse_pol_test)
  496. ##SUB
  497. mse_sub_train = np.linalg.norm(SUB_train_recon - SUB_train_gt) / np.linalg.norm(SUB_train_gt)
  498. mse_sub_test = np.linalg.norm(SUB_test_recon - SUB_test_gt) / np.linalg.norm(SUB_test_gt)
  499. print('sub', mse_sub_train, mse_sub_test)
  500. ## POST
  501. mse_post_train = np.linalg.norm(POST_train_recon - POST_train_gt) / np.linalg.norm(POST_train_gt)
  502. mse_post_test = np.linalg.norm(POST_test_recon - POST_test_gt) / np.linalg.norm(POST_test_gt)
  503. print('post', mse_post_train, mse_post_test)
  504. # %%
  505. #### Linear Regression
  506. ## Probe A
  507. trial = 2
  508. train_recon = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/train_recons_a_{trial}.npy"))
  509. test_recon = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/test_recons_a_{trial}.npy"))
  510. data_recon = np.vstack((train_recon, test_recon))
  511. train_gt = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/train_gt_a_{trial}.npy"))
  512. test_gt = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/test_gt_a_{trial}.npy"))
  513. data_gt= np.vstack((train_gt, test_gt))
  514. neuron_loc = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/neuron_loc_a_{trial}.npy"))
  515. mse = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/mse_a_{trial}.npy"))
  516. ##### Linear Regression
  517. from sklearn.linear_model import LinearRegression
  518. train_in = train_gt[:, neuron_loc]
  519. test_in = test_gt[:, neuron_loc]
  520. data_in = data_gt[:, neuron_loc]
  521. train_out = train_gt
  522. model = LinearRegression()
  523. model.fit(train_in, train_out)
  524. test_out = model.predict(test_in)
  525. print('test out shape',test_out.shape)
  526. data_out = model.predict(data_in)
  527. print('data out shape',data_out.shape)
  528. plt.figure(figsize = (6,2))
  529. plt.plot(test_out)
  530. plt.title('Linear Regression')
  531. plt.show()
  532. plt.figure(figsize = (6,2))
  533. plt.plot(test_gt)
  534. plt.title('Ground Truth')
  535. plt.show()
  536. ######### Plotting the reconstruction
  537. fig, ax = plt.subplots(figsize = (6.5,2))
  538. plt.imshow(data_gt.T, aspect = "auto")
  539. plt.xlabel('Time (s)', fontsize = 12)
  540. plt.ylabel('Neuron', fontsize =12)
  541. # Add red dots for neuron locations on y-axis
  542. for loc in neuron_loc:
  543. ax.plot(0.1, loc, 'ro', markersize=5, clip_on=False, zorder=10)
  544. plt.colorbar()
  545. plt.tight_layout()
  546. # plt.savefig(os.path.join(HOME_DIR, f'worms/data_output/figs/{name}_recon.png'), transparent = True, dpi = 400)
  547. plt.show()
  548. print('out shape',data_out.shape)
  549. fig, ax = plt.subplots(figsize = (6.5,2))
  550. plt.imshow(data_out.T, aspect = "auto")
  551. plt.xlabel('Time (s)', fontsize = 12)
  552. plt.ylabel('Neuron', fontsize = 12)
  553. plt.colorbar()
  554. # Add red dots for neuron locations on y-axis
  555. for loc in neuron_loc:
  556. ax.plot(0.1, loc, 'ro', markersize=5, clip_on=False, zorder=10)
  557. plt.tight_layout()
  558. # plt.savefig(os.path.join(HOME_DIR, f'worms/data_output/figs/{name}_linear.png'), transparent = True, dpi = 400)
  559. plt.show()
  560. n_train = train_out.shape[0]
  561. mask = np.ones(train_out.shape[1], dtype=bool)
  562. mask[neuron_loc] = False
  563. train_ex = data_out[:n_train,:].T[mask]
  564. train_gt_ex = train_gt.T[mask]
  565. mask = np.ones(test_out.shape[1], dtype=bool)
  566. mask[neuron_loc] = False
  567. test_ex = test_out.T[mask]
  568. test_gt_ex = test_gt.T[mask]
  569. mse_train = np.linalg.norm(train_ex - train_gt_ex) / np.linalg.norm(train_gt_ex)
  570. mse_test = np.linalg.norm(test_ex - test_gt_ex) / np.linalg.norm(test_gt_ex)
  571. print('mse train, test', [mse_train, mse_test])
  572. mask = np.ones(train_out.shape[1], dtype=bool)
  573. mask[neuron_loc] = False
  574. data_ex = data_out.T[mask]
  575. data_gt_ex = data_gt.T[mask]
  576. mse_whole = np.linalg.norm(data_ex - data_gt_ex) / np.linalg.norm(data_gt_ex)
  577. print('mse whole', mse_whole)
  578. # %%
  579. # Compute power spectral density for each dataset
  580. from scipy import signal
  581. import numpy as np
  582. name = 'expB_i_dg'
  583. # Parameters for PSD computation
  584. nperseg = min(256, data_out.shape[0] // 4) # Window length for Welch's method
  585. noverlap = nperseg // 2
  586. time = np.linspace(0, 2*data_out.shape[0]/3000, data_out.shape[0])
  587. # Initialize lists to store PSDs
  588. psds_linear = []
  589. psds_gt = []
  590. psds_shred = []
  591. freqs = None
  592. # Compute PSD for each neuron/channel
  593. for i in range(data_out.shape[1]):
  594. # Linear regression PSD
  595. f, psd_linear = signal.welch(data_out[:, i], fs=1/(time[1]-time[0]),
  596. nperseg=nperseg, noverlap=noverlap)
  597. psds_linear.append(psd_linear)
  598. # Ground truth PSD
  599. f, psd_gt = signal.welch(data_gt[:, i], fs=1/(time[1]-time[0]),
  600. nperseg=nperseg, noverlap=noverlap)
  601. psds_gt.append(psd_gt)
  602. # SHRED PSD
  603. f, psd_shred = signal.welch(data_recon[:, i], fs=1/(time[1]-time[0]),
  604. nperseg=nperseg, noverlap=noverlap)
  605. psds_shred.append(psd_shred)
  606. if freqs is None:
  607. freqs = f
  608. # Convert to arrays
  609. psds_linear = np.array(psds_linear)
  610. psds_gt = np.array(psds_gt)
  611. psds_shred = np.array(psds_shred)
  612. # Compute mean and standard deviation across neurons
  613. mean_psd_linear = np.mean(psds_linear, axis=0)
  614. std_psd_linear = np.std(psds_linear, axis=0)
  615. mean_psd_gt = np.mean(psds_gt, axis=0)
  616. std_psd_gt = np.std(psds_gt, axis=0)
  617. mean_psd_shred = np.mean(psds_shred, axis=0)
  618. std_psd_shred = np.std(psds_shred, axis=0)
  619. # Plot the power spectral densities with smooth distributions
  620. fig, ax = plt.subplots(figsize=(3.5, 2))
  621. neuron = 10
  622. # Plot mean with shaded regions for standard deviation
  623. ax.loglog(freqs, mean_psd_linear, 'b-', linewidth=2, label='Linear Regression', alpha=0.8)
  624. ax.fill_between(freqs, mean_psd_linear - std_psd_linear, mean_psd_linear + std_psd_linear,
  625. alpha=0.3, color='blue')
  626. ax.loglog(freqs, mean_psd_gt, 'k-', linewidth=2, label='Ground Truth', alpha=0.8)
  627. ax.fill_between(freqs, mean_psd_gt - std_psd_gt, mean_psd_gt + std_psd_gt,
  628. alpha=0.3, color='black')
  629. ax.loglog(freqs, mean_psd_shred, 'r-', linewidth=2, label='SHRED', alpha=0.8)
  630. ax.fill_between(freqs, mean_psd_shred - std_psd_shred, mean_psd_shred + std_psd_shred,
  631. alpha=0.3, color='red')
  632. ax.set_xlabel('Frequency (Hz)', fontsize=12)
  633. ax.set_ylabel('PSD', fontsize=12)
  634. # ax.legend(fontsize=12)
  635. plt.tight_layout()
  636. plt.savefig(os.path.join(HOME_DIR, f'mice/data_output/figs/{name}_psd.png'), transparent = True, dpi = 400)
  637. plt.show()
  638. # %%
  639. mse_linear = np.linalg.norm(mean_psd_linear - mean_psd_gt) / np.linalg.norm(mean_psd_gt)
  640. mse_shred = np.linalg.norm(mean_psd_shred - mean_psd_gt) / np.linalg.norm(mean_psd_gt)
  641. print(mse_shred)
  642. print(mse_linear)
  643. # %%
  644. #### Linear Regression
  645. ## Probe C
  646. trial = 3
  647. train_recon = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/train_recons_c_{trial}.npy"))
  648. test_recon = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/test_recons_c_{trial}.npy"))
  649. data_recon = np.vstack((train_recon, test_recon))
  650. train_gt = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/train_gt_c_{trial}.npy"))
  651. test_gt = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/test_gt_c_{trial}.npy"))
  652. data_gt= np.vstack((train_gt, test_gt))
  653. neuron_loc = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/neuron_loc_c_{trial}.npy"))
  654. mse = np.load(os.path.join(HOME_DIR, f"mice/data_output/exp_Bi/dg/mse_c_{trial}.npy"))
  655. ##### Linear Regression
  656. from sklearn.linear_model import LinearRegression
  657. train_in = train_gt[:, neuron_loc]
  658. test_in = test_gt[:, neuron_loc]
  659. data_in = data_gt[:, neuron_loc]
  660. train_out = train_gt
  661. model = LinearRegression()
  662. model.fit(train_in, train_out)
  663. test_out = model.predict(test_in)
  664. print('test out shape',test_out.shape)
  665. data_out = model.predict(data_in)
  666. print('data out shape',data_out.shape)
  667. plt.figure(figsize = (6,2))
  668. plt.plot(test_out)
  669. plt.title('Linear Regression')
  670. plt.show()
  671. plt.figure(figsize = (6,2))
  672. plt.plot(test_gt)
  673. plt.title('Ground Truth')
  674. plt.show()
  675. ######### Plotting the reconstruction
  676. fig, ax = plt.subplots(figsize = (6.5,2))
  677. plt.imshow(data_gt.T, aspect = "auto")
  678. plt.xlabel('Time (s)', fontsize = 12)
  679. plt.ylabel('Neuron', fontsize =12)
  680. # Add red dots for neuron locations on y-axis
  681. for loc in neuron_loc:
  682. ax.plot(0.1, loc, 'ro', markersize=5, clip_on=False, zorder=10)
  683. plt.colorbar()
  684. plt.tight_layout()
  685. # plt.savefig(os.path.join(HOME_DIR, f'worms/data_output/figs/{name}_recon.png'), transparent = True, dpi = 400)
  686. plt.show()
  687. print('out shape',data_out.shape)
  688. fig, ax = plt.subplots(figsize = (6.5,2))
  689. plt.imshow(data_out.T, aspect = "auto")
  690. plt.xlabel('Time (s)', fontsize = 12)
  691. plt.ylabel('Neuron', fontsize = 12)
  692. plt.colorbar()
  693. # Add red dots for neuron locations on y-axis
  694. for loc in neuron_loc:
  695. ax.plot(0.1, loc, 'ro', markersize=5, clip_on=False, zorder=10)
  696. plt.tight_layout()
  697. # plt.savefig(os.path.join(HOME_DIR, f'worms/data_output/figs/{name}_linear.png'), transparent = True, dpi = 400)
  698. plt.show()
  699. n_train = train_out.shape[0]
  700. mask = np.ones(train_out.shape[1], dtype=bool)
  701. mask[neuron_loc] = False
  702. train_ex = data_out[:n_train,:].T[mask]
  703. train_gt_ex = train_gt.T[mask]
  704. mask = np.ones(test_out.shape[1], dtype=bool)
  705. mask[neuron_loc] = False
  706. test_ex = test_out.T[mask]
  707. test_gt_ex = test_gt.T[mask]
  708. mse_train = np.linalg.norm(train_ex - train_gt_ex) / np.linalg.norm(train_gt_ex)
  709. mse_test = np.linalg.norm(test_ex - test_gt_ex) / np.linalg.norm(test_gt_ex)
  710. print('mse train, test', [mse_train, mse_test])
  711. mask = np.ones(train_out.shape[1], dtype=bool)
  712. mask[neuron_loc] = False
  713. data_ex = data_out.T[mask]
  714. data_gt_ex = data_gt.T[mask]
  715. mse_whole = np.linalg.norm(data_ex - data_gt_ex) / np.linalg.norm(data_gt_ex)
  716. print('mse whole', mse_whole)
  717. # %%
  718. #Compute power spectral density for each dataset
  719. from scipy import signal
  720. import numpy as np
  721. name = 'expB_i_dg_c'
  722. # Parameters for PSD computation
  723. nperseg = min(256, data_out.shape[0] // 4) # Window length for Welch's method
  724. noverlap = nperseg // 2
  725. time = np.linspace(0, 2*data_out.shape[0]/3000, data_out.shape[0])
  726. # Initialize lists to store PSDs
  727. psds_linear = []
  728. psds_gt = []
  729. psds_shred = []
  730. freqs = None
  731. # Compute PSD for each neuron/channel
  732. for i in range(data_out.shape[1]):
  733. # Linear regression PSD
  734. f, psd_linear = signal.welch(data_out[:, i], fs=1/(time[1]-time[0]),
  735. nperseg=nperseg, noverlap=noverlap)
  736. psds_linear.append(psd_linear)
  737. # Ground truth PSD
  738. f, psd_gt = signal.welch(data_gt[:, i], fs=1/(time[1]-time[0]),
  739. nperseg=nperseg, noverlap=noverlap)
  740. psds_gt.append(psd_gt)
  741. # SHRED PSD
  742. f, psd_shred = signal.welch(data_recon[:, i], fs=1/(time[1]-time[0]),
  743. nperseg=nperseg, noverlap=noverlap)
  744. psds_shred.append(psd_shred)
  745. if freqs is None:
  746. freqs = f
  747. # Convert to arrays
  748. psds_linear = np.array(psds_linear)
  749. psds_gt = np.array(psds_gt)
  750. psds_shred = np.array(psds_shred)
  751. # Compute mean and standard deviation across neurons
  752. mean_psd_linear = np.mean(psds_linear, axis=0)
  753. std_psd_linear = np.std(psds_linear, axis=0)
  754. mean_psd_gt = np.mean(psds_gt, axis=0)
  755. std_psd_gt = np.std(psds_gt, axis=0)
  756. mean_psd_shred = np.mean(psds_shred, axis=0)
  757. std_psd_shred = np.std(psds_shred, axis=0)
  758. # Plot the power spectral densities with smooth distributions
  759. fig, ax = plt.subplots(figsize=(3.5, 2))
  760. neuron = 10
  761. # Plot mean with shaded regions for standard deviation
  762. ax.loglog(freqs, mean_psd_linear, 'b-', linewidth=2, label='Linear Regression', alpha=0.8)
  763. ax.fill_between(freqs, mean_psd_linear - std_psd_linear, mean_psd_linear + std_psd_linear,
  764. alpha=0.3, color='blue')
  765. ax.loglog(freqs, mean_psd_gt, 'k-', linewidth=2, label='Ground Truth', alpha=0.8)
  766. ax.fill_between(freqs, mean_psd_gt - std_psd_gt, mean_psd_gt + std_psd_gt,
  767. alpha=0.3, color='black')
  768. ax.loglog(freqs, mean_psd_shred, 'r-', linewidth=2, label='SHRED', alpha=0.8)
  769. ax.fill_between(freqs, mean_psd_shred - std_psd_shred, mean_psd_shred + std_psd_shred,
  770. alpha=0.3, color='red')
  771. ax.set_xlabel('Frequency (Hz)', fontsize=12)
  772. ax.set_ylabel('PSD', fontsize=12)
  773. # ax.legend(fontsize=12)
  774. plt.tight_layout()
  775. plt.savefig(os.path.join(HOME_DIR, f'mice/data_output/figs/{name}_psd.png'), transparent = True, dpi = 400)
  776. plt.show()
  777. # %%
  778. mse_linear = np.linalg.norm(mean_psd_linear - mean_psd_gt) / np.linalg.norm(mean_psd_gt)
  779. mse_shred = np.linalg.norm(mean_psd_shred - mean_psd_gt) / np.linalg.norm(mean_psd_gt)
  780. print(mse_shred)
  781. print(mse_linear)
  782. # %%

shred_mice_expB_i_dg.ipynb at commit 6b86875, no license · at the source

Overview

Authors: Amy Rude1, J Nathan Kutz1,2
ORCID iDs: J Nathan Kutz
  1. Department of Applied Mathematics, University of Washington, Seattle, WA, USA
  2. Autodesk Research, London, UK
Institutions: University of Washington (United States); Autodesk (United Kingdom) (United Kingdom)
Dates: received 12 June 2025; accepted 17 October 2025; published online 17 September 2026; in print September 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1098/rstb.2024.0461 · PMID 42750452 · PMCID PMC13583487 · OpenAlex W7213464516
Open access: hybrid, a free copy (OpenAlex)
Status: code verified
Categories: human (organism), mouse (organism), C. elegans (organism), computational (subfield)
Methods: Spectral & time-frequency, Connectivity, Statistics, Graphs, Machine learning, Single-unit activity, calcium imaging, Physiology & signal measures
Keywords: computational neuroscience, machine learning, population codes, shallow recurrent decoding, sensing
MeSH: Brain*, Machine Learning*, Models, Neurological*, Neurons*, Animals, Caenorhabditis elegans, Humans, Mice, Recurrent Neural Networks (* major topic)
Topic: Genetics, Aging, and Longevity in Model Organisms (Aging, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Funding: Air Force Office of Scientific Research (FA9550-24-1-0141); National Science Foundation AI Institute in Dynamic Systems (2112085)
Citations: not cited yet (Europe PMC); 68 references in the paper

Abstract

Machine learning algorithms are affording new opportunities for building bio-inspired and data-driven models characterizing neural activity. Critical to understanding decision-making and behaviour is quantifying the relationship between the activity of neuronal population codes and individual neurons. We leverage a SHallow REcurrent Decoder (SHRED) architecture for mapping the dynamics of population codes to individual neurons and other proxy measures of neural activity and behaviour. SHRED is constructed from a temporal sequence model, which encodes the temporal dynamics of limited sensor data in multiple scenarios, and a shallow decoder, which reconstructs the corresponding high-dimensional neuronal and/or behavioural states. It is a robust and flexible sensing strategy which allows for decoding the diversity of neural measurements with only a few sensor measurements. Thus, estimates of whole-brain activity, behaviour and individual neurons can be constructed with only a few neural time-series recordings. Several examples in this article further highlight the potential of leveraging non-invasive or minimally invasive measurements to estimate large-scale brain dynamics. We empirically demonstrate the capabilities of the method on a number of model organisms including Caenorhabditis elegans, mouse, zebrafish and human biolocomotion.

This article is part of the discussion meeting issue ‘Digital healthcare for the management of functional neurological disorders’.

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

amysrude/neuralSHRED

License: none: the authors keep all their rights
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: 6b8687501a637ec64ba2ffe65313c6acced4c83a, 15 September 2025
Languages: Jupyter (16), Python (8)
Size: 1,819 files, 24 scripts
Software Heritage: not archived
Found in: “Data accessibility”
Holds: README, environment (sindy-shred-main/requirements.txt), 16 notebooks
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (19 files), PyTorch (19 files), SciPy (16 files), scikit-learn (14 files), Matplotlib (12 files), MNE-Python (11 files), AllenSDK (7 files), pandas (7 files), Pillow (7 files), xarray (6 files)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
21 files

meganebers/mobileSHRED

License: none: the authors keep all their rights
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: 66e85893bea3a463e18aaaf8b087387c64561018, 7 March 2024
Languages: Python (5)
Size: 12 files, 5 scripts
Software Heritage: not archived
Found in: “Data accessibility”
Holds: README, environment (SHRED_base/requirements.txt)
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (5 files), PyTorch (5 files), scikit-learn (3 files), h5py (1 file), imageio (1 file), Matplotlib (1 file), SciPy (1 file)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
6 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:

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

All code and data are publicly available. Figures and experiments can be reproduced from the following GitHub: http://github.com/amysrude/neuralSHRED. The datasets used in the analyses are from the following sites. C. elegans: http://osf.io/2395t/. Mouse: download data: http://allensdk.readthedocs.io/en/latest/_static/examples/nb/ecephys_data_access.html. Documentation: http://allensdk.readthedocs.io/en/latest/visual_coding_neuropixels.html. Zebrafish: http://www.nature.com/articles/nmeth.2434. Biolocomotion: http://github.com/meganebers/mobileSHRED, http://simtk.org/projects/ankleexopred and [65,66,67].

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, 5 keywords, 9 MeSH terms, 2 funders, 58 references.

Cite

This paper

Rude, A., & Kutz, J. N. (2026). Shallow recurrent decoders for neural and behavioural dynamics. Philosophical transactions of the Royal Society of London. Series B, Biological sciences, 381(1958), 20240461. https://doi.org/10.1098/rstb.2024.0461

BibTeX

@article{rude2026shallow,
author = {Rude, Amy and Kutz, J Nathan},
title = {{Shallow recurrent decoders for neural and behavioural dynamics}},
journal = {Philosophical transactions of the Royal Society of London. Series B, Biological sciences},
year = {2026},
month = sep,
volume = {381},
number = {1958},
pages = {20240461},
publisher = {Royal Society},
issn = {0962-8436},
doi = {10.1098/rstb.2024.0461},
url = {https://doi.org/10.1098/rstb.2024.0461},
pmid = {42750452},
pmcid = {PMC13583487}
}

RIS

TY - JOUR
AU - Rude, Amy
AU - Kutz, J Nathan
TI - Shallow recurrent decoders for neural and behavioural dynamics
T2 - Philosophical transactions of the Royal Society of London. Series B, Biological sciences
J2 - Philos Trans R Soc Lond B Biol Sci
PY - 2026
DA - 2026/09/01
VL - 381
IS - 1958
SP - 20240461
SN - 0962-8436
PB - Royal Society
DO - 10.1098/rstb.2024.0461
UR - https://doi.org/10.1098/rstb.2024.0461
LA - en
ER -

CSL-JSON

{
"id": "10.1098/rstb.2024.0461",
"type": "article-journal",
"title": "Shallow recurrent decoders for neural and behavioural dynamics",
"container-title": "Philosophical transactions of the Royal Society of London. Series B, Biological sciences",
"author": [
{
"family": "Rude",
"given": "Amy"
},
{
"family": "Kutz",
"given": "J Nathan"
}
],
"container-title-short": "Philos Trans R Soc Lond B Biol Sci",
"volume": "381",
"issue": "1958",
"page": "20240461",
"DOI": "10.1098/rstb.2024.0461",
"PMID": "42750452",
"PMCID": "PMC13583487",
"ISSN": "0962-8436",
"publisher": "Royal Society",
"URL": "https://doi.org/10.1098/rstb.2024.0461",
"language": "en",
"issued": {
"date-parts": [
[
2026,
9,
1
]
]
}
}

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

Similar papers

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

[1] doi:10.1038/s41593-026-02257-5 [code]
Neural sequences underlying directed turning in Caenorhabditis elegans.
Journal: Nature neuroscience
In common: h5py, Pillow, PyTorch, 4 other tools, C. elegans, 3 references
[2] doi:10.1038/s41467-026-72709-w [code]
An epifluorescence microscope design for naturalistic behavior and cellular activity in freely moving Caenorhabditis elegans.
Journal: Nature communications
In common: imageio, h5py, Pillow, 5 other tools, C. elegans, 1 reference
[3] doi:10.3389/frai.2026.1771088 [code]
Few-shot deployment of pretrained MRI transformers in brain imaging tasks.
Journal: Frontiers in artificial intelligence
In common: imageio, h5py, Pillow, 6 other tools, 1 reference
[4] doi:10.1038/s41467-026-76939-w [code]
HIPPIE: a generative model for electrophysiological analysis across species, technologies, and modalities.
Journal: Nature communications
In common: AllenSDK, h5py, Pillow, 6 other tools, mouse
[5] doi:10.3389/fncom.2026.1786996 [code]
Schumann-anchored golden ratio organization of human neural oscillations.
Journal: Frontiers in computational neuroscience
In common: imageio, MNE-Python, h5py, 6 other tools
[6] 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, pandas, 3 other tools, computational, 3 references
[7] doi:10.1038/s41593-026-02232-0 [code]
Entorhinal cortex represents task-relevant remote locations independently of CA1.
Journal: Nature neuroscience
In common: xarray, h5py, Pillow, 6 other tools, mouse
[8] doi:10.1093/bioinformatics/btag540 [code]
Deciphering spatial heterogeneity by multimodal spatial transcriptomics modelling with SpatialModal.
Journal: Bioinformatics (Oxford, England)
In common: imageio, h5py, Pillow, 6 other tools, mouse
[9] doi:10.1038/s41467-026-76098-y [code]
A single computational objective can produce specialization of streams in visual cortex.
Journal: Nature communications
In common: xarray, h5py, Pillow, 6 other tools
[10] doi:10.1371/journal.pcbi.1014152 [code]
Synchronization properties in C. elegans: Relating behavioral circuits to structural and functional neuronal connectivity.
Journal: PLoS computational biology
In common: pandas, SciPy, Matplotlib, 1 other tool, C. elegans, 4 references

Contribute

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

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

Request its removal

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

Discussion, reproductions, activity

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

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

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