OSCR

A genetic algorithm for self-supervised models of oscillatory neurodynamics.

Code ↔ Paper

7 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 7 matches
  1. [1] § Materials and methods ↔ misc/Jaxley_Mechanisms.ipynb, lines 473–522 · score 0.65 · stochastic delta rule, model parameters, 0–1, logic, genetic, supervised
  2. [2] § Materials and methods ↔ gsdr/optimizers.py, lines 222–352 · score 0.56 · deselection threshold, MCDP factors, lambda, exploration, stochastic, optimization
  3. [3] § Materials and methods ↔ misc/Jaxley_Mechanisms.ipynb, lines 524–564 · score 0.55 · exploration factor, deselection threshold, reverts, optimization, training, loss
  4. [4] § Results ↔ Biophys_SX.ipynb, lines 3117–3178 · score 0.55 · beta band power, gamma band power, synaptic weight, connectivity, model
  5. [5] § Materials and methods ↔ misc/Jaxley_Mechanisms.ipynb, lines 473–522 · score 0.54 · Genetic Stochastic Delta, stochastic delta rule, model parameter, deselects, dynamics, optimization
  6. [6] § Materials and methods ↔ gsdr/analysis.py, lines 85–96 · score 0.53 · Mutual correlation dependent, plasticity
  7. [7] § Materials and methods ↔ misc/Jaxley_Mechanisms.ipynb, lines 49–139 · score 0.52 · pre synaptic, post synaptic, tau, activity, spike, model

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 · 638 lines · 21 KB · no license · 4 matches

  1. # %% [markdown]
  2. # <a href="https://colab.research.google.com/github/HNXJ/GSDR/blob/main/Jaxley_Mechanisms.ipynb" target="_parent"><img src="https://colab.research.google.com/assets/colab-badge.svg" alt="Open In Colab"/></a>
  3. # %% [markdown]
  4. # Jaxley Biophysical model mechanisms
  5. #
  6. # This notebook contains definitions of biophysical mechanisms to be used in Jaxley https://jaxley.readthedocs.io/en/
  7. #
  8. # @HNyXJ (Assisted by Vanderbilt's amplify AI and Google's Gemini)
  9. #
  10. # Mechanisms included or to be implemented in this notebook :
  11. #
  12. # GABAa
  13. #
  14. # GABAb
  15. #
  16. # AMPA
  17. #
  18. # NMDA
  19. #
  20. # ACh
  21. #
  22. # D1
  23. #
  24. # D2
  25. #
  26. # 5HT
  27. #
  28. # %%
  29. %pip install jaxley
  30. from jax import config
  31. config.update("jax_enable_x64", True)
  32. config.update("jax_platform_name", "cpu")
  33. import matplotlib.pyplot as plt
  34. import numpy as np
  35. import jax
  36. import jax.numpy as jnp
  37. from jax import jit, vmap, value_and_grad
  38. import jaxley as jx
  39. from jaxley.channels import Leak, HH
  40. from jaxley.synapses import IonotropicSynapse
  41. from jaxley.connect import fully_connect
  42. # %% [markdown]
  43. # # Neuronal mechanisms
  44. # %% [markdown]
  45. # ## GABAa
  46. # %%
  47. import jax.numpy as jnp
  48. from jaxley.synapses import Synapse
  49. class GradedGABAa(Synapse):
  50. """
  51. A graded (non-spiking) GABAa synapse model based on high-threshold
  52. graded transmission dynamics.
  53. Unlike standard exponential synapses which are triggered by spike events,
  54. this mechanism's gating variable 's' is continuously driven by the
  55. presynaptic voltage via a hyperbolic tangent transfer function.
  56. Dynamics:
  57. ds/dt = -s/tauD + (1/2) * (1 + tanh((V_pre - V_th) / slope)) * ((1-s)/tauR)
  58. I_syn = gGABAa * s * (V_post - EGABAa)
  59. Parameters:
  60. gGABAa (float): Peak synaptic conductance [uS or mS/cm^2 depending on context].
  61. EGABAa (float): Reversal potential [mV]. Default: -80.0.
  62. tauD (float): Decay time constant [ms]. Default: 10.0.
  63. tauR (float): Rise time constant [ms]. Default: 0.2.
  64. V_th (float): Half-activation voltage [mV]. Default: 0.0 (centered tanh).
  65. slope (float): Sensitivity of the activation curve [mV]. Default: 10.0.
  66. References:
  67. 1. Golowasch, J., Casey, M., Abbott, L. F., & Marder, E. (1999).
  68. Network stability from neuronal insight. Journal of Neurobiology, 41(3), 331-348.
  69. 2. Prinz, A. A., Bucher, D., & Marder, E. (2004).
  70. Similar network activity from disparate circuit parameters.
  71. Nature Neuroscience, 7(12), 1345-1352.
  72. """
  73. def __init__(self, name: str = "GradedGABAa"):
  74. super().__init__(name)
  75. # Parameter definitions matching DynaSim defaults
  76. self.synapse_params = {
  77. "gGABAa": 0.25,
  78. "EGABAa": -80.0,
  79. "tauD": 10.0,
  80. "tauR": 0.2,
  81. "slope": 10.0, # Denominator inside the tanh
  82. "V_th": 0.0 # Midpoint of the tanh (implicit 0 in DynaSim)
  83. }
  84. # Initial Condition (IC) matching DynaSim IC=[0.1]
  85. self.synapse_states = {"s": 0.1}
  86. def update_states(self, states, dt, pre_v, post_v, params):
  87. """
  88. Updates the gating variable 's' based on PRE-synaptic voltage.
  89. """
  90. s = states["s"]
  91. tauD = params["tauD"]
  92. tauR = params["tauR"]
  93. slope = params["slope"]
  94. v_th = params["V_th"]
  95. # The transfer function (Activation)
  96. # Corresponds to DynaSim: 1/2 * (1 + tanh(X_pre / 10))
  97. # We added V_th to make it more robust, but it defaults to 0.
  98. activation = 0.5 * (1 + jnp.tanh((pre_v - v_th) / slope))
  99. # Differential equation:
  100. # s' = -s/tauD + activation * (1-s)/tauR
  101. d_s = (-s / tauD) + activation * ((1 - s) / tauR)
  102. # Forward Euler integration
  103. new_s = s + d_s * dt
  104. return {"s": new_s}
  105. def compute_current(self, states, pre_v, post_v, params):
  106. """
  107. Calculates the synaptic current flowing into the POST-synaptic cell.
  108. """
  109. s = states["s"]
  110. g = params["gGABAa"]
  111. e_rev = params["EGABAa"]
  112. # Ohm's Law for the synapse
  113. current = g * s * (post_v - e_rev)
  114. return current
  115. # %%
  116. import jaxley as jx
  117. # Instantiate the channel
  118. gaba_mech = GABAa()
  119. # Create a cell and add the mechanism
  120. cell = jx.Cell()
  121. cell.insert(gaba_mech)
  122. # If you need to modify parameters specifically for this cell:
  123. cell.GABAa.gGABAa = 0.5 # Override default conductance
  124. # %%
  125. net.delete_recordings()
  126. net.cell(0).branch(0).loc(0.0).record()
  127. net.cell(1).branch(0).loc(0.0).record()
  128. net.cell(2).branch(0).loc(0.0).record()
  129. # %%
  130. inputs = jnp.asarray(np.random.rand(100, 2))
  131. labels = jnp.asarray((inputs[:, 0] + inputs[:, 1]) > 1.0)
  132. # %%
  133. fig, ax = plt.subplots(1, 1, figsize=(3, 2))
  134. _ = ax.scatter(inputs[labels, 0], inputs[labels, 1])
  135. _ = ax.scatter(inputs[~labels, 0], inputs[~labels, 1])
  136. # %%
  137. labels = labels.astype(float)
  138. # net.edges
  139. # %%
  140. net.delete_trainables()
  141. net.make_trainable("radius")
  142. net.cell("all").branch("all").loc("all").make_trainable("Leak_gLeak")
  143. net.IonotropicSynapse.edge("all").make_trainable("IonotropicSynapse_gS")
  144. # %% [markdown]
  145. #
  146. #
  147. # * II
  148. #
  149. # %%
  150. params = net.get_parameters()
  151. s = jx.integrate(net, params=params, t_max=5.0)
  152. # %%
  153. def simulate(params, inputs):
  154. currents = jx.datapoint_to_step_currents(i_delay=10.0, i_dur=40.0, i_amp=10*inputs, delta_t=0.025, t_max=100.0)
  155. data_stimuli = None
  156. data_stimuli = net.cell(0).branch(2).loc(1.0).data_stimulate(currents[0], data_stimuli=data_stimuli)
  157. data_stimuli = net.cell(1).branch(2).loc(1.0).data_stimulate(currents[1], data_stimuli=data_stimuli)
  158. # data_stimuli = net.cell(3).branch(2).loc(1.0).data_stimulate(currents[1], data_stimuli=data_stimuli)
  159. # data_stimuli = net.cell(4).branch(2).loc(1.0).data_stimulate(currents[1], data_stimuli=data_stimuli)
  160. return jx.integrate(net, params=params, data_stimuli=data_stimuli, delta_t=0.025)
  161. batched_simulate = vmap(simulate, in_axes=(None, 0))
  162. # %%
  163. traces = batched_simulate(params, inputs[:4])
  164. fig, ax = plt.subplots(1, 1, figsize=(4, 2))
  165. _ = ax.plot(traces[:, 2, :].T)
  166. # %%
  167. def rasterPlot(params, inputs):
  168. """
  169. Plots a raster from the output of batched_simulate.
  170. Args:
  171. params: The network parameters.
  172. inputs: The input data for simulation (expected to be a batch, e.g., inputs[0:1]).
  173. """
  174. # Get simulation traces for the provided input batch.
  175. # If inputs has shape (1, D), then traces_batch will be (1, num_recordings, timepoints).
  176. traces_batch = batched_simulate(params, inputs)
  177. # Select all neuron traces for the first input in the batch.
  178. all_neuron_traces = traces_batch[0, :, :] # Shape: (num_recordings, timepoints)
  179. # Define simulation time parameters (from cell xPLGrxkxUVUi)
  180. t_max = 100.0
  181. dt = 0.025
  182. time_axis = jnp.arange(0, t_max, dt)
  183. # Spike detection: detect when voltage crosses a threshold from below
  184. spike_threshold = -20.0 # Assuming a spike threshold of -20mV
  185. spike_times_list = []
  186. num_neurons = all_neuron_traces.shape[0]
  187. for i in range(num_neurons): # Iterate over each neuron (recording)
  188. neuron_trace = all_neuron_traces[i]
  189. # Find indices where voltage crosses threshold from below
  190. spikes = (neuron_trace[:-1] < spike_threshold) & (neuron_trace[1:] >= spike_threshold)
  191. spike_indices = jnp.where(spikes)[0]
  192. spike_times = time_axis[spike_indices + 1] # +1 because we are checking neuron_trace[1:]
  193. spike_times_list.append(spike_times)
  194. # Plotting the raster
  195. fig, ax = plt.subplots(1, 1, figsize=(10, 5))
  196. for i, spk_times in enumerate(spike_times_list):
  197. ax.vlines(spk_times, i - 0.4, i + 0.4, colors='blue') # Plot vertical lines for each spike
  198. ax.set_xlabel("Time (s)")
  199. ax.set_ylabel("Neuron Index") # Changed label to reflect all neurons
  200. ax.set_title("Raster Plot of All Neurons for the First Input") # Changed title
  201. ax.set_yticks(jnp.arange(num_neurons)) # Set y-ticks to correspond to neuron indices
  202. ax.set_ylim(-0.5, num_neurons - 0.5)
  203. ax.set_xlim(0, t_max)
  204. plt.grid(axis='x', linestyle='--', alpha=0.7)
  205. plt.tight_layout()
  206. plt.show()
  207. # %%
  208. # Call the rasterPlot function with the trained parameters and only the first input (as a batch)
  209. rasterPlot(final_params, inputs[0:1])
  210. # %%
  211. t_max = 100.0
  212. dt = 0.025
  213. levels = 2
  214. time_points = t_max // dt + 2
  215. checkpoints = [int(np.ceil(time_points**(1/levels))) for _ in range(levels)]
  216. def simulate(params, inputs):
  217. currents = jx.datapoint_to_step_currents(i_delay=1.0, i_dur=1.0, i_amp=inputs / 10.0, delta_t=dt, t_max=t_max)
  218. data_stimuli = None
  219. data_stimuli = net.cell(0).branch(2).loc(1.0).data_stimulate(currents[0], data_stimuli=data_stimuli)
  220. data_stimuli = net.cell(1).branch(2).loc(1.0).data_stimulate(currents[1], data_stimuli=data_stimuli)
  221. return jx.integrate(net, params=params, data_stimuli=data_stimuli, checkpoint_lengths=checkpoints)
  222. batched_simulate = vmap(simulate, in_axes=(None, 0))
  223. def predict(params, inputs):
  224. traces = simulate(params, inputs) # Shape `(batchsize, num_recordings, timepoints)`.
  225. prediction = jnp.mean(traces[2]) # Use the average over time of the output neuron (2) as prediction.
  226. return prediction + 72.0 # Such that the prediction is around 0.
  227. batched_predict = vmap(predict, in_axes=(None, 0))
  228. def predict_2(params, inputs):
  229. """
  230. Calculates the mean Power Spectral Density (PSD) of the average signal of all neurons
  231. response for frequencies in the range [0-100Hz].
  232. """
  233. traces = simulate(params, inputs) # Shape `(num_recordings, timepoints)` when called by vmap
  234. signal = jnp.mean(traces, axis=0) # Average across all neurons (axis=0)
  235. N = signal.shape[-1] # Number of time points
  236. fs = 1.0 / dt # Sampling frequency
  237. # Compute one-sided FFT and corresponding frequencies for real signals
  238. signal_fft = jnp.fft.rfft(signal)
  239. freqs = jnp.fft.rfftfreq(N, d=dt)
  240. # Compute Power Spectral Density (PSD)
  241. # PSD = (1/(N*fs)) * |FFT(signal)|^2
  242. psd = (1.0 / (N * fs)) * jnp.abs(signal_fft)**2
  243. # Filter for frequencies in the range [0-100Hz]
  244. mask = (freqs >= 0) & (freqs <= 100)
  245. filtered_psd = psd[mask]
  246. # Return the mean of the filtered PSD as a scalar prediction.
  247. # Handle case where filtered_psd might be empty to avoid error in jnp.mean.
  248. return jnp.mean(filtered_psd) if filtered_psd.size > 0 else 0.0
  249. batched_predict_2 = vmap(predict_2, in_axes=(None, 0))
  250. def loss(opt_params, inputs, labels):
  251. params = transform.forward(opt_params)
  252. # Use the new predict_2 function for loss calculation
  253. predictions = batched_predict(params, inputs)
  254. losses = jnp.abs(predictions - labels) # Mean absolute error loss.
  255. return jnp.mean(losses) # Average across the batch.
  256. jitted_grad = jit(value_and_grad(loss, argnums=0))
  257. # %%
  258. params
  259. # %%
  260. import jaxley.optimize.transforms as jt
  261. # The structure passed to `jx.ParamTransform` should match the structure of `params`.
  262. transform = jx.ParamTransform([
  263. {"radius": jt.SigmoidTransform(0.1, 5.0)},
  264. {"Leak_gLeak":jt.SigmoidTransform(1e-5, 1e-3)},
  265. {"IonotropicSynapse_gS" : jt.SigmoidTransform(1e-5, 1e-2)}
  266. ])
  267. opt_params = transform.inverse(params)
  268. # %%
  269. jitted_grad = jit(value_and_grad(loss, argnums=0))
  270. value, gradient = jitted_grad(params, inputs[:4], labels[:4])
  271. # %%
  272. import optax
  273. key = jax.random.PRNGKey(42)
  274. initial_params = net.get_parameters()
  275. # Inner optimizer (Adam) handles the gradient descent part
  276. optimizer_inner = optax.adam(learning_rate=0.01)
  277. # Create GSDR wrapper
  278. optimizer = GSDR.GSDR(
  279. inner_optimizer=optimizer_inner,
  280. delta_distribution=jax.random.normal,
  281. deselection_threshold=2.0,
  282. a_init=0.4,
  283. a_dynamic=True
  284. )
  285. # Initialize State
  286. opt_state = optimizer.init(initial_params)
  287. # %%
  288. class Dataset:
  289. """A simple Dataloader which returns batches of the data.
  290. Instead of using this simple dataloader, you can also just use one from
  291. PyTorch or Tensorflow. You do not have to understand what is going on here
  292. to follow this tutorial.
  293. """
  294. def __init__(self, inputs: np.ndarray, labels: np.ndarray):
  295. """Initialize the dataloader.
  296. Args:
  297. inputs: Array of shape (num_samples, num_dim)
  298. labels: Array of shape (num_samples,)
  299. """
  300. assert len(inputs) == len(labels), "Inputs and labels must have same length"
  301. self.inputs = inputs
  302. self.labels = labels
  303. self.num_samples = len(inputs)
  304. self._rng_state = None
  305. self.batch_size = 1
  306. def shuffle(self, seed=None):
  307. """Shuffle the dataset in-place"""
  308. self._rng_state = np.random.get_state()[1][0] if seed is None else seed
  309. np.random.seed(self._rng_state)
  310. indices = np.random.permutation(self.num_samples)
  311. self.inputs = self.inputs[indices]
  312. self.labels = self.labels[indices]
  313. return self
  314. def batch(self, batch_size):
  315. """Create batches of the data."""
  316. self.batch_size = batch_size
  317. return self
  318. def __iter__(self):
  319. """Iterate over the dataset."""
  320. self.shuffle(seed=self._rng_state)
  321. for start in range(0, self.num_samples, self.batch_size):
  322. end = min(start + self.batch_size, self.num_samples)
  323. yield self.inputs[start:end], self.labels[start:end]
  324. self._rng_state += 1
  325. # %%
  326. batch_size = 4
  327. dataloader = Dataset(inputs, labels)
  328. dataloader = dataloader.shuffle(seed=0).batch(batch_size)
  329. # --- 5. Loop ---
  330. key = jax.random.PRNGKey(0)
  331. print("Starting training...")
  332. for epoch in range(20):
  333. key, step_key = jax.random.split(key)
  334. epoch_loss = 0.0
  335. for batch_ind, batch in enumerate(dataloader):
  336. current_batch, label_batch = batch
  337. loss_val, gradient = jitted_grad(opt_params, current_batch, label_batch)
  338. updates, opt_state = optimizer.update(gradient, opt_state,
  339. params=params, # Required for GSDR
  340. value=loss_val, # Required for GSDR
  341. key=step_key) # Required for GSDR
  342. opt_params = optax.apply_updates(opt_params, updates)
  343. epoch_loss += loss_val
  344. print(f"epoch {epoch}, loss {epoch_loss}, alpha {opt_state.a}")
  345. final_params = transform.forward(opt_params)
  346. # %%
  347. ntest = 32
  348. # predictions = batched_predict(final_params, inputs[:4])
  349. # %%
  350. fig, ax = plt.subplots(1, 1, figsize=(3, 2))
  351. _ = ax.scatter(labels[:ntest], predictions)
  352. _ = ax.set_xlabel("Label")
  353. _ = ax.set_ylabel("Prediction")
  354. # %%
  355. traces = batched_simulate(final_params, inputs[:4])
  356. fig, ax = plt.subplots(1, 1, figsize=(4, 2))
  357. _ = ax.plot(traces[:, 2, :].T)
  358. # %% [markdown]
  359. # # GSDR optimizer (optax standard format)
  360. # %%
  361. import jax
  362. import jax.numpy as jnp
  363. import optax
  364. from flax.struct import dataclass
  365. from typing import Any, Callable, NamedTuple, Optional
  366. # State (using flax dataclass for JIT compatibility)
  367. @dataclass
  368. class GSDRState:
  369. inner_state: Any # State of the inner optimizer (e.g., SGD, AdaGrad, Adam ... state)
  370. params_opt: Any # Best parameters so far (Optimal parameters)
  371. inner_state_opt: Any # Inner optimizer state corresponding to the optimal parameters
  372. loss_opt: float # Optimal loss
  373. a: float # Current self-supervision factor (alpha)
  374. a_opt: float # Optimal self-supervision factor
  375. def GSDR(
  376. inner_optimizer: optax.GradientTransformation,
  377. delta_distribution: Callable = jax.random.normal,
  378. deselection_threshold: float = 10.0,
  379. a_init: float = 0.5,
  380. a_dynamic: bool = True
  381. ) -> optax.GradientTransformation:
  382. """
  383. Optax-compliant implementation of the Genetic-Stochastic Delta Rule.
  384. Args:
  385. inner_optimizer: The gradient-based optimizer (e.g., optax.adam).
  386. delta_distribution: Function (key, shape) -> tensor for generating noise.
  387. deselection_threshold: Threshold factor to trigger genetic deselection.
  388. a_init: Initial self-supervision factor (0 to 1).
  389. a_dynamic: Whether 'a' should be stochastic/learnable.
  390. Returns:
  391. An optax.GradientTransformation (init_fn, update_fn).
  392. """
  393. def init_fn(params):
  394. inner_state = inner_optimizer.init(params)
  395. return GSDRState(
  396. inner_state=inner_state,
  397. params_opt=params,
  398. inner_state_opt=inner_state,
  399. loss_opt=jnp.inf,
  400. a=a_init,
  401. a_opt=a_init
  402. )
  403. def update_fn(updates, state, params=None, value=None, key=None):
  404. """
  405. Args:
  406. updates: Gradients from loss_fn (standard Optax naming).
  407. state: Current GSDRState.
  408. params: Current model parameters (Required).
  409. value: Current Loss value (Required for GSDR logic).
  410. key: JAX PRNGKey (Required for stochastic Delta).
  411. """
  412. if params is None:
  413. raise ValueError("GSDR requires 'params' to be passed to update().")
  414. if value is None:
  415. raise ValueError("GSDR requires current loss 'value' to be passed to update().")
  416. if key is None:
  417. raise ValueError("GSDR requires a random 'key' to be passed to update().")
  418. grads = updates
  419. loss = value
  420. # Split keys for delta noise and 'a' (exploration factor)
  421. delta_key, a_key = jax.random.split(key)
  422. # --- 1. Genetic Logic (Selection & Deselection) ---
  423. # Optimal loss selection
  424. is_new_opt = loss < state.loss_opt
  425. # Update Optimal State Candidates
  426. new_params_opt = jax.tree.map(
  427. lambda cur, opt: jnp.where(is_new_opt, cur, opt),
  428. params, state.params_opt
  429. )
  430. new_loss_opt = jnp.where(is_new_opt, loss, state.loss_opt)
  431. new_a_opt = jnp.where(is_new_opt, state.a, state.a_opt)
  432. # Keep the inner optimizer optimal state
  433. new_inner_state_opt = jax.tree.map(
  434. lambda cur, opt: jnp.where(is_new_opt, cur, opt),
  435. state.inner_state, state.inner_state_opt
  436. )
  437. # Deselection (Backtrack from the Catastrophic Failure)
  438. # If loss > threshold * best_loss, revert back to the optimal state
  439. # Exclude the case where loss_opt is infinity (start of training)
  440. is_deselect = (loss > (new_loss_opt * deselection_threshold)) & (new_loss_opt != jnp.inf)
  441. # --- 2. Determine Next Step Variables ---
  442. # If Deselecting: Revert 'a' to 'a_opt'. Else: Explore new 'a' (if a is dynamic)
  443. if a_dynamic:
  444. a_random = jax.random.uniform(a_key, minval=0.0, maxval=1.0)
  445. next_a = jnp.where(is_deselect, new_a_opt, a_random)
  446. else:
  447. next_a = state.a # Constant
  448. # If Deselecting: Revert inner_state to optimal. Else: Keep current.
  449. next_inner_state = jax.tree.map(
  450. lambda opt, cur: jnp.where(is_deselect, opt, cur),
  451. new_inner_state_opt, state.inner_state
  452. )
  453. # --- 3. Calculate Updates ---
  454. # A. Inner Optimizer Update (Gradient Descent)
  455. # Use the *potentially reverted* inner state
  456. inner_updates, updated_inner_state = inner_optimizer.update(grads, next_inner_state, params)
  457. # B. Stochastic Delta Update
  458. # Generate noise matching params structure
  459. param_leaves, treedef = jax.tree_util.tree_flatten(params)
  460. subkeys = jax.random.split(delta_key, len(param_leaves))
  461. param_keys_tree = jax.tree_util.tree_unflatten(treedef, subkeys)
  462. delta_noise = jax.tree.map(
  463. lambda p, k: delta_distribution(k, p.shape),
  464. params, param_keys_tree
  465. )
  466. # Scale noise by parameter magnitude (per paper/pseudocode)
  467. delta = jax.tree.map(lambda n, p: n * p, delta_noise, params)
  468. # C. Combine Updates: a * Grads + (1-a) * Delta
  469. # Note: 'next_a' is the 'a' for THIS step.
  470. combined_updates = jax.tree_map(
  471. lambda d, g: next_a * d + (1 - next_a) * g,
  472. delta, inner_updates
  473. )
  474. # --- 4. Handle Reset (The Revert Step) ---
  475. # If is_deselect is True, we want the FINAL params to be params_opt.
  476. # Optax applies: params_new = params + final_updates
  477. # So if reset: params_new = params_opt
  478. # Therefore: params + reset_update = params_opt
  479. # reset_update = params_opt - params
  480. reset_updates = jax.tree.map(
  481. lambda opt, cur: opt - cur,
  482. new_params_opt, params
  483. )
  484. # Select between Reset Update or Calculated Update
  485. final_updates = jax.tree.map(
  486. lambda reset, calc: jnp.where(is_deselect, reset, calc),
  487. reset_updates, combined_updates
  488. )
  489. # If deselected, must NOT advance the inner optimizer state
  490. # (use the reverted state).
  491. # If didn't deselect, use the state returned by inner_optimizer.update
  492. final_inner_state = jax.tree.map(
  493. lambda reset_st, advanced_st: jnp.where(is_deselect, reset_st, advanced_st),
  494. new_inner_state_opt, updated_inner_state
  495. )
  496. # Create new GSDR state
  497. new_state = GSDRState(
  498. inner_state=final_inner_state,
  499. params_opt=new_params_opt,
  500. inner_state_opt=new_inner_state_opt,
  501. loss_opt=new_loss_opt,
  502. a=next_a,
  503. a_opt=new_a_opt
  504. )
  505. return final_updates, new_state
  506. return optax.GradientTransformation(init_fn, update_fn)
  507. # %%
  508. from google.colab import drive
  509. drive.mount('/content/drive')
  510. from drive.MyDrive.Colab import GSDR

Jaxley_Mechanisms.ipynb at commit b4a96fa, no license · at the source

Overview

Authors: Hamed Nejat1, Jason Sherfey2, André M Bastos1,3
  1. Department of Psychology, Vanderbilt University, Nashville, Tennessee, United States of America
  2. Department of Psychological and Brain Sciences, Boston University, Boston, Massachusetts, United States of America
  3. Vanderbilt Brain Institute, Vanderbilt University, Nashville, Tennessee, United States of America
Institutions: Vanderbilt University (United States); Boston University (United States)
Journal: PloS one, volume 21, issue 8, article e0354021
Dates: received 23 March 2026; accepted 2 July 2026; published online 5 August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pone.0354021 · PMID 42555667 · PMCID PMC13440876 · OpenAlex W7172522179
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: computational modeling (no new data) (modality), systems (subfield)
Methods: Machine learning, Spectral & time-frequency, Single-unit activity, calcium imaging
MeSH: Genetic Algorithms*, Models, Neurological*, Action Potentials, Animals, Computer Simulation, Neurons, Stochastic Processes, Visual Cortex (* major topic)
Topic: Neural dynamics and brain function (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: Vanderbilt University Startup Funds; NIMH NIH HHS (R00 MH116100); National Institute of Mental Health (R00MH11610); National Science Foundation (2449277)
Citations: not cited yet (Europe PMC); 132 references in the paper

Abstract

Predictive processing theories propose that the brain builds internal models of its environment by reducing the discrepancy between internally generated predictions and external sensory signals. Prior work has linked these processes to oscillatory activity in gamma (40–100 Hz) and alpha/beta (10–30 Hz) frequency ranges. Current computational approaches face a trade-off: abstract predictive-processing models can implement self-supervised computations but often omit oscillatory spiking dynamics, whereas biophysically constrained spiking models can generate neural rhythms but often require extensive manual tuning. Here, we introduce the Genetic Stochastic Delta Rule (GSDR), an evolutionary optimization framework for fitting nonlinear neural models to electrophysiological objectives. We first evaluate GSDR in simplified optimization settings, then apply it to spiking-network objectives involving firing rates, beta/gamma spectral ratios, and empirical macaque stimulus-evoked gamma dynamics from visual cortex. We show that GSDR can search constrained synaptic parameter spaces, reduce reliance on manual tuning, and reproduce spectral and circuit-level phenotypes associated with predictive routing. We also used Izhikevich simulations as a model-class robustness analysis, showing that the approach is not limited to the original Hodgkin-Huxley-style implementation. These results position GSDR as a methodological framework for automated, multi-objective exploration of oscillatory neural models, with predictive routing serving as a motivating case study rather than as a completed functional proof.

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

Repositories

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

HNXJ/GSDR

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: b4a96fae5a4d00cb15889168dfc6aca76c3037ee, 6 July 2026
Languages: Python (12), Jupyter (5)
Size: 19 files, 17 scripts
Software Heritage: not archived
Found in: the supplementary material
Holds: README, 5 notebooks
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (12 files), JAX (11 files), Matplotlib (8 files), SciPy (5 files)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
18 files
At the source: github.com/HNXJ/GSDR

supp:PMC13440876/pone.0354021.s002.zip

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Languages: Python (12), Jupyter (5)
Size: 18 files, 17 scripts
Software Heritage: not checked
Found in: the supplementary material
Holds: README, 5 notebooks
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (12 files), JAX (11 files), Matplotlib (8 files), SciPy (5 files)
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)
18 files

The paper's code and data availability statement is in the Data section.

Tracing map

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

What the map holds:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 34 scripts, each with its path and the digest of its content;
  • 7 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

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

Data

No dataset and no data link were found in the paper.

Data Availability

https://github.com/HNXJ/GSDR https://github.com/DynaSim/DynaSim/tree/devDL.

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

Versions

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

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 3 authors, 8 MeSH terms, 4 funders, 129 references.

Cite

This paper

Nejat, H., Sherfey, J., & Bastos, A. M. (2026). A genetic algorithm for self-supervised models of oscillatory neurodynamics. PloS one, 21(8), e0354021. https://doi.org/10.1371/journal.pone.0354021

BibTeX

@article{nejat2026genetic,
author = {Nejat, Hamed and Sherfey, Jason and Bastos, André M},
title = {{A genetic algorithm for self-supervised models of oscillatory neurodynamics}},
journal = {PloS one},
year = {2026},
month = aug,
volume = {21},
number = {8},
pages = {e0354021},
publisher = {PLOS},
issn = {1932-6203},
doi = {10.1371/journal.pone.0354021},
url = {https://doi.org/10.1371/journal.pone.0354021},
pmid = {42555667},
pmcid = {PMC13440876}
}

RIS

TY - JOUR
AU - Nejat, Hamed
AU - Sherfey, Jason
AU - Bastos, André M
TI - A genetic algorithm for self-supervised models of oscillatory neurodynamics
T2 - PloS one
J2 - PLoS One
PY - 2026
DA - 2026/08/05
VL - 21
IS - 8
SP - e0354021
SN - 1932-6203
PB - PLOS
DO - 10.1371/journal.pone.0354021
UR - https://doi.org/10.1371/journal.pone.0354021
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pone.0354021",
"type": "article-journal",
"title": "A genetic algorithm for self-supervised models of oscillatory neurodynamics",
"container-title": "PloS one",
"author": [
{
"family": "Nejat",
"given": "Hamed"
},
{
"family": "Sherfey",
"given": "Jason"
},
{
"family": "Bastos",
"given": "André M"
}
],
"container-title-short": "PLoS One",
"volume": "21",
"issue": "8",
"page": "e0354021",
"DOI": "10.1371/journal.pone.0354021",
"PMID": "42555667",
"PMCID": "PMC13440876",
"ISSN": "1932-6203",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pone.0354021",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
5
]
]
}
}

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

Similar papers

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

[1] doi:10.1371/journal.pcbi.1014164 [code]
'Backpropagation and the brain' realized in cortical error neuron microcircuits.
Journal: PLoS computational biology
In common: SciPy, Matplotlib, NumPy, computational modeling (no new data), 6 references
[2] doi:10.1038/s41593-026-02345-6 [code]
Human hippocampal ripples tune cortical responses based on predicted uncertainty.
Journal: Nature neuroscience
In common: 8 references
[3] doi:10.1002/advs.77857 [code]
Brain Network Dynamics of Local and Global Predictive Processing in Aging.
Journal: Advanced science (Weinheim, Baden-Wurttemberg, Germany)
In common: NumPy, 7 references
[4] doi:10.1371/journal.pcbi.1014378 [code]
A mean-field model of neural networks with PV and SOM interneurons reveals connectivity-based mechanisms of gamma oscillations.
Journal: PLoS computational biology
In common: SciPy, Matplotlib, NumPy, 5 references
[5] doi:10.1038/s41467-026-73540-z [code]
Predictive acoustical processing in human cortical layers.
Journal: Nature communications
In common: 7 references
[6] doi:10.1126/sciadv.aed6417 [code]
Intrinsic timing, not temporal prediction, underlies ramping dynamics in visual and parietal cortex during passive behavior.
Journal: Science advances
In common: SciPy, Matplotlib, NumPy, systems, 5 references
[7] doi:10.1371/journal.pcbi.1014458 [code]
Neuronal excitability and parameter variability in the Hodgkin-Huxley model.
Journal: PLoS computational biology
In common: JAX, SciPy, Matplotlib, 1 other tool, computational modeling (no new data), 3 references
[8] doi:10.1371/journal.pcbi.1014022 [code]
Emergence of multifrequency activity in a laminar neural mass model.
Journal: PLoS computational biology
In common: SciPy, Matplotlib, NumPy, computational modeling (no new data), 4 references
[9] doi:10.1038/s41467-026-75359-0 [code]
Neural mechanisms of time-forward predictions for naturalistic auditory tone sequences.
Journal: Nature communications
In common: 7 references
[10] doi:10.1371/journal.pcbi.1014391 [code]
Multi-stable oscillations in cortical networks with two classes of inhibition.
Journal: PLoS computational biology
In common: SciPy, Matplotlib, NumPy, 4 references

Contribute

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

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

Request its removal

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

Discussion, reproductions, activity

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

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

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