Fast efficient coding and sensory adaptation in gain-adaptive recurrent networks.
The 4 matches
- [1] § Methods › Network optimization and simulations ↔ model.ipynb, lines 424–483 · score 0.75 · ReduceLROnPlateau, learning rate schedule, Adam, threshold, gradient, patience
- [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] § 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] § 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
- # %%
- import numpy as np
- import torch
- import itertools
- # from utils import gaussian, probs_mean_and_var, sample_from_poisson
- from matplotlib.pyplot import *
- from matplotlib.colors import to_hex
- from matplotlib.lines import Line2D
- from scipy import integrate, stats, optimize, special
- # %%
- from scipy.interpolate import interp1d
- # %%
- import sys
- print("Python version:", sys.version)
- print("Version info:", sys.version_info)
- # %%
- !pip list
- # %%
- !pip freeze
- # %%
- from cachetools import cached
- # %%
- torch.set_printoptions(linewidth=150, )
- torch.autograd.set_detect_anomaly(True);
- # %%
- # Max Range
- LOWER_BOUND = -200
- UPPER_BOUND = 200
- OVERALL_WIDTH = UPPER_BOUND - LOWER_BOUND
- MID_VAL = (UPPER_BOUND + LOWER_BOUND) / 2
- # Stimuli
- NUM_STIMULI = UPPER_BOUND - LOWER_BOUND + 1
- STIMULI = torch.linspace(LOWER_BOUND, UPPER_BOUND, NUM_STIMULI)
- #
- SQRT_2_PI = torch.sqrt(torch.tensor(2 * torch.pi))
- #
- COLORS_WIDTH = {10:'C2', 20:'C0', 30:'C1', 8:cm.tab20c(9), 6:cm.tab20c(9),
- 5:cm.tab20c(12), 4:cm.tab20(12), 2:cm.tab20(12), 1:cm.tab20(7), 0:cm.tab20(13),
- 10000:'C0' } # 6:cm.tab20c(12),
- # %% [markdown]
- # # Useful functions
- # %%
- # Utility functions for optimization
- def get_best_res(*args):
- best_res = optimize.OptimizeResult( fun=np.inf )
- for res in args:
- if res is None:
- continue
- if res.fun < best_res.fun:
- best_res = res
- return best_res
- def min_objf(f, x0s, v=True, return_all_res=False, **kwargs):
- all_res = []
- for x0 in x0s:
- res = optimize.minimize(f, x0=x0, **kwargs)
- if v: print(x0, '->', res.x, res.fun)
- all_res.append(res)
- best_res = get_best_res(all_res)
- if v: print('=>', best_res.x, best_res.fun)
- if return_all_res:
- return best_res, all_res
- return best_res
- def min_objf_grid(objf, grid_vals, v=True, finish=True, compare_with_res=None, objf_args=(), x0s=None, **kwargs):
- if v:
- ns = [ len(vals) for vals in grid_vals ]
- print(' x '.join([str(n) for n in ns]), ' = ', np.prod( ns ) )
- ### Method 1: grid search
- rez = []
- for x0 in itertools.product(*grid_vals):
- fun = objf(x0, *objf_args)
- rez.append( optimize.OptimizeResult( x=x0, fun=fun, success=False ) )
- res = get_best_res(*rez)
- if v: print('Grid', res.x, res.fun, flush=True)
- rez = [res,]
- # Optimization method(s)
- methods = [None,]
- if 'method' in kwargs:
- method = kwargs.pop('method')
- if method == 'all':
- methods = ['Nelder-Mead', 'Powell', 'L-BFGS-B',] # 'SLSQP']
- elif isinstance(method, list):
- methods = method
- else:
- methods = [method,]
- # Finish after grid
- if finish:
- for method in methods:
- res_m = optimize.minimize(objf, x0=res.x, method=method, args=objf_args, **kwargs )
- if v: print(('-' if method is None else method).ljust(12), res_m.x, res_m.fun, flush=True)
- rez.append(res_m)
- # List of start points
- if x0s is not None:
- for method in methods:
- res_x0s = min_objf(objf, x0s, v=v, args=objf_args, method=method, **kwargs)
- if v: print(('-' if method is None else method).ljust(12), res_x0s.x, res_x0s.fun, flush=True)
- rez.append(res_x0s)
- # Compare
- if compare_with_res is not None:
- rez.append( compare_with_res )
- res = get_best_res(*rez)
- return res
- # %%
- def values_to_indices(values, size=len(STIMULI), min_val=LOWER_BOUND, max_val=UPPER_BOUND ):
- """Convert values in [min_val, max_val] to indices in a vector of length `size`."""
- values = np.asarray(values)
- step = (max_val - min_val) / (size - 1)
- indices = np.round((values - min_val) / step).astype(int)
- return indices
- # %%
- def make_gamma_G_φs(gains, hh, stims):
- '''Given a vector of gains, a matrix hh of the distance function h(s_i-s_j), and a vector of feedforward stimuli,
- return the weights gamma_ij, the denominator G, and the effective locations φ_s.'''
- gamma_num = torch.diag(gains) + hh * gains.unsqueeze(0) # shape = (num_neurons, num_neurons)
- gamma_den = gamma_num.sum(axis=1) # shape = (num_neurons,) ### This is G(s)
- gammas = gamma_num / gamma_den.unsqueeze(1) # shape = (num_neurons, num_neurons)
- φs = ( gammas * stims ).sum(axis=1) # (num_neurons,)
- return gammas, gamma_den, φs
- # %%
- def effective_tc_center_and_squaredwidth_given_rates(rates):
- '''Returns the effective locations φ_s and squared widths σ_r^2.'''
- denom = rates.sum(axis=1)
- true_φs = ( rates * STIMULI ).sum(axis=1) / denom
- true_sigmarsSq = ( rates * (STIMULI-true_φs.unsqueeze(1))**2 ).sum(axis=1) / denom
- return true_φs, true_sigmarsSq
- # %%
- def make_recurrent_weights(stims, rec_width, rec_magnitude, inhib):
- if rec_magnitude == 0:
- return torch.zeros((len(stims), len(stims)))
- distances = torch.abs(stims.unsqueeze(0) - stims.unsqueeze(1))
- gaussian_kernel = torch.exp(-distances ** 2 / (2 * rec_width ** 2))
- recurrent_weights = gaussian_kernel #+ uniform_density
- eigvals = torch.linalg.eigvals(recurrent_weights)
- max_eigenvalue = eigvals.abs().max()
- scaling_factor = rec_magnitude / max_eigenvalue
- recurrent_weights = recurrent_weights * scaling_factor
- if inhib > 0: # uniform inhibition
- recurrent_weights -= (inhib/len(stims)) * torch.ones_like(recurrent_weights)
- return recurrent_weights
- # %%
- def second_derivative_smoothness(g, dx, normalized=False):
- """Computes second derivative penalty for smoothness regularization"""
- gg = (g/g.mean()) if normalized else g
- g_i_minus_1 = gg[:-2] # g[i-1]
- g_i = gg[1:-1] # g[i]
- g_i_plus_1 = gg[2:] # g[i+1]
- second_deriv = ( g_i_plus_1 - 2 * g_i + g_i_minus_1 ) / dx**2
- return (second_deriv**2).mean() # Mean squared second derivative
- # %%
- def make_Fa_and_Fam1_fcts(prior_xs, prior_ys, a, min_x=LOWER_BOUND, max_x=UPPER_BOUND):
- """
- Given a pdf f(x) and an exponent a, F_a(x) is defined as
- F_a(x) = ( ∫[min_x,x] f(t)^a dt ) / ( ∫[min_x,max_x] f(t)^a dt ).
- If a=1 it is the usual CDF. This returns two functions, corresponding to F_a(x) and its inverse F_a⁻¹(y).
- The discretized support prior_xs is assumed to represent the centers of bins covering the interval.
- The functions obey:
- - F_a(prior_xs[0]) = 0 and F_a(prior_xs[-1]) = 1.
- - For any x <= min_x, fct(x) = 0.
- - For any x >= max_x, fct(x) = 1.
- min_x and max_x serve as the theoretical integration bounds.
- """
- # Ensure that the sampled values lie within the prescribed support.
- assert all(prior_xs >= min_x)
- assert all(prior_xs <= max_x)
- # Evaluate the powered pdf
- prior_ys_pow_a = prior_ys**a
- # Compute the cumulative integral using the trapezoidal rule.
- # 'initial=0' ensures that the cumulative value at the first point is 0.
- num = integrate.cumulative_trapezoid(prior_ys_pow_a, prior_xs, initial=0)
- # Normalize so that the cumulative value at prior_xs[-1] becomes 1.
- num /= num[-1]
- # Create extended arrays that include the full bounds if needed.
- xs_full = np.asarray( prior_xs ).copy()
- nums_full = num.copy()
- # If the left end of the discretization is strictly inside [min_x, max_x], insert the minimum.
- if xs_full[0] > min_x:
- xs_full = np.insert(xs_full, 0, min_x)
- nums_full = np.insert(nums_full, 0, 0)
- # Likewise, if the right end does not reach max_x, append the maximum.
- if xs_full[-1] < max_x:
- xs_full = np.append(xs_full, max_x)
- nums_full = np.append(nums_full, 1)
- # Define the F_a function: for any input x, linearly interpolate the integration values,
- # forcing x <= min_x to return 0 and x >= max_x to return 1.
- def fct(x):
- return np.interp(x, xs_full, nums_full, left=0, right=1)
- # The inverse F_a⁻¹ is defined by inverting the above mapping.
- def fct_m1(y):
- return np.interp(y, nums_full, xs_full, left=min_x, right=max_x)
- return fct, fct_m1
- # %%
- def gaussian(x, mu, sigma, incl_norm=False):
- if incl_norm:
- return (1 / (torch.sqrt(2 * torch.pi * sigma**2))) * torch.exp(-0.5 * ((x - mu) / sigma) ** 2)
- else:
- return torch.exp(-0.5 * ((x - mu) / sigma) ** 2)
- # %%
- def probs_mean_and_var(pxs, xs):
- mu = torch.sum(pxs * xs)
- v = torch.sum(pxs * (xs - mu) ** 2)
- return mu, v
- # %%
- def solve_lyapunov(rates: torch.Tensor, matrix_W: torch.Tensor) -> torch.Tensor:
- """
- Solve (W - I) Σ + Σ (W - I) + Γ = 0 with Γ = diag(rates).
- Args
- ----
- rates : (N,) tensor
- matrix_W : (N, N) tensor (assumed symmetric)
- Returns
- -------
- Sigma : (N, N) tensor
- """
- # Make sure everything is on same device/dtype
- matrix_W = matrix_W
- rates = rates.to(matrix_W)
- n = matrix_W.shape[0]
- I = torch.eye(n, dtype=matrix_W.dtype, device=matrix_W.device)
- A = matrix_W - I
- Gamma = torch.diag(rates)
- # Eigen-decomposition of A (symmetric)
- evals, evecs = torch.linalg.eigh(A) # A = evecs @ diag(evals) @ evecs.T
- # Transform Gamma into eigenbasis: B = U^T Γ U
- B = evecs.T @ Gamma @ evecs
- # Solve elementwise: (λ_i + λ_j) Y_ij + B_ij = 0 => Y_ij = -B_ij / (λ_i + λ_j)
- lam_i = evals.unsqueeze(0) # (1, N)
- lam_j = evals.unsqueeze(1) # (N, 1)
- denom = lam_i + lam_j
- # In a well-posed Lyapunov problem denom != 0; add tiny epsilon for numerical safety.
- eps = 1e-12
- denom = denom + eps * (denom == 0)
- Y = -B / denom
- # Transform back: Σ = U Y U^T
- Sigma = evecs @ Y @ evecs.T
- # Enforce symmetry numerically
- Sigma = 0.5 * (Sigma + Sigma.T)
- return Sigma
- # %%
- def lyapunov_approx_solution(rates: torch.Tensor, matrix_M: torch.Tensor) -> torch.Tensor:
- rates = rates.to(matrix_M)
- Gamma = torch.diag(rates)
- Sigma = 0.25 * (matrix_M @ Gamma + Gamma @ matrix_M)
- Sigma = 0.5 * (Sigma + Sigma.T)
- return Sigma
- # %%
- def lyapunov_diagonal_solution(rates: torch.Tensor, h0) -> torch.Tensor:
- return .5 * (1 + h0) * rates
- # %%
- def sample_from_poisson(rate_parameters, num_samples=1000):
- samples = np.random.poisson(rate_parameters, (num_samples, *rate_parameters.shape))
- return torch.tensor(samples)
- def sample_from_gaussian_rates_then_poisson(base_rates, matrix_M, num_samples=100, xbounds=None):
- '''For the stimuli in xbounds, sample spike counts from a two-step process:
- 1) sample rates from a Gaussian with mean=base_rates and covariance from Lyapunov approx.
- 2) sample spike counts from Poisson with these rates.
- The output shape is (num_samples, num_neurons, num_stimuli_in_xbounds), where num_neurons = base_rates.shape[0].'''
- rng = np.random.default_rng()
- if xbounds is None: # then it's the default, STIMULI
- indices = np.arange(STIMULI.shape[0])
- else:
- xmin, xmax = xbounds
- idx_min, idx_max = values_to_indices([xmin, xmax])
- indices = np.arange(idx_min, idx_max+1)
- samples = torch.zeros((num_samples, base_rates.shape[0], len(indices)) )
- for i, idx in enumerate(indices): # for each stimulus
- this_base_rates = base_rates[:,idx] # shape = (num_neurons,)
- Sigma = lyapunov_approx_solution( rates = this_base_rates, matrix_M=matrix_M )
- # Sample from Gaussian with mean=this_base_rates and covariance=Sigma
- gaussian_samples = rng.multivariate_normal(mean=this_base_rates, cov=Sigma, size=num_samples)
- # Rectify to non-negative rates
- gaussian_samples = np.clip(gaussian_samples, a_min=0, a_max=None)
- # Sample from Poisson with these rates
- samples[:,:,i] = sample_from_poisson(gaussian_samples, 1)
- return samples
- # %% [markdown]
- # # Network class
- # %%
- class Network:
- def __init__(self, num_neurons, ff_sigma, rec_magnitude, rec_width, inhib=0.):
- # neurons ff preferred stimuli
- self.num_neurons = num_neurons
- self.neurons_pref_stims = torch.linspace(LOWER_BOUND, UPPER_BOUND, self.num_neurons)
- self.δ = OVERALL_WIDTH / self.num_neurons
- # feedforward tuning curves
- self.ff_sigma = ff_sigma
- self.ff_tuning_curves = gaussian( self.neurons_pref_stims.unsqueeze(1), STIMULI, self.ff_sigma ) # (num_neurons, NUM_STIMULI)
- # recurrence
- self.rec_magnitude = rec_magnitude
- self.rec_width = rec_width
- self.inhib = inhib
- self.update_W_and_M()
- self.c = self.ff_sigma * np.sqrt(2*np.pi) / ( self.δ * (1 - self.rec_magnitude) )
- # gains
- self.gains = torch.ones(self.num_neurons)
- self.update_r_star()
- def __repr__(self):
- return f"Network(N={self.num_neurons}, ffσ={self.ff_sigma}, rec_mag={self.rec_magnitude}, rec_w={self.rec_width}, inhib={self.inhib})"
- def update_W_and_M(self):
- '''Also updates νrec, hh, hbar, sigma_h_Sq'''
- self.recurrent_weights = make_recurrent_weights(self.neurons_pref_stims, self.rec_width, self.rec_magnitude, self.inhib) # shape = (num_neurons, num_neurons)
- if self.rec_magnitude > 0:
- self.matrix_M = torch.inverse(torch.eye(self.num_neurons) - self.recurrent_weights) # M = (I - W)^-1 # (num_neurons, num_neurons)
- self.νrec = self.rec_width / np.sqrt(2*np.log(1./self.rec_magnitude))
- si_minus_sj = self.neurons_pref_stims.unsqueeze(0)-self.neurons_pref_stims.unsqueeze(1)
- 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) )
- 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)
- self.hbar = (1-self.rec_magnitude)*self.rec_magnitude + (self.rec_magnitude/np.log(1./self.rec_magnitude))
- 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
- self.mu4 = 3. * self.rec_width**4 * self.rec_magnitude * ( 1 - self.rec_magnitude + 2/ (np.log(1./self.rec_magnitude))**3 )
- else:
- self.matrix_M = torch.eye(self.num_neurons)
- self.νrec = 0
- self.hh = torch.zeros((self.num_neurons, self.num_neurons))
- self.hbar = 0
- self.sigma_h_Sq = 0
- self.mu4 = 0
- self.sigma_h_Sq_tilde = self.sigma_h_Sq / (1.+self.hbar)
- self.mu4_tilde = self.mu4 / (1.+self.hbar)
- self.beta = 1 + .5*( 1. + self.hh[0,0] )
- def compute_r_star(self, gains=None):
- '''"r_star" is just r(s) in the paper, here computed as (I - W)^-1 x (gains * ff_tuning_curves).'''
- modulated_ff_tuning_curves = (self.gains if gains is None else gains).unsqueeze(1) * self.ff_tuning_curves # shape = (num_neurons, NUM_STIMULI)
- r_lin = torch.matmul(self.matrix_M, modulated_ff_tuning_curves) # shape = (num_neurons, NUM_STIMULI)
- if self.inhib == 0.: # no inhibition, linear solution
- return r_lin, modulated_ff_tuning_curves
- # if self.inhib > 0.: # with inhibition
- r_star = torch.clamp(r_lin, min=0.0)
- for _ in range(10):
- r_star = torch.clamp( modulated_ff_tuning_curves + torch.matmul(self.recurrent_weights, r_star), min=0.0 )
- return r_star, modulated_ff_tuning_curves
- def update_r_star(self):
- self.r_star, self.modulated_ff_tuning_curves = self.compute_r_star(self.gains)
- def fisher_info(self, gains=None):
- r_star, _ = self.compute_r_star(gains)
- rprime = torch.gradient(r_star, dim=1, spacing=[STIMULI,], edge_order=2)[0]
- FI = torch.zeros_like(r_star)
- mask = r_star > 1e-30
- FI[mask] = (rprime[mask]**2) / r_star[mask]
- FItot = FI.sum(axis=0) # shape = (NUM_STIMULI,)
- return FItot / self.beta
- def effective_tc_center_and_squaredwidth(self):
- '''Returns the effective locations φ_s and squared widths σ_r^2.'''
- return effective_tc_center_and_squaredwidth_given_rates(self.r_star)
- def spike_cost(self, rates, prior_stims, spikes_power):
- powered_exp_spikes_for_each_stimulus = (rates**spikes_power).sum(axis=0)
- return (prior_stims * powered_exp_spikes_for_each_stimulus).sum()
- def expSqE(self, rates, prior_stims, prior_var=None):
- if prior_var is None:
- _, prior_var = probs_mean_and_var(prior_stims, STIMULI)
- _, true_sigmarsSq = effective_tc_center_and_squaredwidth_given_rates(rates)
- expectedSqErr = (prior_stims/(1./prior_var + (1./self.beta)*(rates.T/true_sigmarsSq).sum(axis=1)) ).sum()
- return expectedSqErr
- def expSqE_v2_I(self, rates, prior_stims, prior_var=None):
- # version with I(s) = (1/β) sum_i r_i(s) (s-φ_i)^2 / σ_r_i^4
- if prior_var is None:
- _, prior_var = probs_mean_and_var(prior_stims, STIMULI)
- true_φs, true_sigmarsSq = effective_tc_center_and_squaredwidth_given_rates(rates) # true_sigmarsSq.shape = (num_neurons,)
- fi = (1./self.beta)*(rates.T * (STIMULI.unsqueeze(1)-true_φs)**2/true_sigmarsSq**2).sum(axis=1)
- expectedSqErr = (prior_stims/(1./prior_var + fi ) ).sum()
- return expectedSqErr
- 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,
- regularizer_lambda=0., return_losses=False, eps_rel=0.):
- if start_gains is None: start_gains = prior_neurons
- N = start_gains.numel()
- patience_lr = patience_lr if patience_lr is not None else patience // 2
- log_gains = torch.log(start_gains).detach().requires_grad_(True)
- optimizer = torch.optim.Adam([log_gains], lr=lr)
- scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=patience_lr, factor=0.5, threshold=1e-4 if eps_rel==0. else eps_rel)
- _, prior_var = probs_mean_and_var(prior_stims, STIMULI)
- this_prior = prior_neurons if self.expSqE_kind == 'v2_g' else prior_stims
- rates, _ = self.compute_r_star(start_gains)
- best_expectedSqErr = self.expSqE(rates, this_prior, prior_var)
- best_cost = self.spike_cost(rates, prior_stims, spikes_power=1.)
- best_val_loss_no_reg = best_expectedSqErr + alpha_cost * best_cost
- best_val_loss = best_val_loss_no_reg + regularizer_lambda * second_derivative_smoothness(start_gains, self.δ, normalized=True)
- best_gains = start_gains.clone().detach()
- patience_counter = 0
- now = datetime.datetime.now()
- if v:
- print('Time\t\tIter\t<g>\t\tExpSqErr\tLoss\tpatience')
- for i in range(epochs):
- optimizer.zero_grad() # Reset gradients
- gains = torch.exp(log_gains)
- rates, _ = self.compute_r_star(gains)
- cost = self.spike_cost(rates, prior_stims, spikes_power=1.)
- expectedSqErr = self.expSqE(rates, this_prior, prior_var)
- loss_no_reg = expectedSqErr + alpha_cost*cost
- loss = loss_no_reg + regularizer_lambda * second_derivative_smoothness(gains, self.δ, normalized=True)
- loss.backward() # Backpropagate to compute gradients
- optimizer.step() # Update log_gains
- val_loss_no_reg = loss_no_reg.item()
- val_loss = loss.item()
- scheduler.step(val_loss)
- improved = (best_val_loss - val_loss) > max(0., eps_rel * abs(best_val_loss))
- if improved:
- best_val_loss_no_reg = val_loss_no_reg
- best_val_loss = val_loss
- best_gains = gains.clone().detach()
- best_cost = cost.item()
- best_expectedSqErr = expectedSqErr.item()
- patience_counter = 0
- else:
- patience_counter += 1
- if v and i%500 == 0: #
- 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}')
- if (patience_counter >= patience):
- print(f'Patience reached ({patience_counter}). Stopping optimization.')
- break
- gains = best_gains
- if v:
- 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}')
- self.gains = gains
- self.update_r_star()
- if return_losses:
- return gains, best_val_loss_no_reg, best_val_loss
- return gains
- def optimize_gains_quadratic_g0Delta_expSqErr_wCost_simple(self, prior_stims, alpha_cost, v=True):
- gains = torch.zeros(self.num_neurons)
- prior_mean, prior_var = probs_mean_and_var(prior_stims, STIMULI)
- prior_sd = torch.sqrt(prior_var)
- z = (self.neurons_pref_stims-prior_mean)/prior_sd
- zSq = z**2
- gains = torch.zeros_like(z)
- def objf(x):
- g0, Delta = x
- ind = z.abs() < Delta
- gains[:] = 0.
- gains[ind] = g0 * (1 - zSq[ind]/Delta**2)
- rates, _ = self.compute_r_star(gains)
- cost = self.spike_cost(rates, prior_stims, spikes_power=1.)
- expectedSqErr = self.expSqE(rates, prior_stims, prior_var)
- loss = expectedSqErr + alpha_cost*cost
- #print(x, loss)
- return loss
- 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)
- g0, Delta = res.x
- ind = z.abs() < Delta
- gains[:] = 0.
- gains[ind] = g0 * (1 - zSq[ind]/Delta**2)
- self.gains = gains
- self.update_r_star()
- return g0, Delta, res.fun
- def rec_effective_tc_width(self):
- return np.sqrt(self.ff_sigma**2 + self.sigma_h_Sq_tilde)
- def sample_spike_counts_nocovariance(self, num_samples, xbounds=None):
- if xbounds is None: # then it's the default, STIMULI
- return sample_from_poisson(self.r_star, num_samples) # r_star.shape = (num_neurons, NUM_STIMULI)
- else:
- xmin, xmax = xbounds
- ind = (STIMULI >= xmin) & (STIMULI <= xmax)
- 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)
- def sample_spike_counts(self, num_samples, xbounds=None):
- return sample_from_gaussian_rates_then_poisson(self.r_star, self.matrix_M, num_samples=num_samples, xbounds=xbounds)
- def simulate_recurrent_poisson_many(self, num_samples, T=1.0, tau=0.01, dt=None, xbounds=None, burn_in=0.0):
- """
- Recurrent Poisson simulation for many stimuli and many trials.
- Returns spike counts summed over time (after burn-in), shape:
- (num_samples, num_neurons, num_stimuli_in_xbounds)
- """
- if xbounds is None:
- indices = torch.arange(STIMULI.shape[0])
- else:
- xmin, xmax = xbounds
- idx_min, idx_max = values_to_indices([xmin, xmax])
- indices = torch.arange(idx_min, idx_max + 1)
- if dt is None:
- dt = tau / 10.0
- steps = int(T / dt)
- burn_steps = int(burn_in / dt)
- total_steps = burn_steps + steps
- samples = torch.zeros((num_samples, self.num_neurons, len(indices)) )
- for j, idx in enumerate(indices):
- ff_tc = self.ff_tuning_curves[:, idx]
- ff = self.gains * ff_tc # (num_neurons,)
- r = ff.unsqueeze(0).expand(num_samples, -1).clone() # start at ff rates
- counts = torch.zeros((num_samples, self.num_neurons) )
- for t in range(total_steps):
- lam = torch.clamp(r, min=0.) * dt
- k = torch.poisson(lam) # float tensor
- if t >= burn_steps:
- counts += k
- r = r + (dt / tau) * (-r + ff.unsqueeze(0) + r @ self.recurrent_weights.T) + (1.0 / tau) * (k - lam)
- samples[:, :, j] = counts
- return samples
- def decode(self, spike_counts, prior_mean, prior_var):
- '''spike_counts.shape = (num_samples, num_neurons, num_stimuli).
- Returns estimates of the stimuli, shape = (num_samples, num_stimuli)'''
- assert spike_counts.shape[1] == self.num_neurons
- corrected_spike_counts = (1./self.beta)*spike_counts
- #assert spike_counts.shape[2] == NUM_STIMULI
- φs, sigmarsSq = self.effective_tc_center_and_squaredwidth()
- oo_sigmaSq = 1./sigmarsSq
- # spike_counts has shape (num_samples, num_neurons, num_stimuli), and oo_sigmaSq has shape (num_neurons,).
- denominator = 1./prior_var + torch.einsum('ijk,j->ik', corrected_spike_counts, oo_sigmaSq)
- # spike_counts has shape (num_samples, num_neurons, num_stimuli), and φs*oo_sigmaSq has shape (num_neurons,).
- numerator = prior_mean/prior_var + torch.einsum('ijk,j->ik', corrected_spike_counts, φs*oo_sigmaSq)
- return numerator / denominator
- def posterior_precision(self, spike_counts, prior_var):
- '''spike_counts.shape = (num_samples, num_neurons, num_stimuli).
- Returns the posterior precision (1/variance) for each estimate, shape = (num_samples, num_stimuli)'''
- assert spike_counts.shape[1] == self.num_neurons
- corrected_spike_counts = (1./self.beta)*spike_counts
- #assert spike_counts.shape[2] == NUM_STIMULI
- _, sigmarsSq = self.effective_tc_center_and_squaredwidth()
- oo_sigmaSq = 1./sigmarsSq
- # spike_counts has shape (num_samples, num_neurons, num_stimuli), and oo_sigmaSq has shape (num_neurons,).
- precision = 1./prior_var + torch.einsum('ijk,j->ik', corrected_spike_counts, oo_sigmaSq)
- return precision
- # %%
- def make_gaussian_priors(pref_stims, sd, m=MID_VAL):
- '''This returns two arrays, one for the STIMULI and one for "pref_stims" (e.g., the neurons' preferred stimuli).'''
- # special case: point mass at closest stimulus
- if sd == 0:
- # for STIMULI
- idx0 = torch.argmin((STIMULI - m).abs())
- prior_stimuli = torch.zeros_like(STIMULI)
- prior_stimuli[idx0] = 1.
- # for pref_stims
- idx1 = torch.argmin((pref_stims - m).abs())
- prior_pref_stims = torch.zeros_like(pref_stims)
- prior_pref_stims[idx1] = 1.
- return prior_stimuli, prior_pref_stims
- prior_stimuli = gaussian( STIMULI, m, sd )
- prior_stimuli /= prior_stimuli.sum()
- prior_pref_stims = gaussian( pref_stims, m, sd )
- prior_pref_stims /= prior_pref_stims.sum()
- return prior_stimuli, prior_pref_stims
- # %%
- def make_mixture_uniform_gaussian_priors(pref_stims, sd, m=MID_VAL, uniform_part=0.1):
- '''This returns two arrays, one for the STIMULI and one for "pref_stims" (e.g., the neurons' preferred stimuli).'''
- gaussian_stimuli, gaussian_pref_stims = make_gaussian_priors(pref_stims, sd, m)
- uniform_stimuli = torch.ones_like(STIMULI) / NUM_STIMULI
- uniform_pref_stims = torch.ones_like(pref_stims) / len(pref_stims)
- prior_stimuli = (1-uniform_part)*gaussian_stimuli + uniform_part*uniform_stimuli
- prior_pref_stims = (1-uniform_part)*gaussian_pref_stims + uniform_part*uniform_pref_stims
- return prior_stimuli, prior_pref_stims
- # %% [markdown]
- # # Optimizing gains to minimize expected squared error
- # %%
- #CACHE_find_optimal_gains_for_parameters = {}
- # np.savez_compressed('CACHE_find_optimal_gains_for_parameters.npz', cache=CACHE_find_optimal_gains_for_parameters)
- CACHE_find_optimal_gains_for_parameters = np.load('CACHE_find_optimal_gains_for_parameters.npz', allow_pickle=True)['cache'].item()
- len(CACHE_find_optimal_gains_for_parameters)
- # %%
- @cached(cache=CACHE_find_optimal_gains_for_parameters)
- def find_optimal_gains_for_parameters(num_neurons, ff_sigma, rec_magnitude, rec_width, alpha_cost, prior_w, regularizer_lambda=1000, inhib=0.):
- print( f'prior width = {prior_w}' )
- mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width, inhib=inhib)
- a_prior, a_prior_neurons = make_gaussian_priors(mo.neurons_pref_stims, prior_w)
- if inhib != 0.:
- print('Solution with inhib = 0 ...')
- gains_0, _, _, _ = find_optimal_gains_for_parameters(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width,
- alpha_cost=alpha_cost, prior_w=prior_w, regularizer_lambda=regularizer_lambda, inhib=0.)
- print('Solution with inhib = 0 ... OK')
- else:
- mo.optimize_gains_quadratic_g0Delta_expSqErr_wCost_simple(a_prior, alpha_cost=alpha_cost, v=True)
- gains_0 = mo.gains.clone()
- # smoothen the initial gain solution
- kernel_width = int( np.ceil( mo.rec_effective_tc_width()*mo.δ ))
- x = np.linspace(-3*kernel_width, 3*kernel_width, 6*kernel_width + 1)
- kernel = np.exp(-x**2/(2*kernel_width**2))
- kernel = kernel/kernel.sum()
- gains_1 = torch.Tensor(np.convolve(gains_0, kernel, mode='same'))
- #
- # two additional starting points
- gbar = gains_0.mean()
- gains_const = torch.full_like(gains_0, gbar)
- gains_const_times_prior = gbar * a_prior_neurons * num_neurons
- # same optimization from 3 starts
- 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,
- start_gains=gains_1, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
- print(f'Intermediate loss (smooth start): {loss_smooth:.6g}')
- 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,
- start_gains=gains_const, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
- print(f'Intermediate loss (const start): {loss_const:.6g}')
- 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,
- start_gains=gains_const_times_prior, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
- print(f'Intermediate loss (const x prior start): {loss_const_prior:.6g}')
- # pick best
- loss = loss_smooth
- gains_best = gains_smooth
- best_is = 'smooth'
- if loss_const < loss:
- loss = loss_const
- gains_best = gains_const_opt
- best_is = 'const'
- if loss_const_prior < loss:
- loss = loss_const_prior
- gains_best = gains_const_prior_opt
- best_is = 'const x prior'
- print(f'Best final solution is from "{best_is}" start.')
- print(f'\t----------> loss = {loss:.6g}')
- 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]
- # %% [markdown]
- # ### Peaked prior
- # %%
- #CACHE_find_optimal_gains_for_parameters_uniform_mixture = {}
- CACHE_find_optimal_gains_for_parameters_uniform_mixture = np.load('CACHE_find_optimal_gains_for_parameters_uniform_mixture.npz', allow_pickle=True)['cache'].item()
- # %%
- @cached(cache=CACHE_find_optimal_gains_for_parameters_uniform_mixture)
- 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.):
- print( f'prior width = {prior_w}' )
- mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width, inhib=inhib)
- a_prior, a_prior_neurons = make_mixture_uniform_gaussian_priors(mo.neurons_pref_stims, prior_w, uniform_part=uniform_part)
- if inhib != 0.:
- print('Solution with inhib = 0 ...')
- 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,
- alpha_cost=alpha_cost, prior_w=prior_w, regularizer_lambda=regularizer_lambda, uniform_part=uniform_part, inhib=0.)
- print('Solution with inhib = 0 ... OK')
- else:
- mo.optimize_gains_quadratic_g0Delta_expSqErr_wCost_simple(a_prior, alpha_cost=alpha_cost, v=True)
- gains_0 = mo.gains.clone()
- # smoothen the initial gain solution
- kernel_width = int(np.ceil(mo.rec_effective_tc_width() * mo.δ))
- x = np.linspace(-3 * kernel_width, 3 * kernel_width, 6 * kernel_width + 1)
- kernel = np.exp(-x**2 / (2 * kernel_width**2))
- kernel = kernel / kernel.sum()
- gains_1 = torch.Tensor(np.convolve(gains_0, kernel, mode='same'))
- # two additional starting points
- gbar = gains_0.mean()
- gains_const = torch.full_like(gains_0, gbar)
- gains_const_times_prior = gbar * a_prior_neurons * num_neurons
- # same optimization from 3 starts
- 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,
- start_gains=gains_1, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
- print(f'Intermediate loss (smooth start): {loss_smooth:.6g}')
- 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,
- start_gains=gains_const, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
- print(f'Intermediate loss (const start): {loss_const:.6g}')
- 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,
- start_gains=gains_const_times_prior, v=True, regularizer_lambda=regularizer_lambda, return_losses=True, eps_rel=1e-6)
- print(f'Intermediate loss (const x prior start): {loss_const_prior:.6g}')
- # pick best
- loss = loss_smooth
- gains_best = gains_smooth
- best_is = 'smooth'
- if loss_const < loss:
- loss = loss_const
- gains_best = gains_const_opt
- best_is = 'const'
- if loss_const_prior < loss:
- loss = loss_const_prior
- gains_best = gains_const_prior_opt
- best_is = 'const x prior'
- print(f'Best final solution is from "{best_is}" start.')
- print(f'\t----------> loss = {loss:.6g}')
- 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],
- # %%
- # Compute Full Width at Half Maximum (FWHM) for the tuning curve
- def tc_fwhm(tc):
- tc_half_max = tc.max() / 2.0
- above_half = np.where(tc >= tc_half_max)[0]
- if above_half.size > 0:
- fwhm = STIMULI[above_half[-1]] - STIMULI[above_half[0]]
- else:
- fwhm = np.nan
- return fwhm
- def tc_fwhm_interp(tc):
- # Interpolate the tuning curve
- interp_func = interp1d(STIMULI, tc, kind='cubic', bounds_error=False, fill_value=0.0)
- tc_max = tc.detach().numpy().max()
- half_max = tc_max / 2.0
- # Find where the interpolated curve crosses half max
- fine_x = np.linspace(STIMULI[0], STIMULI[-1], 5000)
- fine_y = interp_func(fine_x)
- above_half = np.where(fine_y >= half_max)[0]
- if above_half.size > 0:
- fwhm = fine_x[above_half[-1]] - fine_x[above_half[0]]
- else:
- fwhm = np.nan
- return fwhm
- # %% [markdown]
- # # Figures
- # %% [markdown]
- # ## FIG. 3 Illustration of mechanism
- # %%
- default_colors = rcParams['axes.prop_cycle'].by_key()['color'] # Default color cycle
- # %%
- ff_sigma = 20.
- rec_magnitude = 0.85
- rec_width = 30.
- mo = Network(num_neurons=100, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
- ired = 55
- iorg = None
- gs = torch.ones(mo.num_neurons)
- gs2 = gs.clone()
- gs2[:mo.num_neurons//2] = 2.
- gs2[mo.num_neurons//2:] = 0
- # Number of colors to extract from cmap
- cmap = cm.viridis_r #Spectral # viridis_r
- num_colors = 19
- cmap_colors = [to_hex(cmap(i / (num_colors - 1))) for i in range(num_colors)]
- colors = rcParams['axes.prop_cycle'].by_key()['color'] # Default color cycle
- # set axes.prop_cycle to use cmap_colors
- rcParams['axes.prop_cycle'] = cycler(color=cmap_colors)
- fig, axs = subplots(nrows=3, ncols=4, figsize=(16, 4), gridspec_kw={'height_ratios': [5, 7, 13]}, sharex=True)
- for jj in [0, 1, 2]: # this merges row 2 and 3 for these columns
- gds = axs[1, jj].get_gridspec()
- for ax in axs[1:, jj]:
- ax.remove()
- axs[1, jj] = fig.add_subplot(gds[1:, jj])
- for jj in [0, 1, 2]:
- this_gs = gs if jj<=1 else gs2 #[gs, gs2][jj-1]
- mo.gains = this_gs
- mo.update_r_star()
- if jj==0:
- axs[1,jj].plot(STIMULI, mo.ff_tuning_curves[5::5].T, lw=.5, alpha=1);
- axs[1,jj].plot(STIMULI, mo.ff_tuning_curves[ired], lw=2, alpha=1, c=colors[3], zorder=1000);
- else:
- axs[1,jj].plot(STIMULI, mo.r_star[5::5].T, lw=.5, alpha=1, clip_on=False);
- if iorg: axs[1,jj].plot(STIMULI, mo.r_star[iorg], lw=2, alpha=1, c=colors[1]);
- axs[1,jj].plot(STIMULI, mo.r_star[ired], lw=2, alpha=1, c=colors[3]);
- axs[1,jj].plot(STIMULI, mo.ff_tuning_curves[ired], lw=1, alpha=1, c=colors[3], zorder=1000, ls=':');
- if jj==1:
- tc_1 = mo.r_star[ired]
- elif jj==2:
- axs[1,jj].plot(STIMULI, tc_1, lw=1, alpha=1, c=colors[3], zorder=1000, ls='--');
- 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))
- 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)
- if iorg: axs[0,jj].scatter(mo.neurons_pref_stims[iorg], this_gs[iorg], color=colors[1], s=12, clip_on=False, zorder=1000)
- axs[0,jj].scatter(mo.neurons_pref_stims[ired], this_gs[ired], color=colors[3], s=18, clip_on=False, zorder=1000)
- if jj==1:
- axs[1,jj].set_prop_cycle(None)
- if iorg: axs[1,jj].plot(STIMULI, mo.ff_tuning_curves[iorg], lw=2, alpha=1, c=colors[1], ls='--');
- else:
- mo.gains = gs
- mo.update_r_star()
- if iorg: axs[1,jj].plot(STIMULI, mo.r_star[iorg], lw=.5, alpha=1, c=colors[1], label=label, ls='--');
- axs[0,jj].set_ylim(0,2.);
- axs[1,jj].set_ylim(0,6);
- axs[1,jj].set_xticks([])
- axs[0,jj].set_xticks([])
- axs[0,jj].set_xlim(LOWER_BOUND, UPPER_BOUND)
- axs[1,jj].set_xlim(LOWER_BOUND, UPPER_BOUND)
- axs[0,jj].set_yticks([0,1,2])
- axs[1,jj].set_yticks([0,2,4,6])
- axs[0,jj].spines['top'].set_visible(False); axs[0,jj].spines['right'].set_visible(False)
- axs[1,jj].spines['top'].set_visible(False); axs[1,jj].spines['right'].set_visible(False)
- axs[1,jj].set_xlabel('Stimulus s', labelpad=8)
- if jj>0:
- axs[0,jj].set_yticklabels([])
- axs[1,jj].set_yticklabels([])
- ####
- 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)
- if iorg: axs[0,3].scatter(mo.neurons_pref_stims[iorg], this_gs[iorg], color=colors[1], s=12, clip_on=False, zorder=1000)
- axs[0,3].scatter(mo.neurons_pref_stims[ired], this_gs[ired], color=colors[3], s=18, clip_on=False, zorder=1000)
- for ii in range(3):
- axs[ii,3].spines['top'].set_visible(False); axs[ii,3].spines['right'].set_visible(False)
- axs[0,3].set_yticks([0,1,2])
- axs[0,3].set_ylim(0,2.);
- axs[0,3].set_yticklabels([])
- axs[0,3].annotate("$g_j$", (.8, 0.6), ha='center', fontsize=9, xycoords='axes fraction')
- 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)
- axs[1,3].scatter(mo.neurons_pref_stims[ired], mo.hh[ired,ired], color=colors[3], s=8, clip_on=False, zorder=1000)
- gammas, gamma_den, φs = make_gamma_G_φs(gs2, mo.hh, mo.neurons_pref_stims)
- 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)
- axs[2,3].scatter(mo.neurons_pref_stims[ired], gammas[ired,ired], color=colors[3], s=8, clip_on=False, zorder=1000)
- axs[1,3].scatter(mo.neurons_pref_stims[ired], 0., color=colors[3], marker='|', clip_on=False, zorder=1000)
- axs[1,3].text(x=mo.neurons_pref_stims[ired], y=0.025, s="$s_i$", ha='center', color=colors[3], fontsize=12)
- axs[1,3].set_title(x=.8, y=0.7, label="$h(s_j-s_i)$", ha='center', fontsize=9)
- axs[1,3].set_ylabel("$h(s_j-s_i)$")
- 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)
- axs[2,3].set_ylabel("Weight $\\gamma_{ij}$")
- axs[2,3].scatter(mo.neurons_pref_stims[ired], 0., color=colors[3], marker='|', clip_on=False, zorder=1000)
- axs[2,3].text(x=mo.neurons_pref_stims[ired], y=0.0075, s="$s_i$", ha='center', color=colors[3], fontsize=12)
- axs[2,3].scatter(φs[ired], 0., color=colors[3], marker='o', fc='w', clip_on=False, zorder=1000)
- axs[2,3].text(x=φs[ired], y=0.0075, s="$\\varphi(s_i)$", ha='center', color=colors[3], fontsize=12)
- 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))
- for ii in range(3):
- axs[ii,3].axvline(0, lw=1, color='lightgrey', ls=':')
- axs[2,3].set_xlabel('Feedforward preferred stimulus $s_j$', labelpad=8)
- axs[1,3].set_ylim(0,None)
- axs[2,3].set_ylim(0,None)
- axs[1,3].set_yticks([])
- axs[2,3].set_yticks([])
- axs[0,0].set_ylabel('Gain g')
- axs[1,0].set_ylabel('Firing rate r(s)')
- fsz = 10.
- axs[0,0].set_title('No recurrence\n+ Uniform gains', fontsize=fsz)
- axs[0,1].set_title('Recurrent connections\n+ Uniform gains', fontsize=fsz)
- axs[0,2].set_title('Recurrent connections\n+ Non-uniform gains', fontsize=fsz)
- axs[1,0].text(0,4.5,'Gains $\\times$ Feedforward tuning curves\n$r(s) = g \\circ f(s)$', ha='center')
- axs[1,1].text(0,4.5,'Effective tuning curves\n$r(s) = (I-W)^{-1} ( g \\circ f(s) )$', ha='center')
- fig.subplots_adjust(wspace=.1, hspace=.15)
- #show()
- #fig.savefig('rnn_tc.pdf', bbox_inches='tight')
- # set axes.prop_cycle back to default
- rcParams['axes.prop_cycle'] = cycler(color=default_colors)
- # %% [markdown]
- # ## FIG. 4 Optimized network for priors 30, 20, 10
- # %%
- num_neurons = 801
- ff_sigma = 5.
- rec_magnitude = 0.95
- rec_width = 6.
- alpha_cost = .5
- reg_lambda = 1000
- inhib = 0.
- fig = figure(figsize=(11, 8))
- # 5 rows: row0(ax0), spacer, row1, row2, row3
- gs = GridSpec(5, 2, figure=fig, height_ratios=[3, 0.3, 1, 1, 1])
- axPriors = fig.add_subplot(gs[0,0])
- axGains = fig.add_subplot(gs[2:,0])
- axPhi = fig.add_subplot(gs[0,1])
- axs_tc = [fig.add_subplot(gs[2+i, 1]) for i in range(3)]
- g_solutions = {}
- g_solutions_approx = {}
- for ip,prior_w in enumerate( [30, 20, 10] ):
- 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,
- alpha_cost=alpha_cost, prior_w=prior_w, regularizer_lambda=reg_lambda, inhib=inhib)
- g_solutions_approx[prior_w] = gains_list[0]
- g_solutions[prior_w] = g_solution
- mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width, inhib=inhib)
- φs_perprior = {}
- argmaxs_perprior = {}
- tuning_curves = {}
- for ip,prior_w in enumerate([10, 20, 30]):
- print( f'prior width = {prior_w}', end=' ')
- c = COLORS_WIDTH.get(prior_w)
- a_prior, a_prior_neurons = make_gaussian_priors(mo.neurons_pref_stims, prior_w)
- prior_mean, prior_var = probs_mean_and_var(a_prior_neurons, mo.neurons_pref_stims)
- 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}$)')
- ####
- g_solution = torch.clamp( g_solutions[prior_w], min=0.)
- mo.gains = g_solution
- mo.update_r_star()
- #
- axGains.plot( mo.neurons_pref_stims, g_solution, c=c, ls='-')
- g_sol_approx = torch.clamp( g_solutions_approx[prior_w], min=0.)
- axGains.plot(mo.neurons_pref_stims, g_sol_approx, c=c, ls='--' )
- #
- true_φs, true_sigmarsSq = mo.effective_tc_center_and_squaredwidth()
- φs_perprior[prior_w] = true_φs
- argmaxs_perprior[prior_w] = STIMULI[ mo.r_star.argmax(axis=1) ]
- tuning_curves[prior_w] = mo.r_star.clone().detach().numpy()
- idxs = [355,368,379,388,395,400,405,412,421,432,445]
- # Number of colors to extract from cmap
- cmap = cm.Spectral #viridis_r #Spectral # viridis_r
- num_colors = len(idxs) # Adjust based on the number of lines you need
- cmap_colors = [to_hex(cmap(i / (num_colors - 1))) for i in range(num_colors)]
- wref = 30
- a_prior_ref, _ = make_gaussian_priors(mo.neurons_pref_stims, wref)
- aa = 1.
- ref_Fa, ref_Fa_m1 = make_Fa_and_Fam1_fcts(STIMULI, a_prior_ref, a=aa)
- for iν, ν in enumerate([30, 20, 10]):
- c = COLORS_WIDTH.get(ν)
- a_prior, _ = make_gaussian_priors(mo.neurons_pref_stims, ν)
- g_sol = g_solutions[ν]
- Fa, Fa_m1 = make_Fa_and_Fam1_fcts(STIMULI, a_prior, a=aa)
- axPhi.plot(φs_perprior[wref], φs_perprior[ν], c=c, ls='-', lw=1)
- axPhi.plot(STIMULI, Fa_m1(ref_Fa(STIMULI)), c=c, ls=':', lw=1)
- axPhi.scatter(φs_perprior[wref][idxs], φs_perprior[ν][idxs], color=cmap_colors, zorder=10, s=20, clip_on=True, ec='k', linewidth=.5)
- xlim_val = 39
- for ip,prior_w in enumerate( [30, 20, 10] ):
- ax = axs_tc[ip]
- ax.set_prop_cycle(cycler(color=cmap_colors))
- a_prior, _ = make_gaussian_priors(mo.neurons_pref_stims, prior_w)
- rrr = tuning_curves[prior_w][idxs]
- ax.plot(STIMULI, (rrr.T/rrr.T.max(axis=0)), lw=.5)
- xs = φs_perprior[prior_w][idxs]
- ys = np.zeros_like(xs)
- n = len(idxs)
- mid = n//2
- # 1) build the front-to-back order of positions in idxs
- order = [mid]
- for k in range(1, n):
- if mid - k >= 0:
- order.append(mid - k)
- if mid + k < n:
- order.append(mid + k)
- # 2) now scatter each in that order, giving the first (the true front) the highest zorder
- base = 1000 # same base you were using
- for layer, pos in enumerate(order):
- z = base - layer
- ax.scatter(xs[pos], ys[pos], color=cmap_colors[pos], s=20, ec='k', linewidth=.5, clip_on=False, zorder=z)
- ax.set_xlim(-xlim_val,xlim_val)
- ax.fill_between(STIMULI, .4*a_prior/a_prior.max(), color=COLORS_WIDTH[prior_w], zorder=-1000, alpha=.2)
- ax.set_ylim(0,None)
- ax.set_ylabel({10:'Narrow', 20:'Medium', 30:'Wide'}[prior_w] + '\nprior', color=COLORS_WIDTH[prior_w], rotation=0, y=.3, labelpad=24)
- axPhi.set_title('Effective preferred stimuli')
- axPhi.set_xlim(-xlim_val, xlim_val); axPhi.set_ylim(-xlim_val, xlim_val)
- axPhi.set_xlabel('Preferred stimulus with Wide prior, $\\varphi_{30}$')
- axPhi.set_ylabel('Preferred stimulus, $\\varphi_{\\sigma_p}$', labelpad=-5)
- axPhi.legend(labels=[f'$\\varphi_{{\\sigma_p}}$ vs. $\\varphi_{{ {wref} }}$', f'$P_{{\\sigma_p}}^{{-1}}(P_{{ {wref} }}(\\varphi_{{ {wref} }}))$',],
- handles=[Line2D([],[],ls='-',c='grey'),Line2D([],[],ls=':',c='grey'),], loc='upper left')
- axPriors.set_title('Priors')
- axPriors.set_xlim(-110, 110); axPriors.set_ylim(0, None)
- axPriors.legend(loc=(.6,.6))
- axPriors.set_yticks([])
- axPriors.set_xlabel('Stimulus')
- axPriors.set_ylabel('pdf')
- axGains.set_title('Optimal gains')
- ###
- axGains.set_xlim(-110, 110); axGains.set_ylim(0, .1099)
- axGains.set_xlabel('Feedforward preferred stimulus',)
- axGains.set_ylabel('$g(s)$', labelpad=0)
- axGains.legend(labels=['Numerical optimization', 'Analytical approximation',],
- handles=[Line2D([],[],ls='-',c='grey'),Line2D([],[],ls='--',c='grey'),], loc=(.45,.83))
- axGains.ticklabel_format( axis='y', style='sci', scilimits=(-2, -2), useMathText=True )
- axGains.yaxis.get_offset_text().set_fontsize(8)
- axGains.yaxis.get_offset_text().set_x(-.025)
- axs_tc[0].set_title('Normalized effective tuning curves')
- axs_tc[0].set_xticks([]); axs_tc[1].set_xticks([])
- axs_tc[0].set_yticks([]); axs_tc[1].set_yticks([]); axs_tc[2].set_yticks([])
- axs_tc[2].set_xlabel('Stimulus')
- # Remove top and right spines for all axes
- for ax in [axPriors, axGains, axPhi, *axs_tc]:
- ax.spines['top'].set_visible(False)
- ax.spines['right'].set_visible(False)
- fig.subplots_adjust(hspace=.35)
- #fig.savefig('opt_gains.pdf', bbox_inches='tight')
- #show()
- # %%
- # np.savez_compressed('CACHE_find_optimal_gains_for_parameters.npz', cache=CACHE_find_optimal_gains_for_parameters)
- # %% [markdown]
- # ## FIG. 5 g, d, sigma, etc.
- # %%
- samples_and_estimates_cache = {}
- # %%
- num_neurons = 801
- ff_sigma = 5.
- rec_magnitude = 0.95
- rec_width = 6.
- alpha_cost = .5
- reg_lambda = 1000
- inhib = 0.
- T = 1
- tau = .01
- burn_in = 0.5
- dt = 0.1 #0.001
- num_samples = 1000
- if dt != 0.001:
- print(f"Warning: dt = {dt} != 0.001. The response variances may be incorrect.")
- prior_ws = [10, 20, 30]
- g_solutions = {}
- for ip,prior_w in enumerate( prior_ws ):
- 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,
- alpha_cost=alpha_cost, prior_w=prior_w, regularizer_lambda=reg_lambda, inhib=inhib)
- g_solutions[prior_w] = g_solution
- mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
- vars = []
- for ip,prior_w in enumerate(prior_ws):
- print( f'prior width = {prior_w}')
- c = COLORS_WIDTH.get(prior_w)
- a_prior, a_prior_neurons = make_gaussian_priors(mo.neurons_pref_stims, prior_w) #, prior_center_val)
- prior_mean, prior_var = probs_mean_and_var(a_prior_neurons, mo.neurons_pref_stims)
- ####
- g_solution = torch.clamp( g_solutions[prior_w], min=0.)
- mo.gains = g_solution
- mo.update_r_star()
- ###
- xlo = - prior_w / 2
- xhi = - xlo
- if prior_w in samples_and_estimates_cache:
- samples, estimates = samples_and_estimates_cache[prior_w]
- else:
- samples = mo.simulate_recurrent_poisson_many(num_samples=num_samples, xbounds=(xlo, xhi), T=T, tau=tau, dt=dt, burn_in=burn_in)
- estimates = mo.decode(samples, MID_VAL, prior_w**2)
- samples_and_estimates_cache[prior_w] = (samples, estimates)
- #
- excursions = estimates - estimates.mean(axis=0)
- vars.append( excursions.var(axis=0).mean() )
- vars = np.array(vars)
- # %%
- num_neurons = 801
- ff_sigma = 5.
- rec_magnitude = 0.95
- rec_width = 6.
- alpha_cost = .5
- reg_lambda = 1000 #600
- inhib = 0.
- T = 1
- tau = .01
- burn_in = 0.5
- dt = 0.1 #0.001
- num_samples = 1000
- prior_ws = [10, 20, 30]
- fig, axs = subplots(nrows=2, ncols=3, figsize=(9,5), sharex=False, )
- axG = axs[0,0]
- axW = axs[0,1]
- axd = axW.twinx()
- axWd = axs[0,2]
- axI = axs[1,0]
- axNbSpikes = axs[1,1]
- axV = axs[1,2]
- g_solutions = {}
- tuning_curves = {}
- φs_perprior = {}
- centergs = {}
- centersigmas = {}; sigmas = {}
- centerdensities = {}; densities = {}
- exp_num_spikes = {}
- center_fish_infos = {}
- fish_infos = {}
- exp_sq_error = {}
- mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
- prior_wparams = np.array([10, 20, 30])
- for ip,prior_w in enumerate( prior_wparams ):
- 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,
- alpha_cost=alpha_cost, prior_w=prior_w, regularizer_lambda=reg_lambda, inhib=inhib)
- mo.gains = g_solution
- mo.update_r_star()
- a_prior, a_prior_neurons = make_gaussian_priors(mo.neurons_pref_stims, prior_w)
- prior_mean, prior_var = probs_mean_and_var(a_prior_neurons, mo.neurons_pref_stims)
- true_φs, true_sigmarsSq = mo.effective_tc_center_and_squaredwidth()
- φs_perprior[prior_w] = true_φs
- tuning_curves[prior_w] = mo.r_star.clone().detach().numpy()
- ### Gains
- centergs[prior_w] = g_solution[len(g_solution)//2]
- # Tuning curves width
- centersigmas[prior_w] = np.sqrt( true_sigmarsSq[len(g_solution)//2] )
- #sigmas[prior_w] = np.sqrt( true_sigmarsSq )
- # Density
- ind = np.abs(mo.neurons_pref_stims - 0.) <= prior_w*10000
- φm1prime = np.gradient(mo.neurons_pref_stims[ind], true_φs[ind])
- # print(true_φs[ind][len(φm1prime)//2])
- centerdensities[prior_w] = φm1prime[len(φm1prime)//2]/mo.δ
- #densities[prior_w] = φm1prime / mo.δ
- # Expected total number of spikes vs width
- exp_num_spikes[prior_w] = (a_prior * mo.r_star.sum(axis=0)).sum()
- # Fisher information - center
- center_fish_infos[prior_w] = SQRT_2_PI * centergs[prior_w] * centerdensities[prior_w] / centersigmas[prior_w]
- exp_sq_error[prior_w] = (a_prior.numpy()/(1/prior_var + (1./mo.beta)*(mo.r_star.T/true_sigmarsSq).sum(axis=1)) ).sum()
- def strip_leading_zero(x, pos):
- return '0' if x==0 else f"{x:.2f}".lstrip("0") if abs(x) < 1 else f"{x:.2f}"
- axG.plot( prior_wparams, [centergs[ν] for ν in prior_wparams], '.-', c='lightgrey', lw=1)
- 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)
- axG.set_title('Gain $g$')
- axG.yaxis.set_major_formatter(FuncFormatter(strip_leading_zero))
- axG.set_ylim(0, None)
- axW.plot( prior_wparams, [centersigmas[ν] for ν in prior_wparams], '.-', c='lightgrey', lw=1)
- 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)
- axW.set_title('Tuning curve width and density')
- axW.set_ylim(0, None)
- axW.legend(handles=[Line2D([], [], color='lightgrey', label='Width $\\sigma$ (left)'),
- Line2D([], [], color='lightgrey', ls='--', label='Density $d$ (right)')],)
- axd.plot( prior_wparams, np.array([centerdensities[ν] for ν in prior_wparams]), '.--', c='lightgrey', lw=1)
- 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)
- axd.set_ylim(0, None)
- # I \propto g d / sigma
- axI.plot( prior_wparams, np.array([center_fish_infos[ν] for ν in prior_wparams]), '.-', c='lightgrey', lw=1)
- 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)
- axI.set_title('Fisher $I \\propto g d / \\sigma$')
- axI.set_xlabel('Prior width $\\sigma_p$')
- xx = prior_wparams
- yy = np.array([center_fish_infos[ν] for ν in prior_wparams])
- a_fit, _ = optimize.curve_fit(lambda x,a: a/x, xx, yy)
- axI.plot(np.linspace(10,30), a_fit[0]/np.linspace(10,30), 'r:', lw=1, label='Fit $1/\\sigma_p$', zorder=-1000)
- axI.legend(fontsize=10, loc='lower left')
- axI.set_ylim(0, None);
- axI.yaxis.set_major_formatter(FuncFormatter(strip_leading_zero))
- axI.set_ylim(0, None)
- # # Var / std dev
- axWd.plot( 1./np.array([centerdensities[ν] for ν in prior_wparams]), [centersigmas[ν] for ν in prior_wparams], '.-', c='lightgrey', lw=1)
- 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)
- axWd.set_title('Width vs. spacing')
- axWd.set_xlabel('Spacing (1/density)')
- axWd.set_ylabel('Width $\\sigma$', labelpad=-7)
- axWd.set_yticks([10, 20])
- axNbSpikes.plot( prior_wparams, [exp_sq_error[ν] for ν in prior_wparams], '.-', c='lightgrey', lw=1, label='$L \\simeq $MSE (left)')
- 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)
- axNbSpikes.set_title('Loss (MSE) and Cost (spikes)')
- axNbSpikes.set_xlabel('Prior width $\\sigma_p$')
- axNbSpikes.set_ylim(0, 24)
- axT = axNbSpikes.twinx()
- axT.plot( prior_wparams, [exp_num_spikes[ν] for ν in prior_wparams], 'x--', c='lightgrey', lw=1, label='$C=$Spiking activity (right)')
- 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)
- axT.set_ylim(0, 42)
- axNbSpikes.legend(loc='upper left', borderpad=0, handletextpad=.5)
- axT.legend(loc='lower left', borderpad=0, handletextpad=.5)
- 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 ])
- line, caps, bars = axV.errorbar([10,20,30], vs, (vs-vs_5, vs_95-vs), c='lightgrey', lw=1, ls=':', label='Subjects')
- axV.plot(prior_wparams, vars, '.-', c='lightgrey', lw=1, label='Network model')
- axV.scatter([10,20,30], vars, color=[COLORS_WIDTH[w] for w in [10,20,30]], s=18, zorder=1000, clip_on=False)
- axV.set_xlabel('Prior width $\\sigma_p$')
- axV.set_title('Response variance')
- axV.legend()
- axV.set_ylim(0, 65)
- for ax in axs.flat:
- ax.spines['top'].set_visible(False)
- ax.spines['right'].set_visible(False)
- axT.spines['top'].set_visible(False)
- axT.spines['right'].set_ls('--')
- axd.spines['top'].set_visible(False)
- axd.spines['right'].set_ls('--')
- for ax in [axG, axW, axd, axI, axNbSpikes, axV]:
- ax.set_xticks([10,20,30])
- fig.tight_layout(h_pad=.5)
- if dt != 0.001:
- print(f"Warning: dt = {dt} != 0.001. The response variances may be incorrect.")
- #fig.savefig('fisher_comp.pdf', bbox_inches='tight')
- #show()
- # %% [markdown]
- # ## Adapter repulsion
- # %% [markdown]
- # ### FIG. 2 Illustration of adapter repulsion
- # %% [markdown]
- # #### Gratings
- # %%
- from scipy.ndimage import rotate
- from matplotlib.offsetbox import OffsetImage, AnnotationBbox
- # %%
- def generate_square_wave_grating(size=256, spatial_freq=10, orientation=0):
- """Return a high-contrast square-wave grating as a NumPy array."""
- x = np.linspace(-0.5, 0.5, size)
- y = np.linspace(-0.5, 0.5, size)
- xv, yv = np.meshgrid(x, y)
- theta = np.deg2rad(orientation)
- xt = xv * np.cos(theta) + yv * np.sin(theta)
- grating = np.sign(np.sin(2 * np.pi * spatial_freq * xt))
- return grating
- def add_grating_to_ax(ax, grating, xy, zoom=1.0, **kwargs):
- """
- Adds a grating image to a given Axes at specified coordinates.
- Parameters:
- - ax: matplotlib Axes object
- - grating: 2D NumPy array, the grating image
- - xy: tuple (x, y), coordinates in data units where the image will be placed
- - zoom: float, scale of the image
- - **kwargs: additional keyword arguments passed to AnnotationBbox (e.g., xycoords)
- """
- imagebox = OffsetImage(grating, cmap='gray', zoom=zoom)
- ab = AnnotationBbox(imagebox, xy, frameon=False, **kwargs)
- ax.add_artist(ab)
- # fig, ax = subplots()
- # grating = generate_square_wave_grating(size=256, spatial_freq=10, orientation=45)
- # add_grating_to_ax(ax, grating, xy=(0.5, 0.5), zoom=0.2)
- # add_grating_to_ax(ax, grating, xy=(0.2, 0.8), zoom=0.3)
- # ax.set_xlim(0, 1)
- # ax.set_ylim(0, 1)
- #show()
- # %%
- def save_grating_to_file(grating, filename, dpi=300, format='png'): # or 'tiff
- """
- Parameters:
- - grating: 2D numpy array with values in [-1, 1]
- - filename: output path, should end in .tif or .tiff
- - dpi: resolution metadata for image (useful for print layout)
- """
- # Normalize to [0, 255] for 8-bit grayscale image
- grating_uint8 = ((grating + 1) / 2 * 255).astype(np.uint8)
- imsave(fname=filename, arr=grating_uint8, cmap='gray', dpi=dpi, format=format)
- for ori in [0, 22.5, 45, 67.5, 90, 112.5, 135, 157.5, ]:
- g = generate_square_wave_grating(size=512, spatial_freq=5, orientation=ori)
- #save_grating_to_file(g, f'figures/grating_{ori}deg.png', dpi=300)
- figure(figsize=(1,1))
- imshow(g, cmap='gray')
- # %% [markdown]
- # #### Fig. 2B
- # %%
- from matplotlib.patches import FancyArrowPatch
- fig, axRep = subplots(figsize=(6, 4))
- xx = np.linspace(-8,95,1001)
- cC = 'C0'; cA = cm.tab20(7)
- axRep.plot(xx, stats.norm(loc=10., scale=8).pdf(xx)/stats.norm(scale=8).pdf(0), c=cC, label='Control', lw=2)
- axRep.plot(xx, stats.norm(loc=15., scale=12).pdf(xx)/stats.norm(scale=12).pdf(0), c=cA, label='Adaptation', lw=2)
- #
- axRep.plot(xx, stats.norm(loc=45., scale=8).pdf(xx)/stats.norm(scale=8).pdf(0), c=cC, ls='--', lw=2)
- axRep.plot(xx, stats.norm(loc=47., scale=7).pdf(xx)/stats.norm(scale=7).pdf(0), c=cA, ls='--', lw=2)
- # Add double arrows to indicate the FWHM of the two normalized Gaussians above
- for (mu, sigma, color) in [(10., 8., cC), (15., 12., cA), (45., 8., cC), (47., 7., cA)]:
- y = stats.norm(loc=mu, scale=sigma).pdf(xx) / stats.norm(scale=sigma).pdf(0)
- half_max = y.max() / 2.0
- above_half = np.where(y >= half_max)[0]
- if above_half.size > 0:
- fwhm_left = xx[above_half[0]]
- fwhm_right = xx[above_half[-1]]
- fwhm_y = y.max() * 0.5 + (0.025 if color == cC else 0.)
- axRep.annotate(
- '', xy=(fwhm_right, fwhm_y), xytext=(fwhm_left, fwhm_y),
- arrowprops=dict(arrowstyle='<->', color=color, lw=1)
- )
- axRep.plot(xx, stats.norm(loc=82, scale=8).pdf(xx)/stats.norm(scale=8).pdf(0), c=cC, ls=':', lw=2)
- axRep.plot(xx, stats.norm(loc=80, scale=7).pdf(xx)/stats.norm(scale=7).pdf(0), c=cA, ls=':', lw=2)
- #
- axRep.vlines(0, 0, 1., color='lightgrey', ls='-', zorder=-100)
- axRep.annotate(" Adapter", xytext=(0, 1.25), xy=(0, 1.0), arrowprops=dict(arrowstyle="-|>", color='r'), ha='center')
- axRep.annotate("", xytext=(10, 1+.03), xy=(15, 1+.03), arrowprops=dict(arrowstyle="-|>", color='k', ls='-', lw=.5, shrinkA=0, shrinkB=0))
- axRep.annotate("", xytext=(44.8, 1+.03), xy=(47.2, 1+.03), arrowprops=dict(arrowstyle="-|>", color='k', ls='-', lw=.5, shrinkA=0, shrinkB=0))
- axRep.annotate("", xytext=(82.2, 1+.03), xy=(79.8, 1+.03), arrowprops=dict(arrowstyle="-|>", color='k', ls='-', lw=.5, shrinkA=0, shrinkB=0))
- axRep.set_ylim(0,1.375)
- axRep.set_xlim(xx[0],xx[-1])
- axRep.set_xlabel('Orientation (difference from adapter)')
- axRep.set_yticks([])
- axRep.set_ylabel('Firing rate')
- axRep.legend(loc=(.275,.87) )
- axRep.spines['top'].set_visible(False)
- axRep.spines['right'].set_visible(False)
- fig.subplots_adjust(wspace=.25)
- #fig.savefig('illus_adapter_repulsion.pdf', bbox_inches='tight')
- # %% [markdown]
- # ### FIG. 6 Adapter repulsion - Model
- # %%
- num_neurons = 801
- ff_sigma = 5.
- rec_magnitude = 0.95
- rec_width = 6.
- inhib = 0.
- alpha_cost = .5
- regularizer_lambda = 1000
- uniform_part = 0.8
- prior_wparams=[10000, 1]
- g_solutions = {}
- # g_solutions_approx = {}
- # g_params_approx = {}
- for ip,prior_w in enumerate( prior_wparams ):
- this_uniform_part = uniform_part if prior_w < 9999 else 1.
- 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,
- prior_w=prior_w, regularizer_lambda=regularizer_lambda, uniform_part=this_uniform_part, inhib=inhib)
- # g_solutions_approx[prior_w] = g_solution_approx
- # g_params_approx[prior_w] = (g0, Delta)
- g_solutions[prior_w] = gains
- # %%
- mo = Network(num_neurons, ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width, inhib=inhib)
- fig, axs = subplots(nrows=2, ncols=5, figsize=(18,6))
- tuning_curves = {}
- φs_perprior = {}
- sigmas = {}
- fwhms = {}
- axg = axs[1,0]
- these_prior_wparams = [10000, 1]
- for ip,prior_w in enumerate(these_prior_wparams):
- print( f'prior width = {prior_w}', end=' ')
- c = COLORS_WIDTH.get(prior_w)
- this_uniform_part = uniform_part if prior_w < 9999 else 1.
- xxx = torch.linspace(LOWER_BOUND, UPPER_BOUND, 10001)
- a_prior, a_prior_xxx = make_mixture_uniform_gaussian_priors(xxx, prior_w, uniform_part=this_uniform_part)
- #prior_mean, prior_var = probs_mean_and_var(a_prior_xxx, xxx)
- label = 'Control' if prior_w==10000 else 'Adaptation'
- axs[0,0].plot( xxx, a_prior_xxx, c=c, ls='-', label=label)
- ####
- g_solution = torch.clamp( g_solutions[prior_w], min=0.)
- mo.gains = g_solution
- mo.update_r_star()
- #
- axg.plot( mo.neurons_pref_stims, g_solution, c=c, ls='-')
- tuning_curves[prior_w] = mo.r_star.clone().detach().numpy()
- fwhms[prior_w] = np.array( [tc_fwhm_interp(tc) for tc in mo.r_star] )
- axs[0,0].set_title('Priors')
- axs[0,0].set_xlim(-40, 40);
- axs[0,0].set_ylim(0, .0045)
- axs[0,0].legend(loc='upper left')
- axs[0,0].set_yticks([])
- g_solution = torch.clamp( g_solutions[1], min=0.)
- ind50 = mo.neurons_pref_stims.abs() <= 50
- imax = g_solution[ind50].argmax()
- print( '***', mo.neurons_pref_stims[ind50][imax], g_solution[ind50][imax] )
- xmax1 = mo.neurons_pref_stims[ind50][imax]
- xmax2 = -xmax1
- axg.scatter([xmax1, xmax2], [0,0], marker='*', clip_on=False, color=COLORS_WIDTH.get(1), zorder=10, s=100, ec='k', lw=.5)
- axg.set_title('Optimal gains')
- axg.set_xlim(-40, 40);
- axg.set_ylim(0, .19)
- axg.set_xlabel('Stimulus $s$')
- for iν, ν in enumerate(these_prior_wparams):
- c = COLORS_WIDTH.get(ν)
- axs[1,4].plot(mo.neurons_pref_stims, fwhms[ν], c=c, ls='-')
- axs[1,4].set_xlim(0, 23)
- axs[1,4].set_ylim(0, 45)
- axs[1,4].set_xlabel('Distance from adapter')
- axs[1,4].set_title('Full Width at Half Maximum')
- axs[1,4].scatter(abs(xmax1), 0, marker='*', clip_on=False, color=COLORS_WIDTH.get(1), zorder=10, s=100, ec='k', lw=.5)
- ######## Tuning curves #######
- # Plot individual neuron tuning curves (normalized) across priors
- neuron_indices = [400, 402, 405, 407, 412, 428, 446]
- axtcs = axs.flatten()[[1,2,3,4,6,7,8]]
- for ii,idx in enumerate(neuron_indices):
- ax = axtcs[ii]
- axs[1,4].scatter(mo.neurons_pref_stims[idx], 0, marker='v', clip_on=False, c='C0', zorder=10, s=20)
- ax.scatter([xmax1, xmax2], [0,0], marker='*', clip_on=False, color=COLORS_WIDTH.get(1), zorder=10, s=100, ec='k', lw=.5)
- for iw,prior_w in enumerate(these_prior_wparams):
- c = COLORS_WIDTH.get(prior_w)
- tc = tuning_curves[prior_w][idx, :] # tuning curve for neuron idx under prior prior_w
- tc_normalized = tc / tc.max() # normalize
- ax.plot(STIMULI, tc_normalized, label=f'prior {prior_w}', c=c)
- # Draw a double arrow to indicate FWHM of the tuning curve
- fwhm_val = fwhms[prior_w][idx]
- above_half = np.where(tc_normalized >= .499)[0]
- # delta_argmax = STIMULI[tc.argmax()] - STIMULI[ tuning_curves[10000][idx, :].argmax() ]
- # 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)
- if idx in [400, 428]:
- if above_half.size > 0:
- x_left = STIMULI[above_half[0]]
- x_right = STIMULI[above_half[-1]]
- y_arrow = 0.5 + 0.02*iw # vertical position for the arrow (normalized units)
- ax.annotate(
- '', xy=(x_left, y_arrow), xytext=(x_right, y_arrow),
- arrowprops=dict(arrowstyle='<->', color=c, lw=1, relpos=0, shrinkA=.5, shrinkB=.5),
- annotation_clip=False
- )
- axs[1,4].scatter(mo.neurons_pref_stims[idx], fwhms[prior_w][idx], marker='o', clip_on=False, c=c, zorder=10, s=20)
- ax.set_xlim(-30, 30)
- ax.set_ylim(0,None)
- ax.axvline(0, lw=1, c='lightgrey', zorder=-1000)
- ax.set_xticks([-30, -20, -10, 0, 10, 20, 30])
- for ax in axs[:1,1]:
- ax.set_ylabel('Normalized response')
- for ax in axs[0,:]:
- ax.set_xticklabels([])
- for ax in axs[1,:-1]:
- ax.set_xlabel('Stimulus $s$')
- axs[0,1].set_title('Tuning curves')
- for ax in axs.flat:
- ax.spines['top'].set_visible(False)
- ax.spines['right'].set_visible(False)
- #fig.savefig('adapter_repulsion.pdf', bbox_inches='tight')
- #show()
- # %% [markdown]
- # ## Approximations (Supplementary Information)
- # %% [markdown]
- # ### FIG. 8 Eigenvalues of W
- # %%
- mo = Network(num_neurons=801, ff_sigma=5, rec_magnitude=.85, rec_width=10)
- figure(figsize=(14,4))
- for rec_width in [.1, .5, 1., 2., 4., 6., 8.]:
- mo.rec_width = rec_width
- mo.update_W_and_M()
- lines = plot(np.arange(801), sorted( torch.linalg.eigvals(mo.recurrent_weights).abs(), reverse=True ), label=f'$\\sigma_{{rec}} = {rec_width}$' )
- 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='--')
- leg = legend(loc=(.88, .15))
- title('Eigenvalues $\\lambda_k$ of $W$')
- xlim(0,801)
- ylim(0,None)
- xlabel('$k$')
- ylabel('$\\lambda_k$')
- legend(handles=[Line2D([], [], color='grey', label='Correct eigenvalue'), Line2D([], [], color='grey', linestyle='--', label='Approximation')], loc=(.65,.5))
- gca().add_artist(leg)
- gca().spines['top'].set_visible(False)
- gca().spines['right'].set_visible(False)
- #savefig('eigenvalues_W.pdf', bbox_inches='tight')
- # %% [markdown]
- # ### FIG. 9-10 Matrix M
- # %% [markdown]
- # $$
- # M_{ij} \approx \frac{1}{1-\lambda_0} \frac{1}{n}
- # + \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)
- # $$
- # %%
- ff_sigma = -1. # shouldn't matter here
- approx_Ms = {}
- for imag, rec_magnitude in enumerate( [.2, .95, .98, .99] ):
- for iw,rec_width in enumerate( [6.,] ):
- mo = Network(num_neurons=801, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
- mo.update_W_and_M()
- # eigenvalues
- lambdas = mo.rec_magnitude * np.exp(-.5*(np.pi*(mo.rec_width)*np.arange(mo.num_neurons)/OVERALL_WIDTH)**2 )
- eigenvalues_of_M = 1./(1-lambdas)
- # matrix of eigenvectors
- # matrix_Q = (factor) np.cos( (np.pi * np.arange(mo.num_neurons) / mo.num_neurons) * ( np.arange(mo.num_neurons)[np.newaxis] + .5) )
- approx_M = np.ones_like(mo.matrix_M) * eigenvalues_of_M[0] / mo.num_neurons
- for i in range(mo.num_neurons):
- for j in range(mo.num_neurons):
- 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))
- approx_M[i,j] += (2./mo.num_neurons) * np.sum( eigenvalues_of_M[1:] * coscos )
- approx_Ms[(rec_magnitude, rec_width)] = approx_M
- # %% [markdown]
- # As an infinite sum of Gaussian:
- # $$
- # M_{ij} \approx \delta_{ij} + \delta \sum_{m=1}^{\infty} \lambda_0^m \frac{1}{\sigma_{rec} \sqrt{m}\sqrt{2 \pi}}
- # \exp{\left( -\frac{1}{2} \frac{(i-j)^2 \delta^2}{m \sigma_{rec}^2} \right)}
- # $$
- # %%
- fig, axs = subplots(nrows=3, ncols=2, figsize=(12,12))
- ff_sigma = -1. # shouldn't matter here
- rec_width = 6.
- for irec, rec_magnitude in enumerate( [0.2, 0.95, 0.99] ):
- these_axs = axs[irec]
- mo = Network(num_neurons=801, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
- mo.update_W_and_M()
- approx_M = approx_Ms[(rec_magnitude, rec_width)]
- these_axs[0].plot( (mo.matrix_M-np.eye(801))[:,40::80], lw=1)
- these_axs[1].plot( (mo.matrix_M-np.eye(801))[:,40::80], lw=1)
- these_axs[0].set_prop_cycle(None)
- these_axs[1].set_prop_cycle(None)
- # cos cos
- approx_M = approx_Ms[(rec_magnitude, rec_width)]
- these_axs[0].plot( (approx_M-np.eye(801))[:,40::80], ls='--', lw=2)
- # infinite sum of Gaussians
- number_of_gaussian = 200
- δ = OVERALL_WIDTH / mo.num_neurons
- approx_M = torch.diag(torch.ones(mo.num_neurons))
- for m in range(1, number_of_gaussian+1):
- sigma_m = np.sqrt(m) * rec_width
- 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))
- these_axs[1].plot( (approx_M-np.eye(801))[:,40::80], ls='--', lw=2)
- for ax in these_axs:
- ax.spines['top'].set_visible(False)
- ax.spines['right'].set_visible(False)
- ax.set_xlim(0,800)
- these_axs[0].set_ylabel(f'$\\lambda_0 = {rec_magnitude}$', rotation=90, fontsize=14)
- for ax in axs[0]:
- ax.set_ylim(0, .0099)
- for ax in axs[1]:
- ax.set_ylim(0, .249)
- for ax in axs[2]:
- ax.set_ylim(0, .8)
- ax.set_xlabel('$i$')
- axs[0,0].set_title('Approximation of $M-I$ with cosines')
- axs[0,1].set_title('Approximation of $M-I$ with sum of Gaussian functions')
- axs[0,0].legend(handles=[Line2D([0], [0], color='grey', lw=1, label='$M_{ij}-\\delta_{ij}$'),
- Line2D([0], [0], color='grey', lw=2, ls='--', label='Approximation')], loc='upper center', ncols=2)
- #fig.savefig('approx_M_cos_gauss.pdf', bbox_inches='tight')
- # %% [markdown]
- # As a Laplace function:
- # $$
- # 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)}
- # $$
- # where
- # $$
- # \nu_{rec} = \frac{\sigma_{rec}}{\sqrt{2\ln{(1/\lambda_0)}}}
- # $$
- # %% [markdown]
- # As a mix of both
- # %%
- fig, axs = subplots(nrows=3, ncols=2, figsize=(12,12))
- ff_sigma = -1. # shouldn't matter here
- rec_width = 6.
- for irec, rec_magnitude in enumerate( [0.2, 0.95, .99] ):
- these_axs = axs[irec]
- mo = Network(num_neurons=801, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
- mo.update_W_and_M()
- approx_M = approx_Ms[(rec_magnitude, rec_width)]
- these_axs[0].plot( (mo.matrix_M-np.eye(801))[:,40::80], lw=1)
- these_axs[1].plot( (mo.matrix_M-np.eye(801))[:,40::80], lw=1)
- these_axs[0].set_prop_cycle(None)
- these_axs[1].set_prop_cycle(None)
- # Gaussian
- approx_M_G = torch.diag(torch.ones(mo.num_neurons))
- m = 1
- sigma_m = np.sqrt(m) * rec_width
- 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))
- these_axs[0].plot( (approx_M_G-np.eye(801))[:,40::80], ls='--', lw=2)
- # Laplace
- ν_rec = rec_width / np.sqrt(2*np.log(1/rec_magnitude))
- approx_M_L = torch.diag(torch.ones(mo.num_neurons)) # delta_ij
- 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 )
- these_axs[1].plot( (approx_M_L-np.eye(801))[:,40::80], ls='--', lw=2)
- these_axs[0].set_prop_cycle(None)
- these_axs[1].set_prop_cycle(None)
- #
- these_axs[0].set_ylabel(f'$\\lambda_0 = {rec_magnitude}$', rotation=90, fontsize=14)
- for ax in these_axs:
- ax.spines['top'].set_visible(False)
- ax.spines['right'].set_visible(False)
- for ax in axs[0]:
- ax.set_ylim(0, .05)
- ax.set_yticks([0, .01, .02, .03, .04])
- for ax in axs[1]:
- ax.set_ylim(0, .27)
- for ax in axs[2]:
- ax.set_ylim(0, .65)
- ax.set_xlabel('$i$')
- for ax in axs.flat:
- ax.set_xlim(0,800)
- axs[0,0].set_title('Gaussian approximation of $M-I$')
- axs[0,1].set_title('Laplace approximation of $M-I$')
- axs[0,0].legend(handles=[Line2D([0], [0], color='grey', lw=1, label='$M_{ij}-\\delta_{ij}$'),
- Line2D([0], [0], color='grey', lw=2, ls='--', label='Approximation')], loc='center')
- #fig.savefig('approx_M_gauss_laplace.pdf', bbox_inches='tight')
- # %% [markdown]
- # ### FIG. 12 Gaussian approxº of rates, and MSE
- # %%
- num_neurons = 801
- ff_sigma = 5.
- rec_magnitude = 0.95 #0.85
- rec_width = 6. # 10.
- mo = Network(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
- mo.update_W_and_M()
- mse_per_stim_perprior = {}
- for prior_w in [10., 20., 30.]:
- samples, estimates = samples_and_estimates_cache[prior_w]
- xmin = MID_VAL-prior_w/2; xmax = MID_VAL+prior_w/2
- ind = (STIMULI>=xmin) & (STIMULI<=xmax)
- these_stims = STIMULI[ind]
- mse_per_stim = ((estimates - these_stims)**2).mean(axis=0)
- mse_per_stim_perprior[prior_w] = mse_per_stim
- # %%
- num_neurons = 801
- ff_sigma = 5.
- rec_magnitude = 0.95
- rec_width = 6.
- if dt != 0.001:
- print(f"Warning: dt = {dt} != 0.001. The MSEs may be incorrect.")
- fig, axs = subplots(nrows=2, ncols=2, figsize=(10,8))
- #
- ss = axs[1,1].get_subplotspec()
- axs[1,1].remove()
- subaxs_gs = ss.subgridspec(1, 3)
- ax11a = fig.add_subplot(subaxs_gs[0])
- ax11b = fig.add_subplot(subaxs_gs[1])
- ax11c = fig.add_subplot(subaxs_gs[2])
- subaxs = [ax11a, ax11b, ax11c]
- #
- mo = Network(num_neurons=num_neurons, ff_sigma=ff_sigma, rec_magnitude=rec_magnitude, rec_width=rec_width)
- mo.update_W_and_M()
- for ii,prior_w in enumerate([10., 20., 30.]):
- c = COLORS_WIDTH.get(prior_w)
- 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,
- alpha_cost=.5, prior_w=prior_w, regularizer_lambda=1000., inhib=0.)
- mo.gains = g_solution
- mo.update_r_star()
- gammas, gamma_den, φs = make_gamma_G_φs(mo.gains, mo.hh, mo.neurons_pref_stims)
- sigmarsSq = ff_sigma**2 + ( gammas * (mo.neurons_pref_stims-φs.unsqueeze(1))**2 ).sum(axis=1) # (num_neurons,)
- sigmars = np.sqrt(sigmarsSq)
- xmax = 60
- # rates
- if prior_w==20.:
- ax = axs[1,0]
- ax.plot(STIMULI, mo.r_star[225:576:25].T, lw=1)
- ax.set_prop_cycle(None)
- amplitudes = (ff_sigma/sigmars) * gamma_den # (num_neurons,)
- approx_r_2 = amplitudes.unsqueeze(1) * np.exp(-(STIMULI-φs.unsqueeze(1))**2 / (2*sigmarsSq.unsqueeze(1)))
- ax.plot(STIMULI, approx_r_2[225:576:25].T, ls='--', lw=2);
- ax.legend(handles=[Line2D([0], [0], color='grey', lw=1, label='Tuning curve $r_i(s)$'),
- Line2D([0], [0], color='grey', lw=2, ls='--', label='Gaussian approximation')], loc='upper left', ncols=1)
- ax.set_title('Gaussian approximations of tuning curves $r_i(s)$')
- ax.set_xlabel('Stimulus $s$')
- ax.set_ylabel('Rate $r_i(s)$')
- ax.set_ylim(0, .54)
- ax.set_xlim(-xmax, xmax)
- # phi
- ax = axs[0,0]
- true_sstar, true_sigmarsSq = mo.effective_tc_center_and_squaredwidth()
- ax.plot(mo.neurons_pref_stims, true_sstar, c=c, zorder=200, label={10:'Narrow', 20:'Medium', 30:'Wide'}.get(prior_w, '') )
- ax.plot(mo.neurons_pref_stims, φs, '--', c=c, zorder=100, )
- ax.axline((0,0), slope=1, c='lightgrey', zorder=-100)
- ax.set_ylim(-xmax, xmax)
- ax.set_xlim(-xmax, xmax)
- ax.set_xlabel('Feedforward preferred stimulus $s_i$')
- ax.set_ylabel('Effective preferred stimulus $\\varphi(s_i)$')
- ax.set_title('Effective preferred stimulus $\\varphi(s_i)$')
- # sigmaSq
- ax = axs[0,1]
- true_sigmars = np.sqrt(true_sigmarsSq)
- ax.plot(mo.neurons_pref_stims, true_sigmars, c=c, zorder=200)
- ax.plot(mo.neurons_pref_stims, sigmars, c=c, ls='--', zorder=100)
- ax.set_ylim(0, 25)
- ax.set_xlim(-xmax, xmax)
- ax.set_xlabel('Feedforward preferred stimulus $s_i$')
- ax.set_ylabel('Effective width $\\sigma_r(s_i)$')
- ax.set_title('Effective width $\\sigma_r(s_i)$')
- ax.legend()
- # MSE
- ax = subaxs[ii]
- xmin = MID_VAL-.5*prior_w; xmax = MID_VAL+.5*prior_w
- ind = (STIMULI>=xmin) & (STIMULI<=xmax)
- these_stims = STIMULI[ind]
- mse_per_stim = mse_per_stim_perprior[prior_w]
- ###
- approx_mse_per_stim = 1./(1./prior_var + (1./mo.beta)*(mo.r_star.T/true_sigmarsSq).sum(axis=1))
- ###
- ax.plot(these_stims, np.sqrt(mse_per_stim), c=c, label='$\\sqrt{MSE}$')
- ax.plot(these_stims, np.sqrt(approx_mse_per_stim[ind]), c=c, ls='--', label='Approximº', lw=2)
- ax.set_xticks([-prior_w/2, 0, prior_w/2])
- ax.set_ylim(0, 11.9)
- ax.spines['top'].set_visible(False)
- ax.spines['right'].set_visible(False)
- if ii > 0:
- ax.set_yticks([])
- subaxs[0].set_ylabel('$\\sqrt{MSE}$')
- subaxs[1].set_xlabel('Stimulus $s$')
- subaxs[1].set_title('Square-root of MSE')
- subaxs[0].legend(loc=(.05, .8))
- leg = axs[0,0].legend()
- axs[0,0].legend( handles=[Line2D([0], [0], color='grey', lw=1, label='$\\varphi(s_i)$'),
- Line2D([0], [0], color='grey', lw=2, ls='--', label='Approximation $\\sum \\gamma_{ij} s_j$')], loc='lower right')
- axs[0,0].add_artist(leg)
- for ax in axs.flat:
- ax.spines['top'].set_visible(False)
- ax.spines['right'].set_visible(False)
- fig.subplots_adjust(hspace=.35)
- #fig.savefig('approx_rates_mses.pdf', bbox_inches='tight')
- # %%
- # %%
model.ipynb, no license · at the source
Overview
- Department of Psychology and Center for Brain Science, Harvard University,Cambridge, MA USA
- Neuroscience Center Zurich, University of Zurich and ETH Zurich,Zurich, Switzerland
- TUM School of Computation, Information and Technology, Technical University of Munich,Garching, Germany
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
Availability: 1 check, the latest on 28 September 2026: the link answers (HTTP 200)
- 28 September 2026: the link answers (HTTP 200)
2 files
- model.ipynb, Jupyter, 1,850 lines, 4 matches
- README.md, Text, 22 lines
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:
- it points to the authors' code: OSF wa7hm
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://
BibTeX
@article{pratcarrabin202
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/
url = {https://
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/
VL - 17
IS - 1
SP - 5554
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1038/
"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":
"volume": "17",
"issue": "1",
"page": "5554",
"DOI": "10.1038/
"PMID": "42140911",
"PMCID": "PMC13291360",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://
"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: eLifeIn 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 advancesIn 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: eLifeIn 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 advancesIn 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 neuroscienceIn 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 biologyIn 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 dataIn 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 biologyIn 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 healthIn 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: NatureIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 1 script, and 4 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:29458952d1d922e3…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[.
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.
