OSCR

Fast efficient coding and sensory adaptation in gain-adaptive recurrent networks.

Code ↔ Paper

4 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 4 matches
  1. [1] § Methods › Network optimization and simulations ↔ model.ipynb, lines 424–483 · score 0.75 · ReduceLROnPlateau, learning rate schedule, Adam, threshold, gradient, patience
  2. [2] § Results › Gain-adaptive recurrent neural network › Gain adaptation yields prior attraction ↔ model.ipynb, lines 930–1073 · score 0.57 · numerical optimization, Normalized effective tuning, Gaussian priors, Optimal gains, Medium, analytical
  3. [3] § Results › Gain-adaptive recurrent neural network › Gain adaptation yields prior attraction ↔ model.ipynb, lines 930–1073 · score 0.52 · analytical approximation, effective tuning curves, optimal gains, numerically, Medium, zero
  4. [4] § Results › Gain-adaptive recurrent neural network › Effective tuning curves ↔ model.ipynb, lines 791–928 · score 0.52 · uniform gains, feedforward tuning curves, effective tuning curves, connectivity, ij, weighted

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Jupyter notebook · 1,850 lines · 79 KB · no license · 4 matches

  1. # %%
  2. import numpy as np
  3. import torch
  4. import itertools
  5. # from utils import gaussian, probs_mean_and_var, sample_from_poisson
  6. from matplotlib.pyplot import *
  7. from matplotlib.colors import to_hex
  8. from matplotlib.lines import Line2D
  9. from scipy import integrate, stats, optimize, special
  10. # %%
  11. from scipy.interpolate import interp1d
  12. # %%
  13. import sys
  14. print("Python version:", sys.version)
  15. print("Version info:", sys.version_info)
  16. # %%
  17. !pip list
  18. # %%
  19. !pip freeze
  20. # %%
  21. from cachetools import cached
  22. # %%
  23. torch.set_printoptions(linewidth=150, )
  24. torch.autograd.set_detect_anomaly(True);
  25. # %%
  26. # Max Range
  27. LOWER_BOUND = -200
  28. UPPER_BOUND = 200
  29. OVERALL_WIDTH = UPPER_BOUND - LOWER_BOUND
  30. MID_VAL = (UPPER_BOUND + LOWER_BOUND) / 2
  31. # Stimuli
  32. NUM_STIMULI = UPPER_BOUND - LOWER_BOUND + 1
  33. STIMULI = torch.linspace(LOWER_BOUND, UPPER_BOUND, NUM_STIMULI)
  34. #
  35. SQRT_2_PI = torch.sqrt(torch.tensor(2 * torch.pi))
  36. #
  37. COLORS_WIDTH = {10:'C2', 20:'C0', 30:'C1', 8:cm.tab20c(9), 6:cm.tab20c(9),
  38. 5:cm.tab20c(12), 4:cm.tab20(12), 2:cm.tab20(12), 1:cm.tab20(7), 0:cm.tab20(13),
  39. 10000:'C0' } # 6:cm.tab20c(12),
  40. # %% [markdown]
  41. # # Useful functions
  42. # %%
  43. # Utility functions for optimization
  44. def get_best_res(*args):
  45. best_res = optimize.OptimizeResult( fun=np.inf )
  46. for res in args:
  47. if res is None:
  48. continue
  49. if res.fun < best_res.fun:
  50. best_res = res
  51. return best_res
  52. def min_objf(f, x0s, v=True, return_all_res=False, **kwargs):
  53. all_res = []
  54. for x0 in x0s:
  55. res = optimize.minimize(f, x0=x0, **kwargs)
  56. if v: print(x0, '->', res.x, res.fun)
  57. all_res.append(res)
  58. best_res = get_best_res(all_res)
  59. if v: print('=>', best_res.x, best_res.fun)
  60. if return_all_res:
  61. return best_res, all_res
  62. return best_res
  63. def min_objf_grid(objf, grid_vals, v=True, finish=True, compare_with_res=None, objf_args=(), x0s=None, **kwargs):
  64. if v:
  65. ns = [ len(vals) for vals in grid_vals ]
  66. print(' x '.join([str(n) for n in ns]), ' = ', np.prod( ns ) )
  67. ### Method 1: grid search
  68. rez = []
  69. for x0 in itertools.product(*grid_vals):
  70. fun = objf(x0, *objf_args)
  71. rez.append( optimize.OptimizeResult( x=x0, fun=fun, success=False ) )
  72. res = get_best_res(*rez)
  73. if v: print('Grid', res.x, res.fun, flush=True)
  74. rez = [res,]
  75. # Optimization method(s)
  76. methods = [None,]
  77. if 'method' in kwargs:
  78. method = kwargs.pop('method')
  79. if method == 'all':
  80. methods = ['Nelder-Mead', 'Powell', 'L-BFGS-B',] # 'SLSQP']
  81. elif isinstance(method, list):
  82. methods = method
  83. else:
  84. methods = [method,]
  85. # Finish after grid
  86. if finish:
  87. for method in methods:
  88. res_m = optimize.minimize(objf, x0=res.x, method=method, args=objf_args, **kwargs )
  89. if v: print(('-' if method is None else method).ljust(12), res_m.x, res_m.fun, flush=True)
  90. rez.append(res_m)
  91. # List of start points
  92. if x0s is not None:
  93. for method in methods:
  94. res_x0s = min_objf(objf, x0s, v=v, args=objf_args, method=method, **kwargs)
  95. if v: print(('-' if method is None else method).ljust(12), res_x0s.x, res_x0s.fun, flush=True)
  96. rez.append(res_x0s)
  97. # Compare
  98. if compare_with_res is not None:
  99. rez.append( compare_with_res )
  100. res = get_best_res(*rez)
  101. return res
  102. # %%
  103. def values_to_indices(values, size=len(STIMULI), min_val=LOWER_BOUND, max_val=UPPER_BOUND ):
  104. """Convert values in [min_val, max_val] to indices in a vector of length `size`."""
  105. values = np.asarray(values)
  106. step = (max_val - min_val) / (size - 1)
  107. indices = np.round((values - min_val) / step).astype(int)
  108. return indices
  109. # %%
  110. def make_gamma_G_φs(gains, hh, stims):
  111. '''Given a vector of gains, a matrix hh of the distance function h(s_i-s_j), and a vector of feedforward stimuli,
  112. return the weights gamma_ij, the denominator G, and the effective locations φ_s.'''
  113. gamma_num = torch.diag(gains) + hh * gains.unsqueeze(0) # shape = (num_neurons, num_neurons)
  114. gamma_den = gamma_num.sum(axis=1) # shape = (num_neurons,) ### This is G(s)
  115. gammas = gamma_num / gamma_den.unsqueeze(1) # shape = (num_neurons, num_neurons)
  116. φs = ( gammas * stims ).sum(axis=1) # (num_neurons,)
  117. return gammas, gamma_den, φs
  118. # %%
  119. def effective_tc_center_and_squaredwidth_given_rates(rates):
  120. '''Returns the effective locations φ_s and squared widths σ_r^2.'''
  121. denom = rates.sum(axis=1)
  122. true_φs = ( rates * STIMULI ).sum(axis=1) / denom
  123. true_sigmarsSq = ( rates * (STIMULI-true_φs.unsqueeze(1))**2 ).sum(axis=1) / denom
  124. return true_φs, true_sigmarsSq
  125. # %%
  126. def make_recurrent_weights(stims, rec_width, rec_magnitude, inhib):
  127. if rec_magnitude == 0:
  128. return torch.zeros((len(stims), len(stims)))
  129. distances = torch.abs(stims.unsqueeze(0) - stims.unsqueeze(1))
  130. gaussian_kernel = torch.exp(-distances ** 2 / (2 * rec_width ** 2))
  131. recurrent_weights = gaussian_kernel #+ uniform_density
  132. eigvals = torch.linalg.eigvals(recurrent_weights)
  133. max_eigenvalue = eigvals.abs().max()
  134. scaling_factor = rec_magnitude / max_eigenvalue
  135. recurrent_weights = recurrent_weights * scaling_factor
  136. if inhib > 0: # uniform inhibition
  137. recurrent_weights -= (inhib/len(stims)) * torch.ones_like(recurrent_weights)
  138. return recurrent_weights
  139. # %%
  140. def second_derivative_smoothness(g, dx, normalized=False):
  141. """Computes second derivative penalty for smoothness regularization"""
  142. gg = (g/g.mean()) if normalized else g
  143. g_i_minus_1 = gg[:-2] # g[i-1]
  144. g_i = gg[1:-1] # g[i]
  145. g_i_plus_1 = gg[2:] # g[i+1]
  146. second_deriv = ( g_i_plus_1 - 2 * g_i + g_i_minus_1 ) / dx**2
  147. return (second_deriv**2).mean() # Mean squared second derivative
  148. # %%
  149. def make_Fa_and_Fam1_fcts(prior_xs, prior_ys, a, min_x=LOWER_BOUND, max_x=UPPER_BOUND):
  150. """
  151. Given a pdf f(x) and an exponent a, F_a(x) is defined as
  152. F_a(x) = ( ∫[min_x,x] f(t)^a dt ) / ( ∫[min_x,max_x] f(t)^a dt ).
  153. If a=1 it is the usual CDF. This returns two functions, corresponding to F_a(x) and its inverse F_a⁻¹(y).
  154. The discretized support prior_xs is assumed to represent the centers of bins covering the interval.
  155. The functions obey:
  156. - F_a(prior_xs[0]) = 0 and F_a(prior_xs[-1]) = 1.
  157. - For any x <= min_x, fct(x) = 0.
  158. - For any x >= max_x, fct(x) = 1.
  159. min_x and max_x serve as the theoretical integration bounds.
  160. """
  161. # Ensure that the sampled values lie within the prescribed support.
  162. assert all(prior_xs >= min_x)
  163. assert all(prior_xs <= max_x)
  164. # Evaluate the powered pdf
  165. prior_ys_pow_a = prior_ys**a
  166. # Compute the cumulative integral using the trapezoidal rule.
  167. # 'initial=0' ensures that the cumulative value at the first point is 0.
  168. num = integrate.cumulative_trapezoid(prior_ys_pow_a, prior_xs, initial=0)
  169. # Normalize so that the cumulative value at prior_xs[-1] becomes 1.
  170. num /= num[-1]
  171. # Create extended arrays that include the full bounds if needed.
  172. xs_full = np.asarray( prior_xs ).copy()
  173. nums_full = num.copy()
  174. # If the left end of the discretization is strictly inside [min_x, max_x], insert the minimum.
  175. if xs_full[0] > min_x:
  176. xs_full = np.insert(xs_full, 0, min_x)
  177. nums_full = np.insert(nums_full, 0, 0)
  178. # Likewise, if the right end does not reach max_x, append the maximum.
  179. if xs_full[-1] < max_x:
  180. xs_full = np.append(xs_full, max_x)
  181. nums_full = np.append(nums_full, 1)
  182. # Define the F_a function: for any input x, linearly interpolate the integration values,
  183. # forcing x <= min_x to return 0 and x >= max_x to return 1.
  184. def fct(x):
  185. return np.interp(x, xs_full, nums_full, left=0, right=1)
  186. # The inverse F_a⁻¹ is defined by inverting the above mapping.
  187. def fct_m1(y):
  188. return np.interp(y, nums_full, xs_full, left=min_x, right=max_x)
  189. return fct, fct_m1
  190. # %%
  191. def gaussian(x, mu, sigma, incl_norm=False):
  192. if incl_norm:
  193. return (1 / (torch.sqrt(2 * torch.pi * sigma**2))) * torch.exp(-0.5 * ((x - mu) / sigma) ** 2)
  194. else:
  195. return torch.exp(-0.5 * ((x - mu) / sigma) ** 2)
  196. # %%
  197. def probs_mean_and_var(pxs, xs):
  198. mu = torch.sum(pxs * xs)
  199. v = torch.sum(pxs * (xs - mu) ** 2)
  200. return mu, v
  201. # %%
  202. def solve_lyapunov(rates: torch.Tensor, matrix_W: torch.Tensor) -> torch.Tensor:
  203. """
  204. Solve (W - I) Σ + Σ (W - I) + Γ = 0 with Γ = diag(rates).
  205. Args
  206. ----
  207. rates : (N,) tensor
  208. matrix_W : (N, N) tensor (assumed symmetric)
  209. Returns
  210. -------
  211. Sigma : (N, N) tensor
  212. """
  213. # Make sure everything is on same device/dtype
  214. matrix_W = matrix_W
  215. rates = rates.to(matrix_W)
  216. n = matrix_W.shape[0]
  217. I = torch.eye(n, dtype=matrix_W.dtype, device=matrix_W.device)
  218. A = matrix_W - I
  219. Gamma = torch.diag(rates)
  220. # Eigen-decomposition of A (symmetric)
  221. evals, evecs = torch.linalg.eigh(A) # A = evecs @ diag(evals) @ evecs.T
  222. # Transform Gamma into eigenbasis: B = U^T Γ U
  223. B = evecs.T @ Gamma @ evecs
  224. # Solve elementwise: (λ_i + λ_j) Y_ij + B_ij = 0 => Y_ij = -B_ij / (λ_i + λ_j)
  225. lam_i = evals.unsqueeze(0) # (1, N)
  226. lam_j = evals.unsqueeze(1) # (N, 1)
  227. denom = lam_i + lam_j
  228. # In a well-posed Lyapunov problem denom != 0; add tiny epsilon for numerical safety.
  229. eps = 1e-12
  230. denom = denom + eps * (denom == 0)
  231. Y = -B / denom
  232. # Transform back: Σ = U Y U^T
  233. Sigma = evecs @ Y @ evecs.T
  234. # Enforce symmetry numerically
  235. Sigma = 0.5 * (Sigma + Sigma.T)
  236. return Sigma
  237. # %%
  238. def lyapunov_approx_solution(rates: torch.Tensor, matrix_M: torch.Tensor) -> torch.Tensor:
  239. rates = rates.to(matrix_M)
  240. Gamma = torch.diag(rates)
  241. Sigma = 0.25 * (matrix_M @ Gamma + Gamma @ matrix_M)
  242. Sigma = 0.5 * (Sigma + Sigma.T)
  243. return Sigma
  244. # %%
  245. def lyapunov_diagonal_solution(rates: torch.Tensor, h0) -> torch.Tensor:
  246. return .5 * (1 + h0) * rates
  247. # %%
  248. def sample_from_poisson(rate_parameters, num_samples=1000):
  249. samples = np.random.poisson(rate_parameters, (num_samples, *rate_parameters.shape))
  250. return torch.tensor(samples)
  251. def sample_from_gaussian_rates_then_poisson(base_rates, matrix_M, num_samples=100, xbounds=None):
  252. '''For the stimuli in xbounds, sample spike counts from a two-step process:
  253. 1) sample rates from a Gaussian with mean=base_rates and covariance from Lyapunov approx.
  254. 2) sample spike counts from Poisson with these rates.
  255. The output shape is (num_samples, num_neurons, num_stimuli_in_xbounds), where num_neurons = base_rates.shape[0].'''
  256. rng = np.random.default_rng()
  257. if xbounds is None: # then it's the default, STIMULI
  258. indices = np.arange(STIMULI.shape[0])
  259. else:
  260. xmin, xmax = xbounds
  261. idx_min, idx_max = values_to_indices([xmin, xmax])
  262. indices = np.arange(idx_min, idx_max+1)
  263. samples = torch.zeros((num_samples, base_rates.shape[0], len(indices)) )
  264. for i, idx in enumerate(indices): # for each stimulus
  265. this_base_rates = base_rates[:,idx] # shape = (num_neurons,)
  266. Sigma = lyapunov_approx_solution( rates = this_base_rates, matrix_M=matrix_M )
  267. # Sample from Gaussian with mean=this_base_rates and covariance=Sigma
  268. gaussian_samples = rng.multivariate_normal(mean=this_base_rates, cov=Sigma, size=num_samples)
  269. # Rectify to non-negative rates
  270. gaussian_samples = np.clip(gaussian_samples, a_min=0, a_max=None)
  271. # Sample from Poisson with these rates
  272. samples[:,:,i] = sample_from_poisson(gaussian_samples, 1)
  273. return samples
  274. # %% [markdown]
  275. # # Network class
  276. # %%
  277. class Network:
  278. def __init__(self, num_neurons, ff_sigma, rec_magnitude, rec_width, inhib=0.):
  279. # neurons ff preferred stimuli
  280. self.num_neurons = num_neurons
  281. self.neurons_pref_stims = torch.linspace(LOWER_BOUND, UPPER_BOUND, self.num_neurons)
  282. self.δ = OVERALL_WIDTH / self.num_neurons
  283. # feedforward tuning curves
  284. self.ff_sigma = ff_sigma
  285. self.ff_tuning_curves = gaussian( self.neurons_pref_stims.unsqueeze(1), STIMULI, self.ff_sigma ) # (num_neurons, NUM_STIMULI)
  286. # recurrence
  287. self.rec_magnitude = rec_magnitude
  288. self.rec_width = rec_width
  289. self.inhib = inhib
  290. self.update_W_and_M()
  291. self.c = self.ff_sigma * np.sqrt(2*np.pi) / ( self.δ * (1 - self.rec_magnitude) )
  292. # gains
  293. self.gains = torch.ones(self.num_neurons)
  294. self.update_r_star()
  295. def __repr__(self):
  296. return f"Network(N={self.num_neurons}, ffσ={self.ff_sigma}, rec_mag={self.rec_magnitude}, rec_w={self.rec_width}, inhib={self.inhib})"
  297. def update_W_and_M(self):
  298. '''Also updates νrec, hh, hbar, sigma_h_Sq'''
  299. self.recurrent_weights = make_recurrent_weights(self.neurons_pref_stims, self.rec_width, self.rec_magnitude, self.inhib) # shape = (num_neurons, num_neurons)
  300. if self.rec_magnitude > 0:
  301. self.matrix_M = torch.inverse(torch.eye(self.num_neurons) - self.recurrent_weights) # M = (I - W)^-1 # (num_neurons, num_neurons)
  302. self.νrec = self.rec_width / np.sqrt(2*np.log(1./self.rec_magnitude))
  303. si_minus_sj = self.neurons_pref_stims.unsqueeze(0)-self.neurons_pref_stims.unsqueeze(1)
  304. self.hh = (1-self.rec_magnitude)*self.rec_magnitude*self.δ/(self.rec_width*SQRT_2_PI) * torch.exp(-(si_minus_sj)**2/(2*self.rec_width**2) )
  305. self.hh += (self.rec_magnitude/np.log(1./self.rec_magnitude)) * self.δ/(2*self.νrec) * torch.exp(-torch.abs(si_minus_sj)/self.νrec ) # shape = (num_neurons, num_neurons)
  306. self.hbar = (1-self.rec_magnitude)*self.rec_magnitude + (self.rec_magnitude/np.log(1./self.rec_magnitude))
  307. self.sigma_h_Sq = (1-self.rec_magnitude)*self.rec_magnitude*self.rec_width**2 + (self.rec_magnitude/np.log(1./self.rec_magnitude)) * self.νrec**2
  308. self.mu4 = 3. * self.rec_width**4 * self.rec_magnitude * ( 1 - self.rec_magnitude + 2/ (np.log(1./self.rec_magnitude))**3 )
  309. else:
  310. self.matrix_M = torch.eye(self.num_neurons)
  311. self.νrec = 0
  312. self.hh = torch.zeros((self.num_neurons, self.num_neurons))
  313. self.hbar = 0
  314. self.sigma_h_Sq = 0
  315. self.mu4 = 0
  316. self.sigma_h_Sq_tilde = self.sigma_h_Sq / (1.+self.hbar)
  317. self.mu4_tilde = self.mu4 / (1.+self.hbar)
  318. self.beta = 1 + .5*( 1. + self.hh[0,0] )
  319. def compute_r_star(self, gains=None):
  320. '''"r_star" is just r(s) in the paper, here computed as (I - W)^-1 x (gains * ff_tuning_curves).'''
  321. modulated_ff_tuning_curves = (self.gains if gains is None else gains).unsqueeze(1) * self.ff_tuning_curves # shape = (num_neurons, NUM_STIMULI)
  322. r_lin = torch.matmul(self.matrix_M, modulated_ff_tuning_curves) # shape = (num_neurons, NUM_STIMULI)
  323. if self.inhib == 0.: # no inhibition, linear solution
  324. return r_lin, modulated_ff_tuning_curves
  325. # if self.inhib > 0.: # with inhibition
  326. r_star = torch.clamp(r_lin, min=0.0)
  327. for _ in range(10):
  328. r_star = torch.clamp( modulated_ff_tuning_curves + torch.matmul(self.recurrent_weights, r_star), min=0.0 )
  329. return r_star, modulated_ff_tuning_curves
  330. def update_r_star(self):
  331. self.r_star, self.modulated_ff_tuning_curves = self.compute_r_star(self.gains)
  332. def fisher_info(self, gains=None):
  333. r_star, _ = self.compute_r_star(gains)
  334. rprime = torch.gradient(r_star, dim=1, spacing=[STIMULI,], edge_order=2)[0]
  335. FI = torch.zeros_like(r_star)
  336. mask = r_star > 1e-30
  337. FI[mask] = (rprime[mask]**2) / r_star[mask]
  338. FItot = FI.sum(axis=0) # shape = (NUM_STIMULI,)
  339. return FItot / self.beta
  340. def effective_tc_center_and_squaredwidth(self):
  341. '''Returns the effective locations φ_s and squared widths σ_r^2.'''
  342. return effective_tc_center_and_squaredwidth_given_rates(self.r_star)
  343. def spike_cost(self, rates, prior_stims, spikes_power):
  344. powered_exp_spikes_for_each_stimulus = (rates**spikes_power).sum(axis=0)
  345. return (prior_stims * powered_exp_spikes_for_each_stimulus).sum()
  346. def expSqE(self, rates, prior_stims, prior_var=None):
  347. if prior_var is None:
  348. _, prior_var = probs_mean_and_var(prior_stims, STIMULI)
  349. _, true_sigmarsSq = effective_tc_center_and_squaredwidth_given_rates(rates)
  350. expectedSqErr = (prior_stims/(1./prior_var + (1./self.beta)*(rates.T/true_sigmarsSq).sum(axis=1)) ).sum()
  351. return expectedSqErr
  352. def expSqE_v2_I(self, rates, prior_stims, prior_var=None):
  353. # version with I(s) = (1/β) sum_i r_i(s) (s-φ_i)^2 / σ_r_i^4
  354. if prior_var is None:
  355. _, prior_var = probs_mean_and_var(prior_stims, STIMULI)
  356. true_φs, true_sigmarsSq = effective_tc_center_and_squaredwidth_given_rates(rates) # true_sigmarsSq.shape = (num_neurons,)
  357. fi = (1./self.beta)*(rates.T * (STIMULI.unsqueeze(1)-true_φs)**2/true_sigmarsSq**2).sum(axis=1)
  358. expectedSqErr = (prior_stims/(1./prior_var + fi ) ).sum()
  359. return expectedSqErr
  360. def optimize_gains_expSqErr_wCost(self, prior_stims, prior_neurons, alpha_cost, lr=1e-3, epochs=5000, patience=100, patience_lr=None, start_gains=None, v=True,
  361. regularizer_lambda=0., return_losses=False, eps_rel=0.):
  362. if start_gains is None: start_gains = prior_neurons
  363. N = start_gains.numel()
  364. patience_lr = patience_lr if patience_lr is not None else patience // 2
  365. log_gains = torch.log(start_gains).detach().requires_grad_(True)
  366. optimizer = torch.optim.Adam([log_gains], lr=lr)
  367. scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=patience_lr, factor=0.5, threshold=1e-4 if eps_rel==0. else eps_rel)
  368. _, prior_var = probs_mean_and_var(prior_stims, STIMULI)
  369. this_prior = prior_neurons if self.expSqE_kind == 'v2_g' else prior_stims
  370. rates, _ = self.compute_r_star(start_gains)
  371. best_expectedSqErr = self.expSqE(rates, this_prior, prior_var)
  372. best_cost = self.spike_cost(rates, prior_stims, spikes_power=1.)
  373. best_val_loss_no_reg = best_expectedSqErr + alpha_cost * best_cost
  374. best_val_loss = best_val_loss_no_reg + regularizer_lambda * second_derivative_smoothness(start_gains, self.δ, normalized=True)
  375. best_gains = start_gains.clone().detach()
  376. patience_counter = 0
  377. now = datetime.datetime.now()
  378. if v:
  379. print('Time\t\tIter\t<g>\t\tExpSqErr\tLoss\tpatience')
  380. for i in range(epochs):
  381. optimizer.zero_grad() # Reset gradients
  382. gains = torch.exp(log_gains)
  383. rates, _ = self.compute_r_star(gains)
  384. cost = self.spike_cost(rates, prior_stims, spikes_power=1.)
  385. expectedSqErr = self.expSqE(rates, this_prior, prior_var)
  386. loss_no_reg = expectedSqErr + alpha_cost*cost
  387. loss = loss_no_reg + regularizer_lambda * second_derivative_smoothness(gains, self.δ, normalized=True)
  388. loss.backward() # Backpropagate to compute gradients
  389. optimizer.step() # Update log_gains
  390. val_loss_no_reg = loss_no_reg.item()
  391. val_loss = loss.item()
  392. scheduler.step(val_loss)
  393. improved = (best_val_loss - val_loss) > max(0., eps_rel * abs(best_val_loss))
  394. if improved:
  395. best_val_loss_no_reg = val_loss_no_reg
  396. best_val_loss = val_loss
  397. best_gains = gains.clone().detach()
  398. best_cost = cost.item()
  399. best_expectedSqErr = expectedSqErr.item()
  400. patience_counter = 0
  401. else:
  402. patience_counter += 1
  403. if v and i%500 == 0: #
  404. print(f'{(-(now - (now:=datetime.datetime.now())))} \t{i}\t<{gains.mean().item():.4g}>\t{best_expectedSqErr:.4g}\t{best_val_loss:.8g}\t{patience_counter}')
  405. if (patience_counter >= patience):
  406. print(f'Patience reached ({patience_counter}). Stopping optimization.')
  407. break
  408. gains = best_gains
  409. if v:
  410. print(f'{(-(now - (now:=datetime.datetime.now())))} \t{i}\t<{gains.mean().item():.4g}>\t{best_expectedSqErr:.4g}\t{best_val_loss:.8g}\t{patience_counter}')
  411. self.gains = gains
  412. self.update_r_star()
  413. if return_losses:
  414. return gains, best_val_loss_no_reg, best_val_loss
  415. return gains
  416. def optimize_gains_quadratic_g0Delta_expSqErr_wCost_simple(self, prior_stims, alpha_cost, v=True):
  417. gains = torch.zeros(self.num_neurons)
  418. prior_mean, prior_var = probs_mean_and_var(prior_stims, STIMULI)
  419. prior_sd = torch.sqrt(prior_var)
  420. z = (self.neurons_pref_stims-prior_mean)/prior_sd
  421. zSq = z**2
  422. gains = torch.zeros_like(z)
  423. def objf(x):
  424. g0, Delta = x
  425. ind = z.abs() < Delta
  426. gains[:] = 0.
  427. gains[ind] = g0 * (1 - zSq[ind]/Delta**2)
  428. rates, _ = self.compute_r_star(gains)
  429. cost = self.spike_cost(rates, prior_stims, spikes_power=1.)
  430. expectedSqErr = self.expSqE(rates, prior_stims, prior_var)
  431. loss = expectedSqErr + alpha_cost*cost
  432. #print(x, loss)
  433. return loss
  434. res = min_objf_grid(objf, grid_vals=[np.arange(.005,.321,.005), np.arange(.5,4.1,.125)], bounds=[(1e-6,None), (1e-3,5.)], method='all', v=v)
  435. g0, Delta = res.x
  436. ind = z.abs() < Delta
  437. gains[:] = 0.
  438. gains[ind] = g0 * (1 - zSq[ind]/Delta**2)
  439. self.gains = gains
  440. self.update_r_star()
  441. return g0, Delta, res.fun
  442. def rec_effective_tc_width(self):
  443. return np.sqrt(self.ff_sigma**2 + self.sigma_h_Sq_tilde)
  444. def sample_spike_counts_nocovariance(self, num_samples, xbounds=None):
  445. if xbounds is None: # then it's the default, STIMULI
  446. return sample_from_poisson(self.r_star, num_samples) # r_star.shape = (num_neurons, NUM_STIMULI)
  447. else:
  448. xmin, xmax = xbounds
  449. ind = (STIMULI >= xmin) & (STIMULI <= xmax)
  450. return sample_from_poisson(self.r_star[:,ind], num_samples) # r_star.shape = (num_neurons, num_stimuli), Result's shape = (num_samples, num_neurons, num_stimuli)
  451. def sample_spike_counts(self, num_samples, xbounds=None):
  452. return sample_from_gaussian_rates_then_poisson(self.r_star, self.matrix_M, num_samples=num_samples, xbounds=xbounds)
  453. def simulate_recurrent_poisson_many(self, num_samples, T=1.0, tau=0.01, dt=None, xbounds=None, burn_in=0.0):
  454. """
  455. Recurrent Poisson simulation for many stimuli and many trials.
  456. Returns spike counts summed over time (after burn-in), shape:
  457. (num_samples, num_neurons, num_stimuli_in_xbounds)
  458. """
  459. if xbounds is None:
  460. indices = torch.arange(STIMULI.shape[0])
  461. else:
  462. xmin, xmax = xbounds
  463. idx_min, idx_max = values_to_indices([xmin, xmax])
  464. indices = torch.arange(idx_min, idx_max + 1)
  465. if dt is None:
  466. dt = tau / 10.0
  467. steps = int(T / dt)
  468. burn_steps = int(burn_in / dt)
  469. total_steps = burn_steps + steps
  470. samples = torch.zeros((num_samples, self.num_neurons, len(indices)) )
  471. for j, idx in enumerate(indices):
  472. ff_tc = self.ff_tuning_curves[:, idx]
  473. ff = self.gains * ff_tc # (num_neurons,)
  474. r = ff.unsqueeze(0).expand(num_samples, -1).clone() # start at ff rates
  475. counts = torch.zeros((num_samples, self.num_neurons) )
  476. for t in range(total_steps):
  477. lam = torch.clamp(r, min=0.) * dt
  478. k = torch.poisson(lam) # float tensor
  479. if t >= burn_steps:
  480. counts += k
  481. r = r + (dt / tau) * (-r + ff.unsqueeze(0) + r @ self.recurrent_weights.T) + (1.0 / tau) * (k - lam)
  482. samples[:, :, j] = counts
  483. return samples
  484. def decode(self, spike_counts, prior_mean, prior_var):
  485. '''spike_counts.shape = (num_samples, num_neurons, num_stimuli).
  486. Returns estimates of the stimuli, shape = (num_samples, num_stimuli)'''
  487. assert spike_counts.shape[1] == self.num_neurons
  488. corrected_spike_counts = (1./self.beta)*spike_counts
  489. #assert spike_counts.shape[2] == NUM_STIMULI
  490. φs, sigmarsSq = self.effective_tc_center_and_squaredwidth()
  491. oo_sigmaSq = 1./sigmarsSq
  492. # spike_counts has shape (num_samples, num_neurons, num_stimuli), and oo_sigmaSq has shape (num_neurons,).
  493. denominator = 1./prior_var + torch.einsum('ijk,j->ik', corrected_spike_counts, oo_sigmaSq)
  494. # spike_counts has shape (num_samples, num_neurons, num_stimuli), and φs*oo_sigmaSq has shape (num_neurons,).
  495. numerator = prior_mean/prior_var + torch.einsum('ijk,j->ik', corrected_spike_counts, φs*oo_sigmaSq)
  496. return numerator / denominator
  497. def posterior_precision(self, spike_counts, prior_var):
  498. '''spike_counts.shape = (num_samples, num_neurons, num_stimuli).
  499. Returns the posterior precision (1/variance) for each estimate, shape = (num_samples, num_stimuli)'''
  500. assert spike_counts.shape[1] == self.num_neurons
  501. corrected_spike_counts = (1./self.beta)*spike_counts
  502. #assert spike_counts.shape[2] == NUM_STIMULI
  503. _, sigmarsSq = self.effective_tc_center_and_squaredwidth()
  504. oo_sigmaSq = 1./sigmarsSq
  505. # spike_counts has shape (num_samples, num_neurons, num_stimuli), and oo_sigmaSq has shape (num_neurons,).
  506. precision = 1./prior_var + torch.einsum('ijk,j->ik', corrected_spike_counts, oo_sigmaSq)
  507. return precision
  508. # %%
  509. def make_gaussian_priors(pref_stims, sd, m=MID_VAL):
  510. '''This returns two arrays, one for the STIMULI and one for "pref_stims" (e.g., the neurons' preferred stimuli).'''
  511. # special case: point mass at closest stimulus
  512. if sd == 0:
  513. # for STIMULI
  514. idx0 = torch.argmin((STIMULI - m).abs())
  515. prior_stimuli = torch.zeros_like(STIMULI)
  516. prior_stimuli[idx0] = 1.
  517. # for pref_stims
  518. idx1 = torch.argmin((pref_stims - m).abs())
  519. prior_pref_stims = torch.zeros_like(pref_stims)
  520. prior_pref_stims[idx1] = 1.
  521. return prior_stimuli, prior_pref_stims
  522. prior_stimuli = gaussian( STIMULI, m, sd )
  523. prior_stimuli /= prior_stimuli.sum()
  524. prior_pref_stims = gaussian( pref_stims, m, sd )
  525. prior_pref_stims /= prior_pref_stims.sum()
  526. return prior_stimuli, prior_pref_stims
  527. # %%
  528. def make_mixture_uniform_gaussian_priors(pref_stims, sd, m=MID_VAL, uniform_part=0.1):
  529. '''This returns two arrays, one for the STIMULI and one for "pref_stims" (e.g., the neurons' preferred stimuli).'''
  530. gaussian_stimuli, gaussian_pref_stims = make_gaussian_priors(pref_stims, sd, m)
  531. uniform_stimuli = torch.ones_like(STIMULI) / NUM_STIMULI
  532. uniform_pref_stims = torch.ones_like(pref_stims) / len(pref_stims)
  533. prior_stimuli = (1-uniform_part)*gaussian_stimuli + uniform_part*uniform_stimuli
  534. prior_pref_stims = (1-uniform_part)*gaussian_pref_stims + uniform_part*uniform_pref_stims
  535. return prior_stimuli, prior_pref_stims
  536. # %% [markdown]
  537. # # Optimizing gains to minimize expected squared error
  538. # %%
  539. #CACHE_find_optimal_gains_for_parameters = {}
  540. # np.savez_compressed('CACHE_find_optimal_gains_for_parameters.npz', cache=CACHE_find_optimal_gains_for_parameters)
  541. CACHE_find_optimal_gains_for_parameters = np.load('CACHE_find_optimal_gains_for_parameters.npz', allow_pickle=True)['cache'].item()
  542. len(CACHE_find_optimal_gains_for_parameters)
  543. # %%
  544. @cached(cache=CACHE_find_optimal_gains_for_parameters)
  545. def find_optimal_gains_for_parameters(num_neurons, ff_sigma, rec_magnitude, rec_width, alpha_cost, prior_w, regularizer_lambda=1000, inhib=0.):
  546. print( f'prior width = {prior_w}' )
  547. mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width, inhib=inhib)
  548. a_prior, a_prior_neurons = make_gaussian_priors(mo.neurons_pref_stims, prior_w)
  549. if inhib != 0.:
  550. print('Solution with inhib = 0 ...')
  551. gains_0, _, _, _ = find_optimal_gains_for_parameters(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width,
  552. alpha_cost=alpha_cost, prior_w=prior_w, regularizer_lambda=regularizer_lambda, inhib=0.)
  553. print('Solution with inhib = 0 ... OK')
  554. else:
  555. mo.optimize_gains_quadratic_g0Delta_expSqErr_wCost_simple(a_prior, alpha_cost=alpha_cost, v=True)
  556. gains_0 = mo.gains.clone()
  557. # smoothen the initial gain solution
  558. kernel_width = int( np.ceil( mo.rec_effective_tc_width()*mo.δ ))
  559. x = np.linspace(-3*kernel_width, 3*kernel_width, 6*kernel_width + 1)
  560. kernel = np.exp(-x**2/(2*kernel_width**2))
  561. kernel = kernel/kernel.sum()
  562. gains_1 = torch.Tensor(np.convolve(gains_0, kernel, mode='same'))
  563. #
  564. # two additional starting points
  565. gbar = gains_0.mean()
  566. gains_const = torch.full_like(gains_0, gbar)
  567. gains_const_times_prior = gbar * a_prior_neurons * num_neurons
  568. # same optimization from 3 starts
  569. gains_smooth, _, loss_smooth = mo.optimize_gains_expSqErr_wCost(a_prior, a_prior_neurons, alpha_cost=alpha_cost, lr=0.01, epochs=50000, patience=500, patience_lr=50,
  570. start_gains=gains_1, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
  571. print(f'Intermediate loss (smooth start): {loss_smooth:.6g}')
  572. gains_const_opt, _, loss_const = mo.optimize_gains_expSqErr_wCost(a_prior, a_prior_neurons, alpha_cost=alpha_cost, lr=0.01, epochs=50000, patience=500, patience_lr=50,
  573. start_gains=gains_const, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
  574. print(f'Intermediate loss (const start): {loss_const:.6g}')
  575. gains_const_prior_opt, _, loss_const_prior = mo.optimize_gains_expSqErr_wCost(a_prior, a_prior_neurons, alpha_cost=alpha_cost, lr=0.01, epochs=50000, patience=500, patience_lr=50,
  576. start_gains=gains_const_times_prior, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
  577. print(f'Intermediate loss (const x prior start): {loss_const_prior:.6g}')
  578. # pick best
  579. loss = loss_smooth
  580. gains_best = gains_smooth
  581. best_is = 'smooth'
  582. if loss_const < loss:
  583. loss = loss_const
  584. gains_best = gains_const_opt
  585. best_is = 'const'
  586. if loss_const_prior < loss:
  587. loss = loss_const_prior
  588. gains_best = gains_const_prior_opt
  589. best_is = 'const x prior'
  590. print(f'Best final solution is from "{best_is}" start.')
  591. print(f'\t----------> loss = {loss:.6g}')
  592. return gains_best.clone(), loss, [gains_0, gains_1, gains_smooth, gains_const_opt, gains_const_prior_opt], [loss_smooth, loss_const, loss_const_prior]
  593. # %% [markdown]
  594. # ### Peaked prior
  595. # %%
  596. #CACHE_find_optimal_gains_for_parameters_uniform_mixture = {}
  597. CACHE_find_optimal_gains_for_parameters_uniform_mixture = np.load('CACHE_find_optimal_gains_for_parameters_uniform_mixture.npz', allow_pickle=True)['cache'].item()
  598. # %%
  599. @cached(cache=CACHE_find_optimal_gains_for_parameters_uniform_mixture)
  600. def find_optimal_gains_for_parameters_uniform_mixture(num_neurons, ff_sigma, rec_magnitude, rec_width, alpha_cost, prior_w, regularizer_lambda, uniform_part=0.8, inhib=0.):
  601. print( f'prior width = {prior_w}' )
  602. mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width, inhib=inhib)
  603. a_prior, a_prior_neurons = make_mixture_uniform_gaussian_priors(mo.neurons_pref_stims, prior_w, uniform_part=uniform_part)
  604. if inhib != 0.:
  605. print('Solution with inhib = 0 ...')
  606. gains_0, _, _, _ = find_optimal_gains_for_parameters_uniform_mixture(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width,
  607. alpha_cost=alpha_cost, prior_w=prior_w, regularizer_lambda=regularizer_lambda, uniform_part=uniform_part, inhib=0.)
  608. print('Solution with inhib = 0 ... OK')
  609. else:
  610. mo.optimize_gains_quadratic_g0Delta_expSqErr_wCost_simple(a_prior, alpha_cost=alpha_cost, v=True)
  611. gains_0 = mo.gains.clone()
  612. # smoothen the initial gain solution
  613. kernel_width = int(np.ceil(mo.rec_effective_tc_width() * mo.δ))
  614. x = np.linspace(-3 * kernel_width, 3 * kernel_width, 6 * kernel_width + 1)
  615. kernel = np.exp(-x**2 / (2 * kernel_width**2))
  616. kernel = kernel / kernel.sum()
  617. gains_1 = torch.Tensor(np.convolve(gains_0, kernel, mode='same'))
  618. # two additional starting points
  619. gbar = gains_0.mean()
  620. gains_const = torch.full_like(gains_0, gbar)
  621. gains_const_times_prior = gbar * a_prior_neurons * num_neurons
  622. # same optimization from 3 starts
  623. gains_smooth, _, loss_smooth = mo.optimize_gains_expSqErr_wCost(a_prior, a_prior_neurons, alpha_cost=alpha_cost, lr=0.01, epochs=50000, patience=500, patience_lr=50,
  624. start_gains=gains_1, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
  625. print(f'Intermediate loss (smooth start): {loss_smooth:.6g}')
  626. gains_const_opt, _, loss_const = mo.optimize_gains_expSqErr_wCost(a_prior, a_prior_neurons, alpha_cost=alpha_cost, lr=0.01, epochs=50000, patience=500, patience_lr=50,
  627. start_gains=gains_const, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
  628. print(f'Intermediate loss (const start): {loss_const:.6g}')
  629. gains_const_prior_opt, _, loss_const_prior = mo.optimize_gains_expSqErr_wCost(a_prior, a_prior_neurons, alpha_cost=alpha_cost, lr=0.01, epochs=50000, patience=500, patience_lr=50,
  630. start_gains=gains_const_times_prior, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
  631. print(f'Intermediate loss (const x prior start): {loss_const_prior:.6g}')
  632. # pick best
  633. loss = loss_smooth
  634. gains_best = gains_smooth
  635. best_is = 'smooth'
  636. if loss_const < loss:
  637. loss = loss_const
  638. gains_best = gains_const_opt
  639. best_is = 'const'
  640. if loss_const_prior < loss:
  641. loss = loss_const_prior
  642. gains_best = gains_const_prior_opt
  643. best_is = 'const x prior'
  644. print(f'Best final solution is from "{best_is}" start.')
  645. print(f'\t----------> loss = {loss:.6g}')
  646. return gains_best.clone(), loss, [gains_0, gains_1, gains_smooth, gains_const_opt, gains_const_prior_opt], [loss_smooth, loss_const, loss_const_prior],
  647. # %%
  648. # Compute Full Width at Half Maximum (FWHM) for the tuning curve
  649. def tc_fwhm(tc):
  650. tc_half_max = tc.max() / 2.0
  651. above_half = np.where(tc >= tc_half_max)[0]
  652. if above_half.size > 0:
  653. fwhm = STIMULI[above_half[-1]] - STIMULI[above_half[0]]
  654. else:
  655. fwhm = np.nan
  656. return fwhm
  657. def tc_fwhm_interp(tc):
  658. # Interpolate the tuning curve
  659. interp_func = interp1d(STIMULI, tc, kind='cubic', bounds_error=False, fill_value=0.0)
  660. tc_max = tc.detach().numpy().max()
  661. half_max = tc_max / 2.0
  662. # Find where the interpolated curve crosses half max
  663. fine_x = np.linspace(STIMULI[0], STIMULI[-1], 5000)
  664. fine_y = interp_func(fine_x)
  665. above_half = np.where(fine_y >= half_max)[0]
  666. if above_half.size > 0:
  667. fwhm = fine_x[above_half[-1]] - fine_x[above_half[0]]
  668. else:
  669. fwhm = np.nan
  670. return fwhm
  671. # %% [markdown]
  672. # # Figures
  673. # %% [markdown]
  674. # ## FIG. 3 Illustration of mechanism
  675. # %%
  676. default_colors = rcParams['axes.prop_cycle'].by_key()['color'] # Default color cycle
  677. # %%
  678. ff_sigma = 20.
  679. rec_magnitude = 0.85
  680. rec_width = 30.
  681. mo = Network(num_neurons=100, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
  682. ired = 55
  683. iorg = None
  684. gs = torch.ones(mo.num_neurons)
  685. gs2 = gs.clone()
  686. gs2[:mo.num_neurons//2] = 2.
  687. gs2[mo.num_neurons//2:] = 0
  688. # Number of colors to extract from cmap
  689. cmap = cm.viridis_r #Spectral # viridis_r
  690. num_colors = 19
  691. cmap_colors = [to_hex(cmap(i / (num_colors - 1))) for i in range(num_colors)]
  692. colors = rcParams['axes.prop_cycle'].by_key()['color'] # Default color cycle
  693. # set axes.prop_cycle to use cmap_colors
  694. rcParams['axes.prop_cycle'] = cycler(color=cmap_colors)
  695. fig, axs = subplots(nrows=3, ncols=4, figsize=(16, 4), gridspec_kw={'height_ratios': [5, 7, 13]}, sharex=True)
  696. for jj in [0, 1, 2]: # this merges row 2 and 3 for these columns
  697. gds = axs[1, jj].get_gridspec()
  698. for ax in axs[1:, jj]:
  699. ax.remove()
  700. axs[1, jj] = fig.add_subplot(gds[1:, jj])
  701. for jj in [0, 1, 2]:
  702. this_gs = gs if jj<=1 else gs2 #[gs, gs2][jj-1]
  703. mo.gains = this_gs
  704. mo.update_r_star()
  705. if jj==0:
  706. axs[1,jj].plot(STIMULI, mo.ff_tuning_curves[5::5].T, lw=.5, alpha=1);
  707. axs[1,jj].plot(STIMULI, mo.ff_tuning_curves[ired], lw=2, alpha=1, c=colors[3], zorder=1000);
  708. else:
  709. axs[1,jj].plot(STIMULI, mo.r_star[5::5].T, lw=.5, alpha=1, clip_on=False);
  710. if iorg: axs[1,jj].plot(STIMULI, mo.r_star[iorg], lw=2, alpha=1, c=colors[1]);
  711. axs[1,jj].plot(STIMULI, mo.r_star[ired], lw=2, alpha=1, c=colors[3]);
  712. axs[1,jj].plot(STIMULI, mo.ff_tuning_curves[ired], lw=1, alpha=1, c=colors[3], zorder=1000, ls=':');
  713. if jj==1:
  714. tc_1 = mo.r_star[ired]
  715. elif jj==2:
  716. axs[1,jj].plot(STIMULI, tc_1, lw=1, alpha=1, c=colors[3], zorder=1000, ls='--');
  717. axs[1,jj].annotate("", xytext=(STIMULI[ tc_1.argmax() ], 2.), xy=(STIMULI[ mo.r_star[ired].argmax() ], 2.), arrowprops=dict(arrowstyle="-|>", color=colors[3], shrinkA=0, shrinkB=0))
  718. axs[0,jj].scatter(mo.neurons_pref_stims, this_gs, c=mo.neurons_pref_stims, cmap=cm.viridis_r, s=6, clip_on=False, zorder=1000)
  719. if iorg: axs[0,jj].scatter(mo.neurons_pref_stims[iorg], this_gs[iorg], color=colors[1], s=12, clip_on=False, zorder=1000)
  720. axs[0,jj].scatter(mo.neurons_pref_stims[ired], this_gs[ired], color=colors[3], s=18, clip_on=False, zorder=1000)
  721. if jj==1:
  722. axs[1,jj].set_prop_cycle(None)
  723. if iorg: axs[1,jj].plot(STIMULI, mo.ff_tuning_curves[iorg], lw=2, alpha=1, c=colors[1], ls='--');
  724. else:
  725. mo.gains = gs
  726. mo.update_r_star()
  727. if iorg: axs[1,jj].plot(STIMULI, mo.r_star[iorg], lw=.5, alpha=1, c=colors[1], label=label, ls='--');
  728. axs[0,jj].set_ylim(0,2.);
  729. axs[1,jj].set_ylim(0,6);
  730. axs[1,jj].set_xticks([])
  731. axs[0,jj].set_xticks([])
  732. axs[0,jj].set_xlim(LOWER_BOUND, UPPER_BOUND)
  733. axs[1,jj].set_xlim(LOWER_BOUND, UPPER_BOUND)
  734. axs[0,jj].set_yticks([0,1,2])
  735. axs[1,jj].set_yticks([0,2,4,6])
  736. axs[0,jj].spines['top'].set_visible(False); axs[0,jj].spines['right'].set_visible(False)
  737. axs[1,jj].spines['top'].set_visible(False); axs[1,jj].spines['right'].set_visible(False)
  738. axs[1,jj].set_xlabel('Stimulus s', labelpad=8)
  739. if jj>0:
  740. axs[0,jj].set_yticklabels([])
  741. axs[1,jj].set_yticklabels([])
  742. ####
  743. axs[0,3].scatter(mo.neurons_pref_stims, this_gs, c=mo.neurons_pref_stims, cmap='viridis_r', s=6, clip_on=False, zorder=1000)
  744. if iorg: axs[0,3].scatter(mo.neurons_pref_stims[iorg], this_gs[iorg], color=colors[1], s=12, clip_on=False, zorder=1000)
  745. axs[0,3].scatter(mo.neurons_pref_stims[ired], this_gs[ired], color=colors[3], s=18, clip_on=False, zorder=1000)
  746. for ii in range(3):
  747. axs[ii,3].spines['top'].set_visible(False); axs[ii,3].spines['right'].set_visible(False)
  748. axs[0,3].set_yticks([0,1,2])
  749. axs[0,3].set_ylim(0,2.);
  750. axs[0,3].set_yticklabels([])
  751. axs[0,3].annotate("$g_j$", (.8, 0.6), ha='center', fontsize=9, xycoords='axes fraction')
  752. axs[1,3].scatter(mo.neurons_pref_stims, mo.hh[ired], c=mo.neurons_pref_stims, cmap='viridis_r', s=6, clip_on=False, zorder=1000)
  753. axs[1,3].scatter(mo.neurons_pref_stims[ired], mo.hh[ired,ired], color=colors[3], s=8, clip_on=False, zorder=1000)
  754. gammas, gamma_den, φs = make_gamma_G_φs(gs2, mo.hh, mo.neurons_pref_stims)
  755. axs[2,3].scatter(mo.neurons_pref_stims, gammas[ired], c=mo.neurons_pref_stims, cmap='viridis_r', s=6, clip_on=False, zorder=1000)
  756. axs[2,3].scatter(mo.neurons_pref_stims[ired], gammas[ired,ired], color=colors[3], s=8, clip_on=False, zorder=1000)
  757. axs[1,3].scatter(mo.neurons_pref_stims[ired], 0., color=colors[3], marker='|', clip_on=False, zorder=1000)
  758. axs[1,3].text(x=mo.neurons_pref_stims[ired], y=0.025, s="$s_i$", ha='center', color=colors[3], fontsize=12)
  759. axs[1,3].set_title(x=.8, y=0.7, label="$h(s_j-s_i)$", ha='center', fontsize=9)
  760. axs[1,3].set_ylabel("$h(s_j-s_i)$")
  761. axs[2,3].set_title(x=.775, y=0.7, label="$\\gamma_{ij}\\propto g_j h(s_j-s_i) + g_i \\delta_{ij}$", ha='center', va='center', fontsize=9)
  762. axs[2,3].set_ylabel("Weight $\\gamma_{ij}$")
  763. axs[2,3].scatter(mo.neurons_pref_stims[ired], 0., color=colors[3], marker='|', clip_on=False, zorder=1000)
  764. axs[2,3].text(x=mo.neurons_pref_stims[ired], y=0.0075, s="$s_i$", ha='center', color=colors[3], fontsize=12)
  765. axs[2,3].scatter(φs[ired], 0., color=colors[3], marker='o', fc='w', clip_on=False, zorder=1000)
  766. axs[2,3].text(x=φs[ired], y=0.0075, s="$\\varphi(s_i)$", ha='center', color=colors[3], fontsize=12)
  767. axs[2,3].annotate("", xytext=(mo.neurons_pref_stims[ired], 0.01), xy=(φs[ired], 0.01), arrowprops=dict(arrowstyle="-|>", color=colors[3], shrinkA=7, shrinkB=14, lw=1))
  768. for ii in range(3):
  769. axs[ii,3].axvline(0, lw=1, color='lightgrey', ls=':')
  770. axs[2,3].set_xlabel('Feedforward preferred stimulus $s_j$', labelpad=8)
  771. axs[1,3].set_ylim(0,None)
  772. axs[2,3].set_ylim(0,None)
  773. axs[1,3].set_yticks([])
  774. axs[2,3].set_yticks([])
  775. axs[0,0].set_ylabel('Gain g')
  776. axs[1,0].set_ylabel('Firing rate r(s)')
  777. fsz = 10.
  778. axs[0,0].set_title('No recurrence\n+ Uniform gains', fontsize=fsz)
  779. axs[0,1].set_title('Recurrent connections\n+ Uniform gains', fontsize=fsz)
  780. axs[0,2].set_title('Recurrent connections\n+ Non-uniform gains', fontsize=fsz)
  781. axs[1,0].text(0,4.5,'Gains $\\times$ Feedforward tuning curves\n$r(s) = g \\circ f(s)$', ha='center')
  782. axs[1,1].text(0,4.5,'Effective tuning curves\n$r(s) = (I-W)^{-1} ( g \\circ f(s) )$', ha='center')
  783. fig.subplots_adjust(wspace=.1, hspace=.15)
  784. #show()
  785. #fig.savefig('rnn_tc.pdf', bbox_inches='tight')
  786. # set axes.prop_cycle back to default
  787. rcParams['axes.prop_cycle'] = cycler(color=default_colors)
  788. # %% [markdown]
  789. # ## FIG. 4 Optimized network for priors 30, 20, 10
  790. # %%
  791. num_neurons = 801
  792. ff_sigma = 5.
  793. rec_magnitude = 0.95
  794. rec_width = 6.
  795. alpha_cost = .5
  796. reg_lambda = 1000
  797. inhib = 0.
  798. fig = figure(figsize=(11, 8))
  799. # 5 rows: row0(ax0), spacer, row1, row2, row3
  800. gs = GridSpec(5, 2, figure=fig, height_ratios=[3, 0.3, 1, 1, 1])
  801. axPriors = fig.add_subplot(gs[0,0])
  802. axGains = fig.add_subplot(gs[2:,0])
  803. axPhi = fig.add_subplot(gs[0,1])
  804. axs_tc = [fig.add_subplot(gs[2+i, 1]) for i in range(3)]
  805. g_solutions = {}
  806. g_solutions_approx = {}
  807. for ip,prior_w in enumerate( [30, 20, 10] ):
  808. g_solution, loss, gains_list, losses_list = find_optimal_gains_for_parameters(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width,
  809. alpha_cost=alpha_cost, prior_w=prior_w, regularizer_lambda=reg_lambda, inhib=inhib)
  810. g_solutions_approx[prior_w] = gains_list[0]
  811. g_solutions[prior_w] = g_solution
  812. mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width, inhib=inhib)
  813. φs_perprior = {}
  814. argmaxs_perprior = {}
  815. tuning_curves = {}
  816. for ip,prior_w in enumerate([10, 20, 30]):
  817. print( f'prior width = {prior_w}', end=' ')
  818. c = COLORS_WIDTH.get(prior_w)
  819. a_prior, a_prior_neurons = make_gaussian_priors(mo.neurons_pref_stims, prior_w)
  820. prior_mean, prior_var = probs_mean_and_var(a_prior_neurons, mo.neurons_pref_stims)
  821. axPriors.plot( mo.neurons_pref_stims, a_prior_neurons, c=c, ls='-', label={10:'Narrow', 20:'Medium', 30:'Wide'}[prior_w] + f' ($\\sigma_p={prior_w}$)')
  822. ####
  823. g_solution = torch.clamp( g_solutions[prior_w], min=0.)
  824. mo.gains = g_solution
  825. mo.update_r_star()
  826. #
  827. axGains.plot( mo.neurons_pref_stims, g_solution, c=c, ls='-')
  828. g_sol_approx = torch.clamp( g_solutions_approx[prior_w], min=0.)
  829. axGains.plot(mo.neurons_pref_stims, g_sol_approx, c=c, ls='--' )
  830. #
  831. true_φs, true_sigmarsSq = mo.effective_tc_center_and_squaredwidth()
  832. φs_perprior[prior_w] = true_φs
  833. argmaxs_perprior[prior_w] = STIMULI[ mo.r_star.argmax(axis=1) ]
  834. tuning_curves[prior_w] = mo.r_star.clone().detach().numpy()
  835. idxs = [355,368,379,388,395,400,405,412,421,432,445]
  836. # Number of colors to extract from cmap
  837. cmap = cm.Spectral #viridis_r #Spectral # viridis_r
  838. num_colors = len(idxs) # Adjust based on the number of lines you need
  839. cmap_colors = [to_hex(cmap(i / (num_colors - 1))) for i in range(num_colors)]
  840. wref = 30
  841. a_prior_ref, _ = make_gaussian_priors(mo.neurons_pref_stims, wref)
  842. aa = 1.
  843. ref_Fa, ref_Fa_m1 = make_Fa_and_Fam1_fcts(STIMULI, a_prior_ref, a=aa)
  844. for iν, ν in enumerate([30, 20, 10]):
  845. c = COLORS_WIDTH.get(ν)
  846. a_prior, _ = make_gaussian_priors(mo.neurons_pref_stims, ν)
  847. g_sol = g_solutions[ν]
  848. Fa, Fa_m1 = make_Fa_and_Fam1_fcts(STIMULI, a_prior, a=aa)
  849. axPhi.plot(φs_perprior[wref], φs_perprior[ν], c=c, ls='-', lw=1)
  850. axPhi.plot(STIMULI, Fa_m1(ref_Fa(STIMULI)), c=c, ls=':', lw=1)
  851. axPhi.scatter(φs_perprior[wref][idxs], φs_perprior[ν][idxs], color=cmap_colors, zorder=10, s=20, clip_on=True, ec='k', linewidth=.5)
  852. xlim_val = 39
  853. for ip,prior_w in enumerate( [30, 20, 10] ):
  854. ax = axs_tc[ip]
  855. ax.set_prop_cycle(cycler(color=cmap_colors))
  856. a_prior, _ = make_gaussian_priors(mo.neurons_pref_stims, prior_w)
  857. rrr = tuning_curves[prior_w][idxs]
  858. ax.plot(STIMULI, (rrr.T/rrr.T.max(axis=0)), lw=.5)
  859. xs = φs_perprior[prior_w][idxs]
  860. ys = np.zeros_like(xs)
  861. n = len(idxs)
  862. mid = n//2
  863. # 1) build the front-to-back order of positions in idxs
  864. order = [mid]
  865. for k in range(1, n):
  866. if mid - k >= 0:
  867. order.append(mid - k)
  868. if mid + k < n:
  869. order.append(mid + k)
  870. # 2) now scatter each in that order, giving the first (the true front) the highest zorder
  871. base = 1000 # same base you were using
  872. for layer, pos in enumerate(order):
  873. z = base - layer
  874. ax.scatter(xs[pos], ys[pos], color=cmap_colors[pos], s=20, ec='k', linewidth=.5, clip_on=False, zorder=z)
  875. ax.set_xlim(-xlim_val,xlim_val)
  876. ax.fill_between(STIMULI, .4*a_prior/a_prior.max(), color=COLORS_WIDTH[prior_w], zorder=-1000, alpha=.2)
  877. ax.set_ylim(0,None)
  878. ax.set_ylabel({10:'Narrow', 20:'Medium', 30:'Wide'}[prior_w] + '\nprior', color=COLORS_WIDTH[prior_w], rotation=0, y=.3, labelpad=24)
  879. axPhi.set_title('Effective preferred stimuli')
  880. axPhi.set_xlim(-xlim_val, xlim_val); axPhi.set_ylim(-xlim_val, xlim_val)
  881. axPhi.set_xlabel('Preferred stimulus with Wide prior, $\\varphi_{30}$')
  882. axPhi.set_ylabel('Preferred stimulus, $\\varphi_{\\sigma_p}$', labelpad=-5)
  883. axPhi.legend(labels=[f'$\\varphi_{{\\sigma_p}}$ vs. $\\varphi_{{ {wref} }}$', f'$P_{{\\sigma_p}}^{{-1}}(P_{{ {wref} }}(\\varphi_{{ {wref} }}))$',],
  884. handles=[Line2D([],[],ls='-',c='grey'),Line2D([],[],ls=':',c='grey'),], loc='upper left')
  885. axPriors.set_title('Priors')
  886. axPriors.set_xlim(-110, 110); axPriors.set_ylim(0, None)
  887. axPriors.legend(loc=(.6,.6))
  888. axPriors.set_yticks([])
  889. axPriors.set_xlabel('Stimulus')
  890. axPriors.set_ylabel('pdf')
  891. axGains.set_title('Optimal gains')
  892. ###
  893. axGains.set_xlim(-110, 110); axGains.set_ylim(0, .1099)
  894. axGains.set_xlabel('Feedforward preferred stimulus',)
  895. axGains.set_ylabel('$g(s)$', labelpad=0)
  896. axGains.legend(labels=['Numerical optimization', 'Analytical approximation',],
  897. handles=[Line2D([],[],ls='-',c='grey'),Line2D([],[],ls='--',c='grey'),], loc=(.45,.83))
  898. axGains.ticklabel_format( axis='y', style='sci', scilimits=(-2, -2), useMathText=True )
  899. axGains.yaxis.get_offset_text().set_fontsize(8)
  900. axGains.yaxis.get_offset_text().set_x(-.025)
  901. axs_tc[0].set_title('Normalized effective tuning curves')
  902. axs_tc[0].set_xticks([]); axs_tc[1].set_xticks([])
  903. axs_tc[0].set_yticks([]); axs_tc[1].set_yticks([]); axs_tc[2].set_yticks([])
  904. axs_tc[2].set_xlabel('Stimulus')
  905. # Remove top and right spines for all axes
  906. for ax in [axPriors, axGains, axPhi, *axs_tc]:
  907. ax.spines['top'].set_visible(False)
  908. ax.spines['right'].set_visible(False)
  909. fig.subplots_adjust(hspace=.35)
  910. #fig.savefig('opt_gains.pdf', bbox_inches='tight')
  911. #show()
  912. # %%
  913. # np.savez_compressed('CACHE_find_optimal_gains_for_parameters.npz', cache=CACHE_find_optimal_gains_for_parameters)
  914. # %% [markdown]
  915. # ## FIG. 5 g, d, sigma, etc.
  916. # %%
  917. samples_and_estimates_cache = {}
  918. # %%
  919. num_neurons = 801
  920. ff_sigma = 5.
  921. rec_magnitude = 0.95
  922. rec_width = 6.
  923. alpha_cost = .5
  924. reg_lambda = 1000
  925. inhib = 0.
  926. T = 1
  927. tau = .01
  928. burn_in = 0.5
  929. dt = 0.1 #0.001
  930. num_samples = 1000
  931. if dt != 0.001:
  932. print(f"Warning: dt = {dt} != 0.001. The response variances may be incorrect.")
  933. prior_ws = [10, 20, 30]
  934. g_solutions = {}
  935. for ip,prior_w in enumerate( prior_ws ):
  936. g_solution, loss, gains_list, losses_list = find_optimal_gains_for_parameters(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width,
  937. alpha_cost=alpha_cost, prior_w=prior_w, regularizer_lambda=reg_lambda, inhib=inhib)
  938. g_solutions[prior_w] = g_solution
  939. mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
  940. vars = []
  941. for ip,prior_w in enumerate(prior_ws):
  942. print( f'prior width = {prior_w}')
  943. c = COLORS_WIDTH.get(prior_w)
  944. a_prior, a_prior_neurons = make_gaussian_priors(mo.neurons_pref_stims, prior_w) #, prior_center_val)
  945. prior_mean, prior_var = probs_mean_and_var(a_prior_neurons, mo.neurons_pref_stims)
  946. ####
  947. g_solution = torch.clamp( g_solutions[prior_w], min=0.)
  948. mo.gains = g_solution
  949. mo.update_r_star()
  950. ###
  951. xlo = - prior_w / 2
  952. xhi = - xlo
  953. if prior_w in samples_and_estimates_cache:
  954. samples, estimates = samples_and_estimates_cache[prior_w]
  955. else:
  956. samples = mo.simulate_recurrent_poisson_many(num_samples=num_samples, xbounds=(xlo, xhi), T=T, tau=tau, dt=dt, burn_in=burn_in)
  957. estimates = mo.decode(samples, MID_VAL, prior_w**2)
  958. samples_and_estimates_cache[prior_w] = (samples, estimates)
  959. #
  960. excursions = estimates - estimates.mean(axis=0)
  961. vars.append( excursions.var(axis=0).mean() )
  962. vars = np.array(vars)
  963. # %%
  964. num_neurons = 801
  965. ff_sigma = 5.
  966. rec_magnitude = 0.95
  967. rec_width = 6.
  968. alpha_cost = .5
  969. reg_lambda = 1000 #600
  970. inhib = 0.
  971. T = 1
  972. tau = .01
  973. burn_in = 0.5
  974. dt = 0.1 #0.001
  975. num_samples = 1000
  976. prior_ws = [10, 20, 30]
  977. fig, axs = subplots(nrows=2, ncols=3, figsize=(9,5), sharex=False, )
  978. axG = axs[0,0]
  979. axW = axs[0,1]
  980. axd = axW.twinx()
  981. axWd = axs[0,2]
  982. axI = axs[1,0]
  983. axNbSpikes = axs[1,1]
  984. axV = axs[1,2]
  985. g_solutions = {}
  986. tuning_curves = {}
  987. φs_perprior = {}
  988. centergs = {}
  989. centersigmas = {}; sigmas = {}
  990. centerdensities = {}; densities = {}
  991. exp_num_spikes = {}
  992. center_fish_infos = {}
  993. fish_infos = {}
  994. exp_sq_error = {}
  995. mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
  996. prior_wparams = np.array([10, 20, 30])
  997. for ip,prior_w in enumerate( prior_wparams ):
  998. g_solution, loss, gains_list, losses_list = find_optimal_gains_for_parameters(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width,
  999. alpha_cost=alpha_cost, prior_w=prior_w, regularizer_lambda=reg_lambda, inhib=inhib)
  1000. mo.gains = g_solution
  1001. mo.update_r_star()
  1002. a_prior, a_prior_neurons = make_gaussian_priors(mo.neurons_pref_stims, prior_w)
  1003. prior_mean, prior_var = probs_mean_and_var(a_prior_neurons, mo.neurons_pref_stims)
  1004. true_φs, true_sigmarsSq = mo.effective_tc_center_and_squaredwidth()
  1005. φs_perprior[prior_w] = true_φs
  1006. tuning_curves[prior_w] = mo.r_star.clone().detach().numpy()
  1007. ### Gains
  1008. centergs[prior_w] = g_solution[len(g_solution)//2]
  1009. # Tuning curves width
  1010. centersigmas[prior_w] = np.sqrt( true_sigmarsSq[len(g_solution)//2] )
  1011. #sigmas[prior_w] = np.sqrt( true_sigmarsSq )
  1012. # Density
  1013. ind = np.abs(mo.neurons_pref_stims - 0.) <= prior_w*10000
  1014. φm1prime = np.gradient(mo.neurons_pref_stims[ind], true_φs[ind])
  1015. # print(true_φs[ind][len(φm1prime)//2])
  1016. centerdensities[prior_w] = φm1prime[len(φm1prime)//2]/mo.δ
  1017. #densities[prior_w] = φm1prime / mo.δ
  1018. # Expected total number of spikes vs width
  1019. exp_num_spikes[prior_w] = (a_prior * mo.r_star.sum(axis=0)).sum()
  1020. # Fisher information - center
  1021. center_fish_infos[prior_w] = SQRT_2_PI * centergs[prior_w] * centerdensities[prior_w] / centersigmas[prior_w]
  1022. exp_sq_error[prior_w] = (a_prior.numpy()/(1/prior_var + (1./mo.beta)*(mo.r_star.T/true_sigmarsSq).sum(axis=1)) ).sum()
  1023. def strip_leading_zero(x, pos):
  1024. return '0' if x==0 else f"{x:.2f}".lstrip("0") if abs(x) < 1 else f"{x:.2f}"
  1025. axG.plot( prior_wparams, [centergs[ν] for ν in prior_wparams], '.-', c='lightgrey', lw=1)
  1026. axG.scatter( [30, 20, 10], [centergs[ν] for ν in [30, 20, 10]], color=[COLORS_WIDTH[w] for w in [30,20,10]], s=18, zorder=1000)
  1027. axG.set_title('Gain $g$')
  1028. axG.yaxis.set_major_formatter(FuncFormatter(strip_leading_zero))
  1029. axG.set_ylim(0, None)
  1030. axW.plot( prior_wparams, [centersigmas[ν] for ν in prior_wparams], '.-', c='lightgrey', lw=1)
  1031. axW.scatter( [30, 20, 10], [centersigmas[ν] for ν in [30, 20, 10]], color=[COLORS_WIDTH[w] for w in [30,20,10]], s=18, zorder=1000)
  1032. axW.set_title('Tuning curve width and density')
  1033. axW.set_ylim(0, None)
  1034. axW.legend(handles=[Line2D([], [], color='lightgrey', label='Width $\\sigma$ (left)'),
  1035. Line2D([], [], color='lightgrey', ls='--', label='Density $d$ (right)')],)
  1036. axd.plot( prior_wparams, np.array([centerdensities[ν] for ν in prior_wparams]), '.--', c='lightgrey', lw=1)
  1037. axd.scatter( np.array([30, 20, 10]), np.array([centerdensities[ν] for ν in [30, 20, 10]]), color=[COLORS_WIDTH[w] for w in [30,20,10]], s=18, zorder=1000)
  1038. axd.set_ylim(0, None)
  1039. # I \propto g d / sigma
  1040. axI.plot( prior_wparams, np.array([center_fish_infos[ν] for ν in prior_wparams]), '.-', c='lightgrey', lw=1)
  1041. axI.scatter( np.array([30, 20, 10]), np.array([center_fish_infos[ν] for ν in [30, 20, 10]]), color=[COLORS_WIDTH[w] for w in [30,20,10]], s=18, zorder=1000)
  1042. axI.set_title('Fisher $I \\propto g d / \\sigma$')
  1043. axI.set_xlabel('Prior width $\\sigma_p$')
  1044. xx = prior_wparams
  1045. yy = np.array([center_fish_infos[ν] for ν in prior_wparams])
  1046. a_fit, _ = optimize.curve_fit(lambda x,a: a/x, xx, yy)
  1047. axI.plot(np.linspace(10,30), a_fit[0]/np.linspace(10,30), 'r:', lw=1, label='Fit $1/\\sigma_p$', zorder=-1000)
  1048. axI.legend(fontsize=10, loc='lower left')
  1049. axI.set_ylim(0, None);
  1050. axI.yaxis.set_major_formatter(FuncFormatter(strip_leading_zero))
  1051. axI.set_ylim(0, None)
  1052. # # Var / std dev
  1053. axWd.plot( 1./np.array([centerdensities[ν] for ν in prior_wparams]), [centersigmas[ν] for ν in prior_wparams], '.-', c='lightgrey', lw=1)
  1054. axWd.scatter( 1./np.array([centerdensities[ν] for ν in [30,20,10]]), [centersigmas[ν] for ν in [30, 20, 10]], color=[COLORS_WIDTH[w] for w in [30,20,10]], s=18, zorder=1000)
  1055. axWd.set_title('Width vs. spacing')
  1056. axWd.set_xlabel('Spacing (1/density)')
  1057. axWd.set_ylabel('Width $\\sigma$', labelpad=-7)
  1058. axWd.set_yticks([10, 20])
  1059. axNbSpikes.plot( prior_wparams, [exp_sq_error[ν] for ν in prior_wparams], '.-', c='lightgrey', lw=1, label='$L \\simeq $MSE (left)')
  1060. axNbSpikes.scatter( [30, 20, 10], [exp_sq_error[ν] for ν in [30, 20, 10]], color=[COLORS_WIDTH[w] for w in [30,20,10]], s=18, zorder=1000)
  1061. axNbSpikes.set_title('Loss (MSE) and Cost (spikes)')
  1062. axNbSpikes.set_xlabel('Prior width $\\sigma_p$')
  1063. axNbSpikes.set_ylim(0, 24)
  1064. axT = axNbSpikes.twinx()
  1065. axT.plot( prior_wparams, [exp_num_spikes[ν] for ν in prior_wparams], 'x--', c='lightgrey', lw=1, label='$C=$Spiking activity (right)')
  1066. axT.scatter( [30, 20, 10], [exp_num_spikes[ν] for ν in [30, 20, 10]], color=[COLORS_WIDTH[w] for w in [30,20,10]], marker='x', s=18, zorder=1000)
  1067. axT.set_ylim(0, 42)
  1068. axNbSpikes.legend(loc='upper left', borderpad=0, handletextpad=.5)
  1069. axT.legend(loc='lower left', borderpad=0, handletextpad=.5)
  1070. vs, vs_5, vs_95 = np.array([15.4389, 32.3343, 48.1599]), np.array([14.3715, 30.0494, 44.9652]), np.array([16.5432, 34.7697, 51.552 ])
  1071. line, caps, bars = axV.errorbar([10,20,30], vs, (vs-vs_5, vs_95-vs), c='lightgrey', lw=1, ls=':', label='Subjects')
  1072. axV.plot(prior_wparams, vars, '.-', c='lightgrey', lw=1, label='Network model')
  1073. axV.scatter([10,20,30], vars, color=[COLORS_WIDTH[w] for w in [10,20,30]], s=18, zorder=1000, clip_on=False)
  1074. axV.set_xlabel('Prior width $\\sigma_p$')
  1075. axV.set_title('Response variance')
  1076. axV.legend()
  1077. axV.set_ylim(0, 65)
  1078. for ax in axs.flat:
  1079. ax.spines['top'].set_visible(False)
  1080. ax.spines['right'].set_visible(False)
  1081. axT.spines['top'].set_visible(False)
  1082. axT.spines['right'].set_ls('--')
  1083. axd.spines['top'].set_visible(False)
  1084. axd.spines['right'].set_ls('--')
  1085. for ax in [axG, axW, axd, axI, axNbSpikes, axV]:
  1086. ax.set_xticks([10,20,30])
  1087. fig.tight_layout(h_pad=.5)
  1088. if dt != 0.001:
  1089. print(f"Warning: dt = {dt} != 0.001. The response variances may be incorrect.")
  1090. #fig.savefig('fisher_comp.pdf', bbox_inches='tight')
  1091. #show()
  1092. # %% [markdown]
  1093. # ## Adapter repulsion
  1094. # %% [markdown]
  1095. # ### FIG. 2 Illustration of adapter repulsion
  1096. # %% [markdown]
  1097. # #### Gratings
  1098. # %%
  1099. from scipy.ndimage import rotate
  1100. from matplotlib.offsetbox import OffsetImage, AnnotationBbox
  1101. # %%
  1102. def generate_square_wave_grating(size=256, spatial_freq=10, orientation=0):
  1103. """Return a high-contrast square-wave grating as a NumPy array."""
  1104. x = np.linspace(-0.5, 0.5, size)
  1105. y = np.linspace(-0.5, 0.5, size)
  1106. xv, yv = np.meshgrid(x, y)
  1107. theta = np.deg2rad(orientation)
  1108. xt = xv * np.cos(theta) + yv * np.sin(theta)
  1109. grating = np.sign(np.sin(2 * np.pi * spatial_freq * xt))
  1110. return grating
  1111. def add_grating_to_ax(ax, grating, xy, zoom=1.0, **kwargs):
  1112. """
  1113. Adds a grating image to a given Axes at specified coordinates.
  1114. Parameters:
  1115. - ax: matplotlib Axes object
  1116. - grating: 2D NumPy array, the grating image
  1117. - xy: tuple (x, y), coordinates in data units where the image will be placed
  1118. - zoom: float, scale of the image
  1119. - **kwargs: additional keyword arguments passed to AnnotationBbox (e.g., xycoords)
  1120. """
  1121. imagebox = OffsetImage(grating, cmap='gray', zoom=zoom)
  1122. ab = AnnotationBbox(imagebox, xy, frameon=False, **kwargs)
  1123. ax.add_artist(ab)
  1124. # fig, ax = subplots()
  1125. # grating = generate_square_wave_grating(size=256, spatial_freq=10, orientation=45)
  1126. # add_grating_to_ax(ax, grating, xy=(0.5, 0.5), zoom=0.2)
  1127. # add_grating_to_ax(ax, grating, xy=(0.2, 0.8), zoom=0.3)
  1128. # ax.set_xlim(0, 1)
  1129. # ax.set_ylim(0, 1)
  1130. #show()
  1131. # %%
  1132. def save_grating_to_file(grating, filename, dpi=300, format='png'): # or 'tiff
  1133. """
  1134. Parameters:
  1135. - grating: 2D numpy array with values in [-1, 1]
  1136. - filename: output path, should end in .tif or .tiff
  1137. - dpi: resolution metadata for image (useful for print layout)
  1138. """
  1139. # Normalize to [0, 255] for 8-bit grayscale image
  1140. grating_uint8 = ((grating + 1) / 2 * 255).astype(np.uint8)
  1141. imsave(fname=filename, arr=grating_uint8, cmap='gray', dpi=dpi, format=format)
  1142. for ori in [0, 22.5, 45, 67.5, 90, 112.5, 135, 157.5, ]:
  1143. g = generate_square_wave_grating(size=512, spatial_freq=5, orientation=ori)
  1144. #save_grating_to_file(g, f'figures/grating_{ori}deg.png', dpi=300)
  1145. figure(figsize=(1,1))
  1146. imshow(g, cmap='gray')
  1147. # %% [markdown]
  1148. # #### Fig. 2B
  1149. # %%
  1150. from matplotlib.patches import FancyArrowPatch
  1151. fig, axRep = subplots(figsize=(6, 4))
  1152. xx = np.linspace(-8,95,1001)
  1153. cC = 'C0'; cA = cm.tab20(7)
  1154. axRep.plot(xx, stats.norm(loc=10., scale=8).pdf(xx)/stats.norm(scale=8).pdf(0), c=cC, label='Control', lw=2)
  1155. axRep.plot(xx, stats.norm(loc=15., scale=12).pdf(xx)/stats.norm(scale=12).pdf(0), c=cA, label='Adaptation', lw=2)
  1156. #
  1157. axRep.plot(xx, stats.norm(loc=45., scale=8).pdf(xx)/stats.norm(scale=8).pdf(0), c=cC, ls='--', lw=2)
  1158. axRep.plot(xx, stats.norm(loc=47., scale=7).pdf(xx)/stats.norm(scale=7).pdf(0), c=cA, ls='--', lw=2)
  1159. # Add double arrows to indicate the FWHM of the two normalized Gaussians above
  1160. for (mu, sigma, color) in [(10., 8., cC), (15., 12., cA), (45., 8., cC), (47., 7., cA)]:
  1161. y = stats.norm(loc=mu, scale=sigma).pdf(xx) / stats.norm(scale=sigma).pdf(0)
  1162. half_max = y.max() / 2.0
  1163. above_half = np.where(y >= half_max)[0]
  1164. if above_half.size > 0:
  1165. fwhm_left = xx[above_half[0]]
  1166. fwhm_right = xx[above_half[-1]]
  1167. fwhm_y = y.max() * 0.5 + (0.025 if color == cC else 0.)
  1168. axRep.annotate(
  1169. '', xy=(fwhm_right, fwhm_y), xytext=(fwhm_left, fwhm_y),
  1170. arrowprops=dict(arrowstyle='<->', color=color, lw=1)
  1171. )
  1172. axRep.plot(xx, stats.norm(loc=82, scale=8).pdf(xx)/stats.norm(scale=8).pdf(0), c=cC, ls=':', lw=2)
  1173. axRep.plot(xx, stats.norm(loc=80, scale=7).pdf(xx)/stats.norm(scale=7).pdf(0), c=cA, ls=':', lw=2)
  1174. #
  1175. axRep.vlines(0, 0, 1., color='lightgrey', ls='-', zorder=-100)
  1176. axRep.annotate(" Adapter", xytext=(0, 1.25), xy=(0, 1.0), arrowprops=dict(arrowstyle="-|>", color='r'), ha='center')
  1177. axRep.annotate("", xytext=(10, 1+.03), xy=(15, 1+.03), arrowprops=dict(arrowstyle="-|>", color='k', ls='-', lw=.5, shrinkA=0, shrinkB=0))
  1178. axRep.annotate("", xytext=(44.8, 1+.03), xy=(47.2, 1+.03), arrowprops=dict(arrowstyle="-|>", color='k', ls='-', lw=.5, shrinkA=0, shrinkB=0))
  1179. axRep.annotate("", xytext=(82.2, 1+.03), xy=(79.8, 1+.03), arrowprops=dict(arrowstyle="-|>", color='k', ls='-', lw=.5, shrinkA=0, shrinkB=0))
  1180. axRep.set_ylim(0,1.375)
  1181. axRep.set_xlim(xx[0],xx[-1])
  1182. axRep.set_xlabel('Orientation (difference from adapter)')
  1183. axRep.set_yticks([])
  1184. axRep.set_ylabel('Firing rate')
  1185. axRep.legend(loc=(.275,.87) )
  1186. axRep.spines['top'].set_visible(False)
  1187. axRep.spines['right'].set_visible(False)
  1188. fig.subplots_adjust(wspace=.25)
  1189. #fig.savefig('illus_adapter_repulsion.pdf', bbox_inches='tight')
  1190. # %% [markdown]
  1191. # ### FIG. 6 Adapter repulsion - Model
  1192. # %%
  1193. num_neurons = 801
  1194. ff_sigma = 5.
  1195. rec_magnitude = 0.95
  1196. rec_width = 6.
  1197. inhib = 0.
  1198. alpha_cost = .5
  1199. regularizer_lambda = 1000
  1200. uniform_part = 0.8
  1201. prior_wparams=[10000, 1]
  1202. g_solutions = {}
  1203. # g_solutions_approx = {}
  1204. # g_params_approx = {}
  1205. for ip,prior_w in enumerate( prior_wparams ):
  1206. this_uniform_part = uniform_part if prior_w < 9999 else 1.
  1207. gains, loss, gains_list, losses_list = find_optimal_gains_for_parameters_uniform_mixture(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width, alpha_cost=alpha_cost,
  1208. prior_w=prior_w, regularizer_lambda=regularizer_lambda, uniform_part=this_uniform_part, inhib=inhib)
  1209. # g_solutions_approx[prior_w] = g_solution_approx
  1210. # g_params_approx[prior_w] = (g0, Delta)
  1211. g_solutions[prior_w] = gains
  1212. # %%
  1213. mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width, inhib=inhib)
  1214. fig, axs = subplots(nrows=2, ncols=5, figsize=(18,6))
  1215. tuning_curves = {}
  1216. φs_perprior = {}
  1217. sigmas = {}
  1218. fwhms = {}
  1219. axg = axs[1,0]
  1220. these_prior_wparams = [10000, 1]
  1221. for ip,prior_w in enumerate(these_prior_wparams):
  1222. print( f'prior width = {prior_w}', end=' ')
  1223. c = COLORS_WIDTH.get(prior_w)
  1224. this_uniform_part = uniform_part if prior_w < 9999 else 1.
  1225. xxx = torch.linspace(LOWER_BOUND, UPPER_BOUND, 10001)
  1226. a_prior, a_prior_xxx = make_mixture_uniform_gaussian_priors(xxx, prior_w, uniform_part=this_uniform_part)
  1227. #prior_mean, prior_var = probs_mean_and_var(a_prior_xxx, xxx)
  1228. label = 'Control' if prior_w==10000 else 'Adaptation'
  1229. axs[0,0].plot( xxx, a_prior_xxx, c=c, ls='-', label=label)
  1230. ####
  1231. g_solution = torch.clamp( g_solutions[prior_w], min=0.)
  1232. mo.gains = g_solution
  1233. mo.update_r_star()
  1234. #
  1235. axg.plot( mo.neurons_pref_stims, g_solution, c=c, ls='-')
  1236. tuning_curves[prior_w] = mo.r_star.clone().detach().numpy()
  1237. fwhms[prior_w] = np.array( [tc_fwhm_interp(tc) for tc in mo.r_star] )
  1238. axs[0,0].set_title('Priors')
  1239. axs[0,0].set_xlim(-40, 40);
  1240. axs[0,0].set_ylim(0, .0045)
  1241. axs[0,0].legend(loc='upper left')
  1242. axs[0,0].set_yticks([])
  1243. g_solution = torch.clamp( g_solutions[1], min=0.)
  1244. ind50 = mo.neurons_pref_stims.abs() <= 50
  1245. imax = g_solution[ind50].argmax()
  1246. print( '***', mo.neurons_pref_stims[ind50][imax], g_solution[ind50][imax] )
  1247. xmax1 = mo.neurons_pref_stims[ind50][imax]
  1248. xmax2 = -xmax1
  1249. axg.scatter([xmax1, xmax2], [0,0], marker='*', clip_on=False, color=COLORS_WIDTH.get(1), zorder=10, s=100, ec='k', lw=.5)
  1250. axg.set_title('Optimal gains')
  1251. axg.set_xlim(-40, 40);
  1252. axg.set_ylim(0, .19)
  1253. axg.set_xlabel('Stimulus $s$')
  1254. for iν, ν in enumerate(these_prior_wparams):
  1255. c = COLORS_WIDTH.get(ν)
  1256. axs[1,4].plot(mo.neurons_pref_stims, fwhms[ν], c=c, ls='-')
  1257. axs[1,4].set_xlim(0, 23)
  1258. axs[1,4].set_ylim(0, 45)
  1259. axs[1,4].set_xlabel('Distance from adapter')
  1260. axs[1,4].set_title('Full Width at Half Maximum')
  1261. axs[1,4].scatter(abs(xmax1), 0, marker='*', clip_on=False, color=COLORS_WIDTH.get(1), zorder=10, s=100, ec='k', lw=.5)
  1262. ######## Tuning curves #######
  1263. # Plot individual neuron tuning curves (normalized) across priors
  1264. neuron_indices = [400, 402, 405, 407, 412, 428, 446]
  1265. axtcs = axs.flatten()[[1,2,3,4,6,7,8]]
  1266. for ii,idx in enumerate(neuron_indices):
  1267. ax = axtcs[ii]
  1268. axs[1,4].scatter(mo.neurons_pref_stims[idx], 0, marker='v', clip_on=False, c='C0', zorder=10, s=20)
  1269. ax.scatter([xmax1, xmax2], [0,0], marker='*', clip_on=False, color=COLORS_WIDTH.get(1), zorder=10, s=100, ec='k', lw=.5)
  1270. for iw,prior_w in enumerate(these_prior_wparams):
  1271. c = COLORS_WIDTH.get(prior_w)
  1272. tc = tuning_curves[prior_w][idx, :] # tuning curve for neuron idx under prior prior_w
  1273. tc_normalized = tc / tc.max() # normalize
  1274. ax.plot(STIMULI, tc_normalized, label=f'prior {prior_w}', c=c)
  1275. # Draw a double arrow to indicate FWHM of the tuning curve
  1276. fwhm_val = fwhms[prior_w][idx]
  1277. above_half = np.where(tc_normalized >= .499)[0]
  1278. # delta_argmax = STIMULI[tc.argmax()] - STIMULI[ tuning_curves[10000][idx, :].argmax() ]
  1279. # ax.text(0, (.9-iw*0.1), f'r={delta_argmax/fwhms[10000][idx]:.4g}, r={(fwhm_val/fwhms[10000][idx]-1):.4g}', ha='center', fontsize=8, color=c)
  1280. if idx in [400, 428]:
  1281. if above_half.size > 0:
  1282. x_left = STIMULI[above_half[0]]
  1283. x_right = STIMULI[above_half[-1]]
  1284. y_arrow = 0.5 + 0.02*iw # vertical position for the arrow (normalized units)
  1285. ax.annotate(
  1286. '', xy=(x_left, y_arrow), xytext=(x_right, y_arrow),
  1287. arrowprops=dict(arrowstyle='<->', color=c, lw=1, relpos=0, shrinkA=.5, shrinkB=.5),
  1288. annotation_clip=False
  1289. )
  1290. axs[1,4].scatter(mo.neurons_pref_stims[idx], fwhms[prior_w][idx], marker='o', clip_on=False, c=c, zorder=10, s=20)
  1291. ax.set_xlim(-30, 30)
  1292. ax.set_ylim(0,None)
  1293. ax.axvline(0, lw=1, c='lightgrey', zorder=-1000)
  1294. ax.set_xticks([-30, -20, -10, 0, 10, 20, 30])
  1295. for ax in axs[:1,1]:
  1296. ax.set_ylabel('Normalized response')
  1297. for ax in axs[0,:]:
  1298. ax.set_xticklabels([])
  1299. for ax in axs[1,:-1]:
  1300. ax.set_xlabel('Stimulus $s$')
  1301. axs[0,1].set_title('Tuning curves')
  1302. for ax in axs.flat:
  1303. ax.spines['top'].set_visible(False)
  1304. ax.spines['right'].set_visible(False)
  1305. #fig.savefig('adapter_repulsion.pdf', bbox_inches='tight')
  1306. #show()
  1307. # %% [markdown]
  1308. # ## Approximations (Supplementary Information)
  1309. # %% [markdown]
  1310. # ### FIG. 8 Eigenvalues of W
  1311. # %%
  1312. mo = Network(num_neurons=801, ff_sigma=5, rec_magnitude=.85, rec_width=10)
  1313. figure(figsize=(14,4))
  1314. for rec_width in [.1, .5, 1., 2., 4., 6., 8.]:
  1315. mo.rec_width = rec_width
  1316. mo.update_W_and_M()
  1317. lines = plot(np.arange(801), sorted( torch.linalg.eigvals(mo.recurrent_weights).abs(), reverse=True ), label=f'$\\sigma_{{rec}} = {rec_width}$' )
  1318. plot(np.arange(801), mo.rec_magnitude * np.exp(-.5*(np.pi*(mo.rec_width)*np.arange(mo.num_neurons)/OVERALL_WIDTH)**2 ), c=lines[0].get_color(), ls='--')
  1319. leg = legend(loc=(.88, .15))
  1320. title('Eigenvalues $\\lambda_k$ of $W$')
  1321. xlim(0,801)
  1322. ylim(0,None)
  1323. xlabel('$k$')
  1324. ylabel('$\\lambda_k$')
  1325. legend(handles=[Line2D([], [], color='grey', label='Correct eigenvalue'), Line2D([], [], color='grey', linestyle='--', label='Approximation')], loc=(.65,.5))
  1326. gca().add_artist(leg)
  1327. gca().spines['top'].set_visible(False)
  1328. gca().spines['right'].set_visible(False)
  1329. #savefig('eigenvalues_W.pdf', bbox_inches='tight')
  1330. # %% [markdown]
  1331. # ### FIG. 9-10 Matrix M
  1332. # %% [markdown]
  1333. # $$
  1334. # M_{ij} \approx \frac{1}{1-\lambda_0} \frac{1}{n}
  1335. # + \frac{2}{n} \sum_{k=1}^{n-1} \frac{1}{1-\lambda_k} \cos\left( \frac{\pi k}{n} \left( i + \frac{1}{2} \right) \right) \cos\left( \frac{\pi k}{n} \left( j + \frac{1}{2} \right) \right)
  1336. # $$
  1337. # %%
  1338. ff_sigma = -1. # shouldn't matter here
  1339. approx_Ms = {}
  1340. for imag, rec_magnitude in enumerate( [.2, .95, .98, .99] ):
  1341. for iw,rec_width in enumerate( [6.,] ):
  1342. mo = Network(num_neurons=801, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
  1343. mo.update_W_and_M()
  1344. # eigenvalues
  1345. lambdas = mo.rec_magnitude * np.exp(-.5*(np.pi*(mo.rec_width)*np.arange(mo.num_neurons)/OVERALL_WIDTH)**2 )
  1346. eigenvalues_of_M = 1./(1-lambdas)
  1347. # matrix of eigenvectors
  1348. # matrix_Q = (factor) np.cos( (np.pi * np.arange(mo.num_neurons) / mo.num_neurons) * ( np.arange(mo.num_neurons)[np.newaxis] + .5) )
  1349. approx_M = np.ones_like(mo.matrix_M) * eigenvalues_of_M[0] / mo.num_neurons
  1350. for i in range(mo.num_neurons):
  1351. for j in range(mo.num_neurons):
  1352. coscos = np.cos( (np.pi * np.arange(1,mo.num_neurons) / mo.num_neurons) * (i + .5)) * np.cos( (np.pi * np.arange(1,mo.num_neurons) / mo.num_neurons) * (j + .5))
  1353. approx_M[i,j] += (2./mo.num_neurons) * np.sum( eigenvalues_of_M[1:] * coscos )
  1354. approx_Ms[(rec_magnitude, rec_width)] = approx_M
  1355. # %% [markdown]
  1356. # As an infinite sum of Gaussian:
  1357. # $$
  1358. # M_{ij} \approx \delta_{ij} + \delta \sum_{m=1}^{\infty} \lambda_0^m \frac{1}{\sigma_{rec} \sqrt{m}\sqrt{2 \pi}}
  1359. # \exp{\left( -\frac{1}{2} \frac{(i-j)^2 \delta^2}{m \sigma_{rec}^2} \right)}
  1360. # $$
  1361. # %%
  1362. fig, axs = subplots(nrows=3, ncols=2, figsize=(12,12))
  1363. ff_sigma = -1. # shouldn't matter here
  1364. rec_width = 6.
  1365. for irec, rec_magnitude in enumerate( [0.2, 0.95, 0.99] ):
  1366. these_axs = axs[irec]
  1367. mo = Network(num_neurons=801, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
  1368. mo.update_W_and_M()
  1369. approx_M = approx_Ms[(rec_magnitude, rec_width)]
  1370. these_axs[0].plot( (mo.matrix_M-np.eye(801))[:,40::80], lw=1)
  1371. these_axs[1].plot( (mo.matrix_M-np.eye(801))[:,40::80], lw=1)
  1372. these_axs[0].set_prop_cycle(None)
  1373. these_axs[1].set_prop_cycle(None)
  1374. # cos cos
  1375. approx_M = approx_Ms[(rec_magnitude, rec_width)]
  1376. these_axs[0].plot( (approx_M-np.eye(801))[:,40::80], ls='--', lw=2)
  1377. # infinite sum of Gaussians
  1378. number_of_gaussian = 200
  1379. δ = OVERALL_WIDTH / mo.num_neurons
  1380. approx_M = torch.diag(torch.ones(mo.num_neurons))
  1381. for m in range(1, number_of_gaussian+1):
  1382. sigma_m = np.sqrt(m) * rec_width
  1383. approx_M += δ * rec_magnitude**m * (1./(sigma_m*SQRT_2_PI))*torch.exp(-( mo.neurons_pref_stims[np.newaxis] - mo.neurons_pref_stims[:,np.newaxis] )**2 / (2*sigma_m**2))
  1384. these_axs[1].plot( (approx_M-np.eye(801))[:,40::80], ls='--', lw=2)
  1385. for ax in these_axs:
  1386. ax.spines['top'].set_visible(False)
  1387. ax.spines['right'].set_visible(False)
  1388. ax.set_xlim(0,800)
  1389. these_axs[0].set_ylabel(f'$\\lambda_0 = {rec_magnitude}$', rotation=90, fontsize=14)
  1390. for ax in axs[0]:
  1391. ax.set_ylim(0, .0099)
  1392. for ax in axs[1]:
  1393. ax.set_ylim(0, .249)
  1394. for ax in axs[2]:
  1395. ax.set_ylim(0, .8)
  1396. ax.set_xlabel('$i$')
  1397. axs[0,0].set_title('Approximation of $M-I$ with cosines')
  1398. axs[0,1].set_title('Approximation of $M-I$ with sum of Gaussian functions')
  1399. axs[0,0].legend(handles=[Line2D([0], [0], color='grey', lw=1, label='$M_{ij}-\\delta_{ij}$'),
  1400. Line2D([0], [0], color='grey', lw=2, ls='--', label='Approximation')], loc='upper center', ncols=2)
  1401. #fig.savefig('approx_M_cos_gauss.pdf', bbox_inches='tight')
  1402. # %% [markdown]
  1403. # As a Laplace function:
  1404. # $$
  1405. # M_{ij} \approx \delta_{ij} + \frac{1}{\ln{(1/\lambda_0)}} \frac{\delta}{2 \nu_{rec}} \exp{\left( -\frac{|i-j|\delta}{\nu_{rec}} \right)}
  1406. # $$
  1407. # where
  1408. # $$
  1409. # \nu_{rec} = \frac{\sigma_{rec}}{\sqrt{2\ln{(1/\lambda_0)}}}
  1410. # $$
  1411. # %% [markdown]
  1412. # As a mix of both
  1413. # %%
  1414. fig, axs = subplots(nrows=3, ncols=2, figsize=(12,12))
  1415. ff_sigma = -1. # shouldn't matter here
  1416. rec_width = 6.
  1417. for irec, rec_magnitude in enumerate( [0.2, 0.95, .99] ):
  1418. these_axs = axs[irec]
  1419. mo = Network(num_neurons=801, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
  1420. mo.update_W_and_M()
  1421. approx_M = approx_Ms[(rec_magnitude, rec_width)]
  1422. these_axs[0].plot( (mo.matrix_M-np.eye(801))[:,40::80], lw=1)
  1423. these_axs[1].plot( (mo.matrix_M-np.eye(801))[:,40::80], lw=1)
  1424. these_axs[0].set_prop_cycle(None)
  1425. these_axs[1].set_prop_cycle(None)
  1426. # Gaussian
  1427. approx_M_G = torch.diag(torch.ones(mo.num_neurons))
  1428. m = 1
  1429. sigma_m = np.sqrt(m) * rec_width
  1430. approx_M_G += δ * rec_magnitude**m * (1./(sigma_m*SQRT_2_PI))*torch.exp(-( mo.neurons_pref_stims[np.newaxis] - mo.neurons_pref_stims[:,np.newaxis] )**2 / (2*sigma_m**2))
  1431. these_axs[0].plot( (approx_M_G-np.eye(801))[:,40::80], ls='--', lw=2)
  1432. # Laplace
  1433. ν_rec = rec_width / np.sqrt(2*np.log(1/rec_magnitude))
  1434. approx_M_L = torch.diag(torch.ones(mo.num_neurons)) # delta_ij
  1435. approx_M_L += ( δ / ( 2*ν_rec*np.log(1/rec_magnitude) ) ) * torch.exp( -torch.abs( mo.neurons_pref_stims[np.newaxis] - mo.neurons_pref_stims[:,np.newaxis] ) / ν_rec )
  1436. these_axs[1].plot( (approx_M_L-np.eye(801))[:,40::80], ls='--', lw=2)
  1437. these_axs[0].set_prop_cycle(None)
  1438. these_axs[1].set_prop_cycle(None)
  1439. #
  1440. these_axs[0].set_ylabel(f'$\\lambda_0 = {rec_magnitude}$', rotation=90, fontsize=14)
  1441. for ax in these_axs:
  1442. ax.spines['top'].set_visible(False)
  1443. ax.spines['right'].set_visible(False)
  1444. for ax in axs[0]:
  1445. ax.set_ylim(0, .05)
  1446. ax.set_yticks([0, .01, .02, .03, .04])
  1447. for ax in axs[1]:
  1448. ax.set_ylim(0, .27)
  1449. for ax in axs[2]:
  1450. ax.set_ylim(0, .65)
  1451. ax.set_xlabel('$i$')
  1452. for ax in axs.flat:
  1453. ax.set_xlim(0,800)
  1454. axs[0,0].set_title('Gaussian approximation of $M-I$')
  1455. axs[0,1].set_title('Laplace approximation of $M-I$')
  1456. axs[0,0].legend(handles=[Line2D([0], [0], color='grey', lw=1, label='$M_{ij}-\\delta_{ij}$'),
  1457. Line2D([0], [0], color='grey', lw=2, ls='--', label='Approximation')], loc='center')
  1458. #fig.savefig('approx_M_gauss_laplace.pdf', bbox_inches='tight')
  1459. # %% [markdown]
  1460. # ### FIG. 12 Gaussian approxº of rates, and MSE
  1461. # %%
  1462. num_neurons = 801
  1463. ff_sigma = 5.
  1464. rec_magnitude = 0.95 #0.85
  1465. rec_width = 6. # 10.
  1466. mo = Network(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
  1467. mo.update_W_and_M()
  1468. mse_per_stim_perprior = {}
  1469. for prior_w in [10., 20., 30.]:
  1470. samples, estimates = samples_and_estimates_cache[prior_w]
  1471. xmin = MID_VAL-prior_w/2; xmax = MID_VAL+prior_w/2
  1472. ind = (STIMULI>=xmin) & (STIMULI<=xmax)
  1473. these_stims = STIMULI[ind]
  1474. mse_per_stim = ((estimates - these_stims)**2).mean(axis=0)
  1475. mse_per_stim_perprior[prior_w] = mse_per_stim
  1476. # %%
  1477. num_neurons = 801
  1478. ff_sigma = 5.
  1479. rec_magnitude = 0.95
  1480. rec_width = 6.
  1481. if dt != 0.001:
  1482. print(f"Warning: dt = {dt} != 0.001. The MSEs may be incorrect.")
  1483. fig, axs = subplots(nrows=2, ncols=2, figsize=(10,8))
  1484. #
  1485. ss = axs[1,1].get_subplotspec()
  1486. axs[1,1].remove()
  1487. subaxs_gs = ss.subgridspec(1, 3)
  1488. ax11a = fig.add_subplot(subaxs_gs[0])
  1489. ax11b = fig.add_subplot(subaxs_gs[1])
  1490. ax11c = fig.add_subplot(subaxs_gs[2])
  1491. subaxs = [ax11a, ax11b, ax11c]
  1492. #
  1493. mo = Network(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
  1494. mo.update_W_and_M()
  1495. for ii,prior_w in enumerate([10., 20., 30.]):
  1496. c = COLORS_WIDTH.get(prior_w)
  1497. g_solution, loss, gains_list, losses_list = find_optimal_gains_for_parameters(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width,
  1498. alpha_cost=.5, prior_w=prior_w, regularizer_lambda=1000., inhib=0.)
  1499. mo.gains = g_solution
  1500. mo.update_r_star()
  1501. gammas, gamma_den, φs = make_gamma_G_φs(mo.gains, mo.hh, mo.neurons_pref_stims)
  1502. sigmarsSq = ff_sigma**2 + ( gammas * (mo.neurons_pref_stims-φs.unsqueeze(1))**2 ).sum(axis=1) # (num_neurons,)
  1503. sigmars = np.sqrt(sigmarsSq)
  1504. xmax = 60
  1505. # rates
  1506. if prior_w==20.:
  1507. ax = axs[1,0]
  1508. ax.plot(STIMULI, mo.r_star[225:576:25].T, lw=1)
  1509. ax.set_prop_cycle(None)
  1510. amplitudes = (ff_sigma/sigmars) * gamma_den # (num_neurons,)
  1511. approx_r_2 = amplitudes.unsqueeze(1) * np.exp(-(STIMULI-φs.unsqueeze(1))**2 / (2*sigmarsSq.unsqueeze(1)))
  1512. ax.plot(STIMULI, approx_r_2[225:576:25].T, ls='--', lw=2);
  1513. ax.legend(handles=[Line2D([0], [0], color='grey', lw=1, label='Tuning curve $r_i(s)$'),
  1514. Line2D([0], [0], color='grey', lw=2, ls='--', label='Gaussian approximation')], loc='upper left', ncols=1)
  1515. ax.set_title('Gaussian approximations of tuning curves $r_i(s)$')
  1516. ax.set_xlabel('Stimulus $s$')
  1517. ax.set_ylabel('Rate $r_i(s)$')
  1518. ax.set_ylim(0, .54)
  1519. ax.set_xlim(-xmax, xmax)
  1520. # phi
  1521. ax = axs[0,0]
  1522. true_sstar, true_sigmarsSq = mo.effective_tc_center_and_squaredwidth()
  1523. ax.plot(mo.neurons_pref_stims, true_sstar, c=c, zorder=200, label={10:'Narrow', 20:'Medium', 30:'Wide'}.get(prior_w, '') )
  1524. ax.plot(mo.neurons_pref_stims, φs, '--', c=c, zorder=100, )
  1525. ax.axline((0,0), slope=1, c='lightgrey', zorder=-100)
  1526. ax.set_ylim(-xmax, xmax)
  1527. ax.set_xlim(-xmax, xmax)
  1528. ax.set_xlabel('Feedforward preferred stimulus $s_i$')
  1529. ax.set_ylabel('Effective preferred stimulus $\\varphi(s_i)$')
  1530. ax.set_title('Effective preferred stimulus $\\varphi(s_i)$')
  1531. # sigmaSq
  1532. ax = axs[0,1]
  1533. true_sigmars = np.sqrt(true_sigmarsSq)
  1534. ax.plot(mo.neurons_pref_stims, true_sigmars, c=c, zorder=200)
  1535. ax.plot(mo.neurons_pref_stims, sigmars, c=c, ls='--', zorder=100)
  1536. ax.set_ylim(0, 25)
  1537. ax.set_xlim(-xmax, xmax)
  1538. ax.set_xlabel('Feedforward preferred stimulus $s_i$')
  1539. ax.set_ylabel('Effective width $\\sigma_r(s_i)$')
  1540. ax.set_title('Effective width $\\sigma_r(s_i)$')
  1541. ax.legend()
  1542. # MSE
  1543. ax = subaxs[ii]
  1544. xmin = MID_VAL-.5*prior_w; xmax = MID_VAL+.5*prior_w
  1545. ind = (STIMULI>=xmin) & (STIMULI<=xmax)
  1546. these_stims = STIMULI[ind]
  1547. mse_per_stim = mse_per_stim_perprior[prior_w]
  1548. ###
  1549. approx_mse_per_stim = 1./(1./prior_var + (1./mo.beta)*(mo.r_star.T/true_sigmarsSq).sum(axis=1))
  1550. ###
  1551. ax.plot(these_stims, np.sqrt(mse_per_stim), c=c, label='$\\sqrt{MSE}$')
  1552. ax.plot(these_stims, np.sqrt(approx_mse_per_stim[ind]), c=c, ls='--', label='Approximº', lw=2)
  1553. ax.set_xticks([-prior_w/2, 0, prior_w/2])
  1554. ax.set_ylim(0, 11.9)
  1555. ax.spines['top'].set_visible(False)
  1556. ax.spines['right'].set_visible(False)
  1557. if ii > 0:
  1558. ax.set_yticks([])
  1559. subaxs[0].set_ylabel('$\\sqrt{MSE}$')
  1560. subaxs[1].set_xlabel('Stimulus $s$')
  1561. subaxs[1].set_title('Square-root of MSE')
  1562. subaxs[0].legend(loc=(.05, .8))
  1563. leg = axs[0,0].legend()
  1564. axs[0,0].legend( handles=[Line2D([0], [0], color='grey', lw=1, label='$\\varphi(s_i)$'),
  1565. Line2D([0], [0], color='grey', lw=2, ls='--', label='Approximation $\\sum \\gamma_{ij} s_j$')], loc='lower right')
  1566. axs[0,0].add_artist(leg)
  1567. for ax in axs.flat:
  1568. ax.spines['top'].set_visible(False)
  1569. ax.spines['right'].set_visible(False)
  1570. fig.subplots_adjust(hspace=.35)
  1571. #fig.savefig('approx_rates_mses.pdf', bbox_inches='tight')
  1572. # %%
  1573. # %%

model.ipynb, no license · at the source

Overview

Authors: Arthur Prat-Carrabin1, Maximilian V. Harl2,3, Samuel J. Gershman1
  1. Department of Psychology and Center for Brain Science, Harvard University,Cambridge, MA USA
  2. Neuroscience Center Zurich, University of Zurich and ETH Zurich,Zurich, Switzerland
  3. TUM School of Computation, Information and Technology, Technical University of Munich,Garching, Germany
Institutions: Harvard University (United States); University of Zurich (Switzerland); Technical University of Munich (Germany)
Journal: Nature communications, volume 17, issue 1, article 5554
Dates: received 18 July 2025; accepted 28 April 2026; published online 15 May 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1038/s41467-026-73032-0 · PMID 42140911 · PMCID PMC13291360 · OpenAlex W4412872186
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: computational modeling (no new data) (modality), human (organism)
Methods: Single-unit activity, calcium imaging, Machine learning
Keywords: Neural encoding, Network models, Perception, Sensory processing
MeSH: Adaptation, Physiological*, Models, Neurological*, Nerve Net*, Neurons*, Action Potentials, Animals, Humans, Recurrent Neural Networks (* major topic)
Topic: CCD and CMOS Imaging Sensors (Electrical and Electronic Engineering, Engineering), according to OpenAlex
Funding: Kempner Institute for the Study of Natural and Artificial Intelligence Polymath Award from Schmidt Sciences
Citations: not cited yet (Europe PMC); 75 references in the paper

Abstract

The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.

Repository

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

OSF wa7hm

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Languages: Jupyter (1)
Size: 3 files, 1 script
Software Heritage: not checked
Found in: “Code availability”
Holds: README, environment (requirements.txt), 1 notebook
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: Matplotlib (1 file), NumPy (1 file), PyTorch (1 file), SciPy (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers (HTTP 200)
  • 28 September 2026: the link answers (HTTP 200)
2 files
At the source: osf.io/wa7hm/

Code availability statement

The paper has a code availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41467-026-73032-0.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 1 script, each with its path and the digest of its content;
  • 4 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

Code and data availability statement

The paper has a code and data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41467-026-73032-0.

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, 28 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 3 authors, 4 keywords, 8 MeSH terms, 1 funder, 65 references.

Cite

This paper

Prat-Carrabin, A., Harl, M. V., & Gershman, S. J. (2026). Fast efficient coding and sensory adaptation in gain-adaptive recurrent networks. Nature communications, 17(1), 5554. https://doi.org/10.1038/s41467-026-73032-0

BibTeX

@article{pratcarrabin2026fast,
author = {Prat-Carrabin, Arthur and Harl, Maximilian V. and Gershman, Samuel J.},
title = {{Fast efficient coding and sensory adaptation in gain-adaptive recurrent networks}},
journal = {Nature communications},
year = {2026},
month = may,
volume = {17},
number = {1},
pages = {5554},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/s41467-026-73032-0},
url = {https://doi.org/10.1038/s41467-026-73032-0},
pmid = {42140911},
pmcid = {PMC13291360}
}

RIS

TY - JOUR
AU - Prat-Carrabin, Arthur
AU - Harl, Maximilian V.
AU - Gershman, Samuel J.
TI - Fast efficient coding and sensory adaptation in gain-adaptive recurrent networks
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/05/15
VL - 17
IS - 1
SP - 5554
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/s41467-026-73032-0
UR - https://doi.org/10.1038/s41467-026-73032-0
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41467-026-73032-0",
"type": "article-journal",
"title": "Fast efficient coding and sensory adaptation in gain-adaptive recurrent networks",
"container-title": "Nature communications",
"author": [
{
"family": "Prat-Carrabin",
"given": "Arthur"
},
{
"family": "Harl",
"given": "Maximilian V."
},
{
"family": "Gershman",
"given": "Samuel J."
}
],
"container-title-short": "Nat Commun",
"volume": "17",
"issue": "1",
"page": "5554",
"DOI": "10.1038/s41467-026-73032-0",
"PMID": "42140911",
"PMCID": "PMC13291360",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41467-026-73032-0",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
15
]
]
}
}

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.7554/elife.101277 [code]
Endogenous precision of the number sense.
Journal: eLife
In common: SciPy, Matplotlib, NumPy, 11 references, author Arthur Prat-Carrabin
[2] doi:10.1126/sciadv.aeg3535 [code]
Undoing of firing rate adaptation enables invariant population codes.
Journal: Science advances
In common: SciPy, Matplotlib, NumPy, 5 references
[3] doi:10.7554/elife.110685 [code]
Sensory adaptation and pupil-linked arousal support flexible evidence accumulation during perceptual decision making.
Journal: eLife
In common: 6 references
[4] doi:10.1126/sciadv.aed4172 [code]
Locomotion optimizes sensory representations through a computational principle shared by rodents and primates.
Journal: Science advances
In common: PyTorch, SciPy, Matplotlib, 1 other tool, 3 references
[5] doi:10.1038/s41593-026-02255-7 [code]
Neural circuits encode prior knowledge of temporal statistics.
Journal: Nature neuroscience
In common: SciPy, Matplotlib, NumPy, 3 references
[6] doi:10.1371/journal.pbio.3003856 [code]
Aging and metabolism contribute separately to brain-body health.
Journal: PLoS biology
In common: PyTorch, SciPy, Matplotlib, 1 other tool, 3 references
[7] doi:10.1038/s41597-026-07077-7 [code]
Everyday Activity Science and Engineering Table Setting Dataset.
Journal: Scientific data
In common: PyTorch, SciPy, Matplotlib, 1 other tool, 3 references
[8] doi:10.1371/journal.pcbi.1013441 [code]
Large vision model framework for automated C. elegans analysis: From static morphometry to dynamic neural activity.
Journal: PLoS computational biology
In common: PyTorch, SciPy, Matplotlib, 1 other tool, 3 references
[9] doi:10.3389/fdgth.2026.1691088 [code]
Real-world federated learning for brain imaging scientists.
Journal: Frontiers in digital health
In common: PyTorch, SciPy, Matplotlib, 1 other tool, 3 references
[10] doi:10.1038/s41586-026-10528-1 [code]
A critical initialization for biological neural networks.
Journal: Nature
In common: PyTorch, SciPy, Matplotlib, 1 other tool, computational modeling (no new data), 2 references

Contribute

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

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

Request its removal

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

Discussion, reproductions, activity

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

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

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