OSCR

Locomotion optimizes sensory representations through a computational principle shared by rodents and primates.

Code ↔ Paper

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

The 15 matches
  1. [1] § METHODS › Nonlinearity optimization ↔ fig2/optimize_all_gabor.ipynb, lines 655–702 · score 0.86 · mutual information, response entropy, Gaussian distribution, stimulus bin, response bin, optimal parameters
  2. [2] § METHODS › Optimization of inhibitory connections ↔ fig4/network_interaction_analysis.ipynb, lines 341–483 · score 0.83 · enforce positivity, Pearson correlation, loss function, network interaction, gradient, subtracted
  3. [3] § METHODS › Nonlinearity optimization ↔ fig2/compute_ori_tuning.py, lines 24–135 · score 0.74 · logistic nonlinearity, Gaussian distribution, response bin, optimal parameters, entropy, utility
  4. [4] § METHODS › Temporal dynamics ↔ fig3/temporal_filtering_analysis.ipynb, lines 1170–1208 · score 0.73 · inverse Fourier transform, white noise, autocorrelation function, temporal filters, spectra, power
  5. [5] § METHODS › Temporal filtering ↔ fig3/temporal_filtering_analysis.ipynb, lines 1170–1208 · score 0.69 · filtered spectra, white noise, temporal filtering, autocorrelation function, spectral, power
  6. [6] § METHODS › Minimal model of sensing during locomotion ↔ minimal_model/agent_sim.py, lines 74–129 · score 0.69 · starting position, minimal model, agent, circle, uniform, velocity
  7. [7] § METHODS › Orientation tuning curves ↔ fig2/optimize_all_gabor.ipynb, lines 1327–1389 · score 0.64 · stationary tuning curve, moving tuning curve, regressing, additive, Gabor filters, fit
  8. [8] § METHODS › Minimal model of sensing during locomotion ↔ minimal_model/agent_sim_analysis.ipynb, lines 78–177 · score 0.63 · starting position, minimal model, trajectories, velocity, radius, agent
  9. [9] § METHODS › Population coding fidelity ↔ fig2/decoding_error_decreases_gabor.ipynb, lines 415–443 · score 0.62 · decoding error, Gaussian filter, decoder, firing rate, smoothed, concatenated
  10. [10] § METHODS › Orientation tuning curves ↔ fig2/compute_ori_tuning.py, lines 177–226 · score 0.61 · stationary tuning curve, moving tuning curve, regressing, additive, Gabor filters, fit
  11. [11] § RESULTS › Modulation of temporal filtering ensures efficiency of sensory coding during locomotion ↔ fig3/temporal_filtering_analysis.ipynb, lines 228–257 · score 0.59 · frequency domain, temporal frequencies, temporal filter, linear, Figure 3
  12. [12] § METHODS › Analysis of data from freely moving mice ↔ supp2/visualize_data_for_anti_mod.ipynb, lines 48–142 · score 0.58 · visual angle, Gabor filter bank, cpd, resolution
  13. [13] § METHODS › Analysis of data from freely moving mice ↔ supp2/visualize_data_for_anti_mod_no_eye_movement_correction.ipynb, lines 35–128 · score 0.58 · visual angle, Gabor filter bank, cpd, resolution
  14. [14] § METHODS › Surround suppression ↔ fig4/network_interaction_analysis.ipynb, lines 1678–1720 · score 0.53 · circular mask, radii, gratings, radius, Gabor filters, fitting
  15. [15] § RESULTS › Stimulus statistics explain differential modulation of sensory coding by locomotion in rodents and primates ↔ fig5/optimize_all_gabor_species_comp.ipynb, lines 1229–1363 · score 0.51 · confidence interval, firing rate, foveal, species, optimized, modulated

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 · 2,590 lines · 123 KB · no license · 3 matches

  1. # %% [markdown]
  2. # # Load Data
  3. # %%
  4. '''
  5. This script optimizes logistic nonlinearities for a variety of Gaussian distribtuions and plots the optimal parameters.
  6. Author: Jonathan Gant
  7. Date: 29.08.2024
  8. '''
  9. import numpy as np
  10. import matplotlib.pyplot as plt
  11. from utilities import logistic_func, calc_MI, calc_entropy
  12. from tqdm import tqdm
  13. import h5py
  14. import os
  15. import bottleneck as bn
  16. # Set random seed for reproducibility
  17. np.random.seed(0)
  18. # load in the data
  19. # all_gabor_responses = h5py.File('../results/new_nat_videos_gabor_responses_full_res_z_score.h5', 'r')
  20. all_gabor_responses = h5py.File('../results/new_nat_videos_gabor_responses_full_res_z_score_more_freq.h5', 'r')
  21. # all_gabor_responses_low_freq = h5py.File('../results/new_nat_videos_gabor_responses_full_res_more_low_freq_z_score.h5', 'r')
  22. # all_gabor_responses_eye_movements = h5py.File('../results/new_nat_videos_gabor_responses_full_res_more_low_freq_z_score_eye_movements.h5', 'r')
  23. # all_gabor_responses_eye_movements_long = h5py.File('../results/new_nat_videos_gabor_responses_full_res_more_low_freq_z_score_eye_movements_2s_interval_stat_only.h5', 'r')
  24. # all_gabor_responses = h5py.File('../results/new_nat_vids_gabor_responses_full_res.h5', 'r')
  25. # video size
  26. resolution_height = 1080
  27. resolution_width = 1920
  28. # fov
  29. horizontal_fov = 92
  30. vertical_fov = 61
  31. # conversion factor of pixels to degrees
  32. horizontal_pixels_per_degree = resolution_width / horizontal_fov
  33. vertical_pixels_per_degree = resolution_height / vertical_fov
  34. # average of the conversion factors to the nearest integer
  35. pixels_per_degree = np.ceil((horizontal_pixels_per_degree + vertical_pixels_per_degree) / 2)
  36. print(pixels_per_degree)
  37. # data hyperparameters
  38. orientation_arr = all_gabor_responses['orientation_arr'][()]
  39. phase_arr = all_gabor_responses['phase_arr'][()]
  40. position_arr = all_gabor_responses['position_arr'][()]
  41. wavelength_arr = all_gabor_responses['wavelength_arr'][()]
  42. # wavelength_arr_low_freq = all_gabor_responses_low_freq['wavelength_arr'][()]
  43. freq_arr = pixels_per_degree / wavelength_arr
  44. # freq_arr_low_freq = pixels_per_degree / wavelength_arr_low_freq
  45. filter_size = (resolution_height, resolution_width)
  46. print(freq_arr)
  47. low_spatial_freq_idx = np.arange(0, 31)
  48. high_spatial_freq_idx = np.arange(35, 70)
  49. # %%
  50. plt.plot(all_gabor_responses['field']['stationary_1'][0, 0, 0, 4, :], color='tab:red')
  51. plt.plot(all_gabor_responses['field']['moving_1'][0, 0, 0, 4, :], color='tab:blue')
  52. # %% [markdown]
  53. # # Temporal filtering
  54. # %% [markdown]
  55. # ## Load and average PSD
  56. # %%
  57. # parse the data and compute fourier transform for each filter and video
  58. environments = ['field', 'forest', 'orchard', 'tall_grass', 'pond']
  59. num_videos = 10
  60. vid_length = 50*30
  61. low_spatial_freq_idx = np.arange(0, 31)
  62. # low_spatial_freq_idx = np.arange(0, 19)
  63. # high_spatial_freq_idx = np.arange(35, 70)
  64. # low_spatial_freq_idx = np.arange(35, 70)
  65. stationary_stim = np.zeros((len(environments), num_videos, len(orientation_arr), len(phase_arr), len(low_spatial_freq_idx), len(position_arr), vid_length))
  66. moving_stim = np.zeros((len(environments), num_videos, len(orientation_arr), len(phase_arr), len(low_spatial_freq_idx), len(position_arr), vid_length))
  67. # get the responses for each environment
  68. for i, env_key in enumerate(environments):
  69. stationary_count = 0
  70. moving_count = 0
  71. all_gabor_responses_env = all_gabor_responses[env_key]
  72. print(all_gabor_responses_env.keys())
  73. for vid_key in all_gabor_responses_env.keys():
  74. if 'stationary' in vid_key:
  75. stationary_stim[i, stationary_count, :, :, :, :, :] = all_gabor_responses_env[vid_key][()][:, :, low_spatial_freq_idx, :, :vid_length]
  76. stationary_count += 1
  77. if 'moving' in vid_key and 'free_moving' not in vid_key:
  78. moving_stim[i, moving_count, :, :, :, :, :] = all_gabor_responses_env[vid_key][()][:, :, low_spatial_freq_idx, :, :vid_length]
  79. moving_count += 1
  80. # %%
  81. def compute_psd_rfft(x, fs):
  82. """
  83. Compute one-sided PSD using rfft for a real-valued signal.
  84. Parameters:
  85. - x: 1D real-valued input signal
  86. - fs: sampling frequency (Hz)
  87. Returns:
  88. - freqs: frequency bins (Hz)
  89. - psd: power spectral density (power/Hz)
  90. """
  91. N = len(x)
  92. X = np.fft.rfft(x, axis=-1)
  93. X_mag2 = np.abs(X) ** 2
  94. psd = X_mag2 / (N * fs)
  95. # Multiply by 2 to account for negative freqs (except DC and Nyquist if N even)
  96. if N % 2 == 0:
  97. psd[1:-1] *= 2
  98. else:
  99. psd[1:] *= 2
  100. return psd
  101. # %%
  102. # compute the fourier transform for each filter and video using scipy
  103. from scipy.fft import rfftn, rfftfreq
  104. sampling_rate = 30
  105. stationary_stim_psd = np.zeros((len(environments), num_videos, len(orientation_arr), len(phase_arr), len(low_spatial_freq_idx), len(position_arr), vid_length//2+1))
  106. moving_stim_psd = np.zeros((len(environments), num_videos, len(orientation_arr), len(phase_arr), len(low_spatial_freq_idx), len(position_arr), vid_length//2+1))
  107. for i, env_key in enumerate(environments):
  108. for j in range(num_videos):
  109. stationary_stim_psd[i, j] = compute_psd_rfft(stationary_stim[i, j], sampling_rate)
  110. moving_stim_psd[i, j] = compute_psd_rfft(moving_stim[i, j], sampling_rate)
  111. # %%
  112. # compute the temporal frequencies
  113. sampling_rate = 30
  114. num_samples = vid_length
  115. # frequencies = fftfreq(num_samples, 1/sampling_rate)[1:vid_length//2]
  116. frequencies = rfftfreq(num_samples, 1/sampling_rate)
  117. # pick a cutoff frequency
  118. rf_window = vid_length/sampling_rate # .5 # seconds
  119. cutoff_freq = 1/rf_window
  120. # find the first index wheter the frequency is greater than the cutoff frequency
  121. cutoff_idx = np.where(frequencies >= cutoff_freq)[0][0]
  122. # %% [markdown]
  123. # ## Welch method
  124. # %%
  125. from scipy.signal import welch
  126. def compute_psd_welch(x, fs, window_size, overlap):
  127. """
  128. Compute one-sided PSD using Welch's method for a real-valued signal with variable window size and full overlap.
  129. Parameters:
  130. - x: 1D real-valued input signal
  131. - fs: sampling frequency (Hz)
  132. - window_size: length of each segment (samples)
  133. Returns:
  134. - freqs: frequency bins (Hz)
  135. - psd: power spectral density (power/Hz)
  136. """
  137. freqs, psd = welch(
  138. x,
  139. fs=fs,
  140. window='boxcar',
  141. nperseg=window_size,
  142. noverlap=overlap, # full overlap
  143. return_onesided=True,
  144. scaling='density',
  145. detrend='constant'
  146. )
  147. return freqs, psd
  148. # %%
  149. # compute the fourier transform for each filter and video using scipy
  150. from scipy.fft import rfftn, rfftfreq
  151. sampling_rate = 30
  152. window_size = int(5 * sampling_rate) # 5 seconds window size
  153. stationary_stim_psd = np.zeros((len(environments), num_videos, len(orientation_arr), len(phase_arr), len(low_spatial_freq_idx), len(position_arr), window_size//2+1))
  154. moving_stim_psd = np.zeros((len(environments), num_videos, len(orientation_arr), len(phase_arr), len(low_spatial_freq_idx), len(position_arr), window_size//2+1))
  155. for i, env_key in enumerate(environments):
  156. for j in range(num_videos):
  157. _, stationary_stim_psd[i, j] = compute_psd_welch(stationary_stim[i, j], sampling_rate, window_size, window_size//2)
  158. frequencies, moving_stim_psd[i, j] = compute_psd_welch(moving_stim[i, j], sampling_rate, window_size, window_size//2)
  159. # %%
  160. # compute the PSD as the square of the absolute value of the fourier transform, also truncate the DC component and the negative frequencies
  161. # set the DC to zero
  162. # stationary_stim_freq_no_dc = np.zeros_like(stationary_stim_psd)
  163. # stationary_stim_freq_no_dc[:, :, :, :, :, :, cutoff_idx:] = stationary_stim_psd[:, :, :, :, :, :, cutoff_idx:]
  164. # moving_stim_freq_no_dc = np.zeros_like(moving_stim_psd)
  165. # moving_stim_freq_no_dc[:, :, :, :, :, :, cutoff_idx:] = moving_stim_psd[:, :, :, :, :, :, cutoff_idx:]
  166. # stationary_stim_psd_avg = np.mean(stationary_stim_freq_no_dc, axis=(0, 1))
  167. # moving_stim_psd_avg = np.mean(moving_stim_freq_no_dc, axis=(0, 1))
  168. stationary_stim_psd_avg = np.mean(stationary_stim_psd, axis=(0, 1))
  169. moving_stim_psd_avg = np.mean(moving_stim_psd, axis=(0, 1))
  170. # %%
  171. ori_idx = 1
  172. phase_idx = 3
  173. freq_idx = 10
  174. pos_idx = 4
  175. plt.plot(frequencies[1:-1], moving_stim_psd_avg[ori_idx, phase_idx, freq_idx, pos_idx, 1:-1], label='Moving', color='tab:orange')
  176. plt.plot(frequencies[1:-1], stationary_stim_psd_avg[ori_idx, phase_idx, freq_idx, pos_idx, 1:-1], label='Stationary', color='tab:gray')
  177. # plt.plot(frequencies[cutoff_idx:], moving_stim_psd_avg[ori_idx, phase_idx, freq_idx, pos_idx, cutoff_idx:], label='Moving', color='tab:orange')
  178. # plt.plot(frequencies[cutoff_idx:], stationary_stim_psd_avg[ori_idx, phase_idx, freq_idx, pos_idx, cutoff_idx:], label='Stationary', color='tab:gray')
  179. plt.xlabel('Frequency (Hz)')
  180. plt.ylabel('Power')
  181. plt.yscale('log')
  182. plt.savefig(f'../manuscript_figures/fig3_psd_example_ori_{ori_idx}_phase_{phase_idx}_freq_{freq_idx}_pos_{pos_idx}.pdf', format='pdf', bbox_inches='tight')
  183. # %% [markdown]
  184. # ## Apply filter
  185. # %%
  186. # definte the temporal filter in the frequency domain with a linear dropoff at a specific frequency determined by a free parameter theta. Define the filter as a function of a set of temporal frequencies
  187. def temporal_filter(temporal_freqs, theta, low_freq_cutoff=2):
  188. # create the filter
  189. filter = np.zeros(temporal_freqs.shape)
  190. # set the filter to 1 for frequencies below theta
  191. filter[temporal_freqs <= theta] = 1
  192. # filter[temporal_freqs < low_freq_cutoff] = 1 - (low_freq_cutoff - temporal_freqs[temporal_freqs < low_freq_cutoff])
  193. # create a linear dropoff for frequencies above theta with a fixed slope
  194. filter[temporal_freqs > theta] = 1 - (temporal_freqs[temporal_freqs > theta] - theta)
  195. # ensure that the filter is non-negative
  196. filter[filter < 0] = 0
  197. return np.square(filter)
  198. def temporal_filter_gauss(temporal_freqs, theta_u, theta_s):
  199. # create the filter as a Gaussian in the frequency domain with mean theta_u and standard deviation theta_s
  200. filter = np.exp(-0.5 * ((temporal_freqs - theta_u) / theta_s)**2)
  201. # ensure that the filter is non-negative
  202. filter[filter < 0] = 0
  203. return np.square(filter)
  204. def temporal_filter_exp(temporal_freq, theta):
  205. # create the filter as an exponential in the frequency domain with mean theta_u and standard deviation theta_s
  206. filter = np.exp(-temporal_freq/theta)
  207. # ensure that the filter is non-negative
  208. filter[filter < 0] = 0
  209. return np.square(filter)
  210. # %%
  211. # take the inverse fourier transform of the sqrt of the temporal filter
  212. from scipy.fft import irfftn, ifftshift
  213. # temporal_filter_freq = temporal_filter_gauss(frequencies, 0, 5)
  214. temporal_filter_freq = temporal_filter(frequencies, 6)
  215. temporal_filter_freq[0] = 0
  216. # temporal_filter_freq = temporal_filter(freq, 3)
  217. # temporal_filter_freq = temporal_filter_exp(frequencies, 5)
  218. inverse_temporal_filter = irfftn(np.sqrt(temporal_filter_freq), axes=-1)[:len(frequencies)]
  219. # %%
  220. # plot the filter in the frequency domain
  221. plt.figure(figsize=(8, 6))
  222. plt.plot(frequencies[1:], np.sqrt(temporal_filter_freq[1:]), label='Temporal filter', color='tab:blue')
  223. plt.xlabel('Frequency (Hz)')
  224. plt.ylabel('Filter amplitude')
  225. plt.savefig('../manuscript_figures/fig3_temporal_filter_freq.pdf', format='pdf', bbox_inches='tight')
  226. # %%
  227. plt.plot(np.linspace(0, 5000, len(frequencies)), inverse_temporal_filter)
  228. plt.axhline(0, color='k', linestyle='--')
  229. plt.xlim(0, 2000)
  230. # %%
  231. # simulate white noise
  232. noise_level = 100
  233. # noise = np.ones(frequencies.shape) * noise_level
  234. noise = np.ones(freq.shape) * noise_level
  235. # %%
  236. # repeat the process again for a range of theta values
  237. num_theta_values = 100
  238. array_shape = moving_stim_psd_avg.shape
  239. theta_values = np.linspace(2, 15, num_theta_values)
  240. moving_stim_information = np.zeros((num_theta_values, array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  241. stationary_stim_information = np.zeros((num_theta_values, array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  242. moving_stim_energy = np.zeros((num_theta_values, array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  243. stationary_stim_energy = np.zeros((num_theta_values, array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  244. moving_stim_energy_linear = np.zeros((num_theta_values, array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  245. stationary_stim_energy_linear = np.zeros((num_theta_values, array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  246. moving_stim_energy_quadratic = np.zeros((num_theta_values, array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  247. stationary_stim_energy_quadratic = np.zeros((num_theta_values, array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  248. for i in tqdm(range(num_theta_values)):
  249. # compute the filter
  250. filter = temporal_filter(frequencies, theta_values[i])
  251. # compute the filtered spectra
  252. filtered_stationary_stim_psd_avg = np.multiply(stationary_stim_psd_avg, filter)
  253. filtered_moving_stim_psd_avg = np.multiply(moving_stim_psd_avg, filter)
  254. # compute the information
  255. moving_stim_information[i] = np.sum(np.log2(1 + filtered_moving_stim_psd_avg / noise), axis=-1)
  256. stationary_stim_information[i] = np.sum(np.log2(1 + filtered_stationary_stim_psd_avg / noise), axis=-1)
  257. # compute the energy
  258. moving_stim_energy[i] = np.sum(filtered_moving_stim_psd_avg, axis=-1)
  259. stationary_stim_energy[i] = np.sum(filtered_stationary_stim_psd_avg, axis=-1)
  260. # compute the energy weighted by the square of the frequencies
  261. moving_stim_energy_quadratic[i] = np.sum(filtered_moving_stim_psd_avg*frequencies**2, axis=-1)
  262. stationary_stim_energy_quadratic[i] = np.sum(filtered_stationary_stim_psd_avg*frequencies**2, axis=-1)
  263. # compute the energy linearly weighted by frequency
  264. moving_stim_energy_linear[i] = np.sum(filtered_moving_stim_psd_avg*frequencies, axis=-1)
  265. stationary_stim_energy_linear[i] = np.sum(filtered_stationary_stim_psd_avg*frequencies, axis=-1)
  266. # %%
  267. # plot an example of the psd before and after filtering
  268. fig, ax = plt.subplots(1, 2, figsize=(16, 6), sharey=True)
  269. filter = temporal_filter(frequencies, 6)
  270. ori_idx = 0
  271. phase_idx = 0
  272. wavelength_idx = 8
  273. position_idx = 4
  274. ax[0].plot(frequencies, moving_stim_psd_avg[ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Moving', color='tab:blue')
  275. ax[0].plot(frequencies, stationary_stim_psd_avg[ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Stationary', color='tab:red')
  276. ax[1].plot(frequencies, np.multiply(moving_stim_psd_avg[ori_idx, phase_idx, wavelength_idx, position_idx, :], filter), label='Moving', color='tab:blue')
  277. ax[1].plot(frequencies, np.multiply(stationary_stim_psd_avg[ori_idx, phase_idx, wavelength_idx, position_idx, :], filter), label='Stationary', color='tab:red')
  278. ax[1].set_yscale('log')
  279. # %%
  280. moving_test_informations = moving_stim_information[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  281. stationary_test_informations = stationary_stim_information[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  282. moving_test_energy = moving_stim_energy[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  283. stationary_test_energy = stationary_stim_energy[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  284. # create a subplot with 1 row and 3 columns. First column plot the information vs. theta, second column plot the energy vs. theta, third column plot the information vs. energy
  285. fig, axs = plt.subplots(1, 3, figsize=(15, 5))
  286. axs[0].plot(theta_values, moving_test_informations, color='tab:blue', label='Moving test')
  287. axs[0].plot(theta_values, stationary_test_informations, color='tab:red', label='Stationary test')
  288. axs[0].set_xlabel('Theta')
  289. axs[0].set_ylabel('Information')
  290. axs[0].set_title('Information vs. theta')
  291. axs[0].legend()
  292. axs[1].plot(theta_values, moving_test_energy, color='tab:blue', label='Moving test')
  293. axs[1].plot(theta_values, stationary_test_energy, color='tab:red', label='Stationary test')
  294. axs[1].set_xlabel('Theta')
  295. axs[1].set_ylabel('Energy')
  296. axs[1].set_title('Energy vs. theta')
  297. axs[1].legend()
  298. axs[2].plot(moving_test_energy, moving_test_informations, color='tab:blue', label='Moving test')
  299. axs[2].plot(stationary_test_energy, stationary_test_informations, color='tab:red', label='Stationary test')
  300. axs[2].set_xlabel('Energy')
  301. axs[2].set_ylabel('Information')
  302. axs[2].set_title('Information vs. energy')
  303. axs[2].legend()
  304. plt.tight_layout()
  305. plt.show()
  306. # %%
  307. lmbd_arr = np.logspace(-5, -3, 16)
  308. fig, ax = plt.subplots(4, 4, figsize=(16, 16))
  309. for i, lmbd in enumerate(lmbd_arr):
  310. moving_utility = moving_test_informations - lmbd*moving_test_energy
  311. stationary_utility = stationary_test_informations - lmbd*stationary_test_energy
  312. ax[i//4, i%4].plot(theta_values, moving_utility, color='tab:blue', label='Moving test')
  313. ax[i//4, i%4].plot(theta_values, stationary_utility, color='tab:red', label='Stationary test')
  314. ax[i//4, i%4].set_title('Lambda = ' + str(np.round(lmbd, 7)))
  315. # %%
  316. # for each filter, for a given energy find the closest value of theta which has that energy
  317. energy_arr = np.logspace(6, 11, 1000)
  318. optimal_theta_val_mov = np.zeros((len(energy_arr), len(orientation_arr), len(phase_arr), len(wavelength_arr[low_spatial_freq_idx]), len(position_arr)))
  319. optimal_theta_val_stat = np.zeros((len(energy_arr), len(orientation_arr), len(phase_arr), len(wavelength_arr[low_spatial_freq_idx]), len(position_arr)))
  320. optimal_theta_val_mov_linear = np.zeros((len(energy_arr), len(orientation_arr), len(phase_arr), len(wavelength_arr[low_spatial_freq_idx]), len(position_arr)))
  321. optimal_theta_val_stat_linear = np.zeros((len(energy_arr), len(orientation_arr), len(phase_arr), len(wavelength_arr[low_spatial_freq_idx]), len(position_arr)))
  322. optimal_theta_val_mov_quadratic = np.zeros((len(energy_arr), len(orientation_arr), len(phase_arr), len(wavelength_arr[low_spatial_freq_idx]), len(position_arr)))
  323. optimal_theta_val_stat_quadratic = np.zeros((len(energy_arr), len(orientation_arr), len(phase_arr), len(wavelength_arr[low_spatial_freq_idx]), len(position_arr)))
  324. # loop over orientation, phase, wavelength, position
  325. for ori_idx in range(len(orientation_arr)):
  326. for phase_idx in range(len(phase_arr)):
  327. for wavelength_idx in range(len(wavelength_arr[low_spatial_freq_idx])):
  328. for position_idx in range(len(position_arr)):
  329. test_moving_energy = moving_stim_energy[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  330. test_stationary_energy = stationary_stim_energy[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  331. # interpolate the energy values to get a smooth curve
  332. test_moving_theta_interp = np.interp(energy_arr, test_moving_energy, theta_values)
  333. test_stationary_theta_interp = np.interp(energy_arr, test_stationary_energy, theta_values)
  334. optimal_theta_val_mov[:, ori_idx, phase_idx, wavelength_idx, position_idx] = test_moving_theta_interp
  335. optimal_theta_val_stat[:, ori_idx, phase_idx, wavelength_idx, position_idx] = test_stationary_theta_interp
  336. # repeat for the linear energy
  337. test_moving_energy_linear = moving_stim_energy_linear[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  338. test_stationary_energy_linear = stationary_stim_energy_linear[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  339. # interpolate the energy values to get a smooth curve
  340. test_moving_theta_interp_linear = np.interp(energy_arr, test_moving_energy_linear, theta_values)
  341. test_stationary_theta_interp_linear = np.interp(energy_arr, test_stationary_energy_linear, theta_values)
  342. optimal_theta_val_mov_linear[:, ori_idx, phase_idx, wavelength_idx, position_idx] = test_moving_theta_interp_linear
  343. optimal_theta_val_stat_linear[:, ori_idx, phase_idx, wavelength_idx, position_idx] = test_stationary_theta_interp_linear
  344. # repeat for the quadratic energy
  345. test_moving_energy_quadratic = moving_stim_energy_quadratic[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  346. test_stationary_energy_quadratic = stationary_stim_energy_quadratic[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  347. # interpolate the energy values to get a smooth curve
  348. test_moving_theta_interp_quadratic = np.interp(energy_arr, test_moving_energy_quadratic, theta_values)
  349. test_stationary_theta_interp_quadratic = np.interp(energy_arr, test_stationary_energy_quadratic, theta_values)
  350. optimal_theta_val_mov_quadratic[:, ori_idx, phase_idx, wavelength_idx, position_idx] = test_moving_theta_interp_quadratic
  351. optimal_theta_val_stat_quadratic[:, ori_idx, phase_idx, wavelength_idx, position_idx] = test_stationary_theta_interp_quadratic
  352. # %%
  353. ori_idx = 0
  354. phase_idx = 0
  355. wavelength_idx = -1
  356. position_idx = 4
  357. # plot the difference in theta values vs. energy
  358. plt.plot(energy_arr, optimal_theta_val_mov[:, ori_idx, phase_idx, wavelength_idx, position_idx]-optimal_theta_val_stat[:, ori_idx, phase_idx, wavelength_idx, position_idx], color='tab:purple')
  359. # %%
  360. optimal_theta_diff = optimal_theta_val_mov - optimal_theta_val_stat
  361. optimal_theta_diff_avg = np.mean(optimal_theta_diff, axis=(1, 2, 4))
  362. optimal_theta_diff_linear = optimal_theta_val_mov_linear - optimal_theta_val_stat_linear
  363. optimal_theta_diff_linear_avg = np.mean(optimal_theta_diff_linear, axis=(1, 2, 4))
  364. optimal_theta_diff_quadratic = optimal_theta_val_mov_quadratic - optimal_theta_val_stat_quadratic
  365. optimal_theta_diff_quadratic_avg = np.mean(optimal_theta_diff_quadratic, axis=(1, 2, 4))
  366. # %%
  367. # plot the average difference in optimal theta, color them with viridis
  368. from matplotlib import cm
  369. from matplotlib.colors import Normalize
  370. from mpl_toolkits.axes_grid1.inset_locator import inset_axes
  371. colors_map = cm.get_cmap('viridis', len(wavelength_arr[low_spatial_freq_idx]))
  372. fig, ax = plt.subplots(1, 3, figsize=(4*3*1.5, 3*1.5), sharey=True)
  373. for i in range(wavelength_arr[low_spatial_freq_idx].shape[0]):
  374. ax[0].plot(energy_arr, optimal_theta_diff_avg[:, i], color=colors_map(i))
  375. ax[1].plot(energy_arr, optimal_theta_diff_linear_avg[:, i], color=colors_map(i))
  376. ax[2].plot(energy_arr, optimal_theta_diff_quadratic_avg[:, i], color=colors_map(i))
  377. ax[0].set_xscale('log')
  378. ax[1].set_xscale('log')
  379. ax[2].set_xscale('log')
  380. ax[0].set_title('Unweighted energy')
  381. ax[1].set_title('Linearly weighted energy')
  382. ax[2].set_title('Quadratically weighted energy')
  383. ax[0].set_xlabel('Energy constraint')
  384. ax[1].set_xlabel('Energy constraint')
  385. ax[2].set_xlabel('Energy constraint')
  386. ax[0].set_ylabel('$\\Delta\\theta$ [moving - stationary]')
  387. # add something to indicate that the darker the color, the lower the spatial frequency
  388. # add a colorbar
  389. norm = Normalize(vmin=min(freq_arr[low_spatial_freq_idx]), vmax=max(freq_arr[low_spatial_freq_idx]))
  390. sm = plt.cm.ScalarMappable(cmap=colors_map, norm=norm)
  391. sm.set_array([])
  392. axins1 = inset_axes(
  393. ax[0],
  394. width="25%", # width: 50% of parent_bbox width
  395. height="5%", # height: 5%
  396. loc="upper right",
  397. )
  398. axins1.xaxis.set_ticks_position("bottom")
  399. cbar = fig.colorbar(sm, cax=axins1, orientation="horizontal")
  400. cbar.set_label('spatial freq.\n[cycles/degree]')
  401. axins2 = inset_axes(
  402. ax[1],
  403. width="25%", # width: 50% of parent_bbox width
  404. height="5%", # height: 5%
  405. loc="upper right",
  406. )
  407. axins2.xaxis.set_ticks_position("bottom")
  408. cbar = fig.colorbar(sm, cax=axins2, orientation="horizontal")
  409. cbar.set_label('spatial freq.\n[cycles/degree]')
  410. axins2 = inset_axes(
  411. ax[2],
  412. width="25%", # width: 50% of parent_bbox width
  413. height="5%", # height: 5%
  414. loc="upper right",
  415. )
  416. axins2.xaxis.set_ticks_position("bottom")
  417. cbar = fig.colorbar(sm, cax=axins2, orientation="horizontal")
  418. cbar.set_label('spatial freq.\n[cycles/degree]')
  419. # %%
  420. ori_idx = 0
  421. phase_idx = 0
  422. position_idx = 4
  423. wavelength_idices = [0, 15, -1]
  424. moving_test_energy = moving_stim_energy[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  425. stationary_test_energy = stationary_stim_energy[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  426. # create a subplot with 1 row and 3 columns. First column plot the information vs. theta, second column plot the energy vs. theta, third column plot the information vs. energy
  427. fig, axs = plt.subplots(1, 3, figsize=(15, 5), sharex=True)
  428. axs[0].plot(theta_values, moving_stim_energy[:, ori_idx, phase_idx, wavelength_idices[0], position_idx], color='tab:blue', label='Moving test')
  429. axs[0].plot(theta_values, stationary_stim_energy[:, ori_idx, phase_idx, wavelength_idices[0], position_idx], color='tab:red', label='Stationary test')
  430. axs[0].set_xlabel('Theta')
  431. axs[0].set_ylabel('Energy')
  432. axs[0].set_title('Low frequency')
  433. axs[0].legend()
  434. axs[1].plot(theta_values, moving_stim_energy[:, ori_idx, phase_idx, wavelength_idices[1], position_idx], color='tab:blue', label='Moving test')
  435. axs[1].plot(theta_values, stationary_stim_energy[:, ori_idx, phase_idx, wavelength_idices[1], position_idx], color='tab:red', label='Stationary test')
  436. axs[1].set_xlabel('Theta')
  437. axs[1].set_ylabel('Energy')
  438. axs[1].set_title('Mid frequency')
  439. axs[1].legend()
  440. axs[2].plot(theta_values, moving_stim_energy[:, ori_idx, phase_idx, wavelength_idices[2], position_idx], color='tab:blue', label='Moving test')
  441. axs[2].plot(theta_values, stationary_stim_energy[:, ori_idx, phase_idx, wavelength_idices[2], position_idx], color='tab:red', label='Stationary test')
  442. axs[2].set_xlabel('Theta')
  443. axs[2].set_ylabel('Energy')
  444. axs[2].set_title('High frequency')
  445. axs[2].legend()
  446. plt.tight_layout()
  447. plt.show()
  448. # %%
  449. # pick a contraint and then see how a stimulus filtered by the optimal filter looks
  450. constraint_energy = 1e9
  451. constraint_idx = np.argmin(np.abs(energy_arr-constraint_energy))
  452. optimal_theta_mov = optimal_theta_val_mov[constraint_idx, :, :, :, :]
  453. optimal_theta_stat = optimal_theta_val_stat[constraint_idx, :, :, :, :]
  454. ori_idx = 0
  455. phase_idx = 0
  456. wavelength_idx = 3
  457. position_idx = 4
  458. theta_mov = optimal_theta_mov[ori_idx, phase_idx, wavelength_idx, position_idx]
  459. theta_stat = optimal_theta_stat[ori_idx, phase_idx, wavelength_idx, position_idx]
  460. # compute the filter
  461. filter_mov = temporal_filter(frequencies, theta_mov)
  462. filter_stat = temporal_filter(frequencies, theta_stat)
  463. # compute the filtered spectra
  464. filtered_moving_stim_psd = np.multiply(noise, filter_mov)
  465. filtered_stationary_stim_psd = np.multiply(noise, filter_stat)
  466. # compute the energy
  467. test_moving_stim_energy = np.sum(filtered_moving_stim_psd, axis=-1)
  468. test_stationary_stim_energy = np.sum(filtered_stationary_stim_psd, axis=-1)
  469. # compute the energy weighted by the square of the frequencies
  470. test_moving_stim_energy_quadratic = np.sum(filtered_moving_stim_psd*frequencies**2, axis=-1)
  471. test_stationary_stim_energy_quadratic = np.sum(filtered_stationary_stim_psd*frequencies**2, axis=-1)
  472. # compute the energy linearly weighted by frequency
  473. test_moving_stim_energy_linear = np.sum(filtered_moving_stim_psd*frequencies, axis=-1)
  474. test_stationary_stim_energy_linear = np.sum(filtered_stationary_stim_psd*frequencies, axis=-1)
  475. # %%
  476. # create a figure which shows the spectra before filtering, the filter, and the spectra after filtering
  477. fig, ax = plt.subplots(1, 3, figsize=(15, 5))
  478. ax[0].plot(frequencies, noise, label='Noise', color='tab:gray')
  479. ax[1].plot(frequencies, filter_mov, label='Moving', color='tab:blue')
  480. ax[1].plot(frequencies, filter_stat, label='Stationary', color='tab:red')
  481. ax[2].plot(frequencies, filtered_moving_stim_psd, label='Moving', color='tab:blue')
  482. ax[2].plot(frequencies, filtered_stationary_stim_psd, label='Stationary', color='tab:red')
  483. # %%
  484. # compute the autocorrelation function by taking the inverse fourier transform of the PSD
  485. from scipy.fft import irfft
  486. # compute the autocorrelation function
  487. moving_autocorr = irfft(filtered_moving_stim_psd)[:len(filtered_moving_stim_psd)//2]
  488. stationary_autocorr = irfft(filtered_stationary_stim_psd)[:len(filtered_stationary_stim_psd)//2]
  489. # normalize the autocorrelation function
  490. moving_autocorr /= np.max(moving_autocorr)
  491. stationary_autocorr /= np.max(stationary_autocorr)
  492. time = np.arange(0, len(moving_autocorr)) / 30
  493. # plot the autocorrelation function
  494. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  495. ax.plot(time, moving_autocorr, label='Moving', color='tab:blue')
  496. ax.plot(time, stationary_autocorr, label='Stationary', color='tab:red')
  497. ax.set_xlabel('Time (s)')
  498. ax.set_ylabel('Autocorrelation')
  499. ax.set_title('Autocorrelation function')
  500. ax.legend()
  501. ax.set_xlim(0,2)
  502. plt.show()
  503. # %% [markdown]
  504. # ## Heatmap energy landscape and theta
  505. # %% [markdown]
  506. # ### Average over filters
  507. # %%
  508. moving_stim_psd.shape
  509. # %%
  510. # compute the average spectrum over all filters
  511. mean_moving_psd = np.mean(moving_stim_psd, axis=(0, 1, 2, 3, 4, 5))
  512. mean_stationary_psd = np.mean(stationary_stim_psd, axis=(0, 1, 2, 3, 4, 5))
  513. # plot the mean spectrum
  514. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  515. ax.plot(frequencies[1:-1], mean_stationary_psd[1:-1], label='Stationary', color='tab:gray')
  516. ax.plot(frequencies[1:-1], mean_moving_psd[1:-1], label='Moving', color='tab:orange')
  517. ax.set_xlabel('Frequency (Hz)')
  518. ax.set_ylabel('Power')
  519. ax.set_title('Mean power spectrum')
  520. ax.set_yscale('log')
  521. ax.legend()
  522. plt.savefig(f'../manuscript_figures/fig3_psd_mean_full.pdf', format='pdf', bbox_inches='tight')
  523. # %%
  524. # normalize the spectrum such that the are under the curve is 1
  525. freq_spacing = np.diff(frequencies)[0]
  526. norm_factor = np.sum(mean_stationary_psd*freq_spacing)
  527. mean_stationary_psd /= norm_factor
  528. mean_moving_psd /= norm_factor
  529. # %%
  530. np.cumsum(mean_stationary_psd*freq_spacing)
  531. # %%
  532. np.where(np.cumsum(mean_moving_psd*freq_spacing) <= auc_val)
  533. # %%
  534. # compute the average spectrum over all filters
  535. # plot the mean spectrum
  536. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  537. ax.plot(frequencies[1:-1], mean_moving_psd[1:-1], label='Moving', color='tab:orange')
  538. ax.plot(frequencies[1:-1], mean_stationary_psd[1:-1], label='Stationary', color='tab:gray')
  539. ax.set_xlabel('Frequency (Hz)')
  540. ax.set_ylabel('Power')
  541. ax.set_title('Mean power spectrum')
  542. ax.set_yscale('log')
  543. # ax.set_xscale('log')
  544. ax.legend()
  545. # find where the area under the stationary psd is a certain value
  546. # auc_val = 1
  547. # aoc_idx_stat = np.where(np.cumsum(mean_stationary_psd*freq_spacing) <= auc_val)[0][-1]
  548. # # find the frequency at that index
  549. # freq_val_stat = frequencies[aoc_idx_stat]
  550. # auc_idx_mov = np.where(np.cumsum(mean_moving_psd*freq_spacing) <= auc_val)[0][-1]
  551. # # find the frequency at that index
  552. # freq_val_mov = frequencies[auc_idx_mov]
  553. auc_val = .95
  554. cum_stationary = np.cumsum(mean_stationary_psd * freq_spacing)
  555. cum_moving = np.cumsum(mean_moving_psd * freq_spacing)
  556. # Interpolate to find the frequency where the cumulative sum reaches auc_val
  557. freq_val_stat = np.interp(auc_val, cum_stationary, frequencies)
  558. freq_val_mov = np.interp(auc_val, cum_moving, frequencies)
  559. # # plot filled area under the curve
  560. # ax.fill_between(frequencies[1:-1], mean_stationary_psd[1:-1], where=(frequencies[1:-1] <= freq_val_stat), color='tab:gray', alpha=0.5)
  561. # ax.fill_between(frequencies[1:-1], mean_moving_psd[1:-1], where=(frequencies[1:-1] <= freq_val_mov), color='tab:orange', alpha=0.5)
  562. # Interpolate the PSD at the cutoff frequencies
  563. psd_val_stat = np.interp(freq_val_stat, frequencies, mean_stationary_psd)
  564. psd_val_mov = np.interp(freq_val_mov, frequencies, mean_moving_psd)
  565. # For fill_between, extend the frequency and PSD arrays to include the interpolated cutoff point
  566. def extend_for_fill(frequencies, psd, cutoff_freq, cutoff_psd):
  567. mask = frequencies <= cutoff_freq
  568. # Find the last index before the cutoff
  569. last_idx = np.where(mask)[0][-1]
  570. # Insert the cutoff point after last_idx
  571. new_freqs = np.insert(frequencies[mask], last_idx + 1, cutoff_freq)
  572. new_psd = np.insert(psd[mask], last_idx + 1, cutoff_psd)
  573. return new_freqs, new_psd
  574. freqs_stat_fill, psd_stat_fill = extend_for_fill(frequencies[1:-1], mean_stationary_psd[1:-1], freq_val_stat, psd_val_stat)
  575. freqs_mov_fill, psd_mov_fill = extend_for_fill(frequencies[1:-1], mean_moving_psd[1:-1], freq_val_mov, psd_val_mov)
  576. # plot filled area under the curve using the interpolated cutoff
  577. ax.fill_between(freqs_stat_fill, psd_stat_fill, color='tab:gray', alpha=0.5)
  578. ax.fill_between(freqs_mov_fill, psd_mov_fill, color='tab:orange', alpha=0.5)
  579. plt.savefig(f'../manuscript_figures/fig3_psd_mean_normalize_2Hz_crop_auc_{auc_val}.pdf', format='pdf', bbox_inches='tight')
  580. # %%
  581. psd_stat_fill
  582. # %%
  583. # plot the energy vs. theta
  584. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  585. ax.plot(frequencies[1:], np.cumsum(mean_stationary_psd*freq_spacing)[1:], label='Stationary', color='tab:gray')
  586. ax.plot(frequencies[1:], np.cumsum(mean_moving_psd*freq_spacing)[1:], label='Moving', color='tab:orange')
  587. ax.axhline(y=auc_val, color='k', linestyle='--')
  588. ax.set_xlabel('Frequency (Hz)')
  589. ax.set_ylabel('Energy')
  590. ax.set_title('Energy vs. theta')
  591. ax.set_yscale('log')
  592. ax.set_xscale('log')
  593. ax.legend()
  594. plt.savefig(f'../manuscript_figures/fig3_energy_vs_freq_aoc_{auc_val}.pdf', format='pdf', bbox_inches='tight')
  595. # %%
  596. # repeat the process again for a range of theta values
  597. num_theta_values = 1000
  598. theta_values = np.linspace(0, 15, num_theta_values)
  599. moving_stim_energy = np.zeros(num_theta_values)
  600. stationary_stim_energy = np.zeros(num_theta_values)
  601. for i in tqdm(range(num_theta_values)):
  602. # compute the filter
  603. filter = temporal_filter(frequencies, theta_values[i])
  604. # compute the filtered spectra
  605. filtered_mean_stationary_psd = np.multiply(mean_stationary_psd, filter)
  606. filtered_mean_moving_psd = np.multiply(mean_moving_psd, filter)
  607. # compute the energy
  608. moving_stim_energy[i] = np.sum(filtered_mean_moving_psd*freq_spacing, axis=-1)
  609. stationary_stim_energy[i] = np.sum(filtered_mean_stationary_psd*freq_spacing, axis=-1)
  610. # %%
  611. # plot the energy vs. theta
  612. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  613. ax.plot(theta_values[:], moving_stim_energy[:], label='Moving', color='tab:orange')
  614. ax.plot(theta_values[:], stationary_stim_energy[:], label='Stationary', color='tab:gray')
  615. ax.set_xlabel('Theta')
  616. ax.set_ylabel('Energy')
  617. ax.set_title('Energy vs. theta')
  618. ax.set_yscale('log')
  619. ax.set_xscale('log')
  620. ax.legend()
  621. # %%
  622. relative = True
  623. # for given energy capacities in the stationary case (between 0 and 1)
  624. energy_capacity_stat = np.linspace(0, 1, 101)
  625. # energy_capacity_mov = np.linspace(0, 2, 101)
  626. delta_energy_capacity = np.linspace(0.01, 10, 1000)
  627. # energy_capacity_stat = np.logspace(-1, 0, 100)
  628. # delta_energy_capacity = np.logspace(-1, 1, 101)
  629. # interpolate the energy values to get a smooth curve
  630. stationary_theta_interp = np.interp(energy_capacity_stat, stationary_stim_energy, theta_values)
  631. if relative:
  632. energy_capacity_mov = np.outer(delta_energy_capacity, energy_capacity_stat) + energy_capacity_stat[np.newaxis, :] # relative
  633. delta_energy_capacity *= 100 # convert to percent
  634. else:
  635. energy_capacity_mov = delta_energy_capacity[:, np.newaxis] + energy_capacity_stat[np.newaxis, :] # absolute
  636. moving_theta_interp = np.interp(energy_capacity_mov, moving_stim_energy, theta_values)
  637. # %%
  638. import matplotlib.colors as mcolors
  639. from matplotlib import cm
  640. norm = mcolors.Normalize(vmin=np.min(theta_values), vmax=np.max(theta_values))
  641. two_color_norm = mcolors.TwoSlopeNorm(vmin=-5, vcenter=0, vmax=5)
  642. # example_points = [(50, 50), (75, 75), (95, 100)]
  643. example_points = [(75, 9), (90, 19), (95, 29), (99, 49)]
  644. fig, ax = plt.subplots(1, 3, figsize=(16, 5), width_ratios=[5, 5, 5], sharey=True)
  645. im0 = ax[0].imshow(np.repeat(stationary_theta_interp[np.newaxis, :], len(delta_energy_capacity), axis=0), origin='lower', norm=norm, cmap='gray_r', extent=[energy_capacity_stat[0], energy_capacity_stat[-1], delta_energy_capacity[0], delta_energy_capacity[-1]], aspect='auto')
  646. ax[0].set_title('Stationary theta')
  647. im1 = ax[1].imshow(moving_theta_interp, origin='lower', norm=norm, cmap='gray_r', extent=[energy_capacity_stat[0], energy_capacity_stat[-1], delta_energy_capacity[0], delta_energy_capacity[-1]], aspect='auto')
  648. ax[1].set_title('Moving theta')
  649. colors_map = mcolors.LinearSegmentedColormap.from_list('custom_cmap', [(1, 127/255, 14/255), (1, 1, 1), (0.5, 0.5, 0.5)], N=256)
  650. im2 = ax[2].imshow(stationary_theta_interp[np.newaxis, :]-moving_theta_interp, origin='lower', norm=two_color_norm, cmap=colors_map, extent=[energy_capacity_stat[0], energy_capacity_stat[-1], delta_energy_capacity[0], delta_energy_capacity[-1]], aspect='auto')
  651. ax[2].set_title('Difference in theta [stationary - moving]')
  652. # ax[2].set_xlim(0.5, 1)
  653. # ax[2].set_ylim(0.5, 1)
  654. # add axis labels
  655. if relative:
  656. ax[0].set_ylabel('Percentage increase in energy capacity')
  657. else:
  658. ax[0].set_ylabel('Increase in energy capacity')
  659. fig.supxlabel('Energy capacity')
  660. # plot the example points in three different colors which are distinct from red, white and blue
  661. colors_map = cm.get_cmap('Blues', len(example_points)+2)
  662. # plot the example points in the first subplot
  663. for i, point in enumerate(example_points):
  664. ax[0].scatter(energy_capacity_stat[point[0]], delta_energy_capacity[point[1]], color=colors_map(i+1), s=100, edgecolor='black', label='Example point ' + str(i+1))
  665. ax[1].scatter(energy_capacity_stat[point[0]], delta_energy_capacity[point[1]], color=colors_map(i+1), s=100, edgecolor='black')
  666. ax[2].scatter(energy_capacity_stat[point[0]], delta_energy_capacity[point[1]], color=colors_map(i+1), s=100, edgecolor='black')
  667. ax[0].set_yscale('log')
  668. # add a colorbar without disrupting the layout of the other subplots
  669. cbar = fig.colorbar(im2, ax=ax[2], orientation='vertical', pad=0.04)
  670. cbar = fig.colorbar(im1, ax=ax[1], orientation='vertical', pad=0.04)
  671. cbar = fig.colorbar(im0, ax=ax[0], orientation='vertical', pad=0.04)
  672. plt.savefig(f'../manuscript_figures/fig3_diff_theta_stat_energy_vs_delta_energy.pdf', format='pdf', bbox_inches='tight')
  673. # %%
  674. stationary_theta_interp[example_points[-1][0]]
  675. # %%
  676. example_filter_mov.shape
  677. # %%
  678. from scipy.fft import irfft
  679. # compute the autocorrelation function for the example points by extracting the theta values in the stationary and moving conditions and then applying the filter to white noise and then computing the autocorrelation function as the inverse fourier transform of the PSD
  680. alpha_arr = [0.25, 0.5, 0.75, 1]
  681. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  682. for i in range(len(example_points)):
  683. # compute the filter
  684. example_theta_mov = moving_theta_interp[example_points[i][1], example_points[i][0]]
  685. example_theta_stat = stationary_theta_interp[example_points[i][0]]
  686. # compute the filter
  687. example_filter_mov = temporal_filter(frequencies, example_theta_mov)
  688. example_filter_stat = temporal_filter(frequencies, example_theta_stat)
  689. # compute the inverse fourier transform
  690. example_moving_autocorr = irfft(example_filter_mov/freq_spacing)
  691. example_stationary_autocorr = irfft(example_filter_stat/freq_spacing)
  692. # normalize the autocorrelation function
  693. example_moving_autocorr /= np.max(example_moving_autocorr)
  694. example_stationary_autocorr /= np.max(example_stationary_autocorr)
  695. # plot the autocorrelation function
  696. time = np.arange(0, len(example_moving_autocorr)) / 30*1000
  697. ax.plot(time, example_moving_autocorr, label='Moving', color='tab:orange', alpha=alpha_arr[i])
  698. ax.plot(time, example_stationary_autocorr, label='Stationary', color='tab:gray', alpha=alpha_arr[i])
  699. ax.set_xlabel('Time (ms)')
  700. ax.set_ylabel('Autocorrelation')
  701. # ax.legend()
  702. ax.set_xlim(0, 500)
  703. plt.savefig(f'../manuscript_figures/fig3_autocorr_example_{i}.pdf', format='pdf', bbox_inches='tight')
  704. plt.show()
  705. plt.close(fig)
  706. # %% [markdown]
  707. # ### Individual filters
  708. # %%
  709. stationary_stim_psd_avg.shape
  710. # %%
  711. # normalize the spectrum such that the are under the curve is 1
  712. norm_factor = np.sum(stationary_stim_psd_avg[:, :, :, :, :]*freq_spacing, axis=-1, keepdims=True)
  713. stationary_stim_psd_avg_norm = stationary_stim_psd_avg[:, :, :, :, :] / norm_factor
  714. moving_stim_psd_avg_norm = moving_stim_psd_avg[:, :, :, :, :] / norm_factor
  715. mat_shape = stationary_stim_psd_avg_norm.shape
  716. # repeat the process again for a range of theta values
  717. num_theta_values = 1000
  718. theta_values = np.linspace(0, 15, num_theta_values)
  719. moving_stim_energy = np.zeros((mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3], num_theta_values))
  720. stationary_stim_energy = np.zeros((mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3], num_theta_values))
  721. for i in tqdm(range(num_theta_values)):
  722. # compute the filter
  723. filter = temporal_filter(frequencies, theta_values[i])
  724. # compute the filtered spectra
  725. filtered_mean_stationary_psd = np.multiply(stationary_stim_psd_avg_norm, filter)
  726. filtered_mean_moving_psd = np.multiply(moving_stim_psd_avg_norm, filter)
  727. # compute the energy
  728. moving_stim_energy[:, :, :, :, i] = np.sum(filtered_mean_moving_psd*freq_spacing, axis=-1)
  729. stationary_stim_energy[:, :, :, :, i] = np.sum(filtered_mean_stationary_psd*freq_spacing, axis=-1)
  730. # %%
  731. # for given energy capacities in the stationary case (between 0 and 1)
  732. relative = True
  733. num_energy_levels = 100
  734. energy_capacity_stat = np.linspace(0, 1, num_energy_levels+1)
  735. delta_energy_capacity = np.linspace(0.01, 10, 10*num_energy_levels)
  736. if relative:
  737. energy_capacity_mov = np.outer(delta_energy_capacity, energy_capacity_stat) + energy_capacity_stat[np.newaxis, :] # relative
  738. delta_energy_capacity *= 100 # convert to percent
  739. else:
  740. energy_capacity_mov = delta_energy_capacity[:, np.newaxis] + energy_capacity_stat[np.newaxis, :]
  741. # energy_capacity_stat = np.logspace(-1, 0, 100)
  742. # delta_energy_capacity = np.logspace(-1, 0, 100)
  743. stationary_theta_interp = np.zeros((num_energy_levels+1, mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3]))
  744. moving_theta_interp = np.zeros((10*num_energy_levels, num_energy_levels+1, mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3]))
  745. # interpolate the energy values to get a smooth curve
  746. for i in range(mat_shape[0]):
  747. for j in range(mat_shape[1]):
  748. for k in range(mat_shape[2]):
  749. for l in range(mat_shape[3]):
  750. # interpolate the energy values to get a smooth curve
  751. stationary_theta_interp[:, i, j, k, l] = np.interp(energy_capacity_stat, stationary_stim_energy[i, j, k, l], theta_values)
  752. for m in range(num_energy_levels):
  753. # interpolate the energy values to get a smooth curve
  754. moving_theta_interp[:, m, i, j, k, l] = np.interp(energy_capacity_mov[:, m], moving_stim_energy[i, j, k, l], theta_values)
  755. # %%
  756. # for each filter, find the cutoff value of the autocorrelation function for different points in the energy landscape
  757. # create a matrix to store the cutoff values for each filter
  758. cutoff_frequency_mov = np.zeros((len(example_points), mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3]))
  759. cutoff_frequency_stat = np.zeros((len(example_points), mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3]))
  760. autocorr_func_mov = np.zeros((len(example_points), mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3], 2*(len(frequencies)-1)))
  761. autocorr_func_stat = np.zeros((len(example_points), mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3], 2*(len(frequencies)-1)))
  762. cutoff = 0.5
  763. for m in range(len(example_points)):
  764. for i in range(mat_shape[0]):
  765. for j in range(mat_shape[1]):
  766. for k in range(mat_shape[2]):
  767. for l in range(mat_shape[3]):
  768. # compute the filter
  769. example_theta_mov = moving_theta_interp[example_points[m][1], example_points[m][0], i, j, k, l]
  770. example_theta_stat = stationary_theta_interp[example_points[m][0], i, j, k, l]
  771. # compute the filter
  772. example_filter_mov = temporal_filter(frequencies, example_theta_mov)
  773. example_filter_stat = temporal_filter(frequencies, example_theta_stat)
  774. # compute the filtered spectra
  775. example_filtered_mean_stationary_psd = example_filter_stat # np.multiply(stationary_stim_psd_avg_norm[i, j, k, l], example_filter_stat)
  776. example_filtered_mean_moving_psd = example_filter_mov # np.multiply(moving_stim_psd_avg_norm[i, j, k, l], example_filter_mov)
  777. # compute the inverse fourier transform
  778. example_moving_autocorr = irfft(example_filtered_mean_moving_psd)
  779. example_stationary_autocorr = irfft(example_filtered_mean_stationary_psd)
  780. # normalize the autocorrelation function
  781. example_moving_autocorr /= np.max(example_moving_autocorr)
  782. example_stationary_autocorr /= np.max(example_stationary_autocorr)
  783. # store the autocorrelation function
  784. autocorr_func_mov[m, i, j, k, l, :] = example_moving_autocorr
  785. autocorr_func_stat[m, i, j, k, l, :] = example_stationary_autocorr
  786. # find the cutoff frequency
  787. try:
  788. cutoff_frequency_mov[m, i, j, k, l] = np.where(example_moving_autocorr < cutoff)[0][0]
  789. except IndexError:
  790. cutoff_frequency_mov[m, i, j, k, l] = len(example_moving_autocorr)
  791. try:
  792. cutoff_frequency_stat[m, i, j, k, l] = np.where(example_stationary_autocorr < cutoff)[0][0]
  793. except IndexError:
  794. cutoff_frequency_stat[m, i, j, k, l] = len(example_stationary_autocorr)
  795. # %%
  796. # plot and save all autocorrelation functions
  797. # make a directory called example_autocorrelation_sim in ../manuscript_figures/
  798. import os
  799. if not os.path.exists('../manuscript_figures/example_autocorrelation_sim_-1_1_white_noise_band_pass'):
  800. os.makedirs('../manuscript_figures/example_autocorrelation_sim_-1_1_white_noise_band_pass')
  801. point_idx = 0
  802. position_idx = 4
  803. time = np.arange(0, len(example_moving_autocorr)) / 30*1000
  804. for ori_idx in range(mat_shape[0]):
  805. for phase_idx in range(mat_shape[1]):
  806. for wavelength_idx in range(mat_shape[2]):
  807. fig, ax = plt.subplots(1, 1, figsize=(4, 3))
  808. ax.plot(time, autocorr_func_mov[point_idx, ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Moving', color='tab:orange')
  809. ax.plot(time, autocorr_func_stat[point_idx, ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Stationary', color='tab:gray')
  810. ax.set_xlim(0, 500)
  811. ax.set_xlabel('Time (ms)')
  812. ax.set_ylabel('Autocorrelation')
  813. ax.set_title(f'Point {point_idx}, Ori {ori_idx}, Phase {phase_idx}, Wavelength {wavelength_idx}, Position {position_idx}')
  814. ax.legend()
  815. ax.set_ylim(-1, 1)
  816. plt.savefig(f'../manuscript_figures/example_autocorrelation_sim_-1_1_white_noise_band_pass/fig3_autocorr_point_{point_idx}_ori_{ori_idx}_phase_{phase_idx}_wavelength_{wavelength_idx}_position_{position_idx}.pdf', format='pdf', bbox_inches='tight')
  817. # close the figure
  818. plt.close(fig)
  819. # %%
  820. # plot the average autocorrelation function for each point
  821. avg_autocorr_func_mov = np.mean(autocorr_func_mov, axis=(1, 2, 3, 4))
  822. avg_autocorr_func_stat = np.mean(autocorr_func_stat, axis=(1, 2, 3, 4))
  823. std_autocorr_func_mov = np.std(autocorr_func_mov, axis=(1, 2, 3, 4))
  824. std_autocorr_func_stat = np.std(autocorr_func_stat, axis=(1, 2, 3, 4))
  825. time = np.arange(0, len(example_moving_autocorr)) / 30*1000
  826. # plot the average autocorrelation function for each point
  827. for point_idx in range(len(example_points)):
  828. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  829. ax.plot(time, avg_autocorr_func_mov[point_idx, :], label='Moving', color='tab:orange')
  830. ax.plot(time, avg_autocorr_func_stat[point_idx, :], label='Stationary', color='tab:gray')
  831. ax.fill_between(time, avg_autocorr_func_mov[point_idx, :] - std_autocorr_func_mov[point_idx, :], avg_autocorr_func_mov[point_idx, :] + std_autocorr_func_mov[point_idx, :], color='tab:orange', alpha=0.2)
  832. ax.fill_between(time, avg_autocorr_func_stat[point_idx, :] - std_autocorr_func_stat[point_idx, :], avg_autocorr_func_stat[point_idx, :] + std_autocorr_func_stat[point_idx, :], color='tab:gray', alpha=0.2)
  833. ax.set_xlim(0, 500)
  834. ax.set_ylim(-1, 1)
  835. plt.savefig(f'../manuscript_figures/fig3_avg_autocorr_point_{point_idx}.pdf', format='pdf', bbox_inches='tight')
  836. plt.show()
  837. plt.close(fig)
  838. # %%
  839. cutoff_frequency_stat /= 30
  840. cutoff_frequency_mov /= 30
  841. # %%
  842. cutoff_frequency_stat *= 1000
  843. cutoff_frequency_mov *= 1000
  844. # %%
  845. np.unique(diff_cutoff)
  846. # %%
  847. np.linspace(-300, 300, 20)
  848. # %%
  849. # take the difference in cutoffs and plot the histogram
  850. diff_cutoff = cutoff_frequency_mov - cutoff_frequency_stat
  851. bins = np.linspace(-300, 300, 20)
  852. idx = 3
  853. plt.hist(diff_cutoff[idx].flatten(), bins=bins, density=True)
  854. plt.xlabel('Difference of cutoff (moving - stationary) in ms')
  855. plt.axvline(np.mean(diff_cutoff[idx]), color='k', linestyle='--')
  856. plt.savefig('../manuscript_figures/fig3_hist_cutoff_diff_theory.pdf', format='pdf', bbox_inches='tight')
  857. # %%
  858. # plot the cutoff frequencies as a density scatter plot with jitter
  859. point = 3
  860. jitter_strength = 10
  861. # Ensure the x and y axes span the same range
  862. min_val = 0
  863. max_val = max(np.max(cutoff_frequency_stat[point]), np.max(cutoff_frequency_mov[point]))
  864. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  865. ax.set_title('Cutoff frequency')
  866. ax.set_xlabel('Stationary cutoff frequency')
  867. ax.set_ylabel('Moving cutoff frequency')
  868. # add jitter to the points
  869. cutoff_stat_jitter = cutoff_frequency_stat[point, :, :, :, :].flatten() + np.random.normal(0, jitter_strength, cutoff_frequency_stat[point, :, :, :, :].flatten().shape)
  870. cutoff_mov_jitter = cutoff_frequency_mov[point, :, :, :, :].flatten() + np.random.normal(0, jitter_strength, cutoff_frequency_mov[point, :, :, :, :].flatten().shape)
  871. # use gaussian kde to plot the density of points
  872. from scipy.stats import gaussian_kde
  873. kde = gaussian_kde([cutoff_stat_jitter, cutoff_mov_jitter])
  874. # plot scatter plot colored by density
  875. density = kde([cutoff_stat_jitter, cutoff_mov_jitter])
  876. ax.scatter(cutoff_stat_jitter, cutoff_mov_jitter, c=density, cmap='viridis', s=1)
  877. # add a colorbar
  878. cbar = fig.colorbar(ax.collections[0], ax=ax, orientation='vertical')
  879. cbar.set_label('Density')
  880. # add a line with slope 1 which goes from 0 to the max value axes
  881. ax.plot([min_val, max_val], [min_val, max_val], color='black', linestyle='--')
  882. ax.set_xlim(0, 500)
  883. ax.set_ylim(0, 500)
  884. plt.savefig(f'../manuscript_figures/fig3_cutoff_frequency_point_{point}.pdf', format='pdf', bbox_inches='tight')
  885. # %%
  886. # plot the cutoff frequencies as a density scatter plot with jitter
  887. point = 3
  888. # Ensure the x and y axes span the same range
  889. min_val = 0
  890. max_val = max(np.max(cutoff_frequency_stat[point]), np.max(cutoff_frequency_mov[point]))
  891. frame_length = 1/30*1000
  892. bins = np.arange(min_val, max_val+frame_length, frame_length)
  893. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  894. ax.set_title('Cutoff frequency')
  895. ax.set_xlabel('Stationary cutoff lag (ms)')
  896. ax.set_ylabel('Moving cutoff lag (ms)')
  897. ax.hist2d(cutoff_frequency_stat[point, :, :, :, :].flatten(), cutoff_frequency_mov[point, :, :, :, :].flatten(), bins=[bins, bins], cmap='Blues')
  898. # add a colorbar
  899. cbar = fig.colorbar(ax.collections[0], ax=ax, orientation='vertical')
  900. cbar.set_label('Density')
  901. # add a line with slope 1 which goes from 0 to the max value axes
  902. ax.plot([min_val, max_val], [min_val, max_val], color='black', linestyle='--')
  903. plt.savefig(f'../manuscript_figures/fig3_cutoff_frequency_density_hist_point_{point}.pdf', format='pdf', bbox_inches='tight')
  904. # %%
  905. # print the fraction above the diagonal
  906. above_diagonal = np.sum(cutoff_frequency_mov[point] > cutoff_frequency_stat[point]) / (cutoff_frequency_mov[point].flatten().shape[0])
  907. print(f'Fraction of points above the diagonal: {above_diagonal:.2f}')
  908. # %%
  909. import matplotlib.colors as mcolors
  910. norm = mcolors.LogNorm(vmin=np.min(theta_values), vmax=np.max(theta_values))
  911. two_color_norm = mcolors.TwoSlopeNorm(vmin=-5, vcenter=0, vmax=5)
  912. random_draws = 10
  913. for i in range(random_draws):
  914. ori_idx = np.random.randint(0, mat_shape[0])
  915. phase_idx = np.random.randint(0, mat_shape[1])
  916. wavelength_idx = np.random.randint(0, mat_shape[2])
  917. position_idx = np.random.randint(0, mat_shape[3])
  918. stationary_theta_test = stationary_theta_interp[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  919. moving_theta_test = moving_theta_interp[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]
  920. fig, ax = plt.subplots(2, 3, figsize=(16, 10))
  921. im0 = ax[0, 0].imshow(np.repeat(stationary_theta_test[np.newaxis, :], 100, axis=0), origin='lower', norm=norm, cmap='viridis', extent=[energy_capacity_stat[0], energy_capacity_stat[-1], delta_energy_capacity[0], delta_energy_capacity[-1]], aspect='auto')
  922. ax[0, 0].set_title('Stationary theta')
  923. im1 = ax[0, 1].imshow(moving_theta_test, origin='lower', norm=norm, cmap='viridis', extent=[energy_capacity_stat[0], energy_capacity_stat[-1], delta_energy_capacity[0], delta_energy_capacity[-1]], aspect='auto')
  924. ax[0, 1].set_title('Moving theta')
  925. im2 = ax[0, 2].imshow(stationary_theta_test[np.newaxis, :]-moving_theta_test, origin='lower', norm=two_color_norm, cmap='bwr', extent=[energy_capacity_stat[0], energy_capacity_stat[-1], delta_energy_capacity[0], delta_energy_capacity[-1]], aspect='auto')
  926. ax[0, 2].set_title('Difference in theta [stationary - moving]')
  927. # add axis labels
  928. ax[0, 0].set_ylabel('Percentage increase in energy capacity')
  929. ax[0, 1].set_xlabel('Energy capacity')
  930. fig.suptitle(f'Orientation: {orientation_arr[ori_idx]}, Phase: {phase_arr[phase_idx]}, Frequency: {freq_arr[low_spatial_freq_idx][wavelength_idx]}, Position: {position_arr[position_idx]}')
  931. # add a colorbar without disrupting the layout of the other subplots
  932. cbar = fig.colorbar(im2, ax=ax[0, 2], orientation='vertical', fraction=0.02, pad=0.04)
  933. cbar = fig.colorbar(im1, ax=ax[0, 1], orientation='vertical', fraction=0.02, pad=0.04)
  934. cbar = fig.colorbar(im0, ax=ax[0, 0], orientation='vertical', fraction=0.02, pad=0.04)
  935. # plot the spectra and the energy vs. theta
  936. ax[1, 0].plot(frequencies, stationary_stim_psd_avg[ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Stationary', color='tab:red')
  937. ax[1, 0].plot(frequencies, moving_stim_psd_avg[ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Moving', color='tab:blue')
  938. ax[1, 0].set_xlabel('Frequency (Hz)')
  939. ax[1, 0].set_ylabel('Power')
  940. ax[1, 0].set_title('Power spectrum')
  941. ax[1, 0].set_yscale('log')
  942. ax[1, 1].plot(theta_values, stationary_stim_energy[ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Stationary', color='tab:red')
  943. ax[1, 1].plot(theta_values, moving_stim_energy[ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Moving', color='tab:blue')
  944. ax[1, 1].set_xlabel('Theta')
  945. ax[1, 1].set_ylabel('Energy')
  946. ax[1, 1].set_title('Energy vs. theta')
  947. # make last plot empty
  948. ax[1, 2].axis('off')
  949. plt.show()
  950. # %% [markdown]
  951. # ### Block trial stimuli
  952. # %%
  953. import numpy as np
  954. # create white noise block structure, have 1 second of white noise with low variance and then 1 second of white noise with high variance
  955. # the sampling rate is 10 ms.
  956. num_samples = 100 # number of samples in the white noise block
  957. low_noise = 0.1
  958. high_noise = 1
  959. sampling_rate = 100
  960. # repeat this 100 times with new sampling so that in the end I have an array of shape (100, 300)
  961. num_blocks = 100
  962. white_noise_blocks = np.array([np.concatenate([
  963. np.random.normal(0, low_noise, 20), # low variance
  964. np.random.normal(0, high_noise, num_samples), # high variance
  965. np.random.normal(0, low_noise, 80), # low variance
  966. ]) for _ in range(num_blocks)])
  967. # take the fft of the white noise blocks
  968. white_noise_blocks_fft = np.fft.rfft(white_noise_blocks, axis=-1)
  969. white_noise_frequencies = np.fft.rfftfreq(white_noise_blocks.shape[1], d=1/sampling_rate)
  970. # %%
  971. avg_white_noise_psd = np.mean(np.abs(white_noise_blocks_fft)**2, axis=0)[4:]
  972. # %%
  973. import matplotlib.pyplot as plt
  974. plt.plot(white_noise_frequencies[4:], avg_white_noise_psd)
  975. # %%
  976. # for each filter, find the cutoff value of the autocorrelation function for different points in the energy landscape
  977. # create a matrix to store the cutoff values for each filter
  978. cutoff_frequency_mov = np.zeros((len(example_points), mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3]))
  979. cutoff_frequency_stat = np.zeros((len(example_points), mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3]))
  980. autocorr_func_mov = np.zeros((len(example_points), mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3], len(white_noise_frequencies[4:])//2))
  981. autocorr_func_stat = np.zeros((len(example_points), mat_shape[0], mat_shape[1], mat_shape[2], mat_shape[3], len(white_noise_frequencies[4:])//2))
  982. cutoff = 0.5
  983. for m in range(len(example_points)):
  984. for i in range(mat_shape[0]):
  985. for j in range(mat_shape[1]):
  986. for k in range(mat_shape[2]):
  987. for l in range(mat_shape[3]):
  988. # compute the filter
  989. example_theta_mov = moving_theta_interp[example_points[m][1], example_points[m][0], i, j, k, l]
  990. example_theta_stat = stationary_theta_interp[example_points[m][0], i, j, k, l]
  991. # compute the filter
  992. example_filter_mov = temporal_filter(white_noise_frequencies[4:], example_theta_mov)
  993. example_filter_stat = temporal_filter(white_noise_frequencies[4:], example_theta_stat)
  994. # compute the filtered spectra
  995. example_filtered_mean_stationary_psd = np.multiply(avg_white_noise_psd, example_filter_stat)
  996. example_filtered_mean_moving_psd = np.multiply(avg_white_noise_psd, example_filter_mov)
  997. # compute the inverse fourier transform
  998. example_moving_autocorr = irfft(example_filtered_mean_moving_psd)[:len(example_filtered_mean_moving_psd)//2]
  999. example_stationary_autocorr = irfft(example_filtered_mean_stationary_psd)[:len(example_filtered_mean_stationary_psd)//2]
  1000. # normalize the autocorrelation function
  1001. example_moving_autocorr /= np.max(example_moving_autocorr)
  1002. example_stationary_autocorr /= np.max(example_stationary_autocorr)
  1003. # store the autocorrelation function
  1004. autocorr_func_mov[m, i, j, k, l, :] = example_moving_autocorr
  1005. autocorr_func_stat[m, i, j, k, l, :] = example_stationary_autocorr
  1006. # # find the cutoff frequency
  1007. try:
  1008. cutoff_frequency_mov[m, i, j, k, l] = np.where(example_moving_autocorr < cutoff)[0][0]
  1009. cutoff_frequency_stat[m, i, j, k, l] = np.where(example_stationary_autocorr < cutoff)[0][0]
  1010. except IndexError:
  1011. # if the cutoff is not found, set it to nan
  1012. cutoff_frequency_mov[m, i, j, k, l] = np.nan
  1013. cutoff_frequency_stat[m, i, j, k, l] = np.nan
  1014. # %%
  1015. # save the autocorrelation and cutoff frequency matrices
  1016. np.savez('../data/fig3_autocorr_cutoff_frequency_white_noise.npz',
  1017. autocorr_func_mov=autocorr_func_mov,
  1018. autocorr_func_stat=autocorr_func_stat,
  1019. cutoff_frequency_mov=cutoff_frequency_mov,
  1020. cutoff_frequency_stat=cutoff_frequency_stat,
  1021. white_noise_frequencies=white_noise_frequencies[4:],
  1022. example_points=example_points,
  1023. frequencies=frequencies[99:],
  1024. theta_values=theta_values,
  1025. stationary_theta_interp=stationary_theta_interp,
  1026. moving_theta_interp=moving_theta_interp)
  1027. # %%
  1028. time
  1029. # %%
  1030. autocorr_func_mov.shape
  1031. # %%
  1032. point_idx = 1
  1033. position_idx = 5
  1034. ori = 0
  1035. phase = 0
  1036. wavelength = 3
  1037. time = np.arange(0, len(white_noise_frequencies[4:])//2*10, 10)
  1038. plt.plot(time, autocorr_func_mov[point_idx, ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Moving', color='tab:orange')
  1039. plt.plot(time, autocorr_func_stat[point_idx, ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Moving', color='tab:gray')
  1040. plt.xlim(0, 500)
  1041. plt.axhline(0, color='black', linestyle='--')
  1042. # %%
  1043. # plot and save all autocorrelation functions
  1044. # make a directory called example_autocorrelation_sim in ../manuscript_figures/
  1045. import os
  1046. if not os.path.exists('../manuscript_figures/example_autocorrelation_sim_-1_1_block_stim'):
  1047. os.makedirs('../manuscript_figures/example_autocorrelation_sim_-1_1_block_stim')
  1048. point_idx = 1
  1049. position_idx = 4
  1050. for ori_idx in range(mat_shape[0]):
  1051. for phase_idx in range(mat_shape[1]):
  1052. for wavelength_idx in range(mat_shape[2]):
  1053. fig, ax = plt.subplots(1, 1, figsize=(4, 3))
  1054. ax.plot(time, autocorr_func_mov[point_idx, ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Moving', color='tab:orange')
  1055. ax.plot(time, autocorr_func_stat[point_idx, ori_idx, phase_idx, wavelength_idx, position_idx, :], label='Stationary', color='tab:gray')
  1056. ax.set_xlabel('Time (ms)')
  1057. ax.set_ylabel('Autocorrelation')
  1058. ax.set_title(f'Point {point_idx}, Ori {ori_idx}, Phase {phase_idx}, Wavelength {wavelength_idx}, Position {position_idx}')
  1059. ax.legend()
  1060. ax.set_ylim(-1, 1)
  1061. ax.set_xlim(0, 500)
  1062. plt.savefig(f'../manuscript_figures/example_autocorrelation_sim_-1_1_block_stim/fig3_autocorr_point_{point_idx}_ori_{ori_idx}_phase_{phase_idx}_wavelength_{wavelength_idx}_position_{position_idx}_block_stim.pdf', format='pdf', bbox_inches='tight')
  1063. # close the figure
  1064. plt.close(fig)
  1065. # %%
  1066. # plot the average autocorrelation function for each point
  1067. avg_autocorr_func_mov = np.mean(autocorr_func_mov, axis=(1, 2, 3, 4))
  1068. avg_autocorr_func_stat = np.mean(autocorr_func_stat, axis=(1, 2, 3, 4))
  1069. std_autocorr_func_mov = np.std(autocorr_func_mov, axis=(1, 2, 3, 4))
  1070. std_autocorr_func_stat = np.std(autocorr_func_stat, axis=(1, 2, 3, 4))
  1071. time = np.arange(0, len(avg_autocorr_func_mov[1])*10, 10)
  1072. # plot the average autocorrelation function for each point
  1073. for point_idx in range(len(example_points)):
  1074. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  1075. ax.plot(time, avg_autocorr_func_mov[point_idx, :], label='Moving', color='tab:orange')
  1076. ax.plot(time, avg_autocorr_func_stat[point_idx, :], label='Stationary', color='tab:gray')
  1077. ax.fill_between(time, avg_autocorr_func_mov[point_idx, :] - std_autocorr_func_mov[point_idx, :], avg_autocorr_func_mov[point_idx, :] + std_autocorr_func_mov[point_idx, :], color='tab:orange', alpha=0.2)
  1078. ax.fill_between(time, avg_autocorr_func_stat[point_idx, :] - std_autocorr_func_stat[point_idx, :], avg_autocorr_func_stat[point_idx, :] + std_autocorr_func_stat[point_idx, :], color='tab:gray', alpha=0.2)
  1079. # ax.set_xlim(0, 1000)
  1080. ax.set_ylim(-1, 1)
  1081. plt.show()
  1082. plt.savefig(f'../manuscript_figures/fig3_avg_autocorr_point_{point_idx}_block_stimuli.pdf', format='pdf', bbox_inches='tight')
  1083. plt.close(fig)
  1084. # %%
  1085. # plot the cutoff frequencies as a density scatter plot with jitter
  1086. point = 2
  1087. jitter_strength = 10
  1088. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  1089. ax.set_title('Cutoff frequency')
  1090. ax.set_xlabel('Stationary cutoff frequency')
  1091. ax.set_ylabel('Moving cutoff frequency')
  1092. # add jitter to the points and remove nan values
  1093. cutoff_frequency_stat_temp = cutoff_frequency_stat[point, :, :, :, :].flatten()
  1094. cutoff_frequency_mov_temp = cutoff_frequency_mov[point, :, :, :, :].flatten()
  1095. # check what percent of values in cutoff_frequency_mov are nan
  1096. print(f'Percent of NaN values in cutoffs: {np.sum(np.isnan(cutoff_frequency_mov_temp))/ cutoff_frequency_mov_temp.size * 100:.2f}%')
  1097. cutoff_frequency_stat_temp = cutoff_frequency_stat_temp[~np.isnan(cutoff_frequency_stat_temp)]*10
  1098. cutoff_frequency_mov_temp = cutoff_frequency_mov_temp[~np.isnan(cutoff_frequency_mov_temp)]*10
  1099. # Ensure the x and y axes span the same range
  1100. min_val = 0
  1101. max_val = max(np.max(cutoff_frequency_stat_temp), np.max(cutoff_frequency_mov_temp))
  1102. cutoff_stat_jitter = cutoff_frequency_stat_temp + np.random.normal(0, jitter_strength, cutoff_frequency_stat_temp.shape)
  1103. cutoff_mov_jitter = cutoff_frequency_mov_temp + np.random.normal(0, jitter_strength, cutoff_frequency_mov_temp.shape)
  1104. # use gaussian kde to plot the density of points
  1105. from scipy.stats import gaussian_kde
  1106. kde = gaussian_kde([cutoff_stat_jitter, cutoff_mov_jitter])
  1107. # plot scatter plot colored by density
  1108. density = kde([cutoff_stat_jitter, cutoff_mov_jitter])
  1109. ax.scatter(cutoff_stat_jitter, cutoff_mov_jitter, c=density, cmap='viridis', s=1)
  1110. # add a colorbar
  1111. cbar = fig.colorbar(ax.collections[0], ax=ax, orientation='vertical')
  1112. cbar.set_label('Density')
  1113. # add a line with slope 1 which goes from 0 to the max value axes
  1114. ax.plot([min_val, max_val], [min_val, max_val], color='black', linestyle='--')
  1115. ax.set_xlim(0, 500)
  1116. ax.set_ylim(0, 500)
  1117. plt.savefig(f'../manuscript_figures/fig3_cutoff_frequency_density_scatter_block_stim_point_{point}.pdf', format='pdf', bbox_inches='tight')
  1118. # %%
  1119. # plot the cutoff frequencies as a density scatter plot with jitter
  1120. point = 2
  1121. # Ensure the x and y axes span the same range
  1122. # add jitter to the points and remove nan values
  1123. cutoff_frequency_stat_temp = cutoff_frequency_stat[point, :, :, :, :].flatten()
  1124. cutoff_frequency_mov_temp = cutoff_frequency_mov[point, :, :, :, :].flatten()
  1125. # check what percent of values in cutoff_frequency_mov are nan
  1126. print(f'Percent of NaN values in cutoffs: {np.sum(np.isnan(cutoff_frequency_mov_temp))/ cutoff_frequency_mov_temp.size * 100:.2f}%')
  1127. cutoff_frequency_stat_temp = cutoff_frequency_stat_temp[~np.isnan(cutoff_frequency_stat_temp)]*10
  1128. cutoff_frequency_mov_temp = cutoff_frequency_mov_temp[~np.isnan(cutoff_frequency_mov_temp)]*10
  1129. # Ensure the x and y axes span the same range
  1130. min_val = 0
  1131. max_val = max(np.max(cutoff_frequency_stat_temp), np.max(cutoff_frequency_mov_temp))
  1132. frame_length = 10
  1133. bins = np.arange(min_val, 510, frame_length)
  1134. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  1135. ax.set_title('Cutoff frequency')
  1136. ax.set_xlabel('Stationary cutoff lag (ms)')
  1137. ax.set_ylabel('Moving cutoff lag (ms)')
  1138. ax.hist2d(cutoff_frequency_stat_temp, cutoff_frequency_mov_temp, bins=[bins, bins], cmap='Blues')
  1139. # add a colorbar
  1140. cbar = fig.colorbar(ax.collections[0], ax=ax, orientation='vertical')
  1141. cbar.set_label('Density')
  1142. # add a line with slope 1 which goes from 0 to the max value axes
  1143. ax.plot([min_val, max_val], [min_val, max_val], color='black', linestyle='--')
  1144. plt.savefig(f'../manuscript_figures/fig3_cutoff_frequency_density_scatter_block_stim_point_{point}.pdf', format='pdf', bbox_inches='tight')
  1145. # %% [markdown]
  1146. # ## Normalized Utility
  1147. # %%
  1148. # compute the max info and energy for each filter
  1149. moving_max_info = np.max(moving_stim_information, axis=0)
  1150. stationary_max_info = np.max(stationary_stim_information, axis=0)
  1151. moving_max_energy = np.max(moving_stim_energy, axis=0)
  1152. stationary_max_energy = np.max(stationary_stim_energy, axis=0)
  1153. max_info = np.maximum(moving_max_info, stationary_max_info)
  1154. max_energy = np.maximum(moving_max_energy, stationary_max_energy)
  1155. # normalize
  1156. moving_stim_information_norm = moving_stim_information / max_info
  1157. stationary_stim_information_norm = stationary_stim_information / max_info
  1158. moving_stim_energy_norm = moving_stim_energy / max_energy
  1159. stationary_stim_energy_norm = stationary_stim_energy / max_energy
  1160. # %%
  1161. ori_idx = 0
  1162. phase_idx = 0
  1163. wavelength_idx = 0
  1164. position_idx = 4
  1165. moving_test_informations = moving_stim_information_norm[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  1166. stationary_test_informations = stationary_stim_information_norm[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  1167. moving_test_energy = moving_stim_energy_norm[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  1168. stationary_test_energy = stationary_stim_energy_norm[:, ori_idx, phase_idx, wavelength_idx, position_idx]
  1169. # create a subplot with 1 row and 3 columns. First column plot the information vs. theta, second column plot the energy vs. theta, third column plot the information vs. energy
  1170. fig, axs = plt.subplots(1, 3, figsize=(15, 5))
  1171. axs[0].plot(theta_values, moving_test_informations, 'o-', color='tab:blue', label='Moving test')
  1172. axs[0].plot(theta_values, stationary_test_informations, 'o-', color='tab:red', label='Stationary test')
  1173. axs[0].set_xlabel('Theta')
  1174. axs[0].set_ylabel('Information')
  1175. axs[0].set_title('Information vs. theta')
  1176. axs[0].legend()
  1177. axs[1].plot(theta_values, moving_test_energy, 'o-', color='tab:blue', label='Moving test')
  1178. axs[1].plot(theta_values, stationary_test_energy, 'o-', color='tab:red', label='Stationary test')
  1179. axs[1].set_xlabel('Theta')
  1180. axs[1].set_ylabel('Energy')
  1181. axs[1].set_title('Energy vs. theta')
  1182. axs[1].legend()
  1183. axs[2].plot(moving_test_energy, moving_test_informations, 'o-', color='tab:blue', label='Moving test')
  1184. axs[2].plot(stationary_test_energy, stationary_test_informations, 'o-', color='tab:red', label='Stationary test')
  1185. axs[2].set_xlabel('Energy')
  1186. axs[2].set_ylabel('Information')
  1187. axs[2].set_title('Information vs. energy')
  1188. axs[2].legend()
  1189. plt.tight_layout()
  1190. plt.show()
  1191. # %%
  1192. # compute the utility for a variety of lambdas
  1193. lmbd_arr = np.linspace(0, 100, 100)
  1194. moving_utility = np.zeros((lmbd_arr.shape[0], num_theta_values, len(orientation_arr), len(phase_arr), len(wavelength_arr[low_spatial_freq_idx]), len(position_arr)))
  1195. stationary_utility = np.zeros((lmbd_arr.shape[0], num_theta_values, len(orientation_arr), len(phase_arr), len(wavelength_arr[low_spatial_freq_idx]), len(position_arr)))
  1196. for i, lmbd in enumerate(lmbd_arr):
  1197. moving_utility[i] = moving_stim_information_norm - lmbd*moving_stim_energy_norm
  1198. stationary_utility[i] = stationary_stim_information_norm - lmbd*stationary_stim_energy_norm
  1199. # %%
  1200. # compute the optimal value of theta for each filter
  1201. optimal_theta_idx_moving = np.argmax(moving_utility, axis=1)
  1202. optimal_theta_idx_stationary = np.argmax(stationary_utility, axis=1)
  1203. optimal_theta_moving = theta_values[optimal_theta_idx_moving]
  1204. optimal_theta_stationary = theta_values[optimal_theta_idx_stationary]
  1205. # %%
  1206. optimal_theta_moving.shape
  1207. # %%
  1208. # take the difference and then the average
  1209. optimal_theta_diff = optimal_theta_moving - optimal_theta_stationary
  1210. optimal_theta_diff_avg = np.mean(optimal_theta_diff, axis=(1, 2, 4))
  1211. # %%
  1212. optimal_theta_diff.shape
  1213. # %%
  1214. plt.hist(np.min(optimal_theta_diff, axis=0).flatten(), bins=20)
  1215. # %%
  1216. # plot the average difference in optimal theta, color them with viridis
  1217. from matplotlib import cm
  1218. from matplotlib.colors import Normalize
  1219. colors_map = cm.get_cmap('viridis', len(wavelength_arr[low_spatial_freq_idx]))
  1220. fig, ax = plt.subplots(1, 1, figsize=(8, 6))
  1221. for i in range(wavelength_arr[low_spatial_freq_idx].shape[0]):
  1222. ax.plot(lmbd_arr, optimal_theta_diff_avg[:, i], color=colors_map(i))
  1223. # %%
  1224. ori_idx = 0
  1225. phase_idx = 0
  1226. wavelength_idx = 0
  1227. position_idx = 2
  1228. # plot the delta theta moving - stationary for a specific filter
  1229. fig, ax = plt.subplots(1, 1, figsize=(8, 6))
  1230. ax.plot(lmbd_arr, optimal_theta_moving[:, ori_idx, phase_idx, wavelength_idx, position_idx] - optimal_theta_stationary[:, ori_idx, phase_idx, wavelength_idx, position_idx])
  1231. # %%
  1232. import matplotlib.colors as mcolors
  1233. norm = mcolors.TwoSlopeNorm(vmin=np.min([moving_utility, stationary_utility]), vcenter=0, vmax=np.max([moving_utility, stationary_utility]))
  1234. # create a heatmap of the utility for theta vs. lambda for stationary and moving
  1235. fig, ax = plt.subplots(1, 2, figsize=(16, 6), sharey=True)
  1236. ax[1].imshow(moving_utility[:, :, ori_idx, phase_idx, wavelength_idx, position_idx], aspect='auto', cmap='bwr', origin='lower')
  1237. ax[1].set_title('Moving utility')
  1238. ax[0].imshow(stationary_utility[:, :, ori_idx, phase_idx, wavelength_idx, position_idx], aspect='auto', cmap='bwr', origin='lower')
  1239. ax[0].set_title('Stationary utility')
  1240. ax[0].set_xlabel('Theta')
  1241. ax[1].set_xlabel('Theta')
  1242. ax[0].set_ylabel('Lambda')
  1243. # set the tick labels
  1244. ax[0].set_xticks(np.arange(0, num_theta_values, 10))
  1245. ax[0].set_xticklabels(np.round(theta_values[::10], 2))
  1246. ax[1].set_xticks(np.arange(0, num_theta_values, 10))
  1247. ax[1].set_xticklabels(np.round(theta_values[::10], 2))
  1248. ax[0].set_yticks(np.arange(0, len(lmbd_arr), 10))
  1249. ax[0].set_yticklabels(np.round(lmbd_arr[::10], 2))
  1250. # %% [markdown]
  1251. # ## Gaussian filter
  1252. # %%
  1253. # repeat the process again for a range of theta values
  1254. num_theta_values = 50
  1255. array_shape = moving_stim_psd_avg.shape
  1256. theta_values_u = [0] # np.linspace(0, np.max(frequencies), num_theta_values)
  1257. theta_values_s = np.logspace(-1, 1, num_theta_values)
  1258. moving_stim_information = np.zeros((len(theta_values_u), len(theta_values_s), array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  1259. stationary_stim_information = np.zeros((len(theta_values_u), len(theta_values_s), array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  1260. moving_stim_energy = np.zeros((len(theta_values_u), len(theta_values_s), array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  1261. stationary_stim_energy = np.zeros((len(theta_values_u), len(theta_values_s), array_shape[0], array_shape[1], array_shape[2], array_shape[3]))
  1262. for i in tqdm(range(len(theta_values_u))):
  1263. for j in range(len(theta_values_s)):
  1264. # compute the filter
  1265. filter = temporal_filter_gauss(frequencies, theta_values_u[i], theta_values_s[j])
  1266. # compute the filtered spectra
  1267. filtered_stationary_stim_psd_avg = np.multiply(stationary_stim_psd_avg, filter)
  1268. filtered_moving_stim_psd_avg = np.multiply(moving_stim_psd_avg, filter)
  1269. # compute the information
  1270. moving_stim_information[i, j] = np.sum(np.log2(1 + filtered_moving_stim_psd_avg / noise), axis=-1)
  1271. stationary_stim_information[i, j] = np.sum(np.log2(1 + filtered_stationary_stim_psd_avg / noise), axis=-1)
  1272. # compute the energy
  1273. moving_stim_energy[i, j] = np.sum(filtered_moving_stim_psd_avg, axis=-1)
  1274. stationary_stim_energy[i, j] = np.sum(filtered_stationary_stim_psd_avg, axis=-1)
  1275. # %%
  1276. import matplotlib.colors as mcolors
  1277. norm = mcolors.CenteredNorm()
  1278. # plot the information as a heatmap of theta_u and theta_s
  1279. fig, ax = plt.subplots(1, 3, figsize=(24, 6))
  1280. im = ax[0].imshow(moving_stim_information[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_information[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]), aspect='auto', cmap='coolwarm', origin='lower', norm=norm)
  1281. ax[0].set_xlabel('SD')
  1282. ax[0].set_ylabel('Mean')
  1283. ax[0].set_title('Info')
  1284. ax[1].imshow(moving_stim_energy[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_energy[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]), aspect='auto', cmap='coolwarm', origin='lower', norm=norm)
  1285. ax[1].set_xlabel('SD')
  1286. ax[1].set_ylabel('Mean')
  1287. ax[1].set_title('Energy')
  1288. ax[2].imshow(moving_stim_information[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_information[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]) - 0.1*moving_stim_energy[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_energy[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]), aspect='auto', cmap='coolwarm', origin='lower', norm=norm)
  1289. ax[2].set_xlabel('SD')
  1290. ax[2].set_ylabel('Mean')
  1291. # %%
  1292. import matplotlib.colors as mcolors
  1293. norm = mcolors.CenteredNorm()
  1294. # plot the information as a heatmap of theta_u and theta_s
  1295. fig, ax = plt.subplots(1, 3, figsize=(24, 6))
  1296. im = ax[0].imshow(stationary_stim_information[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_information[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]), aspect='auto', cmap='coolwarm', origin='lower', norm=norm)
  1297. ax[0].set_xlabel('SD')
  1298. ax[0].set_ylabel('Mean')
  1299. ax[0].set_title('Info')
  1300. ax[1].imshow(stationary_stim_energy[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_energy[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]), aspect='auto', cmap='coolwarm', origin='lower', norm=norm)
  1301. ax[1].set_xlabel('SD')
  1302. ax[1].set_ylabel('Mean')
  1303. ax[1].set_title('Energy')
  1304. ax[2].imshow(stationary_stim_information[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_information[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]) - 0.1*stationary_stim_energy[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_energy[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]), aspect='auto', cmap='coolwarm', origin='lower', norm=norm)
  1305. ax[2].set_xlabel('SD')
  1306. ax[2].set_ylabel('Mean')
  1307. # %%
  1308. plt.plot(moving_stim_information[0, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_information[0, :, ori_idx, phase_idx, wavelength_idx, position_idx]) - 2 * moving_stim_energy[0, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_energy[0, :, ori_idx, phase_idx, wavelength_idx, position_idx]), 'o-', color='tab:blue', label='Moving test')
  1309. plt.plot(stationary_stim_information[0, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_information[0, :, ori_idx, phase_idx, wavelength_idx, position_idx]) - 2 * stationary_stim_energy[0, :, ori_idx, phase_idx, wavelength_idx, position_idx]/np.max(moving_stim_energy[0, :, ori_idx, phase_idx, wavelength_idx, position_idx]), 'o-', color='tab:red', label='Stationary test')
  1310. # %%
  1311. # compute the max info and energy for each filter
  1312. moving_max_info = np.max(moving_stim_information[0], axis=0)
  1313. stationary_max_info = np.max(stationary_stim_information[0], axis=0)
  1314. moving_max_energy = np.max(moving_stim_energy[0], axis=0)
  1315. stationary_max_energy = np.max(stationary_stim_energy[0], axis=0)
  1316. max_info = np.maximum(moving_max_info, stationary_max_info)
  1317. max_energy = np.maximum(moving_max_energy, stationary_max_energy)
  1318. # normalize
  1319. moving_stim_information_norm = moving_stim_information[0] / max_info
  1320. stationary_stim_information_norm = stationary_stim_information[0] / max_info
  1321. moving_stim_energy_norm = moving_stim_energy[0] / max_energy
  1322. stationary_stim_energy_norm = stationary_stim_energy[0] / max_energy
  1323. # %%
  1324. moving_stim_information_norm.shape
  1325. # %%
  1326. # compute the utility for a variety of lambdas
  1327. lmbd_arr = np.linspace(0, 2, 50)
  1328. moving_utility = np.zeros((lmbd_arr.shape[0], num_theta_values, len(orientation_arr), len(phase_arr), len(wavelength_arr[low_spatial_freq_idx]), len(position_arr)))
  1329. stationary_utility = np.zeros((lmbd_arr.shape[0], num_theta_values, len(orientation_arr), len(phase_arr), len(wavelength_arr[low_spatial_freq_idx]), len(position_arr)))
  1330. for i, lmbd in enumerate(lmbd_arr):
  1331. moving_utility[i] = moving_stim_information_norm - lmbd*moving_stim_energy_norm
  1332. stationary_utility[i] = stationary_stim_information_norm - lmbd*stationary_stim_energy_norm
  1333. # %%
  1334. moving_utility.shape
  1335. # %%
  1336. # %%
  1337. import matplotlib.colors as mcolors
  1338. norm = mcolors.TwoSlopeNorm(vmin=np.min([moving_utility, stationary_utility]), vcenter=0, vmax=np.max([moving_utility, stationary_utility]))
  1339. ori_idx = 0
  1340. phase_idx = 0
  1341. wavelength_idx = 10
  1342. position_idx = 4
  1343. # utility as as a heatmaps of lambda and theta
  1344. fig, ax = plt.subplots(1, 2, figsize=(16, 7))
  1345. ax[0].imshow(stationary_utility[:, :, ori_idx, phase_idx, wavelength_idx, position_idx].T, aspect='auto', cmap='coolwarm', origin='lower', norm=norm)
  1346. ax[0].set_xlabel('Lambda')
  1347. ax[0].set_ylabel('SD')
  1348. im = ax[1].imshow(moving_utility[:, :, ori_idx, phase_idx, wavelength_idx, position_idx].T, aspect='auto', cmap='coolwarm', origin='lower', norm=norm)
  1349. ax[1].set_xlabel('Lambda')
  1350. ax[1].set_ylabel('SD')
  1351. ax[0].set_title('Stationary Utility')
  1352. ax[1].set_title('Moving Utility')
  1353. # for each lambda find the optimal theta and then plot a line overlaid on the heatmap connecting these values
  1354. optimal_theta_stationary_idx = np.argmin(np.abs(stationary_utility[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]), axis=1)
  1355. optimal_theta_moving_idx = np.argmin(np.abs(moving_utility[:, :, ori_idx, phase_idx, wavelength_idx, position_idx]), axis=1)
  1356. optimal_theta_stationary = theta_values_s[optimal_theta_stationary_idx]
  1357. optimal_theta_moving = theta_values_s[optimal_theta_moving_idx]
  1358. ax[0].plot(np.arange(len(lmbd_arr)), optimal_theta_stationary_idx, color='k')
  1359. ax[1].plot(np.arange(len(lmbd_arr)), optimal_theta_moving_idx, color='k')
  1360. # %%
  1361. optimal_theta_stationary_idx.shape
  1362. # %%
  1363. plt.plot(optimal_theta_moving - optimal_theta_stationary)
  1364. # %% [markdown]
  1365. # ## Streamlined windowed and demeaned
  1366. # %%
  1367. # parse the data and compute fourier transform for each filter and video
  1368. environments = ['field', 'forest', 'orchard', 'tall_grass', 'pond']
  1369. num_videos = 10
  1370. vid_length = 50*30
  1371. low_spatial_freq_idx = np.arange(0, 31)
  1372. # low_spatial_freq_idx = np.arange(0, 19)
  1373. # high_spatial_freq_idx = np.arange(35, 70)
  1374. # low_spatial_freq_idx = np.arange(35, 70)
  1375. stationary_stim = np.zeros((len(environments), num_videos, len(orientation_arr), len(phase_arr), len(low_spatial_freq_idx), len(position_arr), vid_length))
  1376. moving_stim = np.zeros((len(environments), num_videos, len(orientation_arr), len(phase_arr), len(low_spatial_freq_idx), len(position_arr), vid_length))
  1377. # get the responses for each environment
  1378. for i, env_key in enumerate(environments):
  1379. stationary_count = 0
  1380. moving_count = 0
  1381. all_gabor_responses_env = all_gabor_responses[env_key]
  1382. print(all_gabor_responses_env.keys())
  1383. for vid_key in all_gabor_responses_env.keys():
  1384. if 'stationary' in vid_key:
  1385. stationary_stim[i, stationary_count, :, :, :, :, :] = all_gabor_responses_env[vid_key][()][:, :, low_spatial_freq_idx, :, :vid_length]
  1386. stationary_count += 1
  1387. if 'moving' in vid_key and 'free_moving' not in vid_key:
  1388. moving_stim[i, moving_count, :, :, :, :, :] = all_gabor_responses_env[vid_key][()][:, :, low_spatial_freq_idx, :, :vid_length]
  1389. moving_count += 1
  1390. # %%
  1391. from scipy.signal import welch
  1392. def compute_psd_welch(x, fs, window_size, overlap):
  1393. """
  1394. Compute one-sided PSD using Welch's method for a real-valued signal with variable window size and full overlap.
  1395. Parameters:
  1396. - x: 1D real-valued input signal
  1397. - fs: sampling frequency (Hz)
  1398. - window_size: length of each segment (samples)
  1399. Returns:
  1400. - freqs: frequency bins (Hz)
  1401. - psd: power spectral density (power/Hz)
  1402. """
  1403. freqs, psd = welch(
  1404. x,
  1405. fs=fs,
  1406. window='boxcar',
  1407. nperseg=window_size,
  1408. noverlap=overlap, # full overlap
  1409. return_onesided=True,
  1410. scaling='density',
  1411. detrend='constant'
  1412. )
  1413. return freqs, psd
  1414. # %%
  1415. # Compute the spectra for multiple window sizes: 10, 5, 2, and 1 seconds
  1416. sampling_rate = 30
  1417. window_lengths_sec = [5] # [10, 5, 2]
  1418. window_sizes = [int(w * sampling_rate) for w in window_lengths_sec]
  1419. # Prepare arrays to store PSDs for each window size
  1420. stationary_stim_psd_all = []
  1421. moving_stim_psd_all = []
  1422. frequencies_all = []
  1423. for window_size in window_sizes:
  1424. # Pre-allocate arrays for this window size
  1425. psd_shape = (
  1426. len(environments), num_videos, len(orientation_arr), len(phase_arr),
  1427. len(low_spatial_freq_idx), len(position_arr), window_size // 2 + 1
  1428. )
  1429. stationary_stim_psd = np.zeros(psd_shape)
  1430. moving_stim_psd = np.zeros(psd_shape)
  1431. frequencies = None
  1432. for i, env_key in enumerate(environments):
  1433. for j in range(num_videos):
  1434. _, stationary_stim_psd[i, j] = compute_psd_welch(
  1435. stationary_stim[i, j], sampling_rate, window_size, window_size // 2
  1436. )
  1437. frequencies, moving_stim_psd[i, j] = compute_psd_welch(
  1438. moving_stim[i, j], sampling_rate, window_size, window_size // 2
  1439. )
  1440. stationary_stim_psd_all.append(stationary_stim_psd)
  1441. moving_stim_psd_all.append(moving_stim_psd)
  1442. frequencies_all.append(frequencies)
  1443. # %%
  1444. stationary_stim_psd_avg_all = []
  1445. moving_stim_psd_avg_all = []
  1446. for i in range(len(window_sizes)):
  1447. stationary_stim_psd_avg_all.append(np.mean(stationary_stim_psd_all[i], axis=(0, 1)))
  1448. moving_stim_psd_avg_all.append(np.mean(moving_stim_psd_all[i], axis=(0, 1)))
  1449. # %%
  1450. def temporal_filter(temporal_freqs, theta, low_freq_cutoff=2):
  1451. # create the filter
  1452. filter = np.zeros(temporal_freqs.shape)
  1453. # set the filter to 1 for frequencies below theta
  1454. filter[temporal_freqs <= theta] = 1
  1455. # filter[temporal_freqs < low_freq_cutoff] = 1 - (low_freq_cutoff - temporal_freqs[temporal_freqs < low_freq_cutoff])
  1456. # create a linear dropoff for frequencies above theta with a fixed slope
  1457. filter[temporal_freqs > theta] = 1 - (temporal_freqs[temporal_freqs > theta] - theta)
  1458. # ensure that the filter is non-negative
  1459. filter[filter < 0] = 0
  1460. return np.square(filter)
  1461. # %%
  1462. mean_moving_psd_all = []
  1463. mean_stationary_psd_all = []
  1464. sd_moving_psd_all = []
  1465. sd_stationary_psd_all = []
  1466. for win_idx, (moving_stim_psd, stationary_stim_psd, frequencies) in enumerate(zip(moving_stim_psd_all, stationary_stim_psd_all, frequencies_all)):
  1467. mean_moving_psd = np.mean(moving_stim_psd, axis=(0, 1, 2, 3, 4, 5))
  1468. mean_stationary_psd = np.mean(stationary_stim_psd, axis=(0, 1, 2, 3, 4, 5))
  1469. sd_moving_psd = np.std(moving_stim_psd, axis=(0, 1, 2, 3, 4, 5))
  1470. sd_stationary_psd = np.std(stationary_stim_psd, axis=(0, 1, 2, 3, 4, 5))
  1471. mean_moving_psd_all.append(mean_moving_psd)
  1472. mean_stationary_psd_all.append(mean_stationary_psd)
  1473. sd_moving_psd_all.append(sd_moving_psd)
  1474. sd_stationary_psd_all.append(sd_stationary_psd)
  1475. # Plot and save
  1476. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  1477. ax.plot(frequencies[1:-1], mean_stationary_psd[1:-1], label='Stationary', color='tab:gray')
  1478. ax.plot(frequencies[1:-1], mean_moving_psd[1:-1], label='Moving', color='tab:orange')
  1479. # Optionally add SD shading:
  1480. # ax.fill_between(frequencies[1:-1], mean_stationary_psd[1:-1] - sd_stationary_psd[1:-1], mean_stationary_psd[1:-1] + sd_stationary_psd[1:-1], color='tab:gray', alpha=0.2)
  1481. # ax.fill_between(frequencies[1:-1], mean_moving_psd[1:-1] - sd_moving_psd[1:-1], mean_moving_psd[1:-1] + sd_moving_psd[1:-1], color='tab:orange', alpha=0.2)
  1482. ax.set_xlabel('Frequency (Hz)')
  1483. ax.set_ylabel('Power')
  1484. ax.set_title(f'Mean power spectrum ({window_lengths_sec[win_idx]}s window)')
  1485. ax.set_yscale('log')
  1486. ax.legend()
  1487. plt.savefig(f'../manuscript_figures/fig3_psd_mean_full_window_{window_lengths_sec[win_idx]}s.pdf', format='pdf', bbox_inches='tight')
  1488. plt.show()
  1489. plt.close(fig)
  1490. # %%
  1491. # compute the average spectrum over all filters
  1492. # plot the mean spectrum
  1493. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  1494. ax.plot(frequencies[1:-1], mean_moving_psd[1:-1], label='Moving', color='tab:orange')
  1495. ax.plot(frequencies[1:-1], mean_stationary_psd[1:-1], label='Stationary', color='tab:gray')
  1496. ax.set_xlabel('Frequency (Hz)')
  1497. ax.set_ylabel('Power')
  1498. ax.set_yscale('log')
  1499. # ax.set_xscale('log')
  1500. ax.legend()
  1501. # find where the area under the stationary psd is a certain value
  1502. # auc_val = 1
  1503. # aoc_idx_stat = np.where(np.cumsum(mean_stationary_psd*freq_spacing) <= auc_val)[0][-1]
  1504. # # find the frequency at that index
  1505. # freq_val_stat = frequencies[aoc_idx_stat]
  1506. # auc_idx_mov = np.where(np.cumsum(mean_moving_psd*freq_spacing) <= auc_val)[0][-1]
  1507. # # find the frequency at that index
  1508. # freq_val_mov = frequencies[auc_idx_mov]
  1509. auc_val_stat = 1
  1510. auc_val_mov = 2
  1511. cum_stationary = np.cumsum(mean_stationary_psd * freq_spacing)
  1512. cum_moving = np.cumsum(mean_moving_psd * freq_spacing)
  1513. # Interpolate to find the frequency where the cumulative sum reaches auc_val
  1514. freq_val_stat = np.interp(auc_val_stat, cum_stationary, frequencies)
  1515. freq_val_mov = np.interp(auc_val_mov, cum_moving, frequencies)
  1516. # # plot filled area under the curve
  1517. # ax.fill_between(frequencies[1:-1], mean_stationary_psd[1:-1], where=(frequencies[1:-1] <= freq_val_stat), color='tab:gray', alpha=0.5)
  1518. # ax.fill_between(frequencies[1:-1], mean_moving_psd[1:-1], where=(frequencies[1:-1] <= freq_val_mov), color='tab:orange', alpha=0.5)
  1519. # Interpolate the PSD at the cutoff frequencies
  1520. psd_val_stat = np.interp(freq_val_stat, frequencies, mean_stationary_psd)
  1521. psd_val_mov = np.interp(freq_val_mov, frequencies, mean_moving_psd)
  1522. # For fill_between, extend the frequency and PSD arrays to include the interpolated cutoff point
  1523. def extend_for_fill(frequencies, psd, cutoff_freq, cutoff_psd):
  1524. mask = frequencies <= cutoff_freq
  1525. # Find the last index before the cutoff
  1526. last_idx = np.where(mask)[0][-1]
  1527. # Insert the cutoff point after last_idx
  1528. new_freqs = np.insert(frequencies[mask], last_idx + 1, cutoff_freq)
  1529. new_psd = np.insert(psd[mask], last_idx + 1, cutoff_psd)
  1530. return new_freqs, new_psd
  1531. freqs_stat_fill, psd_stat_fill = extend_for_fill(frequencies[1:-1], mean_stationary_psd[1:-1], freq_val_stat, psd_val_stat)
  1532. freqs_mov_fill, psd_mov_fill = extend_for_fill(frequencies[1:-1], mean_moving_psd[1:-1], freq_val_mov, psd_val_mov)
  1533. # plot filled area under the curve using the interpolated cutoff
  1534. ax.fill_between(freqs_stat_fill[:-1], psd_stat_fill[:-1], color='tab:gray', alpha=0.5)
  1535. ax.fill_between(freqs_mov_fill, psd_mov_fill, color='tab:orange', alpha=0.5)
  1536. ax.set_xlim(0, 15)
  1537. ax.set_ylim(1e-3, 4)
  1538. # remove top and right spines
  1539. ax.spines['top'].set_visible(False)
  1540. ax.spines['right'].set_visible(False)
  1541. ax.set_title(f'Mean power spectrum, moving cutoff: {freq_val_mov:.2f} Hz, stationary cutoff: {freq_val_stat:.2f} Hz')
  1542. plt.savefig(f'../manuscript_figures/fig3_psd_mean_normalize_2Hz_crop_auc_stat_{auc_val_stat}_auc_mov_{auc_val_mov}.pdf', format='pdf', bbox_inches='tight')
  1543. # %%
  1544. # Normalize the spectrum for each window so that the area under the curve is 1
  1545. normalized_mean_stationary_psd_all = []
  1546. normalized_mean_moving_psd_all = []
  1547. for win_idx in range(len(window_sizes)):
  1548. frequencies = frequencies_all[win_idx]
  1549. mean_stationary_psd = mean_stationary_psd_all[win_idx]
  1550. mean_moving_psd = mean_moving_psd_all[win_idx]
  1551. freq_spacing = np.diff(frequencies)[0]
  1552. norm_factor = np.sum(mean_stationary_psd * freq_spacing)
  1553. mean_stationary_psd_norm = mean_stationary_psd / norm_factor
  1554. mean_moving_psd_norm = mean_moving_psd / norm_factor
  1555. normalized_mean_stationary_psd_all.append(mean_stationary_psd_norm)
  1556. normalized_mean_moving_psd_all.append(mean_moving_psd_norm)
  1557. # %%
  1558. # Repeat the process for a range of theta values for each window size
  1559. num_theta_values = 1000
  1560. theta_values = np.linspace(0, 15, num_theta_values)
  1561. all_moving_stim_energy = []
  1562. all_stationary_stim_energy = []
  1563. for win_idx in range(len(window_sizes)):
  1564. frequencies = frequencies_all[win_idx]
  1565. mean_stationary_psd = normalized_mean_stationary_psd_all[win_idx]
  1566. mean_moving_psd = normalized_mean_moving_psd_all[win_idx]
  1567. freq_spacing = np.diff(frequencies)[0]
  1568. moving_stim_energy = np.zeros(num_theta_values)
  1569. stationary_stim_energy = np.zeros(num_theta_values)
  1570. for i in tqdm(range(num_theta_values), desc=f"Window {window_lengths_sec[win_idx]}s"):
  1571. # compute the filter
  1572. filter = temporal_filter(frequencies, theta_values[i])
  1573. # compute the filtered spectra
  1574. filtered_mean_stationary_psd = mean_stationary_psd * filter
  1575. filtered_mean_moving_psd = mean_moving_psd * filter
  1576. # compute the energy
  1577. moving_stim_energy[i] = np.sum(filtered_mean_moving_psd * freq_spacing)
  1578. stationary_stim_energy[i] = np.sum(filtered_mean_stationary_psd * freq_spacing)
  1579. all_moving_stim_energy.append(moving_stim_energy)
  1580. all_stationary_stim_energy.append(stationary_stim_energy)
  1581. # %%
  1582. # Plot the energy vs. theta for each window size
  1583. fig, ax = plt.subplots(1, len(window_sizes), figsize=(3*1.5*len(window_sizes), 3*1.5))
  1584. if len(window_sizes) == 1:
  1585. ax = [ax]
  1586. for win_idx in range(len(window_sizes)):
  1587. ax[win_idx].axhline(1, color='k', linestyle='--')
  1588. ax[win_idx].axhline(2, color='tab:orange', linestyle='--')
  1589. ax[win_idx].plot(theta_values, all_moving_stim_energy[win_idx], label='Moving', color='tab:orange')
  1590. ax[win_idx].plot(theta_values, all_stationary_stim_energy[win_idx], label='Stationary', color='tab:gray')
  1591. ax[win_idx].set_xlabel('Theta')
  1592. ax[win_idx].set_title(f'Window: {window_lengths_sec[win_idx]}s')
  1593. ax[win_idx].set_yscale('log')
  1594. ax[win_idx].set_xlim(0, 15)
  1595. ax[win_idx].set_ylim(.3, 4)
  1596. # remove the right and top spines
  1597. ax[win_idx].spines['right'].set_visible(False)
  1598. ax[win_idx].spines['top'].set_visible(False)
  1599. if win_idx == 0:
  1600. ax[win_idx].set_ylabel('Energy')
  1601. ax[win_idx].legend()
  1602. plt.tight_layout()
  1603. plt.savefig(f'../manuscript_figures/fig3_energy_vs_theta_filtered.pdf', format='pdf', bbox_inches='tight')
  1604. plt.show()
  1605. # %%
  1606. freq_val_stat - freq_val_mov
  1607. # %% [markdown]
  1608. # ### Full 2D optimization
  1609. # %%
  1610. relative = True
  1611. # for given energy capacities in the stationary case (between 0 and 1)
  1612. energy_capacity_stat = np.linspace(0, 1, 101)
  1613. # energy_capacity_mov = np.linspace(0, 2, 101)
  1614. delta_energy_capacity = np.linspace(0.01, 10, 1000)
  1615. # energy_capacity_stat = np.logspace(-1, 0, 100)
  1616. # delta_energy_capacity = np.logspace(-1, 1, 101)
  1617. # interpolate the energy values to get a smooth curve
  1618. stationary_theta_interp = np.interp(energy_capacity_stat, stationary_stim_energy, theta_values)
  1619. if relative:
  1620. energy_capacity_mov = np.outer(delta_energy_capacity, energy_capacity_stat) + energy_capacity_stat[np.newaxis, :] # relative
  1621. delta_energy_capacity *= 100 # convert to percent
  1622. else:
  1623. energy_capacity_mov = delta_energy_capacity[:, np.newaxis] + energy_capacity_stat[np.newaxis, :] # absolute
  1624. moving_theta_interp = np.interp(energy_capacity_mov, moving_stim_energy, theta_values)
  1625. # %%
  1626. import matplotlib.colors as mcolors
  1627. from matplotlib import cm
  1628. norm = mcolors.Normalize(vmin=np.min(theta_values), vmax=np.max(theta_values))
  1629. two_color_norm = mcolors.TwoSlopeNorm(vmin=-5, vcenter=0, vmax=5)
  1630. # example_points = [(50, 50), (75, 75), (95, 100)]
  1631. example_points = [(75, 9), (90, 19), (95, 29), (99, 49)]
  1632. fig, ax = plt.subplots(1, 3, figsize=(16, 5), width_ratios=[5, 5, 5], sharey=True)
  1633. im0 = ax[0].imshow(np.repeat(stationary_theta_interp[np.newaxis, :], len(delta_energy_capacity), axis=0), origin='lower', norm=norm, cmap='gray_r', extent=[energy_capacity_stat[0], energy_capacity_stat[-1], delta_energy_capacity[0], delta_energy_capacity[-1]], aspect='auto')
  1634. ax[0].set_title('Stationary theta')
  1635. im1 = ax[1].imshow(moving_theta_interp, origin='lower', norm=norm, cmap='gray_r', extent=[energy_capacity_stat[0], energy_capacity_stat[-1], delta_energy_capacity[0], delta_energy_capacity[-1]], aspect='auto')
  1636. ax[1].set_title('Moving theta')
  1637. colors_map = mcolors.LinearSegmentedColormap.from_list('custom_cmap', [(1, 127/255, 14/255), (1, 1, 1), (0.5, 0.5, 0.5)], N=256)
  1638. im2 = ax[2].imshow(stationary_theta_interp[np.newaxis, :]-moving_theta_interp, origin='lower', norm=two_color_norm, cmap=colors_map, extent=[energy_capacity_stat[0], energy_capacity_stat[-1], delta_energy_capacity[0], delta_energy_capacity[-1]], aspect='auto')
  1639. ax[2].set_title('Difference in theta [stationary - moving]')
  1640. # ax[2].set_xlim(0.5, 1)
  1641. # ax[2].set_ylim(0.5, 1)
  1642. # add axis labels
  1643. if relative:
  1644. ax[0].set_ylabel('Percentage increase in energy capacity')
  1645. else:
  1646. ax[0].set_ylabel('Increase in energy capacity')
  1647. fig.supxlabel('Energy capacity')
  1648. # plot the example points in three different colors which are distinct from red, white and blue
  1649. colors_map = cm.get_cmap('Blues', len(example_points)+2)
  1650. # plot the example points in the first subplot
  1651. for i, point in enumerate(example_points):
  1652. ax[0].scatter(energy_capacity_stat[point[0]], delta_energy_capacity[point[1]], color=colors_map(i+1), s=100, edgecolor='black', label='Example point ' + str(i+1))
  1653. ax[1].scatter(energy_capacity_stat[point[0]], delta_energy_capacity[point[1]], color=colors_map(i+1), s=100, edgecolor='black')
  1654. ax[2].scatter(energy_capacity_stat[point[0]], delta_energy_capacity[point[1]], color=colors_map(i+1), s=100, edgecolor='black')
  1655. ax[0].set_yscale('log')
  1656. # add a colorbar without disrupting the layout of the other subplots
  1657. cbar = fig.colorbar(im2, ax=ax[2], orientation='vertical', pad=0.04)
  1658. cbar = fig.colorbar(im1, ax=ax[1], orientation='vertical', pad=0.04)
  1659. cbar = fig.colorbar(im0, ax=ax[0], orientation='vertical', pad=0.04)
  1660. plt.savefig(f'../manuscript_figures/fig3_diff_theta_stat_energy_vs_delta_energy_{window_lengths_sec[win_idx]}s.pdf', format='pdf', bbox_inches='tight')
  1661. # %%
  1662. from scipy.fft import irfft
  1663. # compute the autocorrelation function for the example points by extracting the theta values in the stationary and moving conditions and then applying the filter to white noise and then computing the autocorrelation function as the inverse fourier transform of the PSD
  1664. alpha_arr = [0.25, 0.5, 0.75, 1]
  1665. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  1666. for i in range(len(example_points)):
  1667. # compute the filter
  1668. example_theta_mov = moving_theta_interp[example_points[i][1], example_points[i][0]]
  1669. example_theta_stat = stationary_theta_interp[example_points[i][0]]
  1670. # compute the filter
  1671. example_filter_mov = temporal_filter(frequencies, example_theta_mov)
  1672. example_filter_stat = temporal_filter(frequencies, example_theta_stat)
  1673. # compute the inverse fourier transform
  1674. example_moving_autocorr = irfft(example_filter_mov/freq_spacing)
  1675. example_stationary_autocorr = irfft(example_filter_stat/freq_spacing)
  1676. # normalize the autocorrelation function
  1677. example_moving_autocorr /= np.max(example_moving_autocorr)
  1678. example_stationary_autocorr /= np.max(example_stationary_autocorr)
  1679. # plot the autocorrelation function
  1680. time = np.arange(0, len(example_moving_autocorr)) / 30*1000
  1681. ax.plot(time, example_moving_autocorr, label='Moving', color='tab:orange', alpha=alpha_arr[i])
  1682. ax.plot(time, example_stationary_autocorr, label='Stationary', color='tab:gray', alpha=alpha_arr[i])
  1683. ax.set_xlabel('Time (ms)')
  1684. ax.set_ylabel('Autocorrelation')
  1685. # ax.legend()
  1686. ax.set_xlim(0, 500)
  1687. plt.savefig(f'../manuscript_figures/fig3_autocorr_example_{i}_{window_lengths_sec[win_idx]}s.pdf', format='pdf', bbox_inches='tight')
  1688. plt.show()
  1689. plt.close(fig)
  1690. # %% [markdown]
  1691. # ### 1D optimization
  1692. # %%
  1693. # Repeat for all windows, but keep energy_capacity_stat fixed at 1
  1694. relative = True
  1695. energy_capacity_stat = 1.0 # Always use 1 for stationary energy capacity
  1696. delta_energy_capacity = np.linspace(0.01, 3, 300)
  1697. all_stationary_theta_interp = []
  1698. all_moving_theta_interp = []
  1699. for win_idx in range(len(window_sizes)):
  1700. stationary_stim_energy = all_stationary_stim_energy[win_idx]
  1701. moving_stim_energy = all_moving_stim_energy[win_idx]
  1702. # theta_values should be defined as before
  1703. # If not, use: theta_values = np.linspace(0, 15, stationary_stim_energy.shape[0])
  1704. # Interpolate stationary theta for energy_capacity_stat = 1
  1705. stationary_theta_interp = np.interp(energy_capacity_stat, stationary_stim_energy, theta_values)
  1706. if relative:
  1707. energy_capacity_mov = delta_energy_capacity*energy_capacity_stat + energy_capacity_stat # relative to stat
  1708. delta_energy_capacity_scaled = delta_energy_capacity * 100 # percent
  1709. else:
  1710. energy_capacity_mov = delta_energy_capacity + energy_capacity_stat # absolute
  1711. # Interpolate moving theta for each delta
  1712. moving_theta_interp = np.interp(energy_capacity_mov, moving_stim_energy, theta_values)
  1713. all_stationary_theta_interp.append(stationary_theta_interp)
  1714. all_moving_theta_interp.append(moving_theta_interp)
  1715. # %%
  1716. import matplotlib.colors as mcolors
  1717. from matplotlib import cm
  1718. # example_points = [(50, 50), (75, 75), (95, 100)]
  1719. example_points = [49, 99, 149]
  1720. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  1721. for win_idx in range(len(window_lengths_sec)):
  1722. ax.plot(delta_energy_capacity*100, all_stationary_theta_interp[win_idx]-all_moving_theta_interp[win_idx], label='Window: ' + str(window_lengths_sec[win_idx]) + 's', color=cm.Reds((win_idx+1)/len(window_lengths_sec)))
  1723. ax.scatter(delta_energy_capacity[example_points]*100, all_stationary_theta_interp[win_idx]-all_moving_theta_interp[win_idx][example_points], color=cm.viridis((np.arange(len(example_points))+1)/len(example_points)), s=50, edgecolor='black')
  1724. ax.set_xlabel('Increase in energy capacity (%)')
  1725. ax.set_ylabel('Theta difference (stationary - moving)')
  1726. ax.legend()
  1727. plt.savefig(f'../manuscript_figures/fig3_diff_theta_fixed_stat_energy_vs_delta_energy_{window_lengths_sec[win_idx]}s.pdf', format='pdf', bbox_inches='tight')
  1728. # %%
  1729. from scipy.fft import irfft
  1730. # Repeat for each window size
  1731. alpha_arr = [0.25, 0.5, 0.75, 1]
  1732. for win_idx in range(len(window_sizes)):
  1733. frequencies = frequencies_all[win_idx]
  1734. freq_spacing = np.diff(frequencies)[0]
  1735. stationary_theta_interp = all_stationary_theta_interp[win_idx]
  1736. moving_theta_interp = all_moving_theta_interp[win_idx]
  1737. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  1738. for i in range(len(example_points)):
  1739. # compute the filter
  1740. example_theta_mov = moving_theta_interp[example_points[i]]
  1741. example_theta_stat = stationary_theta_interp
  1742. # compute the filter
  1743. example_filter_mov = temporal_filter(frequencies, example_theta_mov)
  1744. example_filter_stat = temporal_filter(frequencies, example_theta_stat)
  1745. # set the DC to zero
  1746. example_filter_mov[0] = 0
  1747. example_filter_stat[0] = 0
  1748. # compute the inverse fourier transform
  1749. example_moving_autocorr = irfft(example_filter_mov)
  1750. example_stationary_autocorr = irfft(example_filter_stat)
  1751. # normalize the autocorrelation function
  1752. example_moving_autocorr /= np.max(example_moving_autocorr)
  1753. example_stationary_autocorr /= np.max(example_stationary_autocorr)
  1754. # plot the autocorrelation function
  1755. time = np.arange(0, len(example_moving_autocorr)) / 30*1000
  1756. ax.plot(time, example_moving_autocorr, color=cm.viridis((np.arange(len(example_points))+1)/len(example_points))[i], label=f'$\\Delta E$ {delta_energy_capacity[example_points[i]]*100}')
  1757. ax.set_xlabel('Time (ms)')
  1758. ax.set_ylabel('Autocorrelation')
  1759. ax.set_title(f'Autocorrelation for {window_lengths_sec[win_idx]}s window')
  1760. ax.set_xlim(0, 500)
  1761. ax.plot(time, example_stationary_autocorr, color='tab:gray', label='stationary')
  1762. ax.legend()
  1763. plt.savefig(f'../manuscript_figures/fig3_autocorr_{window_lengths_sec[win_idx]}s.pdf', format='pdf', bbox_inches='tight')
  1764. plt.show()
  1765. plt.close(fig)
  1766. # %% [markdown]
  1767. # ### Individual filters
  1768. # %%
  1769. # Repeat the normalization and energy computation for all window sizes
  1770. all_moving_stim_energy = []
  1771. all_stationary_stim_energy = []
  1772. num_theta_values = 1000
  1773. theta_values = np.linspace(0, 15, num_theta_values)
  1774. all_stationary_theta_interp = []
  1775. all_moving_theta_interp = []
  1776. num_energy_levels = 100
  1777. energy_capacity_stat = np.linspace(0, 1, num_energy_levels+1)
  1778. delta_energy_capacity = np.linspace(0.01, 10, 10*num_energy_levels)
  1779. relative = True
  1780. for win_idx in range(len(window_sizes)):
  1781. stationary_stim_psd_avg = stationary_stim_psd_avg_all[win_idx]
  1782. moving_stim_psd_avg = moving_stim_psd_avg_all[win_idx]
  1783. # Use the correct frequency array for this window size
  1784. frequencies = frequencies_all[win_idx]
  1785. freq_spacing = np.diff(frequencies)[0]
  1786. # Normalize the spectrum such that the area under the curve is 1
  1787. norm_factor = np.sum(stationary_stim_psd_avg * freq_spacing, axis=-1, keepdims=True)
  1788. stationary_stim_psd_avg_norm = stationary_stim_psd_avg / norm_factor
  1789. moving_stim_psd_avg_norm = moving_stim_psd_avg / norm_factor
  1790. mat_shape = stationary_stim_psd_avg_norm.shape
  1791. # Compute energy for a range of theta values
  1792. moving_stim_energy = np.zeros((*mat_shape[:-1], num_theta_values))
  1793. stationary_stim_energy = np.zeros((*mat_shape[:-1], num_theta_values))
  1794. for i in tqdm(range(num_theta_values), desc=f"Window {window_lengths_sec[win_idx]}s"):
  1795. filter = temporal_filter(frequencies, theta_values[i])
  1796. filtered_mean_stationary_psd = stationary_stim_psd_avg_norm * filter
  1797. filtered_mean_moving_psd = moving_stim_psd_avg_norm * filter
  1798. moving_stim_energy[:, :, :, :, i] = np.sum(filtered_mean_moving_psd * freq_spacing, axis=-1)
  1799. stationary_stim_energy[:, :, :, :, i] = np.sum(filtered_mean_stationary_psd * freq_spacing, axis=-1)
  1800. all_moving_stim_energy.append(moving_stim_energy)
  1801. all_stationary_stim_energy.append(stationary_stim_energy)
  1802. # --- Interpolation block (from your prompt) ---
  1803. if relative:
  1804. energy_capacity_mov = np.outer(delta_energy_capacity, energy_capacity_stat) + energy_capacity_stat[np.newaxis, :]
  1805. delta_energy_capacity_scaled = delta_energy_capacity * 100 # convert to percent
  1806. else:
  1807. energy_capacity_mov = delta_energy_capacity[:, np.newaxis] + energy_capacity_stat[np.newaxis, :]
  1808. stationary_theta_interp = np.zeros((num_energy_levels+1, *mat_shape[:-1]))
  1809. moving_theta_interp = np.zeros((10*num_energy_levels, num_energy_levels+1, *mat_shape[:-1]))
  1810. for i in range(mat_shape[0]):
  1811. for j in range(mat_shape[1]):
  1812. for k in range(mat_shape[2]):
  1813. for l in range(mat_shape[3]):
  1814. stationary_theta_interp[:, i, j, k, l] = np.interp(
  1815. energy_capacity_stat, stationary_stim_energy[i, j, k, l], theta_values
  1816. )
  1817. for m in range(num_energy_levels):
  1818. moving_theta_interp[:, m, i, j, k, l] = np.interp(
  1819. energy_capacity_mov[:, m], moving_stim_energy[i, j, k, l], theta_values
  1820. )
  1821. all_stationary_theta_interp.append(stationary_theta_interp)
  1822. all_moving_theta_interp.append(moving_theta_interp)
  1823. # %%
  1824. # For each window size, compute the cutoff value of the autocorrelation function for different points in the energy landscape
  1825. cutoff = 0.5
  1826. example_points = [(75, 9), (90, 19), (95, 29), (99, 49)] # or use your own
  1827. all_cutoff_frequency_mov = []
  1828. all_cutoff_frequency_stat = []
  1829. all_autocorr_func_mov = []
  1830. all_autocorr_func_stat = []
  1831. for win_idx in range(len(all_stationary_theta_interp)):
  1832. stationary_theta_interp = all_stationary_theta_interp[win_idx]
  1833. moving_theta_interp = all_moving_theta_interp[win_idx]
  1834. frequencies = frequencies_all[win_idx]
  1835. mat_shape = stationary_theta_interp.shape[1:] # (ori, phase, freq, pos)
  1836. cutoff_frequency_mov = np.zeros((len(example_points), *mat_shape))
  1837. cutoff_frequency_stat = np.zeros((len(example_points), *mat_shape))
  1838. autocorr_func_mov = np.zeros((len(example_points), *mat_shape, 2*(len(frequencies)-1)))
  1839. autocorr_func_stat = np.zeros((len(example_points), *mat_shape, 2*(len(frequencies)-1)))
  1840. for m, (stat_idx, mov_idx) in enumerate(example_points):
  1841. for i in range(mat_shape[0]):
  1842. for j in range(mat_shape[1]):
  1843. for k in range(mat_shape[2]):
  1844. for l in range(mat_shape[3]):
  1845. # Get theta values for this filter and energy point
  1846. example_theta_stat = stationary_theta_interp[stat_idx, i, j, k, l]
  1847. example_theta_mov = moving_theta_interp[mov_idx, stat_idx, i, j, k, l]
  1848. # Compute temporal filters
  1849. example_filter_stat = temporal_filter(frequencies, example_theta_stat)
  1850. example_filter_mov = temporal_filter(frequencies, example_theta_mov)
  1851. example_filter_stat[0] = 0
  1852. example_filter_mov[0] = 0
  1853. # Compute autocorrelation (inverse FFT of filter)
  1854. example_stationary_autocorr = irfft(example_filter_stat)
  1855. example_moving_autocorr = irfft(example_filter_mov)
  1856. # Normalize
  1857. example_stationary_autocorr /= np.max(example_stationary_autocorr)
  1858. example_moving_autocorr /= np.max(example_moving_autocorr)
  1859. # Store
  1860. autocorr_func_stat[m, i, j, k, l, :] = example_stationary_autocorr
  1861. autocorr_func_mov[m, i, j, k, l, :] = example_moving_autocorr
  1862. # Find cutoff
  1863. try:
  1864. cutoff_frequency_stat[m, i, j, k, l] = np.where(example_stationary_autocorr < cutoff)[0][0]
  1865. except IndexError:
  1866. cutoff_frequency_stat[m, i, j, k, l] = len(example_stationary_autocorr)
  1867. try:
  1868. cutoff_frequency_mov[m, i, j, k, l] = np.where(example_moving_autocorr < cutoff)[0][0]
  1869. except IndexError:
  1870. cutoff_frequency_mov[m, i, j, k, l] = len(example_moving_autocorr)
  1871. all_cutoff_frequency_mov.append(cutoff_frequency_mov)
  1872. all_cutoff_frequency_stat.append(cutoff_frequency_stat)
  1873. all_autocorr_func_mov.append(autocorr_func_mov)
  1874. all_autocorr_func_stat.append(autocorr_func_stat)
  1875. # Now all_cutoff_frequency_mov, all_cutoff_frequency_stat, all_autocorr_func_mov, all_autocorr_func_stat
  1876. # are lists, one per window size, containing the results for each window.
  1877. # %%
  1878. # Plot the average autocorrelation function for each point and window size
  1879. for win_idx in range(len(all_autocorr_func_mov)):
  1880. avg_autocorr_func_mov = np.mean(all_autocorr_func_mov[win_idx], axis=(1, 2, 3, 4))
  1881. avg_autocorr_func_stat = np.mean(all_autocorr_func_stat[win_idx], axis=(1, 2, 3, 4))
  1882. std_autocorr_func_mov = np.std(all_autocorr_func_mov[win_idx], axis=(1, 2, 3, 4))
  1883. std_autocorr_func_stat = np.std(all_autocorr_func_stat[win_idx], axis=(1, 2, 3, 4))
  1884. # Use the correct time axis for this window size
  1885. time = np.arange(0, avg_autocorr_func_mov.shape[1]) / 30 * 1000
  1886. for point_idx in range(len(example_points)):
  1887. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  1888. ax.plot(time, avg_autocorr_func_mov[point_idx, :], label='Moving', color='tab:orange')
  1889. ax.plot(time, avg_autocorr_func_stat[point_idx, :], label='Stationary', color='tab:gray')
  1890. ax.fill_between(
  1891. time,
  1892. avg_autocorr_func_mov[point_idx, :] - std_autocorr_func_mov[point_idx, :],
  1893. avg_autocorr_func_mov[point_idx, :] + std_autocorr_func_mov[point_idx, :],
  1894. color='tab:orange', alpha=0.2
  1895. )
  1896. ax.fill_between(
  1897. time,
  1898. avg_autocorr_func_stat[point_idx, :] - std_autocorr_func_stat[point_idx, :],
  1899. avg_autocorr_func_stat[point_idx, :] + std_autocorr_func_stat[point_idx, :],
  1900. color='tab:gray', alpha=0.2
  1901. )
  1902. ax.set_xlim(0, 500)
  1903. ax.set_ylim(-1, 1)
  1904. ax.set_title(f'Window size: {window_lengths_sec[win_idx]}s, Point {point_idx}')
  1905. plt.savefig(f'../manuscript_figures/fig3_avg_autocorr_point_{point_idx}_window_{window_lengths_sec[win_idx]}s.pdf', format='pdf', bbox_inches='tight')
  1906. plt.show()
  1907. plt.close(fig)
  1908. # ...existing code...
  1909. # %%
  1910. all_cutoff_frequency_mov = np.array(all_cutoff_frequency_mov)
  1911. all_cutoff_frequency_stat = np.array(all_cutoff_frequency_stat)
  1912. all_cutoff_frequency_stat = all_cutoff_frequency_stat / 30 * 1000
  1913. all_cutoff_frequency_mov = all_cutoff_frequency_mov / 30 * 1000
  1914. diff_cutoff_frequency = all_cutoff_frequency_mov - all_cutoff_frequency_stat
  1915. # %%
  1916. # take the difference in cutoffs and plot the histogram
  1917. bins = np.linspace(-300, 300, 20)
  1918. window_idx = 3
  1919. point_idx = 3
  1920. for window_idx in range(len(window_sizes)):
  1921. for point_idx in range(len(example_points)):
  1922. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  1923. ax.hist(diff_cutoff_frequency[window_idx, point_idx].flatten(), bins=bins, density=True)
  1924. ax.set_xlabel('Difference of cutoff (moving - stationary) in ms')
  1925. ax.axvline(np.mean(diff_cutoff_frequency[window_idx, point_idx]), color='k', linestyle='--')
  1926. ax.set_title(f'Window size: {window_sizes[window_idx]/30}s, Point {point_idx}')
  1927. plt.savefig(f'../manuscript_figures/fig3_hist_cutoff_diff_theory_window_{window_sizes[window_idx]}_point_{point_idx}.pdf', format='pdf', bbox_inches='tight')
  1928. plt.show()
  1929. plt.close(fig)
  1930. # %%
  1931. diff_cutoff_frequency.shape
  1932. # %%
  1933. avg_diff_cutoff_frequency = np.mean(diff_cutoff_frequency, axis=(2, 3, 4, 5))
  1934. median_diff_cutoff_frequency = np.median(diff_cutoff_frequency, axis=(2, 3, 4, 5))
  1935. # %%
  1936. norm = mcolors.TwoSlopeNorm(vmin=-100, vcenter=0, vmax=100)
  1937. energy_capacity_stat = np.linspace(0, 1, num_energy_levels+1)
  1938. delta_energy_capacity = np.linspace(0.01, 10, 10*num_energy_levels)
  1939. im = plt.imshow(avg_diff_cutoff_frequency.T, norm=norm, cmap='bwr', aspect='auto')
  1940. plt.xticks(np.arange(len(window_sizes)), np.array(window_sizes)/30)
  1941. plt.xlabel('Window size (s)')
  1942. plt.ylabel('Constraint point')
  1943. plt.yticks(np.arange(len(example_points)), [f'$E_s$ {energy_capacity_stat[example_points[i][0]]:.2f}\n$\\Delta E$ {delta_energy_capacity[example_points[i][1]]:.2f}' for i in range(len(example_points))])
  1944. plt.colorbar(im, label='moving cutoff - stationary cutoff (ms)')
  1945. plt.suptitle('Average difference in cutoff frequency for different energy constraints')
  1946. plt.savefig(f'../manuscript_figures/fig3_avg_diff_cutoff_frequency.pdf', format='pdf', bbox_inches='tight')
  1947. # %%
  1948. norm = mcolors.TwoSlopeNorm(vmin=-100, vcenter=0, vmax=100)
  1949. energy_capacity_stat = np.linspace(0, 1, num_energy_levels+1)
  1950. delta_energy_capacity = np.linspace(0.01, 10, 10*num_energy_levels)
  1951. im = plt.imshow(median_diff_cutoff_frequency.T, norm=norm, cmap='bwr', aspect='auto')
  1952. plt.xticks(np.arange(len(window_sizes)), np.array(window_sizes)/30)
  1953. plt.xlabel('Window size (s)')
  1954. plt.ylabel('Constraint point')
  1955. plt.yticks(np.arange(len(example_points)), [f'$E_s$ {energy_capacity_stat[example_points[i][0]]:.2f}\n$\\Delta E$ {delta_energy_capacity[example_points[i][1]]:.2f}' for i in range(len(example_points))])
  1956. plt.colorbar(im, label='moving cutoff - stationary cutoff (ms)')
  1957. plt.suptitle('Median difference in cutoff frequency for different energy constraints')
  1958. plt.savefig(f'../manuscript_figures/fig3_median_diff_cutoff_frequency.pdf', format='pdf', bbox_inches='tight')
  1959. # %% [markdown]
  1960. # #### 1D optimization
  1961. # %%
  1962. # Repeat the normalization and energy computation for all window sizes
  1963. all_moving_stim_energy = []
  1964. all_stationary_stim_energy = []
  1965. num_theta_values = 1000
  1966. theta_values = np.linspace(0, 15, num_theta_values)
  1967. all_stationary_theta_interp = []
  1968. all_moving_theta_interp = []
  1969. num_energy_levels = 30
  1970. energy_capacity_stat = 1.0
  1971. delta_energy_capacity = np.linspace(0.01, 3, 10*num_energy_levels)
  1972. relative = True
  1973. for win_idx in range(len(window_sizes)):
  1974. stationary_stim_psd_avg = stationary_stim_psd_avg_all[win_idx]
  1975. moving_stim_psd_avg = moving_stim_psd_avg_all[win_idx]
  1976. # Use the correct frequency array for this window size
  1977. frequencies = frequencies_all[win_idx]
  1978. freq_spacing = np.diff(frequencies)[0]
  1979. # Normalize the spectrum such that the area under the curve is 1
  1980. norm_factor = np.sum(stationary_stim_psd_avg * freq_spacing, axis=-1, keepdims=True)
  1981. stationary_stim_psd_avg_norm = stationary_stim_psd_avg / norm_factor
  1982. moving_stim_psd_avg_norm = moving_stim_psd_avg / norm_factor
  1983. mat_shape = stationary_stim_psd_avg_norm.shape
  1984. # Compute energy for a range of theta values
  1985. moving_stim_energy = np.zeros((*mat_shape[:-1], num_theta_values))
  1986. stationary_stim_energy = np.zeros((*mat_shape[:-1], num_theta_values))
  1987. for i in tqdm(range(num_theta_values), desc=f"Window {window_lengths_sec[win_idx]}s"):
  1988. filter = temporal_filter(frequencies, theta_values[i])
  1989. filtered_mean_stationary_psd = stationary_stim_psd_avg_norm * filter
  1990. filtered_mean_moving_psd = moving_stim_psd_avg_norm * filter
  1991. moving_stim_energy[:, :, :, :, i] = np.sum(filtered_mean_moving_psd * freq_spacing, axis=-1)
  1992. stationary_stim_energy[:, :, :, :, i] = np.sum(filtered_mean_stationary_psd * freq_spacing, axis=-1)
  1993. all_moving_stim_energy.append(moving_stim_energy)
  1994. all_stationary_stim_energy.append(stationary_stim_energy)
  1995. # --- Interpolation block (from your prompt) ---
  1996. if relative:
  1997. energy_capacity_mov = delta_energy_capacity*energy_capacity_stat + energy_capacity_stat # relative to stat
  1998. delta_energy_capacity_scaled = delta_energy_capacity * 100 # percent
  1999. else:
  2000. energy_capacity_mov = delta_energy_capacity[:, np.newaxis] + energy_capacity_stat[np.newaxis, :]
  2001. stationary_theta_interp = np.zeros(mat_shape[:-1])
  2002. moving_theta_interp = np.zeros((10*num_energy_levels, *mat_shape[:-1]))
  2003. for i in range(mat_shape[0]):
  2004. for j in range(mat_shape[1]):
  2005. for k in range(mat_shape[2]):
  2006. for l in range(mat_shape[3]):
  2007. stationary_theta_interp[i, j, k, l] = np.interp(
  2008. energy_capacity_stat, stationary_stim_energy[i, j, k, l], theta_values
  2009. )
  2010. moving_theta_interp[:, i, j, k, l] = np.interp(
  2011. energy_capacity_mov, moving_stim_energy[i, j, k, l], theta_values
  2012. )
  2013. all_stationary_theta_interp.append(stationary_theta_interp)
  2014. all_moving_theta_interp.append(moving_theta_interp)
  2015. # %%
  2016. moving_theta_interp[49, 0, -1, :, 4]
  2017. # %%
  2018. from scipy.fft import irfft
  2019. cutoff = 0.5
  2020. example_points = [49, 99, 149] # or your own, for 1D fixed stationary energy
  2021. all_cutoff_frequency_mov = []
  2022. all_cutoff_frequency_stat = []
  2023. all_autocorr_func_mov = []
  2024. all_autocorr_func_stat = []
  2025. for win_idx in range(len(all_stationary_theta_interp)):
  2026. stationary_theta_interp = all_stationary_theta_interp[win_idx] # shape: (filters...)
  2027. moving_theta_interp = all_moving_theta_interp[win_idx] # shape: (delta_energy, filters...)
  2028. frequencies = frequencies_all[win_idx]
  2029. mat_shape = stationary_theta_interp.shape # (ori, phase, freq, pos)
  2030. cutoff_frequency_mov = np.zeros((len(example_points), *mat_shape))
  2031. cutoff_frequency_stat = np.zeros((len(example_points), *mat_shape))
  2032. autocorr_func_mov = np.zeros((len(example_points), *mat_shape, 2*(len(frequencies)-1)))
  2033. autocorr_func_stat = np.zeros((len(example_points), *mat_shape, 2*(len(frequencies)-1)))
  2034. for m, mov_idx in enumerate(example_points):
  2035. for i in range(mat_shape[0]):
  2036. for j in range(mat_shape[1]):
  2037. for k in range(mat_shape[2]):
  2038. for l in range(mat_shape[3]):
  2039. # Get theta values for this filter and energy point
  2040. example_theta_stat = stationary_theta_interp[i, j, k, l]
  2041. example_theta_mov = moving_theta_interp[mov_idx, i, j, k, l]
  2042. # Compute temporal filters
  2043. example_filter_stat = temporal_filter(frequencies, example_theta_stat)
  2044. example_filter_mov = temporal_filter(frequencies, example_theta_mov)
  2045. example_filter_stat[0] = 0
  2046. example_filter_mov[0] = 0
  2047. # Compute autocorrelation (inverse FFT of filter)
  2048. example_stationary_autocorr = irfft(example_filter_stat)
  2049. example_moving_autocorr = irfft(example_filter_mov)
  2050. # Normalize
  2051. example_stationary_autocorr /= np.max(example_stationary_autocorr)
  2052. example_moving_autocorr /= np.max(example_moving_autocorr)
  2053. # Store
  2054. autocorr_func_stat[m, i, j, k, l, :] = example_stationary_autocorr
  2055. autocorr_func_mov[m, i, j, k, l, :] = example_moving_autocorr
  2056. # Find cutoff
  2057. try:
  2058. cutoff_frequency_stat[m, i, j, k, l] = np.where(example_stationary_autocorr < cutoff)[0][0]
  2059. except IndexError:
  2060. cutoff_frequency_stat[m, i, j, k, l] = len(example_stationary_autocorr)
  2061. try:
  2062. cutoff_frequency_mov[m, i, j, k, l] = np.where(example_moving_autocorr < cutoff)[0][0]
  2063. except IndexError:
  2064. cutoff_frequency_mov[m, i, j, k, l] = len(example_moving_autocorr)
  2065. all_cutoff_frequency_mov.append(cutoff_frequency_mov)
  2066. all_cutoff_frequency_stat.append(cutoff_frequency_stat)
  2067. all_autocorr_func_mov.append(autocorr_func_mov)
  2068. all_autocorr_func_stat.append(autocorr_func_stat)
  2069. # %%
  2070. # do a paired t-test on the cutoff frequencies for each point and window size
  2071. from scipy.stats import ttest_rel
  2072. mov_cutoff = all_cutoff_frequency_mov[0, 0]
  2073. stat_cutoff = all_cutoff_frequency_stat[0, 0]
  2074. t_stat, p_value = ttest_rel(mov_cutoff.flatten(), stat_cutoff.flatten())
  2075. print(f"T-statistic: {t_stat}, P-value: {p_value}")
  2076. # %%
  2077. import sys
  2078. print(sys.float_info.min)
  2079. # %%
  2080. # Plot the average autocorrelation function for each point and window size
  2081. for win_idx in range(len(all_autocorr_func_mov)):
  2082. avg_autocorr_func_mov = np.mean(all_autocorr_func_mov[win_idx], axis=(1, 2, 3, 4))
  2083. avg_autocorr_func_stat = np.mean(all_autocorr_func_stat[win_idx], axis=(1, 2, 3, 4))
  2084. std_autocorr_func_mov = np.std(all_autocorr_func_mov[win_idx], axis=(1, 2, 3, 4))
  2085. std_autocorr_func_stat = np.std(all_autocorr_func_stat[win_idx], axis=(1, 2, 3, 4))
  2086. # Use the correct time axis for this window size
  2087. time = np.arange(0, avg_autocorr_func_mov.shape[1]) / 30 * 1000
  2088. for point_idx in range(len(example_points)):
  2089. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  2090. ax.plot(time, avg_autocorr_func_mov[point_idx, :], label='Moving', color='tab:orange')
  2091. ax.plot(time, avg_autocorr_func_stat[point_idx, :], label='Stationary', color='tab:gray')
  2092. ax.fill_between(
  2093. time,
  2094. avg_autocorr_func_mov[point_idx, :] - std_autocorr_func_mov[point_idx, :],
  2095. avg_autocorr_func_mov[point_idx, :] + std_autocorr_func_mov[point_idx, :],
  2096. color='tab:orange', alpha=0.2
  2097. )
  2098. ax.fill_between(
  2099. time,
  2100. avg_autocorr_func_stat[point_idx, :] - std_autocorr_func_stat[point_idx, :],
  2101. avg_autocorr_func_stat[point_idx, :] + std_autocorr_func_stat[point_idx, :],
  2102. color='tab:gray', alpha=0.2
  2103. )
  2104. ax.set_xlim(0, 500)
  2105. ax.set_ylim(-1, 1)
  2106. ax.set_title(f'Window size: {window_lengths_sec[win_idx]}s, Point {point_idx}')
  2107. plt.savefig(f'../manuscript_figures/fig3_avg_autocorr_point_{point_idx}_window_{window_lengths_sec[win_idx]}s.pdf', format='pdf', bbox_inches='tight')
  2108. plt.show()
  2109. plt.close(fig)
  2110. # ...existing code...
  2111. # %%
  2112. all_autocorr_func_mov[1].shape
  2113. # %%
  2114. import os
  2115. # Directory to save figures
  2116. save_dir = "../manuscript_figures/autocorr_individual_corrected"
  2117. os.makedirs(save_dir, exist_ok=True)
  2118. # plot each individual autocorrelation function for a given windows size and position, for all example points, orientations, phases, and wavelengths
  2119. window_idx = 1
  2120. position_idx = 4 # choose a position index, e.g., 0 for the first position
  2121. example_points = [49] # or your own, for 1D fixed stationary energy
  2122. frequencies = frequencies_all[window_idx]
  2123. mat_shape = all_autocorr_func_mov[window_idx].shape[1:5] # (ori, phase, freq, pos)
  2124. i = 0
  2125. point_idx = 49
  2126. for ori_idx in range(mat_shape[0]):
  2127. for phase_idx in range(mat_shape[1]):
  2128. for freq_id in range(mat_shape[2]):
  2129. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  2130. ax.plot(
  2131. np.arange(0, len(all_autocorr_func_mov[window_idx][i, ori_idx, phase_idx, freq_id, position_idx])) / 30 * 1000,
  2132. all_autocorr_func_mov[window_idx][i, ori_idx, phase_idx, freq_id, position_idx],
  2133. label='Moving', color='tab:orange'
  2134. )
  2135. ax.plot(
  2136. np.arange(0, len(all_autocorr_func_stat[window_idx][i, ori_idx, phase_idx, freq_id, position_idx])) / 30 * 1000,
  2137. all_autocorr_func_stat[window_idx][i, ori_idx, phase_idx, freq_id, position_idx],
  2138. label='Stationary', color='tab:gray'
  2139. )
  2140. ax.set_xlim(0, 500)
  2141. ax.set_ylim(-1, 1)
  2142. ax.set_title(f'Window size: {window_lengths_sec[window_idx]}s\n$\\Delta E$ {delta_energy_capacity[point_idx]*100:.2f}, Ori {orientation_arr[ori_idx]:.2f}, Phase {phase_arr[phase_idx]:.2f}, Freq {freq_arr[freq_id]:.2f} cpd')
  2143. ax.set_xlabel('Time (ms)')
  2144. ax.set_ylabel('Autocorrelation')
  2145. ax.legend()
  2146. # Save figure
  2147. fname = (
  2148. f"autocorr_win{window_lengths_sec[window_idx]}s_point{point_idx}_"
  2149. f"ori{ori_idx}_phase{phase_idx}_freq{freq_id}_pos{position_idx}.pdf"
  2150. )
  2151. plt.savefig(os.path.join(save_dir, fname), bbox_inches='tight', format='pdf')
  2152. plt.close(fig)
  2153. # %%
  2154. all_cutoff_frequency_mov = np.array(all_cutoff_frequency_mov)
  2155. all_cutoff_frequency_stat = np.array(all_cutoff_frequency_stat)
  2156. all_cutoff_frequency_stat = all_cutoff_frequency_stat / 30 * 1000
  2157. all_cutoff_frequency_mov = all_cutoff_frequency_mov / 30 * 1000
  2158. diff_cutoff_frequency = all_cutoff_frequency_mov - all_cutoff_frequency_stat
  2159. # %%
  2160. # take the difference in cutoffs and plot the histogram
  2161. bins = np.linspace(-300, 300, 20)
  2162. window_idx = 3
  2163. point_idx = 3
  2164. for window_idx in range(len(window_sizes)):
  2165. for point_idx in range(len(example_points)):
  2166. fig, ax = plt.subplots(1, 1, figsize=(4*1.5, 3*1.5))
  2167. ax.hist(diff_cutoff_frequency[window_idx, point_idx].flatten(), bins=bins, density=True)
  2168. ax.set_xlabel('Difference of cutoff (moving - stationary) in ms')
  2169. ax.axvline(np.mean(diff_cutoff_frequency[window_idx, point_idx]), color='k', linestyle='--')
  2170. ax.set_title(f'Window size: {window_sizes[window_idx]/30}s, Point {point_idx}')
  2171. plt.savefig(f'../manuscript_figures/fig3_hist_cutoff_diff_theory_window_{window_sizes[window_idx]}_point_{point_idx}.pdf', format='pdf', bbox_inches='tight')
  2172. plt.show()
  2173. plt.close(fig)
  2174. # %%

temporal_filtering_analysis.ipynb at commit 24c076f, no license · at the source

Overview

  1. Faculty of Biology, LMU, Munich, Germany
  2. Graduate School of Systemic Neurosciences, Munich, Germany
  3. Bernstein Center for Computational Neuroscience Munich, Munich, Germany
Journal: Science advances, volume 12, issue 35, article eaed4172
Dates: received 27 October 2025; accepted 20 July 2026; published online 28 August 2026; in print August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1126/sciadv.aed4172 · PMID 42664357 · PMCID PMC13524049 · OpenAlex W7154616600
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: cellular / molecular (subfield)
Methods: Smoothing, state filtering, decompositions, Statistics, Connectivity, Machine learning, Single-unit activity, calcium imaging
MeSH: Locomotion*, Models, Neurological*, Primates*, Rodentia*, Sensory Receptor Cells*, Animals (* major topic)
Journal subjects: Neuroscience
Topic: Neural dynamics and brain function (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Citations: not cited yet (Europe PMC); 64 references in the paper

Abstract

Behavior modulates the activity of sensory systems in multiple ways: from gain changes in individual neurons to changing interactions in neural populations. These effects are not universal; while movement has a strong influence on sensory coding in rodents, its impact on primates is less prominent. The diversity of effects that locomotion exerts on sensory neurons, as well as disparities between species, raises questions about the existence of universal principles that may underlie sensation during behavior. We propose that sensory systems are internally modulated to match systematic changes in stimulus statistics caused by locomotion, to facilitate an accurate and efficient sensory code. We find that model neurons, adapted to stimuli recorded during movement in natural environments, predict and reproduce a broad spectrum of experimental observations in rodents and primates. This simple principle of maintaining coding efficiency across behavioral states reconciles the diversity of ways in which locomotion modulates visual coding in different animal species.

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

Zenodo 20624372

License: CC-BY-4.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: “Data, code, and materials availability:”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (31 files), h5py (21 files), Matplotlib (20 files), SciPy (19 files), scikit-learn (7 files), OpenCV (6 files), pandas (2 files), PyTorch (2 files), Numba (1 file), rpy2 (1 file), statsmodels (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)
32 files

mlynarski-group/locomotion-modulation

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 24c076f12556f265925d8408ef9f6bff8131385f, 10 June 2026
Languages: Jupyter (18), Python (13)
Size: 42 files, 31 scripts
Software Heritage: not archived
Found in: the Zenodo archive record
Holds: README, environment (requirements.txt), 18 notebooks
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (31 files), h5py (21 files), Matplotlib (20 files), SciPy (19 files), scikit-learn (7 files), OpenCV (6 files), pandas (2 files), PyTorch (2 files), Numba (1 file), rpy2 (1 file), statsmodels (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
32 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;
  • 62 scripts, each with its path and the digest of its content;
  • 15 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

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

Data

Datasets cited

Data, code, and materials availability

All data and code needed to evaluate and reproduce the results in the paper are present in the paper and/or the Supplementary Materials. The data (natural videos) used in the analysis are available at https://doi.org/10.12751/g-node.14m4gq. The code is publicly available at https://doi.org/10.5281/zenodo.20624372. No new materials were generated in this study.

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, 6 MeSH terms, 59 references.

Cite

This paper

Gant, J. M., & Młynarski, W. F. (2026). Locomotion optimizes sensory representations through a computational principle shared by rodents and primates. Science advances, 12(35), eaed4172. https://doi.org/10.1126/sciadv.aed4172

BibTeX

@article{gant2026locomotion,
author = {Gant, Jonathan M. and Młynarski, Wiktor F.},
title = {{Locomotion optimizes sensory representations through a computational principle shared by rodents and primates}},
journal = {Science advances},
year = {2026},
month = aug,
volume = {12},
number = {35},
pages = {eaed4172},
publisher = {American Association for the Advancement of Science},
issn = {2375-2548},
doi = {10.1126/sciadv.aed4172},
url = {https://doi.org/10.1126/sciadv.aed4172},
pmid = {42664357},
pmcid = {PMC13524049}
}

RIS

TY - JOUR
AU - Gant, Jonathan M.
AU - Młynarski, Wiktor F.
TI - Locomotion optimizes sensory representations through a computational principle shared by rodents and primates
T2 - Science advances
J2 - Sci Adv
PY - 2026
DA - 2026/08/28
VL - 12
IS - 35
SP - eaed4172
SN - 2375-2548
PB - American Association for the Advancement of Science
DO - 10.1126/sciadv.aed4172
UR - https://doi.org/10.1126/sciadv.aed4172
LA - en
ER -

CSL-JSON

{
"id": "10.1126/sciadv.aed4172",
"type": "article-journal",
"title": "Locomotion optimizes sensory representations through a computational principle shared by rodents and primates",
"container-title": "Science advances",
"author": [
{
"family": "Gant",
"given": "Jonathan M."
},
{
"family": "Młynarski",
"given": "Wiktor F."
}
],
"container-title-short": "Sci Adv",
"volume": "12",
"issue": "35",
"page": "eaed4172",
"DOI": "10.1126/sciadv.aed4172",
"PMID": "42664357",
"PMCID": "PMC13524049",
"ISSN": "2375-2548",
"publisher": "American Association for the Advancement of Science",
"URL": "https://doi.org/10.1126/sciadv.aed4172",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
28
]
]
}
}

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.1371/journal.pcbi.1014123 [code]
Subunit-specific behavioral modulation of sensory tuning in the visual cortex.
Journal: PLoS computational biology
In common: pandas, Matplotlib, NumPy, 10 references, author Wiktor F. Młynarski
[2] doi:10.1038/s41467-026-75347-4 [code]
Sleep reveals dynamics integrating and segregating movement and stimulus representations in V1.
Journal: Nature communications
In common: Numba, PyTorch, scikit-learn, 4 other tools, 7 references
[3] doi:10.1038/s41467-026-71667-7
Behavioural states control binocular vision through input-specific mechanisms.
Journal: Nature communications
In common: 9 references
[4] doi:10.1016/j.isci.2026.116055 [code]
Mapping the transcriptional diversity of calcium signaling in the mouse and human brain.
Journal: iScience
In common: rpy2, Numba, h5py, 7 other tools
[5] doi:10.1038/s41586-026-10348-3 [code]
An enteric neuron ionotropic receptor regulates salt stress resistance.
Journal: Nature
In common: Numba, OpenCV, h5py, 7 other tools, cellular / molecular
[6] doi:10.1016/j.xcrm.2026.102766 [code]
A longitudinal single-cell and spatial multiomic atlas of pediatric high-grade glioma.
Journal: Cell reports. Medicine
In common: rpy2, Numba, OpenCV, 6 other tools, cellular / molecular
[7] doi:10.7554/elife.109717 [code]
Retrosplenial cortex enables context-dependent goal-directed sensorimotor transformation.
Journal: eLife
In common: Numba, OpenCV, h5py, 6 other tools, 1 reference
[8] doi:10.1038/s41593-026-02267-3 [code]
Spatial proteomic analysis in human Alzheimer's disease brains enables identification of microenvironment-dependent microglial cell states.
Journal: Nature neuroscience
In common: rpy2, Numba, h5py, 6 other tools, cellular / molecular
[9] doi:10.1016/j.isci.2026.116825 [code]
Social hierarchy shapes behavioral and transcriptional responses to chronic stress and ketamine in male mice.
Journal: iScience
In common: Numba, OpenCV, h5py, 7 other tools
[10] doi:10.1038/s41467-026-72057-9 [code]
Sex-specific behavioral feedback modulates sensorimotor processing and drives flexible social behavior.
Journal: Nature communications
In common: Numba, OpenCV, h5py, 7 other tools

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.