OSCR

Probabilistic inference of homonymous and heteronymous recurrent inhibition in human muscles from large-scale motor neuron recordings.

Code ↔ Paper

21 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 21 matches
  1. [1] § MATERIALS AND METHODS › Simulation model › Simulated recurrent connectivity via Renshaw cells ↔ simulator.py, lines 39–219 · score 0.87 · refractory period, membrane potential, low pass filtered, Renshaw cell, hyperpolarizing, voltage
  2. [2] § MATERIALS AND METHODS › Simulation-based inference › Neural density estimator ↔ Simulation_based_inference/SBI_main_script.ipynb, lines 246–259 · score 0.80 · masked autoregressive flow, training batch, MAF, sbi, network, validation
  3. [3] § MATERIALS AND METHODS › Synchronization cross-histograms and their features › Feature extraction ↔ analyzer.py, lines 400–462 · score 0.73 · rising edge, trapezoidal function, falling edge, backward, plateau, window
  4. [4] § MATERIALS AND METHODS › Simulation model › Inputs to motor neurons ↔ simulator.py, lines 39–219 · score 0.71 · low pass filtered, half width, 0–5 Hz, chosen, motor neuron, noise
  5. [5] § MATERIALS AND METHODS › Synchronization cross-histograms and their features › Constructing and interpreting the synchronization cross-histograms ↔ analyzer.py, lines 26–69 · score 0.68 · homonymous pool, heteronymous pool, probability distribution, cross histogram, resolution, smoothed
  6. [6] § RESULTS › Cross-histogram signatures of recurrent inhibition and higher-frequency common input ↔ R_scripts_figures_results/Proxy_features_VS_ground_truth_plots_FIG2.Rmd, lines 6–101 · score 0.66 · disynaptic delay, ground truth, peak height, proxies, trough area, recurrent inhibition
  7. [7] § MATERIALS AND METHODS › Synchronization cross-histograms and their features › Curve fitting ↔ analyzer.py, lines 464–532 · score 0.65 · generalized logistic, baseline curve, probability distributions, rising, fitted, plateau
  8. [8] § RESULTS › Simulation-based inference to quantify recurrent inhibition ↔ Simulation_based_inference/SBI_main_script.ipynb, lines 2126–2229 · score 0.65 · 0–100 %, posterior density, posterior modes, ground truth, Pearson, accuracy
  9. [9] § MATERIALS AND METHODS › Synchronization cross-histograms and their features › Constructing and interpreting the synchronization cross-histograms ↔ Experimental_data_processing/organize_experimental_data.ipynb, lines 340–457 · score 0.64 · triceps surae, maximized spike, SOL, GM, VM, VL
  10. [10] § MATERIALS AND METHODS › Other data analyses › Recruitment thresholds ↔ Experimental_data_processing/process_experimental_data.ipynb, lines 1028–1084 · score 0.64 · valid ramp, recruitment threshold, derivative, pre, motor unit, Hz
  11. [11] § MATERIALS AND METHODS › Synchronization cross-histograms and their features › Curve fitting ↔ analyzer.py, lines 464–532 · score 0.63 · Mexican hat, vertical offset, decay, Gaussian, curve, width
  12. [12] § MATERIALS AND METHODS › Simulation-based inference › Validation ↔ Simulation_based_inference/SBI_main_script.ipynb, lines 2126–2229 · score 0.60 · posterior density, posterior modes, ground truth, Pearson, accuracy, uniform
  13. [13] § RESULTS › Simulation-based inference to quantify recurrent inhibition ↔ Simulation_based_inference/posterior_prediction_simulated_obs_VS_experimental_obs.ipynb, lines 22–89 · score 0.60 · posterior predictive simulations, posterior predictive checks, peak height, trough area, firing rate, recurrent inhibition
  14. [14] § MATERIALS AND METHODS › Simulation model › Motor neuron model ↔ simulator.py, lines 1512–1608 · score 0.58 · refractory period, leak, capacitance, reset, motor neuron, ms
  15. [15] § MATERIALS AND METHODS › Experimental study › Decoding of motor unit firing activity ↔ Experimental_data_processing/process_experimental_data.ipynb, lines 470–600 · score 0.57 · spike trigger, MUAP, realigned, onset, motor unit, activity
  16. [16] § MATERIALS AND METHODS › Experimental study › Decoding of motor unit firing activity ↔ Experimental_data_processing/process_experimental_data.ipynb, lines 470–600 · score 0.53 · noise ratio, channels, motor unit, EMG, signals, activity
  17. [17] § MATERIALS AND METHODS › Simulation model › Motor neuron model ↔ simulator.py, lines 768–793 · score 0.51 · potassium reversal potential, decayed
  18. [18] § MATERIALS AND METHODS › Other data analyses › Intramuscular coherence ↔ simulator.py, lines 1046–1105 · score 0.51 · frequency domain, Welch, Hann, summing, correlated, windows
  19. [19] § RESULTS › Simulation-based inference to quantify recurrent inhibition ↔ Simulation_based_inference/posterior_prediction_simulated_obs_VS_experimental_obs.ipynb, lines 462–536 · score 0.51 · interquartile range, simulated observations, median, inference, posterior
  20. [20] § MATERIALS AND METHODS › Simulation model › Motor neuron model ↔ simulator.py, lines 1512–1608 · score 0.50 · membrane potential, Euler, voltage, reset, motor neuron, threshold
  21. [21] § MATERIALS AND METHODS › Other data analyses › Intramuscular coherence ↔ analyzer.py, lines 26–69 · score 0.50 · random subsets, magnitude, overlap, iterations, smaller, coherence

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

Python · 2,008 lines · 111 KB · MIT · 6 matches

  1. import numpy as np
  2. import matplotlib.pyplot as plt
  3. import matplotlib.gridspec as gridspec
  4. import seaborn as sns
  5. import os
  6. import sys
  7. import json
  8. import pandas as pd
  9. from brian2 import *
  10. from scipy.signal import windows, butter, filtfilt, sosfiltfilt
  11. from dataclasses import dataclass, field, asdict
  12. from typing import Dict
  13. from scipy.signal import welch
  14. import time
  15. import Cython
  16. import h5py
  17. from h5py import string_dtype
  18. from typing import List
  19. import traceback
  20. from threading import local
  21. import logging
  22. import warnings
  23. # # # IGNORE WARNINGS # # #
  24. # # # Ignore warnings from Brian2 which do not affect simulation behavior
  25. # Silence the “TimedArray uses dt … not aligned” warning
  26. logging.getLogger('brian2.input.timedarray').setLevel(logging.ERROR)
  27. # Silence the “internal variable … exists in the namespace” warning
  28. logging.getLogger('brian2.groups.group').setLevel(logging.ERROR)
  29. # # #
  30. # ignore only the “not compatible with tight_layout” UserWarning
  31. warnings.filterwarnings(
  32. "ignore",
  33. message=r".*not compatible with tight_layout.*",
  34. category=UserWarning,
  35. )
  36. # # # DATA CLASS FOR SIMULATION's PARAMETERS
  37. @dataclass
  38. class SimulationParameters:
  39. # # # OUTPUT PARAMETERS
  40. output_folder_name: str = "simulation_batch_"
  41. make_unique_output_folder: bool = True
  42. output_plots: bool = False
  43. # # # RANDOM SEED
  44. pre_specify_random_seed: bool = True
  45. random_seed: int = field(init=False)
  46. # # # TIME PARAMETERS
  47. fsamp: int = 2048 # in samples per second
  48. # note that this is used to determine the time bins for input and output, but the actual integration time step is determined by 'defaultclock.dt' from Brian2 (0.1ms by default)
  49. duration: float = 5 # in seconds
  50. duration_with_ignored_window: float = field(init=False) # in seconds
  51. edges_ignore_duration: float = 1 # in seconds
  52. # # # NEURON NUMBERS AND SIZE
  53. nb_pools: int = 1
  54. full_pool_sim: bool = True
  55. # ^ If True, will simulate the chosen number of motor neurons as a population going from min soma size to max soma size, with a distribution determined by an exponent
  56. # ^ If True, min_soma_diameter, max_soma_diameter and size_distribution_exponent will be used
  57. # ^ If false, will simulate N motor neurons by sampling from a gaussian distribution specified by mean_soma_diameter and sd_prctl_soma_diameter
  58. # In both cases, the normalized size will depend on min_soma_diameter & max_soma_diameter
  59. nb_motoneurons_per_pool: int = 300
  60. nb_RCs_per_pool_pair: int = 60 # if 10 and 1 pool, only 10 RCs. If 2 pools, then there are 4 pairs of pools, so 40 RCs. Generally, total nb of RCs = nb_pools**2 * nb_RCs_per_pool_pair
  61. # size parameters if full_pool_sim == True ; if
  62. min_soma_diameter: float = 50 # in micrometers, for smallest motor neuron # Manuel et al. 2019 "Scaling of motor output, from Mouse to Humans"
  63. max_soma_diameter: float = 100 # in micrometers, for largest motor neuron # Manuel et al. 2019 "Scaling of motor output, from Mouse to Humans"
  64. size_distribution_exponent: float = 2 # between 0-1 => more large MN than small MN; 1 => uniform distribution (linear relationship between MN index and soma diameter); >1 => more small motoneurons than large MNs
  65. # size parameters if full_pool_sim == False
  66. mean_soma_diameter: float = 60 # in micrometers, for average motor neuron size of the simulated population
  67. sd_prct_soma_diameter: float = 5 # as a % of mean_soma_diameter. If mean_soma_diameter=50 and sd_prct_soma_diameter=10 for example, the distribution will be a gaussian with mean 50 and sd 5 micrometers
  68. # To calculate at initialization
  69. total_nb_motoneurons: int = field(init=False)
  70. total_nb_renshaw_cells: int = field(init=False)
  71. RC_pair_indices: Dict[tuple, np.ndarray] = field(init=False)
  72. # # # NEURON THRESHOLD AND EQUILIBRIUM POTENTIALS
  73. voltage_rest: Quantity = field(default_factory=lambda: 0 * mvolt) # arbitrary; 0 at rest
  74. voltage_thresh: Quantity = field(default_factory=lambda: 10 * mvolt) # arbitrary; 10 for generating a spike
  75. voltage_AHP: Quantity = field(default_factory=lambda: -13.3 * mvolt) # resting potential is typically -70mvolt, spike threshold is typically -55mvolt, and potassium reversal potential is typically -90mvolt.
  76. # To keep the relationships the same despite the arbitrary 0mvolt (resting voltage) and 10mvolt (spike generation threshold), voltage_AHP is set to -13.3 mvolt
  77. # # # COMMON INPUT PARAMETERS
  78. nb_of_common_inputs: int = 2 # Each comon input is distributed to the whole pool, but each has its own frequency content
  79. frequency_range_to_set_input_power: List[float] = field(
  80. default_factory=lambda: [0,5]) # in hz
  81. # Will set the frequency range over which the std scaling (both for excitatory common input and independent input) is done.
  82. # Thus, total power of the input signal will be the one selected by the user within the 'frequency_range_to_set_input_power',
  83. # but the total power increases with the frequency range of the input signal
  84. excitatory_input_baseline: List[float] = field( # in nA (nanoAmperes) # List of length at least equal to nb_pools
  85. default_factory=lambda: [25*1e3, 25*1e3])
  86. common_input_std: np.ndarray = field( # in nA (nanoAmperes) or in % of excitatory_input_baseline (check common_input_std_as_amp_or_prct) => it has to be a numpy array of size = nb_pools (or more) x nb_of_inputs
  87. default_factory=lambda: np.array([
  88. [3.0*1e3, 0.0], # common_input_std for pool 0 [input 0, input 1]
  89. [3.0*1e3, 0.0] # common_input_std for pool 1 [input 0, input 1]
  90. ]))
  91. common_input_std_as_amp_or_prct: List[str] = field( # list of size nb_of_common_inputs => the unit will be the same regardless of the pool
  92. default_factory=lambda: ['nA', 'nA']) # 'percent' or 'nA'
  93. frequency_range_of_common_input: np.ndarray = field( # in Hz. This has to be a numpy array of size = nb_pools (or more) x nb_of_inputs x 2 (lower and upper bounds of frequency bandwidth)
  94. default_factory=lambda: np.array([
  95. [[0.0,5.0],[30.0,40.0]], # frequency range for pool 0 inputs [[low, high] for input 0], [[low, high] for input 1]
  96. [[0.0,5.0],[30.0,40.0]] # frequency range for pool 1 inputs [[low, high] for input 0], [[low, high] for input 1]
  97. ]))
  98. max_frequency_of_any_input: float = 80 # This can be necessary when setting the frequency content through another script
  99. common_input_characteristics: dict = field(init=False) # Those values are set AFTER initialization - they are just nicer to work with for SBI later
  100. # ──── FREQUENCY FILTERING "MASTER PARAMETERS" ────
  101. scale_filter_order_to_frequency: bool = True
  102. filter_order_scaling_coeff: float = 0.5 # Multiplying the cutoff frequency to determine filter order. Used only if scale_filter_order_to_frequency == True
  103. default_freq_filter_order: int = 5 # Used if scale_filter_order_to_frequency = False
  104. lowest_freq_filter_order: int = 5 # # Used if scale_filter_order_to_frequency = True
  105. max_freq_filter_order: int = 100
  106. # ──── INPUTS DISTRIBUTION AND CORRELATION ACROSS POOLS
  107. set_same_excitatory_input_for_all_pools: bool = False # If true, this will override all the generated excitatory inputs to be the same for all pools
  108. set_arbitrary_correlation_between_excitatory_inputs: bool = True
  109. between_pool_excitatory_input_correlation: float = 0.5 # > 0 and < 1 # used only if set_arbitrary_correlation_between_excitatory_inputs == True
  110. # negative correlations are accepted as inputs but the procedure to set arbitrary correlations doesn't work for negative correlations (it results in correlations around 0)
  111. # a rough method has been implemented to deal with this problem, but it can induce spurious reversal of the sign of correlations
  112. # # # INDEPENDENT (NOISY) INPUT PARAMETERS
  113. low_pass_filter_of_MN_independent_input: float = 50 # in Hz
  114. # Remember that the power will be scaled according to the frequency band set by frequency_range_to_set_input_power, so for instance,
  115. # if (MN_independent_input_std = common_input_std) and if max freq of independent input = 50 and max freq of common input = 5,
  116. # then the total power of the independent noise will be 10 times larger than the common input power (because the frequency range is 10 times larger)
  117. # So in order to get the independent input total power to be 3 times larger than the total common input power when
  118. # THE COMMON INPUT FREQUENCY RANGE IS 0-5Hz, and when THE INDEPENDENT INPUT FREQUENCY RANGE IS 0-50Hz,
  119. # we need to set the independent input std to be be sqrt(0.3)
  120. independent_input_absolute_or_ratio: str = 'ratio' # 'absolute' or 'ratio'
  121. independent_input_power: float = 3
  122. # If independent_input_absolute_or_ratio == 'ratio', then independent_input_power is the ratio of the independent input std to the common input std
  123. # Assuming that the independent input frequency range is 10 times larger than the common input frequency range
  124. # ^ This ratio is set relative to the FIRST common input only
  125. # If independent_input_absolute_or_ratio == 'absolute', then independent_input_power is the absolute value of the independent input std
  126. low_pass_filter_of_RC_independent_input: float = 50 # in Hz
  127. # # # MOTOR NEURONS ELECTROPHYSIOLOGICAL PROPERTIES CONSTANTS (Caillet et al 2022)
  128. # Resistance constants, used to generate the motor neuron resistance (in Ohms) according to their size
  129. resistance_constant: float = 9.6*(10**5)
  130. resistance_exponent: float = 2.4*(-1)
  131. # Rheobase constants, used to generate the motor neuron input current offset (in nA) according to their size
  132. rheobase_constant: float = 9.0*(10**-4)
  133. rheobase_exponent: float = 2.5
  134. rheobase_scaling: float = 6.0*1e2 # manually-tuned scaling
  135. # Capacitance constants, used to generate the motor neuron capacitance (in Farads) according to their size
  136. capacitance_constant: float = 1.2
  137. capacitance_exponent: float = 1
  138. # Afterhyperpolarization constants, used to generate the motor neuron AHP duration (in ms) according to their size + refractory period (in ms)
  139. AHP_duration_constant: float = 2.5 * (10**4)
  140. AHP_duration_exponent: float = 1.5 * (-1)
  141. # ^ these variables are used to create the variable 'motoneurons_AHP_conductance_decay_time_constant'
  142. # Caillet et al describe the relationship for the DURATION of the AHP, but for the equations I am using a time constant to control for the AHP conductance decay
  143. # I consider that the duration of the AHP correspond to the time it takes for the peak input at time 0 (x0) to decay to a tenth of its value x(t)<=X0/10
  144. # Thus, motoneurons_AHP_conductance_decay_time_constant = AHP_duration / ln(10)
  145. AHP_conductance_delta_after_spiking: Quantity = field(default_factory=lambda: 1.0 * msiemens) # Hyperpolarizing conductance change after a spike
  146. refractory_period_absolute: float = 5 # in ms
  147. # Axonal conduction velocity constants, used to generate the MN-to-fiber velocity (in m/s) according to their size
  148. # Then, the delay (ms) is calculated from the axonal conduction velocity, assuming a 0.5m axon length => so correspond to the conduction speed from MN to muscle fiber (speed in m/s, so multiply speed by 2)
  149. axonal_conduction_velocity_constant: float = 4.0*2
  150. axonal_conduction_velocity_exponent: float = 0.7
  151. # # # RENSHAW CELLS ELECTROPHYSIOLOGICAL PROPERTIES
  152. tau_Renshaw: Quantity = field(default_factory=lambda: 8*ms)
  153. refractory_period_RC: Quantity = field(default_factory=lambda: 5*ms)
  154. # # # MOTOR NEURONS <=> RENSHAW CELLS CONNECTIVITY
  155. binary_connectivity: bool = False # If true, MN to RC and RC to MN weights are either 0 or 1.
  156. # This makes the distribution of disynaptic connections between MNs (especially the dsitribution's std) less controllable
  157. # 1) Full MN→MN target matrix (size: nb_pools × nb_pools)
  158. # disynpatic_inhib_connections_desired_MN_MN[i,j] is the *desired* mean number of disynaptic synapses from MN pool i to MN pool j
  159. # This is NOT a probability and can thus be > 1
  160. disynpatic_inhib_connections_desired_MN_MN: np.ndarray = field(
  161. default_factory=lambda: np.array([ # Increase the size of the array if nb_pools > 2
  162. [1.0, 0.0], # pool 0→pool 0, pool 0→pool 1
  163. [0.0, 0.0], # pool 1→pool 0, pool 1→pool 1
  164. ])
  165. )
  166. # 2) Split ratio of excitation vs inhibition:
  167. # alpha = fraction of weight allocated to MN→RC vs RC→MN
  168. # (so MN→RC uses p = (disynpatic_inhib_connections_desired_MN_MN)**alpha, RC→MN uses p = (disynpatic_inhib_connections_desired_MN_MN)**(1-alpha) )
  169. split_MN_RC_ratio: float = 0.5
  170. # 3) Distribution rule for sampling each bipartite graph:
  171. # 'binarize' → Bernoulli(p), and max MN->MN inhibition of 1
  172. # 'gaussian' → each pre: k∼N(mean*n_post, std*n_post)
  173. # 'size_gaussian'→ same but mean interpolates by MN size rank
  174. distribution_type: str = 'gaussian'
  175. # 4) Distribution rule parameters:
  176. # - binarize: {} # no extra params
  177. # - gaussian: {'std': 0.1}
  178. # - size_gaussian: {'std': 0.1,
  179. # 'ratio_large_small': 2.0}
  180. distribution_params: Dict[str, float] = field(default_factory=lambda: {
  181. 'std': 0.2, # net std of the MN->MN connectivity is sqrt(2*std**2 + std**4) / sqrt(nb_RCs_per_pool_pair). So with std=0.4 and nb_RCs_per_pool_pair=10, the std of MN->MN connectivity is ~0.20
  182. 'std_is_prct': True, # if 'std_is_prct == True', then std will be interpreted as a % relative to disynpatic_inhib_connections_desired_MN_MN (with std=0.1 being 10% of mean connectivity for instance)
  183. # if 'std_is_prct == False', the value will be interpreted directly as the std of the number of disynaptic connections (weights)
  184. # if 'std=0.2" and if 'std_is_prct=True', the actual std will be 10% and not 20%
  185. 'ratio_large_small': 3.0
  186. })
  187. disynaptic_inhib_received_arbitrary_adjustment: float = 0, # Correspond to the std of connection weights from all RCs to a given motorneuron that will be added. If 0, correspond to the specified connectivity, if 1 for, will be the specified connectivity on average +/- 1 (with a cutoff at 0)
  188. # 5) Enforce heteronymous‐pool non‐overlap?
  189. prevent_heteronymous_pool_overlap: bool = False
  190. # # # MOTOR NEURONS <=> RENSHAW CELLS POST-SYNAPTIC EFFECTS
  191. MN_to_Renshaw_EPSP: Quantity = field(default_factory=lambda: 6.7*mvolt) # 6.7*mvolt # when > 10*mvolt, ensures that 1 MN spike = 1 RC spike # increase in V in renshaw cell when receiving spike from MN - From Moore et al 2015 = MN-RC pair recordings, with 1 MN spike on average resulting in a probability of 0.3 of RC spike
  192. Renshaw_to_MN_IPSP: Quantity = field(default_factory=lambda: 3.0*1e3*nA)
  193. # ^ later in the code, this is turned into Coulomb (Total charge of an IPSP = current per second)
  194. # If scale_initial_IPSP_to_be_same_integral_regardless_of_synaptic_tau == True:
  195. # Then the value (in nA) defined in Renshaw_to_MN_IPSP will be the total charge (in Coulomb)
  196. # If scale_initial_IPSP_to_be_same_integral_regardless_of_synaptic_tau == False:
  197. # Then the value (in nA) defined in Renshaw_to_MN_IPSP will be the initial IPSP current
  198. scale_initial_IPSP_to_be_same_integral_regardless_of_synaptic_tau: bool = False # If True, Renshaw_to_MN_IPSP will be used not as initial current when an IPSP is received, but as a target total current received regardless of the chosen synaptic time constant
  199. synaptic_IPSP_membrane_or_user_defined_time_constant: str = 'membrane' # 'user_defined' or 'membrane'
  200. synaptic_IPSP_decay_time_constant: Quantity = field(default_factory=lambda: 10*ms) # Used only if 'synaptic_IPSP_membrane_or_user_defined_time_constant' == 'user_defined'
  201. # Time constant of the synaptic IPSP decay in the MN # Check figures from Williams & Baker 2009 (simulation); between ~5 and ~10ms membrane tau for Uchiyama & Windhorst 2007 (simulation); ~15-30ms inhibition duration in Ozyurt et al 2019 (experimental data)
  202. # If 10ms for example, the IPSP will decay to ~37% of its peak value after 10ms and to ~14% after 20ms (0.14 is ~0.37²)
  203. # William & Baker 2009 J Neuroscience: RC's IPSP rise time of 5.5 ms and half-width of 18.5 ms
  204. MN_RC_synpatic_delay: Quantity = field(default_factory=lambda: 5*ms)
  205. RC_independent_input_std: float = 3.333 # in mvolt # if 3.33, on average, RC membrane potential will fluctuate at 1/3 of the spike threshold (10mvolt)
  206. # # # POST INIT FUNCTION, EXECUTING AFTER INITIALIZATION
  207. def __post_init__(self):
  208. # Setting variables that are calculated post-initialization
  209. if self.pre_specify_random_seed:
  210. self.random_seed = 42
  211. else:
  212. self.random_seed = np.random.randint(2**10)
  213. self.duration_with_ignored_window = self.duration + (2 * self.edges_ignore_duration)
  214. self.total_nb_motoneurons = self.nb_pools * self.nb_motoneurons_per_pool
  215. self.total_nb_renshaw_cells = (self.nb_pools**2)*self.nb_RCs_per_pool_pair
  216. self.RC_pair_indices = {(i, j): np.arange(
  217. (i*self.nb_pools + j)*self.nb_RCs_per_pool_pair,
  218. (i*self.nb_pools + j + 1)*self.nb_RCs_per_pool_pair)
  219. for i in range(self.nb_pools) for j in range(self.nb_pools)}
  220. #### PARAMETERS CHECK ####
  221. if self.nb_pools < 1:
  222. raise ValueError("nb_pools must be ≥1")
  223. if self.distribution_type not in ("binarize","gaussian","size_gaussian"):
  224. raise ValueError(f"Unknown distribution_type={self.distribution_type!r}")
  225. # — Check excitatory_input_baseline —
  226. if len(self.excitatory_input_baseline) < self.nb_pools:
  227. raise ValueError(
  228. f"excitatory_input_baseline must be a list of length >= nb_pools"
  229. )
  230. # — Check frequency_range_of_common_input —
  231. if not isinstance(self.frequency_range_of_common_input, np.ndarray):
  232. raise TypeError(
  233. f"frequency_range_of_common_input must be of type np.ndarray, not {type(self.frequency_range_of_common_input).__name__}"
  234. )
  235. array_shape_temp = np.shape(self.frequency_range_of_common_input)
  236. if (array_shape_temp[0] < self.nb_pools) and (array_shape_temp[1] < self.nb_of_common_inputs) and (array_shape_temp[2] < 2):
  237. raise ValueError(
  238. f"frequency_range_of_common_input must be np.ndarray of size [nb_pools (or more) x nb_of_inputs x 2], but got shape {np.shape(self.frequency_range_of_common_input)} instead"
  239. )
  240. # — Check limit of frequency_range_of_common_input —
  241. highs = self.frequency_range_of_common_input[..., 1] # grab a view of all the “high” edges
  242. M = self.max_frequency_of_any_input
  243. # Clamp & warn
  244. if np.any(highs > M):
  245. warnings.warn(f"Clamping {np.sum(highs > M)} upper‐bounds down to {M} Hz.")
  246. # write back into the dataclass array
  247. self.frequency_range_of_common_input[..., 1] = np.where(
  248. highs > M,
  249. M, # if above the max, set to M
  250. highs # otherwise leave as-is
  251. )
  252. # — Check common_input_std —
  253. if not isinstance(self.common_input_std, np.ndarray):
  254. raise TypeError(
  255. f"common_input_std must be np.ndarray of size [nb_pools (or more) x nb_of_inputs], not {type(self.common_input_std).__name__}"
  256. )
  257. array_shape_temp = np.shape(self.common_input_std)
  258. if (array_shape_temp[0] < self.nb_pools) and (array_shape_temp[1] < self.nb_of_common_inputs):
  259. raise ValueError(
  260. f"common_input_std must be np.ndarray of size [nb_pools (or more) x nb_of_inputs], but got shape {np.shape(self.common_input_std)} instead"
  261. )
  262. # - Check independent_input_absolute_or_ratio -
  263. if self.independent_input_absolute_or_ratio not in ('absolute', 'ratio'):
  264. raise ValueError(
  265. f"independent_input_absolute_or_ratio must be 'absolute' or 'ratio', not {self.independent_input_absolute_or_ratio!r}"
  266. )
  267. # Set common input std to nA if defined as a % of baseline
  268. for pooli in range(self.nb_pools):
  269. for inputi in range(self.nb_of_common_inputs):
  270. if self.common_input_std_as_amp_or_prct[inputi] == 'percent':
  271. self.common_input_std[pooli][inputi] = (self.common_input_std[pooli][inputi]*self.excitatory_input_baseline[pooli])/100 # common_input_std should be in %
  272. elif self.common_input_std_as_amp_or_prct[inputi] == 'nA':
  273. self.common_input_std[pooli][inputi] = self.common_input_std[pooli][inputi] # no change
  274. else:
  275. raise ValueError(f"Unknown common_input_std_as_amp_or_prct={self.common_input_std_as_amp_or_prct[inputi]!r}\n Should be 'percent' or 'nA'")
  276. # Reframe some of the parameter values to make simulation-based inference easier
  277. self.common_input_characteristics = {"Frequency_middle_of_range": {},
  278. "Frequency_half_width_of_range": {}}
  279. for pooli in range(self.nb_pools):
  280. key_pool = f"pool_{pooli}"
  281. self.common_input_characteristics["Frequency_middle_of_range"][f"{key_pool}"] = {}
  282. self.common_input_characteristics["Frequency_half_width_of_range"][f"{key_pool}"] = {}
  283. for inputi in range(self.nb_of_common_inputs):
  284. key_input = f"input_{inputi}"
  285. self.common_input_characteristics["Frequency_middle_of_range"][f"{key_pool}"][f"{key_input}"] = np.mean(
  286. self.frequency_range_of_common_input[pooli][inputi])
  287. self.common_input_characteristics["Frequency_half_width_of_range"][f"{key_pool}"][f"{key_input}"] = (
  288. self.frequency_range_of_common_input[pooli][inputi][1] - self.frequency_range_of_common_input[pooli][inputi][0]) / 2
  289. # frequency_range_of_common_input = np.array([
  290. # [ [low, high] (input 0), [low, high] (input 1) ], # pool 0
  291. # [ [low, high] (input 0), [low, high] (input 1) ] # pool 1
  292. # ])
  293. ######################## END OF SimulationParameters CLASS
  294. ######################################
  295. ### FETCH THREAD-LOCAL STORAGE FOR "CURRENT" PARAMETERS
  296. ### So that they can be used in the helper functions without explicitely passing them as arguments
  297. ######################################
  298. _params_context = local()
  299. def _set_filter_params(params: SimulationParameters):
  300. """Call this at the top of run_simulation to make params visible to the filters."""
  301. _params_context.params = params
  302. def _get_filter_params() -> SimulationParameters:
  303. try:
  304. return _params_context.params
  305. except AttributeError:
  306. raise RuntimeError("No filter params set; forgot to call _set_filter_params?")
  307. ######################################
  308. ### HELPER FUNCTIONS
  309. ######################################
  310. def lerp(a, b, t):
  311. return a + t * (b - a)
  312. # ───────────────────
  313. # Filtering functions
  314. def butter_lowpass(cutoff, fs, order):
  315. """
  316. Return an SOS filter for a lowpass Butterworth of the given order.
  317. """
  318. nyq = 0.5 * fs
  319. Wn = cutoff / nyq
  320. # design in SOS form
  321. sos = butter(order, Wn,
  322. btype='low',
  323. analog=False,
  324. output='sos')
  325. return sos
  326. def butter_highpass(cutoff, fs, order):
  327. """
  328. Return an SOS filter for a highpass Butterworth of the given order.
  329. """
  330. nyq = 0.5 * fs
  331. Wn = cutoff / nyq
  332. sos = butter(order, Wn,
  333. btype='high',
  334. analog=False,
  335. output='sos')
  336. return sos
  337. # public filter‐wrappers
  338. def lowpass_filter(data, cutoff, fs):
  339. p = _get_filter_params()
  340. order = p.default_freq_filter_order
  341. if p.scale_filter_order_to_frequency:
  342. if cutoff > p.lowest_freq_filter_order:
  343. order = cutoff * p.filter_order_scaling_coeff
  344. if ~np.isfinite(order):
  345. order = p.default_freq_filter_order
  346. elif order > p.max_freq_filter_order:
  347. order = p.max_freq_filter_order
  348. elif order < p.lowest_freq_filter_order:
  349. order = p.lowest_freq_filter_order
  350. order = np.floor(order).astype(int)
  351. sos = butter_lowpass(cutoff, fs, order=order)
  352. return sosfiltfilt(sos, data)
  353. def highpass_filter(data, cutoff, fs, order=None):
  354. p = _get_filter_params()
  355. order = p.default_freq_filter_order
  356. if p.scale_filter_order_to_frequency:
  357. if cutoff > p.lowest_freq_filter_order:
  358. order = cutoff * p.filter_order_scaling_coeff
  359. if ~np.isfinite(order):
  360. order = p.default_freq_filter_order
  361. elif order > p.max_freq_filter_order:
  362. order = p.max_freq_filter_order
  363. elif order < p.lowest_freq_filter_order:
  364. order = p.lowest_freq_filter_order
  365. order = np.floor(order).astype(int)
  366. sos = butter_highpass(cutoff, fs, order=order)
  367. return sosfiltfilt(sos, data)
  368. # ───────────────────────
  369. def filter_artifact_removal(fsamp, edges_ignore_duration):
  370. duration_to_remove = 1/edges_ignore_duration # in second
  371. Wind_s = duration_to_remove * 2
  372. artifact_removal_window = windows.hann(round(fsamp * Wind_s))
  373. artifact_removal_window = artifact_removal_window[:int(np.round(len(artifact_removal_window)/2))]
  374. nb_samples_artifact_removal_window = len(artifact_removal_window)
  375. return artifact_removal_window, nb_samples_artifact_removal_window
  376. def scale_to_band_std(signal, fs, band, desired_std):
  377. """
  378. Scale `signal` so that its *power* (variance) in [fmin,fmax] ⟶ (desired_std)^2.
  379. Uses Welch's PSD + frequency‐bin masking. Avoids filtfilt entirely.
  380. Parameters
  381. ----------
  382. signal : 1D array
  383. fs : float
  384. Sampling rate (Hz)
  385. band : (fmin, fmax)
  386. desired_std : float
  387. Target standard deviation within [fmin,fmax]
  388. nperseg : int
  389. segment length for Welch. If signal_length < nperseg, welch auto‐truncates.
  390. Returns
  391. -------
  392. scaled_signal : ndarray, same shape as `signal`
  393. """
  394. fmin, fmax = band
  395. if fmin < 0 or fmax >= fs/2 or fmin >= fmax:
  396. raise ValueError("band must be [fmin, fmax] with 0≤fmin<fmax<fs/2")
  397. # 1) full‐PSD
  398. freqs, Pxx = welch(signal, fs=fs, window='hann', nperseg=fs, scaling='spectrum')
  399. # 2) mask to only [fmin, fmax]
  400. mask = (freqs >= fmin) & (freqs <= fmax)
  401. if not np.any(mask):
  402. raise ValueError(f"No PSD bins in [{fmin},{fmax}] Hz; choose a smaller nperseg or adjust band.")
  403. current_power = np.trapz(Pxx[mask], freqs[mask]) # = ∫_{fmin}^{fmax} PSD(f) df
  404. if current_power <= 0:
  405. raise ValueError(f"No power in the {fmin}–{fmax} Hz band to scale.")
  406. # 3) we want variance in [fmin,fmax] = (desired_std)^2
  407. desired_power = desired_std**2
  408. scale_factor = np.sqrt(desired_power / current_power)
  409. # 4) apply to entire time-series
  410. return signal * scale_factor
  411. def Generate_filtered_gaussian_noise_input(fsamp,
  412. duration_with_ignored_window, edges_ignore_duration, artifact_removal_window,
  413. low_pass_filter_cutoff, input_mean, scaling_std, high_pass_filter_cutoff=0):
  414. # Generate random input
  415. temp_input = np.random.normal(0, 1, int(duration_with_ignored_window * fsamp))
  416. # Apply artifact removal window to the end of the signal
  417. end_ignore_start = int(duration_with_ignored_window * fsamp) - int(np.round(edges_ignore_duration * fsamp))
  418. temp_input[end_ignore_start:] = temp_input[end_ignore_start:] * np.flip(artifact_removal_window)
  419. # Apply artifact removal window to the beginning of the signal
  420. beginning_ignore_end = int(np.round(edges_ignore_duration * fsamp))
  421. temp_input[:beginning_ignore_end] = temp_input[:beginning_ignore_end] * artifact_removal_window
  422. # Apply low-pass filter
  423. temp_input = lowpass_filter(temp_input, low_pass_filter_cutoff, fsamp)
  424. # Apply high-pass filter
  425. if high_pass_filter_cutoff >= 1:
  426. temp_input = highpass_filter(temp_input, high_pass_filter_cutoff, fsamp)
  427. # Normalize the signal
  428. temp_input = temp_input - np.mean(temp_input)
  429. temp_input = temp_input / np.std(temp_input)
  430. # Scale and add mean
  431. temp_input = temp_input * scaling_std
  432. temp_input = temp_input + input_mean
  433. return temp_input
  434. def is_positive_definite(X):
  435. try:
  436. np.linalg.cholesky(X)
  437. return True
  438. except np.linalg.LinAlgError:
  439. return False
  440. def nearest_positive_definite(A):
  441. B = (A + A.T) / 2
  442. U, s, Vt = np.linalg.svd(B)
  443. H = np.dot(Vt.T * s, Vt)
  444. A2 = (B + H) / 2
  445. A3 = (A2 + A2.T) / 2
  446. if is_positive_definite(A3):
  447. return A3
  448. spacing = np.spacing(np.linalg.norm(A))
  449. I = np.eye(A.shape[0])
  450. k = 1
  451. while not is_positive_definite(A3):
  452. A3 += I * spacing * k
  453. k += 1
  454. return A3
  455. def adjust_correlation(data, R_desired, threshold_for_adjusting_correl=0.05):
  456. """
  457. Adjust correlation of 'data' so its correlation matrix approximates R_desired.
  458. Handles negative correlations by a post-processing sign-flip heuristic.
  459. """
  460. data = np.asarray(data, dtype=float)
  461. n_samples, n_vars = data.shape
  462. if R_desired.shape != (n_vars, n_vars):
  463. raise ValueError("R_desired must be an n_vars x n_vars matrix.")
  464. if not np.allclose(R_desired, R_desired.T):
  465. raise ValueError("R_desired must be symmetric.")
  466. if not np.all(np.diag(R_desired) == 1.0):
  467. raise ValueError("Diagonal elements of R_desired must be 1.")
  468. # Store the sign pattern of desired correlations
  469. sign_pattern = np.sign(R_desired)
  470. # Use absolute values to form a positive definite target
  471. R_abs = np.abs(R_desired)
  472. # Make sure R_abs is positive definite
  473. R_abs_pd = R_abs if is_positive_definite(R_abs) else nearest_positive_definite(R_abs)
  474. # Standardize data
  475. orig_means = np.mean(data, axis=0)
  476. orig_stds = np.std(data, axis=0, ddof=1)
  477. if np.any(orig_stds == 0):
  478. raise ValueError("One or more input variables are constant; cannot adjust correlation.")
  479. data_std = (data - orig_means) / orig_stds
  480. # Compute current correlation
  481. R_current = np.corrcoef(data_std, rowvar=False)
  482. # Make sure R_current is PD
  483. R_current_pd = R_current if is_positive_definite(R_current) else nearest_positive_definite(R_current)
  484. # Cholesky decompositions
  485. L_current = np.linalg.cholesky(R_current_pd)
  486. L_desired = np.linalg.cholesky(R_abs_pd)
  487. # Whiten data
  488. inv_L_current = np.linalg.inv(L_current)
  489. data_white = data_std @ inv_L_current
  490. # Apply desired correlation (absolute values)
  491. data_transformed_std = data_white @ L_desired
  492. # Rescale back
  493. data_transformed = data_transformed_std * orig_stds + orig_means
  494. # Now, adjust signs if needed
  495. # Check the resulting correlations
  496. R_final = np.corrcoef(data_transformed, rowvar=False)
  497. # Check the resulting correlation and impose correlation relative to the first (first column of data) input if necessary
  498. for inputi in range(len(data_transformed[0])):
  499. if inputi==0:
  500. continue # skipping the first input (correlation with itself)
  501. else:
  502. temp_correl = np.corrcoef(data_transformed[:,0],data_transformed[:,inputi])[0, 1]
  503. temp_desired_correl = R_desired[0, inputi]
  504. diff_actual_VS_desired = temp_correl - temp_desired_correl # negative if the correlation is too low, positive if the correlation is too high
  505. if (abs(diff_actual_VS_desired) > threshold_for_adjusting_correl) and (diff_actual_VS_desired < 0): # if the correlation is too low
  506. data_transformed[:, inputi] = data_transformed[:,inputi]*(1-abs(diff_actual_VS_desired)) + (data_transformed[:,0]*abs(diff_actual_VS_desired))
  507. if (abs(diff_actual_VS_desired) > threshold_for_adjusting_correl) and (diff_actual_VS_desired > 0): # if the correlation is too high
  508. data_transformed[:, inputi] = data_transformed[:,inputi]*(1-abs(diff_actual_VS_desired)) - (data_transformed[:,0]*abs(diff_actual_VS_desired))
  509. # If a desired correlation is negative but we got a positive one, flip one column
  510. # This is a heuristic. We try flipping the second column in the pair.
  511. min_R_final_for_flip = 0.2 # adjust this threshold as needed
  512. min_R_desired_for_flip = -0.2 # adjust this threshold as needed
  513. for i in range(n_vars):
  514. for j in range(i+1, n_vars):
  515. if R_desired[i, j] < min_R_desired_for_flip: # desired negative correlation exceeding a threshold
  516. if R_final[i, j] > min_R_final_for_flip: # got a positive correlation exceeding a threshold instead
  517. # Flip the sign of column j
  518. data_transformed[:, j] = -data_transformed[:, j]
  519. # Recalculate R_final after the flip
  520. # R_final = np.corrcoef(data_transformed, rowvar=False)
  521. # Example usage:
  522. # A = np.random.randn(1000, 3)
  523. # R_desired = np.array([[1.0, -0.5, 0.3],
  524. # [-0.5, 1.0, -0.2],
  525. # [0.3, -0.2, 1.0]])
  526. # data_new = adjust_correlation(A, R_desired)
  527. # np.corrcoef(data_new, rowvar=False) should reflect the sign pattern of R_desired.
  528. return data_transformed
  529. def sample_block(n_pre, n_post, p_block, dist, params, ranks=None, binary=False):
  530. # If we want real‐valued weights, create A as float. If binary, keep it integer.
  531. if binary:
  532. A = np.zeros((n_pre, n_post), dtype=int)
  533. else:
  534. A = np.zeros((n_pre, n_post), dtype=float)
  535. if dist == "binarize":
  536. if binary:
  537. # pure 0/1 Bernoulli
  538. A[:] = (np.random.rand(n_pre, n_post) < p_block).astype(int)
  539. else:
  540. # real‐valued in [0, p_block)
  541. A[:] = np.random.rand(n_pre, n_post) * p_block
  542. elif dist == "gaussian":
  543. for u in range(n_pre):
  544. for v in range(n_post):
  545. μ = p_block
  546. if params['std_is_prct']:
  547. σ = p_block * params.get("std", 0.0)
  548. else:
  549. σ = params.get("std", 0.0)
  550. if binary:
  551. # sample a probability p_samp ∼ TruncNormal(μ, σ), then Bernoulli
  552. p_samp = np.clip(np.random.randn() * σ + μ, 0, 1)
  553. A[u, v] = (np.random.rand() < p_samp).astype(int)
  554. else:
  555. # sample a truncated normal weight; keep it as float
  556. raw = np.random.randn() * σ + μ
  557. A[u, v] = max(raw, 0.0)
  558. elif dist == "size_gaussian" and (ranks is not None):
  559. R = params.get("ratio_large_small", 1.0)
  560. for u in range(n_pre):
  561. r = ranks[u]
  562. # interpolate mean between μ_s and μ_l
  563. μ0 = p_block
  564. if params['std_is_prct']:
  565. σ = p_block * params.get("std", 0.0)
  566. else:
  567. σ = params.get("std", 0.0)
  568. μ_s = 2 * μ0 / (1 + R)
  569. μ_l = R * μ_s
  570. μ_ij = μ_s + r * (μ_l - μ_s)
  571. for v in range(n_post):
  572. if binary:
  573. p_samp = np.clip(np.random.randn() * σ + μ_ij, 0, 1)
  574. A[u, v] = (np.random.rand() < p_samp).astype(int)
  575. else:
  576. raw = np.random.randn() * σ + μ_ij
  577. A[u, v] = max(raw, 0.0)
  578. else:
  579. # fallback to Gaussian logic
  580. for u in range(n_pre):
  581. for v in range(n_post):
  582. μ = p_block
  583. if params['std_is_prct']:
  584. σ = p_block * params.get("std", 0.0)
  585. else:
  586. σ = params.get("std", 0.0)
  587. if binary:
  588. p_samp = np.clip(np.random.randn() * σ + μ, 0, 1)
  589. A[u, v] = (np.random.rand() < p_samp).astype(int)
  590. else:
  591. raw = np.random.randn() * σ + μ
  592. A[u, v] = max(raw, 0.0)
  593. return A
  594. def plot_connectivity_matrix(
  595. mat, title,
  596. pool_labels_pre, pool_labels_post,
  597. n_pre_pool, n_post_pool,
  598. cmap='viridis',
  599. add_colorbar=False,
  600. savepath=None
  601. ):
  602. n_pre, n_post = mat.shape
  603. n_pre_blocks = len(pool_labels_pre)
  604. n_post_blocks = len(pool_labels_post)
  605. # compute block means & stds (unchanged) …
  606. means = np.zeros((n_pre_blocks, n_post_blocks))
  607. stds = np.zeros((n_pre_blocks, n_post_blocks))
  608. for i in range(n_pre_blocks):
  609. for j in range(n_post_blocks):
  610. r = slice(i*n_pre_pool, (i+1)*n_pre_pool)
  611. c = slice(j*n_post_pool, (j+1)*n_post_pool)
  612. blk = mat[r,c]
  613. means[i,j] = blk.mean()
  614. stds[i,j] = blk.std()
  615. fig, ax = plt.subplots(figsize=(10,10))
  616. im = ax.imshow(mat, cmap=cmap, aspect='equal')
  617. # grid lines at pool boundaries
  618. for y in np.arange(n_pre_pool, n_pre, n_pre_pool):
  619. ax.axhline(y-0.5, color='white', lw=1)
  620. for x in np.arange(n_post_pool, n_post, n_post_pool):
  621. ax.axvline(x-0.5, color='white', lw=1)
  622. # ### 1) show every cell index tick ###
  623. ax.set_xticks(np.arange(n_post))
  624. ax.set_xticklabels([str(i) for i in range(n_post)], rotation=90, fontsize=6)
  625. ax.set_yticks(np.arange(n_pre))
  626. ax.set_yticklabels([str(i) for i in range(n_pre)], fontsize=6)
  627. ax.invert_yaxis()
  628. # ### 2) overlay pool labels at group centers ###
  629. # x‐axis (pool‐pair labels)
  630. post_centers = [j*n_post_pool + n_post_pool/2 for j in range(n_post_blocks)]
  631. for j, lbl in enumerate(pool_labels_post):
  632. ax.text(
  633. post_centers[j], -0.7, # x at center, y just above top row
  634. lbl,
  635. ha='center', va='bottom',
  636. fontsize=10, fontweight='bold',
  637. rotation=90,
  638. clip_on=False
  639. )
  640. # y‐axis (pre‐pool labels)
  641. pre_centers = [i*n_pre_pool + n_pre_pool/2 for i in range(n_pre_blocks)]
  642. for i, lbl in enumerate(pool_labels_pre):
  643. ax.text(
  644. -0.7, pre_centers[i], # x just before first column, y at center
  645. lbl,
  646. ha='right', va='center',
  647. fontsize=10, fontweight='bold',
  648. clip_on=False
  649. )
  650. # title + subtitle
  651. μ, σ = mat.mean(), mat.std()
  652. subtitle = "\n".join(f"{pool_labels_pre[i]}→{pool_labels_post[j]}: {means[i,j]:.2f}±{stds[i,j]:.2f}"
  653. for i in range(n_pre_blocks)
  654. for j in range(n_post_blocks))
  655. ax.set_title(f"{title}\nOverall μ={μ:.2f}, σ={σ:.2f}\n{subtitle}", pad=50)
  656. ax.set_xlabel(" ".join(pool_labels_post[0].split()[:1]).capitalize() + " cell index (receiving)")
  657. ax.set_ylabel(" ".join(pool_labels_pre[0].split()[:1]).capitalize() + " cell index (delivering)")
  658. if add_colorbar:
  659. cbar = fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
  660. cbar.set_label("# of disynaptic connections", rotation=90, labelpad=20)
  661. plt.tight_layout()
  662. if savepath is not None:
  663. plt.savefig(savepath)
  664. plt.show()
  665. def _ensure_logging(): # Define logger to keep track of each simulation's progress despite parallelization
  666. """Make sure each process has a file+console logger attached."""
  667. root = logging.getLogger()
  668. if not root.handlers:
  669. # file handler
  670. fh = logging.FileHandler("simulations_progress_log.log", mode="a")
  671. fh.setFormatter(logging.Formatter("%(asctime)s %(processName)s %(levelname)s: %(message)s"))
  672. fh.setLevel(logging.INFO)
  673. root.addHandler(fh)
  674. # console handler
  675. ch = logging.StreamHandler()
  676. ch.setLevel(logging.INFO)
  677. ch.setFormatter(logging.Formatter("%(processName)s: %(message)s"))
  678. root.addHandler(ch)
  679. root.setLevel(logging.INFO)
  680. def make_unique_output_dir(parent_folder, prefix="simulation_output_"):
  681. i = 0
  682. while True:
  683. dirname = f"{parent_folder}//{prefix}{i}"
  684. try:
  685. # this call is atomic: if two processes hit the same i,
  686. # only one will succeed, the other will get FileExistsError
  687. os.mkdir(dirname)
  688. return dirname, i
  689. except FileExistsError:
  690. i += 1
  691. def wrap_motoneuron_properties(motoneuron_soma_diameter,
  692. motoneurons_resistance, motoneurons_input_weight,
  693. motoneurons_capacitance, motoneurons_membrane_conductance,
  694. motoneurons_membrane_time_constant,
  695. motoneurons_AHP_duration, motoneurons_AHP_conductance_decay_time_constant,
  696. motoneurons_refractory_periods, motoneurons_rheobases,
  697. motoneurons_spike_transmission_delays):
  698. return {
  699. 'soma_diameter': motoneuron_soma_diameter,
  700. 'resistance': motoneurons_resistance,
  701. 'input_weight': motoneurons_input_weight,
  702. 'capacitance': motoneurons_capacitance,
  703. 'membrane_conductance': motoneurons_membrane_conductance,
  704. 'membrane_time_constant': motoneurons_membrane_time_constant,
  705. 'AHP_duration': motoneurons_AHP_duration,
  706. 'AHP_conductance_decay_time_constant': motoneurons_AHP_conductance_decay_time_constant,
  707. 'refractory_period': motoneurons_refractory_periods,
  708. 'rheobase': motoneurons_rheobases,
  709. 'spike_transmission_delay': motoneurons_spike_transmission_delays
  710. }
  711. ######################################
  712. ### FUNCTIONS TO SET UP THE SIMULATION
  713. ######################################
  714. # # # SET BRIAN2 EQUATIONS
  715. def set_brian2_equations():
  716. MN_equations = Equations('''
  717. dv/dt = (
  718. - g_leak*(v - voltage_rest) # leak current from membrane conductance
  719. - g_ahp*(v - voltage_AHP) # AHP current from AHP conductance (based on potassium reversal potential)
  720. + input_weight*(input_MN_timedarray_amp(t,i) + I_syn) # excitatory + inhibitory synaptic currents (weighted by input_weight, i.e. normalized resistance)
  721. )/C_m : volt (unless refractory)
  722. dI_syn/dt = -I_syn/tau_syn : amp # Input from RC decays exponentially
  723. dg_ahp/dt = -g_ahp/tau_ahp : siemens # AHP conductance decays exponentially
  724. g_leak : siemens
  725. C_m : farad
  726. input_weight : 1
  727. voltage_rest : volt
  728. voltage_AHP : volt
  729. refractory_period : second
  730. tau_syn : second # synaptic input from RC time constant (can be either the membrane time constant or an arbitrarily defined time constant)
  731. tau_ahp : second # AHP time constant
  732. ''')
  733. RC_equations = Equations('''
  734. dv/dt = (input_RC_timedarray_volt(t,i)-v)/tau: volt (unless refractory)
  735. tau : second
  736. ''')
  737. return MN_equations, RC_equations
  738. # # # CREATE MOTOR NEURON SIZES AND ASSIGN THEM TO THEIR POOLS
  739. def generate_motor_neurons(full_pool_sim,
  740. min_soma_diameter, max_soma_diameter,
  741. nb_pools, nb_motoneurons_per_pool, total_nb_motoneurons,
  742. size_distribution_exponent,
  743. mean_soma_diameter, sd_prct_soma_diameter,
  744. generate_figure=False, savepath=None):
  745. motoneuron_soma_diameters = np.zeros(total_nb_motoneurons)
  746. motoneuron_normalized_soma_diameters = np.zeros(total_nb_motoneurons)
  747. pool_list_by_MN, idx_of_MN_by_pool = [], {}
  748. for pooli in range(nb_pools):
  749. poolname_temp = f"pool_{pooli}"
  750. idx_of_MN_by_pool[poolname_temp] = np.arange(pooli*nb_motoneurons_per_pool, (pooli+1)*nb_motoneurons_per_pool)
  751. mni_offset = pooli*nb_motoneurons_per_pool
  752. if full_pool_sim == True:
  753. for mni in range(nb_motoneurons_per_pool):
  754. pool_list_by_MN.append(poolname_temp)
  755. motoneuron_soma_diameters[mni_offset+mni] = lerp(min_soma_diameter, max_soma_diameter, (mni/(nb_motoneurons_per_pool-1))**size_distribution_exponent )
  756. motoneuron_normalized_soma_diameters[mni_offset+mni] = lerp(0, 1, (mni/(nb_motoneurons_per_pool-1))**size_distribution_exponent )
  757. else:
  758. motoneuron_sizes_sampled_from_gaussian = np.random.normal(
  759. loc=mean_soma_diameter, scale=(sd_prct_soma_diameter/100)*mean_soma_diameter, size=nb_motoneurons_per_pool)
  760. motoneuron_sizes_sampled_from_gaussian = np.clip(motoneuron_sizes_sampled_from_gaussian, min_soma_diameter, max_soma_diameter)
  761. motoneuron_sizes_sampled_from_gaussian = np.sort(motoneuron_sizes_sampled_from_gaussian)
  762. for mni in range(nb_motoneurons_per_pool):
  763. pool_list_by_MN.append(poolname_temp)
  764. motoneuron_soma_diameters[mni_offset+mni] = motoneuron_sizes_sampled_from_gaussian[mni]
  765. motoneuron_normalized_soma_diameters[mni_offset+mni] = (motoneuron_sizes_sampled_from_gaussian[mni] - min_soma_diameter) / (max_soma_diameter - min_soma_diameter)
  766. if generate_figure:
  767. fig, axes = plt.subplots(1, 3, figsize=(18, 5))
  768. # 1) Histogram of raw soma diameters
  769. ax = axes[0]
  770. weights = np.ones_like(motoneuron_soma_diameters) * (100 / len(motoneuron_soma_diameters))
  771. ax.hist(
  772. motoneuron_soma_diameters,
  773. density=False,
  774. weights=weights,
  775. edgecolor='white',
  776. color='gray',
  777. alpha=1
  778. )
  779. ymin, ymax = ax.get_ylim()
  780. ax.vlines(min_soma_diameter, ymin, ymax, color='C1', label='Min soma diameter', linewidth=2)
  781. ax.vlines(max_soma_diameter, ymin, ymax, color='C3', label='Max soma diameter', linewidth=2)
  782. ax.set_xlabel("Soma diameter (μm)")
  783. ax.set_ylabel("Proportion (% of MNs)")
  784. ax.set_title("Distribution of MN soma diameters")
  785. ax.legend(loc='upper right')
  786. # 2) Histogram of normalized soma diameters
  787. ax = axes[1]
  788. weights_norm = np.ones_like(motoneuron_normalized_soma_diameters) * (100 / len(motoneuron_normalized_soma_diameters))
  789. ax.hist(
  790. motoneuron_normalized_soma_diameters,
  791. density=False,
  792. weights=weights_norm,
  793. edgecolor='white',
  794. color='gray',
  795. alpha=0.5
  796. )
  797. ymin, ymax = ax.get_ylim()
  798. ax.vlines(0, ymin, ymax, color='C1', label='Min normalized', linewidth=2)
  799. ax.vlines(1, ymin, ymax, color='C3', label='Max normalized', linewidth=2)
  800. ax.set_xlabel("Normalized soma diameter")
  801. ax.set_ylabel("Proportion (% of MNs)")
  802. ax.set_title("Normalized MN size distribution")
  803. ax.legend(loc='upper right')
  804. # 3) Soma diameter vs index
  805. ax = axes[2]
  806. ax.plot(motoneuron_soma_diameters, color='gray')
  807. mean_diameter = np.mean(motoneuron_soma_diameters)
  808. ax.hlines(mean_diameter, 0, len(motoneuron_soma_diameters)-1, color='C2', linestyle='--',
  809. label=f"Mean = {mean_diameter:.1f} μm")
  810. ax.set_xlabel("MN index")
  811. ax.set_ylabel("Soma diameter (μm)")
  812. ax.set_title("MN soma diameter vs. index")
  813. ax.legend(loc='upper right')
  814. # Set layout and save
  815. plt.tight_layout()
  816. if savepath is not None:
  817. save_file = os.path.join(savepath, 'MN_sizes.png')
  818. fig.savefig(save_file)
  819. plt.show()
  820. return (motoneuron_soma_diameters, motoneuron_normalized_soma_diameters,
  821. pool_list_by_MN, idx_of_MN_by_pool)
  822. # # # CREATE MOTOR NEURON ELECTROPHYSIOLOGICAL PROPERTIES
  823. def generate_motor_neuron_electrophysiological_properties(
  824. total_nb_motoneurons, motoneuron_soma_diameters,
  825. resistance_constant, resistance_exponent,
  826. capacitance_constant, capacitance_exponent,
  827. AHP_duration_constant, AHP_duration_exponent,
  828. rheobase_constant, rheobase_exponent, rheobase_scaling,
  829. refractory_period_absolute,
  830. axonal_conduction_velocity_constant, axonal_conduction_velocity_exponent,
  831. generate_figure=False, savepath=None):
  832. motoneurons_resistance = np.zeros(total_nb_motoneurons)
  833. motoneurons_input_weight = np.zeros(total_nb_motoneurons)
  834. motoneurons_capacitance = np.zeros(total_nb_motoneurons)
  835. motoneurons_membrane_conductance = np.zeros(total_nb_motoneurons)
  836. motoneurons_membrane_time_constant = np.zeros(total_nb_motoneurons)
  837. motoneurons_AHP_duration = np.zeros(total_nb_motoneurons)
  838. motoneurons_AHP_conductance_decay_time_constant = np.zeros(total_nb_motoneurons)
  839. motoneurons_refractory_periods = np.zeros(total_nb_motoneurons)
  840. motoneurons_rheobases = np.zeros(total_nb_motoneurons)
  841. motoneurons_spike_transmission_delays = np.zeros(total_nb_motoneurons)
  842. for mni in range(total_nb_motoneurons):
  843. motoneurons_resistance[mni] = resistance_constant*(motoneuron_soma_diameters[mni]**resistance_exponent)
  844. motoneurons_input_weight[mni] = motoneurons_resistance[mni] / motoneurons_resistance[0] # normalized value so that smallest MN has weight of 1
  845. motoneurons_capacitance[mni] = capacitance_constant*(motoneuron_soma_diameters[mni]**capacitance_exponent)
  846. motoneurons_membrane_conductance[mni] = 1/motoneurons_resistance[mni]
  847. motoneurons_membrane_time_constant[mni] = (motoneurons_resistance[mni]*ohm) * (motoneurons_capacitance[mni]*farad) / 100 # using the right units should make it be in seconds already. Diving by 100 because some mismatch of scale somewhere
  848. motoneurons_AHP_duration[mni] = AHP_duration_constant*(motoneuron_soma_diameters[mni]**AHP_duration_exponent)
  849. motoneurons_AHP_conductance_decay_time_constant[mni] = motoneurons_AHP_duration[mni] / np.log(10) # Time constant defined so that the AHP duration corresponds to the duration to reach 1/10th of the hyperpolarizing increase in conductance caused by a spike
  850. motoneurons_refractory_periods[mni] = refractory_period_absolute
  851. motoneurons_rheobases[mni] = rheobase_constant*(motoneuron_soma_diameters[mni]**rheobase_exponent)*rheobase_scaling
  852. motoneurons_spike_transmission_delays[mni] = 0.5/(axonal_conduction_velocity_constant*(motoneuron_soma_diameters[mni]**axonal_conduction_velocity_exponent)) # in s # The delay (s) is calculated from the axonal conduction velocity, assuming a 0.5m axon length => so correspond to the conduction speed from MN to muscle fiber (speed in m/s, so multiply speed by 2 -> numerator is 0.5 meter)
  853. motoneuron_properties_dict = wrap_motoneuron_properties(motoneuron_soma_diameters,
  854. motoneurons_resistance, motoneurons_input_weight,
  855. motoneurons_capacitance, motoneurons_membrane_conductance,
  856. motoneurons_membrane_time_constant,
  857. motoneurons_AHP_duration, motoneurons_AHP_conductance_decay_time_constant,
  858. motoneurons_refractory_periods, motoneurons_rheobases,
  859. motoneurons_spike_transmission_delays)
  860. if generate_figure:
  861. # Create a 5×1 grid of subplots
  862. fig, axes = plt.subplots(6, 1, figsize=(8, 22), sharex=True)
  863. # 1) Resistance & Input weight (dual axis)
  864. ax1 = axes[0]
  865. curve1, = ax1.plot(motoneurons_resistance, label="Resistance (Ω)", color='C1', linewidth=2)
  866. ax2 = ax1.twinx()
  867. curve2, = ax2.plot(
  868. motoneurons_input_weight,
  869. label="Normalized input resistance",
  870. color='C5',
  871. linewidth=2,
  872. linestyle=":"
  873. )
  874. ax1.set_ylabel("Resistance (Ω)", color='C1')
  875. ax2.set_ylabel("Normalized input resistance", color='C5')
  876. ax1.tick_params(axis='y', labelcolor='C1')
  877. ax2.tick_params(axis='y', labelcolor='C5')
  878. ax1.legend([curve1, curve2], [curve1.get_label(), curve2.get_label()], loc='best')
  879. ax1.set_title("Motoneuron Resistance & Input Weight")
  880. # 2) Capacitance
  881. ax3 = axes[1]
  882. ax3.plot(motoneurons_capacitance, label="Capacitance (F)", color='C2')
  883. ax3.set_ylabel("Capacitance (F)")
  884. ax3.legend(loc='best')
  885. ax3.set_title("Motoneuron Capacitance")
  886. # 3) Membrane conductance
  887. ax4 = axes[2]
  888. ax4.plot(motoneurons_membrane_conductance, label="Membrane conductance (mS)", color='C3')
  889. ax4.set_ylabel("Conductance (mS)")
  890. ax4.legend(loc='best')
  891. ax4.set_title("Motoneuron Membrane Conductance")
  892. # 5) Membrane time constant (determines RC IPSP effect)
  893. ax5 = axes[3]
  894. ln1, = ax5.plot(motoneurons_membrane_time_constant, label="Membrane time constant (ms)", color='black', linestyle='--')
  895. ax5.set_ylabel("Time (ms)")
  896. ax5.set_title("Membrane time constant\n(can be used for RC's IPSPs decay rate)")
  897. # 6) AHP & Refractory
  898. ax6 = axes[4]
  899. ln1, = ax6.plot(motoneurons_AHP_duration, label="AHP duration (ms)", color='blue', linestyle='--')
  900. ln2, = ax6.plot(motoneurons_AHP_conductance_decay_time_constant, label="AHP time constant (ms)", color='blue')
  901. ln3, = ax6.plot(motoneurons_refractory_periods, label="Refractory period (ms)", color='C6')
  902. ax6.set_ylabel("Time (ms)")
  903. ax6.legend([ln1, ln2, ln3], [ln1.get_label(), ln2.get_label(), ln3.get_label()], loc='best')
  904. ax6.set_title("AHP & Refractory Properties")
  905. # 7) Rheobase
  906. ax7 = axes[5]
  907. ax7.plot(motoneurons_rheobases, label="Rheobase (nA)", color='C7')
  908. ax7.set_xlabel("MN index")
  909. ax7.set_ylabel("Rheobase (nA)")
  910. ax7.legend(loc='best')
  911. ax7.set_title("Motoneuron Rheobase")
  912. # Adjust layout and save
  913. plt.tight_layout()
  914. if savepath is not None:
  915. new_filename = f'MN_electrophysiological_properties.png'
  916. save_file_path = os.path.join(savepath, new_filename)
  917. plt.savefig(save_file_path)
  918. plt.show()
  919. return (motoneurons_resistance, motoneurons_input_weight, motoneurons_capacitance, motoneurons_membrane_conductance, motoneurons_membrane_time_constant,
  920. motoneurons_AHP_conductance_decay_time_constant, motoneurons_refractory_periods, motoneurons_rheobases, motoneurons_spike_transmission_delays,
  921. motoneuron_properties_dict)
  922. # # # GENERATE AND DISTRIBUTE COMMON INPUT(S)
  923. def generate_and_distribute_common_inputs(fsamp, nb_of_common_inputs,
  924. nb_pools, nb_motoneurons_per_pool,
  925. duration_with_ignored_window, edges_ignore_duration,
  926. frequency_range_of_common_input, # array of shape [nb of pools x nb of common inpts x 2] (each cell contains the low and high frequency boundaries)
  927. common_input_std, # array of size [nb of pools x nb of common inpts] (at this point, should already be input as nA)
  928. frequency_range_to_set_input_power,
  929. set_arbitrary_correlation_between_excitatory_inputs,
  930. between_pool_excitatory_input_correlation,
  931. set_same_excitatory_input_for_all_pools,
  932. generate_figure=False, savepath=None
  933. ):
  934. """
  935. Generate nb_of_common_inputs Gaussian noise inputs per pool,
  936. filtered in the range specified by common_input_std.
  937. Optionally impose a target correlation structure,
  938. scale to desired band std, and plot:
  939. 1) time series
  940. 2) power spectrum
  941. 3) correlation matrix
  942. """
  943. # logger = logging.getLogger(__name__)
  944. # 1) Generate raw inputs
  945. MN_excit_input = {}
  946. for pooli in range(nb_pools):
  947. MN_excit_input[pooli] = {}
  948. # Generate_filtered_gaussian_noise_input returns array shape (n_samples,)
  949. artifact_removal_window, _ = filter_artifact_removal(fsamp, edges_ignore_duration)
  950. for inputi in range(nb_of_common_inputs):
  951. MN_excit_input[pooli][inputi] = Generate_filtered_gaussian_noise_input(fsamp,
  952. duration_with_ignored_window, edges_ignore_duration, artifact_removal_window,
  953. low_pass_filter_cutoff=frequency_range_of_common_input[pooli][inputi][1], # the low pass is the second boundary
  954. input_mean=0, scaling_std=1,
  955. high_pass_filter_cutoff=frequency_range_of_common_input[pooli][inputi][0]) # the high pas is the first boundary
  956. # 2) Impose correlation if requested - correlations are adjusted with the following mapping: input 0 of pool 0 is adjusted to input 0 of pool 1, input 1 of pool 0 is adjusted to input 1 of pool 1, etc. ...
  957. if set_arbitrary_correlation_between_excitatory_inputs and nb_pools > 1:
  958. # stack into shape (n_pools, n_samples)
  959. for inputi in range(nb_of_common_inputs):
  960. data = np.vstack([MN_excit_input[pooli][inputi] for pooli in range(nb_pools)])
  961. data = data.T # shape (n_samples, n_pools)
  962. # build target correlation matrix
  963. C_target = np.zeros((nb_pools, nb_pools))
  964. for i in range(nb_pools):
  965. for j in range(nb_pools):
  966. if i == j:
  967. C_target[i, j] = 1.0
  968. else:
  969. C_target[i, j] = between_pool_excitatory_input_correlation
  970. # adjust
  971. transformed = adjust_correlation(data, C_target)
  972. # unpack
  973. for pooli in range(nb_pools):
  974. MN_excit_input[pooli][inputi] = transformed[:, pooli]
  975. # 3) Scale to desired band std - only for the first (likely low-freq) input
  976. for pooli in range(nb_pools):
  977. for inputi in range(nb_of_common_inputs):
  978. if inputi == 0: # The first input (which should be low freq) is scaled relative to the desired frequency band.
  979. MN_excit_input[pooli][inputi] = scale_to_band_std(
  980. MN_excit_input[pooli][inputi],
  981. fsamp,
  982. frequency_range_to_set_input_power,
  983. common_input_std[pooli][inputi])
  984. else: # The other inputs are scaled directly to their requested stds, relative to the power of the first input in the frequency_range_to_set_input_power
  985. freq_range_of_input_temp_diff = frequency_range_of_common_input[pooli][inputi][1] - frequency_range_of_common_input[pooli][inputi][0]
  986. frequency_range_to_set_input_power_diff = frequency_range_to_set_input_power[1] - frequency_range_to_set_input_power[0]
  987. MN_excit_input[pooli][inputi] *= (
  988. common_input_std[pooli][inputi] * np.sqrt(freq_range_of_input_temp_diff/frequency_range_to_set_input_power_diff))
  989. # 4) Optionally make all pools share the same input
  990. if set_same_excitatory_input_for_all_pools:
  991. for inputi in range(nb_of_common_inputs):
  992. base = MN_excit_input[0][inputi]
  993. for pooli in range(nb_pools):
  994. MN_excit_input[pooli][inputi] = base
  995. # 5) Collapse (merge) all the inputs from a given pool together
  996. for pooli in range(nb_pools):
  997. # Get all the arrays for this pool as a real list
  998. arrays = list(MN_excit_input[pooli].values())
  999. # Start with zeros of the same shape as the first array
  1000. total = np.zeros_like(arrays[0])
  1001. # Sum them all
  1002. for arr in arrays:
  1003. total += arr
  1004. # Replace the dict with the collapsed array
  1005. MN_excit_input[pooli] = total
  1006. # 6) Compute pairwise correlation matrix
  1007. corr_mat = np.eye(nb_pools)
  1008. for i in range(nb_pools):
  1009. for j in range(i+1, nb_pools):
  1010. r = np.corrcoef(MN_excit_input[i], MN_excit_input[j])[0,1]
  1011. corr_mat[i,j] = corr_mat[j,i] = r
  1012. # 7) Build weight matrix: each MN receives 100% of its pool's input
  1013. total_nb_mn = nb_pools * nb_motoneurons_per_pool
  1014. inputs_to_mn_weight_matrix = np.zeros((total_nb_mn, nb_pools))
  1015. for pooli in range(nb_pools):
  1016. row_start = pooli * nb_motoneurons_per_pool
  1017. row_end = row_start + nb_motoneurons_per_pool
  1018. inputs_to_mn_weight_matrix[row_start:row_end, pooli] = 1.0
  1019. # 8) Compute power spectrum of common input
  1020. total_power = {}
  1021. power_within_band_of_interest = {}
  1022. psd = {}
  1023. # psd_within_band = {}
  1024. nperseg = min(fsamp, len(MN_excit_input[0])) # only first pool
  1025. for pooli in range(nb_pools):
  1026. freqs, psd[pooli] = welch( # freqs stay the same each time so no need to have a dict containing them for each pool
  1027. MN_excit_input[pooli],
  1028. fs=fsamp,
  1029. window='hann',
  1030. nperseg=nperseg,
  1031. noverlap=nperseg // 2,
  1032. scaling='density',
  1033. detrend='constant')
  1034. total_power[pooli] = np.trapz(psd[pooli], freqs)
  1035. mask = (freqs >= frequency_range_to_set_input_power[0]) & (freqs <= frequency_range_to_set_input_power[1])
  1036. power_within_band_of_interest[pooli] = np.trapz(psd[pooli][mask], freqs[mask])
  1037. # logger.info(f"Time domain total power of common input = {np.mean(MN_excit_input[0]**2):.2f}")
  1038. # logger.info(f"Frequency domain total power of common input = {total_power[0]:.2f}")
  1039. # logger.info(f"Time domain within-band power of common input = undefined yet")
  1040. # logger.info(f"Frequency domain within-band power of common input = {power_within_band_of_interest[0]:.2f}")
  1041. # 8) Plotting: time series, power spectrum, correlation matrix, distribution of inputs to MNs
  1042. max_freq_lim = 150
  1043. power_per_frequency_band = {"frequencies": freqs[np.array(range(max_freq_lim)).astype(int)],
  1044. "power": psd[0][np.array(range(max_freq_lim)).astype(int)]} # only first pool
  1045. # Save power_per_frequency_band as csv file
  1046. if savepath is not None:
  1047. power_per_frequency_band_df = pd.DataFrame(power_per_frequency_band)
  1048. power_per_frequency_band_df.to_csv(os.path.join(savepath, 'power_per_frequency_band_common_input.csv'), index=False)
  1049. if generate_figure:
  1050. pool_colors = {0: "blue", 1: "red", 2: "orange", 3: "green"}
  1051. # Build figure with GridSpec
  1052. fig = plt.figure(figsize=(14, 10))
  1053. gs = fig.add_gridspec(2, 4, height_ratios=[1, 1], width_ratios=[1, 1, 1, 1], hspace=0.4, wspace=0.3)
  1054. custom_input_labels = [f'Pool {p}' for p in range(nb_pools)]
  1055. # 1) Time series - full (row 0, all columns)
  1056. ax_ts = fig.add_subplot(gs[0, 0:2])
  1057. t = np.linspace(0, duration_with_ignored_window, len(MN_excit_input[0]))
  1058. for pooli in range(nb_pools):
  1059. ax_ts.plot(t, MN_excit_input[pooli]/1000, label=f'Pool {pooli}',
  1060. color=pool_colors[pooli], alpha = 0.5)
  1061. ax_ts.set_xlabel("Time (s)")
  1062. ax_ts.set_ylabel("Input amplitude (microAmperes)")
  1063. ax_ts.set_title("Common Input(s) Time Series (full)")
  1064. ax_ts.legend(loc='upper right')
  1065. # 1) Time series - zoomed-in (first 3 seconds) (row 0, all columns)
  1066. ax_ts_zommed = fig.add_subplot(gs[0, 2:])
  1067. time_to_display = np.min([edges_ignore_duration+3, duration_with_ignored_window])
  1068. for pooli in range(nb_pools):
  1069. ax_ts_zommed.plot(t, MN_excit_input[pooli]/1000, label=f'Pool {pooli}',
  1070. color=pool_colors[pooli], alpha = 0.5)
  1071. ax_ts_zommed.set_xlabel("Time (s)")
  1072. ax_ts_zommed.set_ylabel("Input amplitude (microAmperes)")
  1073. ax_ts_zommed.set_title("Common Input(s) Time Series (Zoomed-in)")
  1074. ax_ts_zommed.set_xlim(left=edges_ignore_duration, right=time_to_display)
  1075. ax_ts_zommed.legend(loc='upper right')
  1076. # 2) Power spectrum (row 1, columns 0-1)
  1077. ax_spec = fig.add_subplot(gs[1, 0:2])
  1078. plt.axvspan(frequency_range_to_set_input_power[0], frequency_range_to_set_input_power[1],
  1079. color='red', # the fill color
  1080. alpha=0.2, # transparency
  1081. label=f"Scaling band: {frequency_range_to_set_input_power[0]}–{frequency_range_to_set_input_power[1]} Hz")
  1082. for pooli in range(nb_pools):
  1083. # first, fill under the PSD curve:
  1084. ax_spec.fill_between(
  1085. freqs, psd[pooli],
  1086. y2=0, color=pool_colors[pooli], alpha=0.3,
  1087. label=f'Pool {pooli}\n - Power={total_power[pooli]:.2f})\n - Power within band to set input={power_within_band_of_interest[pooli]:.2f}')
  1088. # then draw the line on top
  1089. ax_spec.plot(freqs, psd[pooli], color=pool_colors[pooli])
  1090. ax_spec.set_xlabel("Frequency (Hz)")
  1091. ax_spec.set_ylabel("Power spectrum")
  1092. ax_spec.set_xlim([0, max_freq_lim])
  1093. ax_spec.set_title("Common input(s) power Spectrum (Welch)")
  1094. ax_spec.legend(fontsize=8, loc='upper right')
  1095. # 3) Correlation matrix (row 1, column 2)
  1096. ax_corr = fig.add_subplot(gs[1, 2])
  1097. sns.heatmap(
  1098. corr_mat,
  1099. ax=ax_corr,
  1100. annot=True,
  1101. cmap="Spectral_r", # "RdYlBu_r",
  1102. vmin=-1, vmax=1,
  1103. xticklabels=custom_input_labels,
  1104. yticklabels=custom_input_labels
  1105. )
  1106. ax_corr.set_title("Pairwise Correlation Matrix")
  1107. ax_corr.set_xlabel("Pool")
  1108. ax_corr.set_ylabel("Pool")
  1109. # 4) Weight distribution matrix (row 1, column 3)
  1110. ax_wgt = fig.add_subplot(gs[1, 3])
  1111. sns.heatmap(
  1112. inputs_to_mn_weight_matrix,
  1113. ax=ax_wgt,
  1114. cmap= "viridis", # "RdYlBu_r",
  1115. vmin=0, vmax=1,
  1116. xticklabels=custom_input_labels,
  1117. yticklabels=False # turn off Seaborn’s default labels
  1118. )
  1119. # now add just every Nth MN index:
  1120. max_labels = 10
  1121. step = max(1, total_nb_mn // max_labels)
  1122. yticks = np.arange(0, total_nb_mn, step)
  1123. ax_wgt.set_yticks(yticks)
  1124. ax_wgt.set_yticklabels(yticks)
  1125. ax_wgt.set_title("Input distribution to MNs")
  1126. ax_wgt.set_xlabel("Pool Input")
  1127. ax_wgt.set_ylabel("MN Index")
  1128. # Adjust layout and save
  1129. plt.tight_layout()
  1130. if savepath is not None:
  1131. new_filename = f'Common_inputs.png'
  1132. save_file_path = os.path.join(savepath, new_filename)
  1133. plt.savefig(save_file_path)
  1134. plt.show()
  1135. return (MN_excit_input, corr_mat, inputs_to_mn_weight_matrix,
  1136. total_power, power_within_band_of_interest, power_per_frequency_band)
  1137. # # # GENERATE INDEPENDENT INPUTS AND CREATE BRIAN2 TIME ARRAYS
  1138. def generate_independent_inputs(fsamp,
  1139. nb_pools, total_nb_motoneurons, total_nb_renshaw_cells,
  1140. low_pass_filter_of_MN_independent_input,
  1141. low_pass_filter_of_RC_independent_input,
  1142. ref_common_input_power,
  1143. MN_independent_input_absolute_or_ratio,
  1144. MN_independent_input_power,
  1145. RC_independent_input_std,
  1146. MN_excit_input,
  1147. inputs_to_mn_weight_matrix,
  1148. excitatory_input_baseline, # list of size <= nb_pools
  1149. duration_with_ignored_window,
  1150. edges_ignore_duration,
  1151. motoneurons_rheobases,
  1152. generate_figure=False, savepath=None
  1153. ):
  1154. # logger = logging.getLogger(__name__)
  1155. # Motor neuron - generate independent input
  1156. artifact_removal_window, _ = filter_artifact_removal(fsamp, edges_ignore_duration) # Get artifact removal window
  1157. MN_independent_input = []
  1158. nperseg = min(fsamp, len(MN_excit_input[0])) # only first pool
  1159. for mni in range(total_nb_motoneurons):
  1160. temp_input = Generate_filtered_gaussian_noise_input(fsamp,
  1161. duration_with_ignored_window, edges_ignore_duration, artifact_removal_window,
  1162. low_pass_filter_cutoff=low_pass_filter_of_MN_independent_input,
  1163. input_mean=0, scaling_std=1)
  1164. pooli = np.round(mni // (total_nb_motoneurons/nb_pools)).astype(int)
  1165. # Scale it to the desired power
  1166. temp_input_power = np.mean(temp_input**2)
  1167. if MN_independent_input_absolute_or_ratio == 'absolute':
  1168. independent_input_scaling_factor = MN_independent_input_power
  1169. elif MN_independent_input_absolute_or_ratio == 'ratio':
  1170. desired_independent_power = MN_independent_input_power * ref_common_input_power[pooli]
  1171. independent_input_scaling_factor = np.sqrt(desired_independent_power / temp_input_power)
  1172. temp_input *= independent_input_scaling_factor
  1173. # Keep the first independent input as a separate variable for plotting purpose:
  1174. MN_independent_input.append(temp_input)
  1175. if generate_figure and mni==0:
  1176. independent_input_for_plotting = temp_input.copy()
  1177. # # Some logs for debugging
  1178. # logger.info(f"Pool # = {pooli}")
  1179. # logger.info(f"inital temp_input_power = {temp_input_power:.2f}")
  1180. # logger.info(f"ref_common_input_power = {ref_common_input_power[0]:.2f}")
  1181. # logger.info(f"desired_independent_power = {desired_independent_power:.2f}")
  1182. # logger.info(f"Scaling factor = {independent_input_scaling_factor:.2f}")
  1183. # logger.info(f"Transformed temp input power = {np.mean(temp_input**2):.2f}")
  1184. # logger.info(f"independent_input_for_plotting power = {np.mean(independent_input_for_plotting**2):.2f}")
  1185. # logger.info(f"Excitatory input shape = {MN_excit_input[pooli].shape}")
  1186. # # logger.info(f"Excitatory input power = {np.mean((np.array(MN_excit_input[pooli]).flatten() * inputs_to_mn_weight_matrix[mni, pooli])**2)}")
  1187. # logger.info(f"Excitatory input power = {np.mean(np.array(MN_excit_input[pooli])**2)}")
  1188. # logger.info(f"Excitatory input power multiplied by weight matrix = {np.mean((np.array(MN_excit_input[pooli]) * inputs_to_mn_weight_matrix[mni, pooli])**2)}")
  1189. # Get mean power of the independent input
  1190. max_freq_lim = 150
  1191. psd_independent_input = []
  1192. for mni in range(len(MN_independent_input)):
  1193. freqs, psd_temp = welch( # freqs stay the same each time so no need to have a dict containing them for each pool
  1194. MN_independent_input[mni],
  1195. fs=fsamp,
  1196. window='hann',
  1197. nperseg=nperseg,
  1198. noverlap=nperseg // 2,
  1199. scaling='density',
  1200. detrend='constant')
  1201. psd_independent_input.append(psd_temp)
  1202. power_per_frequency_band_independent = {"frequencies": freqs,
  1203. "power": np.mean(np.array(psd_independent_input), axis = 0)}
  1204. total_power_independent = np.trapz(power_per_frequency_band_independent["power"], freqs)
  1205. power_per_frequency_band_independent = {"frequencies": freqs[np.array(range(max_freq_lim)).astype(int)],
  1206. "power": power_per_frequency_band_independent["power"][np.array(range(max_freq_lim)).astype(int)]}
  1207. # Save power_per_frequency_band as csv file
  1208. if savepath is not None:
  1209. power_per_frequency_band_df = pd.DataFrame(power_per_frequency_band_independent)
  1210. power_per_frequency_band_df.to_csv(os.path.join(savepath, 'power_per_frequency_band_independent_input.csv'), index=False)
  1211. temp_timed_array = np.zeros((int(np.round(duration_with_ignored_window/second*fsamp)), total_nb_motoneurons))
  1212. for mni in range(total_nb_motoneurons):
  1213. pooli = np.round(mni // (total_nb_motoneurons/nb_pools)).astype(int)
  1214. temp_timed_array[:, mni] += np.array(
  1215. MN_excit_input[pooli]) * inputs_to_mn_weight_matrix[mni, pooli] # the weights are relative to the distribution of common input to the different pools, not the size of the motor neurons!
  1216. # add excitatory baseline and independent input
  1217. temp_timed_array[:, mni] += excitatory_input_baseline[pooli]
  1218. temp_timed_array[:, mni] += MN_independent_input[mni]
  1219. # Rheobase = clip value to 0 if it is below a given value (in nA)
  1220. temp_timed_array[:, mni] = np.clip(
  1221. temp_timed_array[:, mni]-motoneurons_rheobases[mni],
  1222. a_min=0, a_max=np.inf)
  1223. input_MN_timedarray_amp = TimedArray(temp_timed_array * nA, dt=(1/fsamp)*second) # in nano Ampere
  1224. # Renshaw cell - generate independent input and fill time array
  1225. RC_independent_input = []
  1226. temp_timed_array = np.zeros((int(np.round(duration_with_ignored_window/second*fsamp)), total_nb_renshaw_cells))
  1227. for renshawi in range(total_nb_renshaw_cells):
  1228. RC_independent_input.append(Generate_filtered_gaussian_noise_input(fsamp,
  1229. duration_with_ignored_window, edges_ignore_duration, artifact_removal_window,
  1230. low_pass_filter_cutoff=low_pass_filter_of_RC_independent_input,
  1231. input_mean=0, scaling_std=RC_independent_input_std))
  1232. input_RC_timedarray_volt = TimedArray(temp_timed_array * mvolt, dt=(1/fsamp)*second)
  1233. # Sanity check plot
  1234. if generate_figure:
  1235. plt.figure(figsize=(30,5))
  1236. plt.plot(input_MN_timedarray_amp.values[:,
  1237. np.linspace(0, total_nb_motoneurons-1, 10, dtype=int)],
  1238. alpha=0.2, color = 'C0')
  1239. test = np.mean(input_MN_timedarray_amp.values[:,np.arange(1,total_nb_motoneurons)],axis=1)
  1240. plt.plot(test, color='darkblue')
  1241. plt.xlabel("Time (samples)")
  1242. plt.ylabel("Input (Amperes)")
  1243. if savepath is not None:
  1244. new_filename = f'Input_TimedArray_SanityCheck.png'
  1245. save_file_path = os.path.join(savepath, new_filename)
  1246. plt.savefig(save_file_path)
  1247. plt.show()
  1248. plt.figure()
  1249. duration_to_plot = 3 # in seconds
  1250. nb_samples_to_plot = int(np.round(duration_to_plot*fsamp))
  1251. independent_input_for_plotting_power = np.mean(independent_input_for_plotting**2)
  1252. common_input_for_plotting_power = np.mean(MN_excit_input[0]**2)
  1253. if nb_samples_to_plot > len(independent_input_for_plotting):
  1254. nb_samples_to_plot = len(independent_input_for_plotting)
  1255. plt.plot(np.arange(nb_samples_to_plot)/fsamp,
  1256. MN_excit_input[0][:nb_samples_to_plot] * inputs_to_mn_weight_matrix[0, 0],
  1257. label=f"Common input delivered to MN#0\nPower={common_input_for_plotting_power:.2f}",
  1258. color = 'red')
  1259. plt.plot(np.arange(nb_samples_to_plot)/fsamp,
  1260. independent_input_for_plotting[:nb_samples_to_plot],
  1261. label=f"Independent input delivered to MN#0\nPower={independent_input_for_plotting_power:.2f}",
  1262. color='blue', alpha=0.5, linewidth=0.5)
  1263. plt.xlabel("Time (s)")
  1264. plt.ylabel("Input (nA)")
  1265. plt.title(f"Recalculated ratio of independent VS common input = {independent_input_for_plotting_power/common_input_for_plotting_power:.2f}")
  1266. plt.legend()
  1267. if savepath is not None:
  1268. new_filename = f'Input_common_VS_independent_MN0.png'
  1269. save_file_path = os.path.join(savepath, new_filename)
  1270. plt.savefig(save_file_path)
  1271. plt.show()
  1272. return input_MN_timedarray_amp, input_RC_timedarray_volt, total_power_independent, power_per_frequency_band_independent
  1273. # # # CREATE CONNECTIVITY BETWEEN MOTOR NEURONS AND RENSHAW CELLS
  1274. def create_connectivity(total_nb_motoneurons, total_nb_renshaw_cells,
  1275. nb_pools, nb_motoneurons_per_pool, RC_pair_indices, nb_RCs_per_pool_pair,
  1276. disynpatic_inhib_connections_desired_MN_MN,
  1277. split_MN_RC_ratio, motoneuron_normalized_soma_diameters,
  1278. distribution_type, distribution_params, disynaptic_inhib_received_arbitrary_adjustment,
  1279. distribution_binary_weights,
  1280. generate_figure=False, savepath=None):
  1281. # --- Allocate global adjacency matrices ---
  1282. MN_to_Renshaw_connectivity_matrix = np.zeros(
  1283. (total_nb_motoneurons, total_nb_renshaw_cells), dtype=float
  1284. )
  1285. Renshaw_to_MNs_connectivity_matrix = np.zeros(
  1286. (total_nb_renshaw_cells, total_nb_motoneurons), dtype=float
  1287. )
  1288. # --- Fill in each pool‐pair block independently ---
  1289. motoneuron_size_ranks = np.asarray(motoneuron_normalized_soma_diameters) # need this if using size_gaussian
  1290. for i in range(nb_pools):
  1291. mn_pre = np.arange(i*nb_motoneurons_per_pool,
  1292. (i+1)*nb_motoneurons_per_pool)
  1293. for j in range(nb_pools):
  1294. mn_post = np.arange(j*nb_motoneurons_per_pool,
  1295. (j+1)*nb_motoneurons_per_pool)
  1296. rc_inds = RC_pair_indices[(i, j)]
  1297. nRC = len(rc_inds)
  1298. # desired *mean disynaptic count* for this block
  1299. mean_disyn = disynpatic_inhib_connections_desired_MN_MN[i, j]
  1300. # solve p_mn_rc * p_rc_mn = mean_disyn / nRC
  1301. base_conn = mean_disyn / nRC
  1302. p_mn_rc = base_conn ** split_MN_RC_ratio
  1303. p_rc_mn = base_conn ** (1.0 - split_MN_RC_ratio)
  1304. # sample MN→RC sub‐matrix and write into global
  1305. subA = sample_block(
  1306. n_pre=len(mn_pre),
  1307. n_post=nRC,
  1308. p_block=p_mn_rc,
  1309. dist=distribution_type,
  1310. params=distribution_params,
  1311. ranks=(motoneuron_size_ranks[mn_pre]
  1312. if distribution_type=='size_gaussian' else None),
  1313. binary=distribution_binary_weights
  1314. )
  1315. MN_to_Renshaw_connectivity_matrix[np.ix_(mn_pre, rc_inds)] = subA
  1316. # sample RC→MN sub‐matrix and write into global
  1317. subB = sample_block(
  1318. n_pre=nRC,
  1319. n_post=len(mn_post),
  1320. p_block=p_rc_mn,
  1321. dist=distribution_type,
  1322. params=distribution_params,
  1323. ranks=None,
  1324. binary=distribution_binary_weights
  1325. )
  1326. Renshaw_to_MNs_connectivity_matrix[np.ix_(rc_inds, mn_post)] = subB
  1327. # --- Modify the connectivity from RCs to MNs from the point of view of each MN
  1328. if disynaptic_inhib_received_arbitrary_adjustment > 0:
  1329. for mni in range(total_nb_motoneurons):
  1330. Renshaw_to_MNs_connectivity_matrix[:, mni] += np.random.normal(loc=0, scale=disynaptic_inhib_received_arbitrary_adjustment, size=total_nb_renshaw_cells)
  1331. # Make sure the weights are non-negative (no excitation from Renshaw cells!)
  1332. Renshaw_to_MNs_connectivity_matrix = np.maximum(Renshaw_to_MNs_connectivity_matrix, 0)
  1333. # --- Compute disynaptic MN→MN counts (or binary for 'binarize') ---
  1334. MN_to_MN_counts = MN_to_Renshaw_connectivity_matrix.dot(
  1335. Renshaw_to_MNs_connectivity_matrix)
  1336. # np.fill_diagonal(MN_to_MN_counts, 0)
  1337. if distribution_type == 'binarize':
  1338. # interpret exactly as mean *probability* of ≥1 path
  1339. MN_to_MN_connectivity_matrix = (MN_to_MN_counts > 0).astype(int)
  1340. else:
  1341. # keep raw counts as your “weights”
  1342. MN_to_MN_connectivity_matrix = MN_to_MN_counts
  1343. ### Plotting, if requested
  1344. if generate_figure:
  1345. mn_pool_labels = [f"MN pool {i}" for i in range(nb_pools)]
  1346. rc_pair_list = [(i,j) for i in range(nb_pools) for j in range(nb_pools)]
  1347. rc_pool_pair_labels = [f"RC pool pair {i}-{j}" for (i,j) in rc_pair_list]
  1348. plot_connectivity_matrix(
  1349. MN_to_Renshaw_connectivity_matrix,
  1350. "Monosynaptic MN → RC",
  1351. mn_pool_labels,
  1352. rc_pool_pair_labels,
  1353. n_pre_pool=nb_motoneurons_per_pool,
  1354. n_post_pool=nb_RCs_per_pool_pair,
  1355. cmap='autumn',
  1356. add_colorbar=False,
  1357. savepath=f'{savepath}/Connectivity_MN_to_RC.png'
  1358. )
  1359. plot_connectivity_matrix(
  1360. Renshaw_to_MNs_connectivity_matrix,
  1361. "Monosynaptic RC → MN",
  1362. rc_pool_pair_labels,
  1363. mn_pool_labels,
  1364. n_pre_pool=nb_RCs_per_pool_pair,
  1365. n_post_pool=nb_motoneurons_per_pool,
  1366. cmap='winter',
  1367. add_colorbar=False,
  1368. savepath=f'{savepath}/Connectivity_RC_to_MN.png'
  1369. )
  1370. plot_connectivity_matrix(
  1371. MN_to_MN_connectivity_matrix,
  1372. "Disynaptic MN → MN via RCs",
  1373. mn_pool_labels,
  1374. mn_pool_labels,
  1375. n_pre_pool=nb_motoneurons_per_pool,
  1376. n_post_pool=nb_motoneurons_per_pool,
  1377. cmap='viridis',
  1378. add_colorbar=True,
  1379. savepath=f'{savepath}/Connectivity_MN_to_MN.png'
  1380. )
  1381. # ─── Histograms per pool→pool ─────────────────── #
  1382. nP = nb_pools
  1383. counts_MN_RC = {}
  1384. counts_RC_MN = {}
  1385. counts_MN_MN = {}
  1386. for i in range(nP):
  1387. mn_pre = slice(i*nb_motoneurons_per_pool, (i+1)*nb_motoneurons_per_pool)
  1388. for j in range(nP):
  1389. mn_post = slice(j*nb_motoneurons_per_pool, (j+1)*nb_motoneurons_per_pool)
  1390. rc_inds = RC_pair_indices[(i,j)]
  1391. counts_MN_RC[(i,j)] = (
  1392. MN_to_Renshaw_connectivity_matrix[np.ix_(range(*mn_pre.indices(total_nb_motoneurons)), rc_inds)]
  1393. .sum(axis=1)
  1394. )
  1395. counts_RC_MN[(i,j)] = (
  1396. Renshaw_to_MNs_connectivity_matrix[np.ix_(rc_inds, range(*mn_post.indices(total_nb_motoneurons)))]
  1397. .sum(axis=0)
  1398. )
  1399. counts_MN_MN[(i,j)] = (
  1400. MN_to_MN_connectivity_matrix[mn_pre, mn_post]
  1401. .sum(axis=1)
  1402. )
  1403. fig = plt.figure(figsize=(4*nP, 4*nP))
  1404. outer = gridspec.GridSpec(nP, nP, wspace=0.4, hspace=0.6)
  1405. for i in range(nP):
  1406. for j in range(nP):
  1407. cell = outer[i,j]
  1408. inner = gridspec.GridSpecFromSubplotSpec(2,1, subplot_spec=cell,
  1409. height_ratios=[1,1], hspace=0.2)
  1410. # top
  1411. ax1 = fig.add_subplot(inner[0])
  1412. ax1.hist(counts_MN_RC[(i,j)], bins='auto', alpha=0.6, label='MN→RC', color='C1')
  1413. ax1.hist(counts_RC_MN[(i,j)], bins='auto', alpha=0.6, label='RC→MN', color='C0')
  1414. ax1.set_ylabel("Count")
  1415. ax1.set_title(f"Pools {i}→{j}")
  1416. ax1.legend(fontsize=8, loc="upper left")
  1417. ax1.text(0.95,0.7,
  1418. f"μ₁={counts_MN_RC[(i,j)].mean():.1f}±{counts_MN_RC[(i,j)].std():.1f}\n"
  1419. f"μ₂={counts_RC_MN[(i,j)].mean():.1f}±{counts_RC_MN[(i,j)].std():.1f}",
  1420. transform=ax1.transAxes, ha='right', va='top', fontsize=7)
  1421. # bottom
  1422. ax2 = fig.add_subplot(inner[1])
  1423. ax2.hist(counts_MN_MN[(i,j)], bins='auto', alpha=0.6, color='green', label='MN→MN')
  1424. ax2.set_xlabel("Number of synapses")
  1425. ax2.set_ylabel("Count")
  1426. ax2.legend(fontsize=8, loc="upper left")
  1427. ax2.text(0.95,0.7,
  1428. f"μ₃={counts_MN_MN[(i,j)].mean():.1f}±{counts_MN_MN[(i,j)].std():.1f}",
  1429. transform=ax2.transAxes, ha='right', va='top', fontsize=7)
  1430. plt.suptitle("Number of synapses per pool, per cell type")
  1431. plt.tight_layout()
  1432. if savepath is not None:
  1433. new_filename = f'Connectivity_histogram.png'
  1434. save_file_path = os.path.join(savepath, new_filename)
  1435. plt.savefig(save_file_path)
  1436. plt.show()
  1437. return MN_to_MN_connectivity_matrix, MN_to_Renshaw_connectivity_matrix, Renshaw_to_MNs_connectivity_matrix
  1438. # # # CREATE Brian2 NEURONGROUPS AND SYNAPSES OBJECTS
  1439. def create_neurongroups_and_synapses_objects(
  1440. total_nb_motoneurons, total_nb_renshaw_cells,
  1441. MN_equations, RC_equations, voltage_rest, voltage_thresh,
  1442. motoneurons_membrane_conductance, motoneurons_capacitance,
  1443. motoneurons_input_weight, synaptic_IPSP_decay_time_constant_per_MN,
  1444. motoneurons_AHP_conductance_decay_time_constant, motoneurons_refractory_periods,
  1445. AHP_conductance_delta_after_spiking, MN_to_Renshaw_excit, Renshaw_to_MN_inhib,
  1446. scale_initial_IPSP_to_be_same_integral_regardless_of_synaptic_tau,
  1447. tau_Renshaw, MN_RC_synpatic_delay, refractory_period_RC,
  1448. MN_to_Renshaw_connectivity_matrix, Renshaw_to_MNs_connectivity_matrix):
  1449. logger = logging.getLogger(__name__)
  1450. # # ----------------- BRIAN2 COMMON NAMESPACE FOR VARIABLES
  1451. common_brian2_namespace = {
  1452. 'voltage_thresh': voltage_thresh,
  1453. 'voltage_rest': voltage_rest,
  1454. 'refractory_period_RC': refractory_period_RC,
  1455. 'AHP_conductance_delta_after_spiking': AHP_conductance_delta_after_spiking,
  1456. 'MN_to_Renshaw_excit': MN_to_Renshaw_excit,
  1457. 'Renshaw_to_MN_IPSP_integral': Renshaw_to_MN_inhib * second # Should be in amp * s = Coulomb (total charge). Renshaw_to_MN_inhib is already in amp.
  1458. }
  1459. # ----------------- NEURON GROUPS -----------
  1460. # MOTOR NEURONS
  1461. motoneurons = NeuronGroup(
  1462. total_nb_motoneurons,
  1463. MN_equations,
  1464. threshold='v > voltage_thresh',
  1465. reset='''
  1466. v = voltage_rest
  1467. g_ahp += AHP_conductance_delta_after_spiking
  1468. ''',
  1469. refractory='refractory_period',
  1470. method='euler',
  1471. namespace=common_brian2_namespace
  1472. )
  1473. motoneurons.v = voltage_rest # in mV #
  1474. motoneurons.g_leak = motoneurons_membrane_conductance * msiemens # in milisiemens
  1475. motoneurons.C_m = motoneurons_capacitance * ufarad # in microfarads
  1476. motoneurons.refractory_period = motoneurons_refractory_periods * ms # in milliseconds
  1477. motoneurons.input_weight = motoneurons_input_weight # dimensionless unit
  1478. motoneurons.tau_syn = synaptic_IPSP_decay_time_constant_per_MN # already in millisecond
  1479. motoneurons.tau_ahp = motoneurons_AHP_conductance_decay_time_constant * ms # in millisecond
  1480. motoneurons.g_ahp = 0*siemens # Initialize AHP with a current of 0
  1481. # logger.info(f"motoneurons tau syn = {motoneurons.tau_syn}")
  1482. # logger.info(f"motoneurons input weights = {motoneurons.input_weight}")
  1483. # logger.info(f"IPSP integrals = {common_brian2_namespace['Renshaw_to_MN_IPSP_integral']}")
  1484. # RENSHAW CELLS
  1485. renshaw_cells = NeuronGroup(
  1486. total_nb_renshaw_cells,
  1487. RC_equations,
  1488. threshold='v > voltage_thresh',
  1489. reset='v = voltage_rest',
  1490. refractory='refractory_period_RC',
  1491. method='euler',
  1492. namespace=common_brian2_namespace
  1493. )
  1494. renshaw_cells.v = voltage_rest # Initialize membrane potential
  1495. renshaw_cells.tau = tau_Renshaw
  1496. # ----------------- SYNAPSES -----------
  1497. # Connect motor neurons to Renshaw cells
  1498. synapses_MN_to_Renshaw = Synapses(motoneurons, renshaw_cells, 'w : 1',
  1499. on_pre='v += MN_to_Renshaw_excit*w',
  1500. delay = MN_RC_synpatic_delay,
  1501. namespace=common_brian2_namespace)
  1502. pre_indices, post_indices = np.nonzero(MN_to_Renshaw_connectivity_matrix)
  1503. weights_to_assign = MN_to_Renshaw_connectivity_matrix[pre_indices,post_indices]
  1504. if len(pre_indices)>0 and len(post_indices)>0:
  1505. synapses_MN_to_Renshaw.connect(i=pre_indices, j=post_indices)
  1506. synapses_MN_to_Renshaw.w = weights_to_assign
  1507. else:
  1508. synapses_MN_to_Renshaw.active = False
  1509. # Connect Renshaw cells to motor neurons
  1510. # # ----------------- CHANGE Renshaw_to_MN_inhib IF IT IS USED AS A TARGET RELATIVE TO THE SYNAPTIC TIME CONSTANT
  1511. if scale_initial_IPSP_to_be_same_integral_regardless_of_synaptic_tau:
  1512. on_pre_action = 'I_syn -= (Renshaw_to_MN_IPSP_integral * w) / tau_syn'
  1513. logger.info(f" Defined IPSP is an integral over hyperpolarizing current (hyperpolarizing charge, in Coulomb) = {common_brian2_namespace['Renshaw_to_MN_IPSP_integral']}")
  1514. else:
  1515. on_pre_action = 'I_syn -= (Renshaw_to_MN_IPSP_integral * w) / (1*second)' # Re-interpret the initial IPSP amplitude as an area over 1s
  1516. logger.info(f" Defined IPSP is the initial hyperpolarizing current induced by the Renshaw cell's IPSP (hyperpolarizing current, in Amp) = {common_brian2_namespace['Renshaw_to_MN_IPSP_integral']/second}")
  1517. synapses_Renshaw_to_MN = Synapses(renshaw_cells, motoneurons,
  1518. '''
  1519. w: 1
  1520. Renshaw_to_MN_IPSP_integral: amp*second # the total ∫I(t)dt you want = total charge (Coulomb)
  1521. ''',
  1522. on_pre=on_pre_action, # pick up each post‐cell's tau_syn
  1523. delay = MN_RC_synpatic_delay,
  1524. namespace=common_brian2_namespace)
  1525. pre_indices, post_indices = np.nonzero(Renshaw_to_MNs_connectivity_matrix)
  1526. weights_to_assign = Renshaw_to_MNs_connectivity_matrix[pre_indices, post_indices]
  1527. if len(pre_indices)>0 and len(post_indices)>0:
  1528. synapses_Renshaw_to_MN.connect(i=pre_indices, j=post_indices)
  1529. synapses_Renshaw_to_MN.w = weights_to_assign
  1530. synapses_Renshaw_to_MN.Renshaw_to_MN_IPSP_integral = common_brian2_namespace['Renshaw_to_MN_IPSP_integral']
  1531. else:
  1532. synapses_Renshaw_to_MN.active = False
  1533. return motoneurons, renshaw_cells, synapses_MN_to_Renshaw, synapses_Renshaw_to_MN
  1534. # # # GET SPIKE TRAINS
  1535. def get_spike_trains(spike_monitor_MN, spike_monitor_RC,
  1536. spike_transmission_delay, total_nb_motoneurons, total_nb_renshaws,
  1537. edges_ignore_duration, duration_with_ignored_window, motoneurons_soma_diameters,
  1538. generate_figure=False, savepath=None):
  1539. """
  1540. Pull out and (optionally) plot the post‐delay spike trains from your monitors.
  1541. Parameters
  1542. ----------
  1543. spike_monitor_MN : Brian2 SpikeMonitor recording the motoneurons
  1544. spike_monitor_RC : Brian2 SpikeMonitor recording the Renshaw cells
  1545. spike_transmission_delay : array_like, length total_nb_motoneurons
  1546. delay (in seconds) to add to each MN spike train before returning.
  1547. total_nb_motoneurons : int
  1548. How many MNs in total you expect (so that even silent ones get an empty list).
  1549. generate_figure : bool
  1550. savepath : str or None
  1551. """
  1552. # 1) Extract raw trains from the monitors
  1553. raw_MN = spike_monitor_MN.spike_trains() # { neuron_index: array_of_times_in_s, ... }
  1554. raw_RC = spike_monitor_RC.spike_trains()
  1555. # helper to strip units only if needed
  1556. def to_seconds(qt):
  1557. # If it's a Brian Quantity, divide by second
  1558. if isinstance(qt, Quantity):
  1559. return np.asarray(qt/second, dtype=float)
  1560. # Otherwise assume it's already a float array in seconds
  1561. return np.asarray(qt, dtype=float)
  1562. # build MN list
  1563. spike_trains_MN = []
  1564. for m in range(total_nb_motoneurons):
  1565. qt = raw_MN.get(m, np.array([],float))
  1566. times_s = to_seconds(qt)
  1567. # add your float delays (in seconds)
  1568. times_s = times_s + float(spike_transmission_delay[m])
  1569. spike_trains_MN.append(times_s)
  1570. # build RC list
  1571. # max_r = max(raw_RC.keys())+1 if raw_RC else 0 # Get nb of Renshaw cells directly from the monitor
  1572. spike_trains_RC = []
  1573. for r in range(total_nb_renshaws):
  1574. qt = raw_RC.get(r, np.array([],float))
  1575. times_s = to_seconds(qt)
  1576. spike_trains_RC.append(times_s)
  1577. # (Optional) Plot them
  1578. if generate_figure:
  1579. fig, (ax1, ax2) = plt.subplots(2,1, figsize=(20,10), sharex=True)
  1580. t_min, t_max = 0.0, duration_with_ignored_window
  1581. # shade + dashed lines for ignored windows
  1582. for ax in (ax1, ax2):
  1583. ax.axvspan(t_min, edges_ignore_duration, color='grey', alpha=0.3)
  1584. ax.axvspan(t_max - edges_ignore_duration, t_max, color='grey', alpha=0.3)
  1585. ax.axvline(edges_ignore_duration, color='black', linestyle='--', linewidth=1)
  1586. ax.axvline(t_max - edges_ignore_duration, color='black', linestyle='--', linewidth=1)
  1587. # prepare colormap for MN sizes
  1588. cmap = plt.get_cmap('viridis')
  1589. norm = mpl.colors.Normalize(
  1590. vmin=motoneurons_soma_diameters.min(),
  1591. vmax=motoneurons_soma_diameters.max()
  1592. )
  1593. sm = mpl.cm.ScalarMappable(norm=norm, cmap=cmap)
  1594. sm.set_array([]) # for the colorbar
  1595. # panel 1: MN rasters, colored by soma diameter
  1596. ax1.set_title("Motoneuron Spike Trains")
  1597. for mni, times in enumerate(spike_trains_MN):
  1598. c = cmap(norm(motoneurons_soma_diameters[mni]))
  1599. ax1.eventplot(times, lineoffsets=mni, colors=[c], alpha=0.6)
  1600. ax1.set_ylabel("MN index")
  1601. cbar = fig.colorbar(sm, ax=ax1, pad=0.02)
  1602. cbar.set_label("MN soma diameter (µm)")
  1603. # panel 2: RC rasters (unchanged)
  1604. ax2.set_title("Renshaw Cell Spike Trains")
  1605. for rci, times in enumerate(spike_trains_RC):
  1606. ax2.eventplot(times, lineoffsets=rci, colors='purple', alpha=0.2)
  1607. ax2.set_xlabel("Time (s)")
  1608. ax2.set_ylabel("RC index")
  1609. plt.tight_layout()
  1610. if savepath is not None:
  1611. fn = os.path.join(savepath, "Spike_Trains.png")
  1612. plt.savefig(fn)
  1613. plt.show()
  1614. return spike_trains_MN, spike_trains_RC
  1615. # # # SAVE OUTPUT AS HDF5 FILE WITH H5PY
  1616. def save_output_hdf5(directory_name,
  1617. params, # SimulationParameters dataclass
  1618. motoneurons_and_pools_idx, # dict of dicts, see generate_motor_neurons() (saving {pool_list_by_MN, idx_of_MN_by_pool})
  1619. motoneurons_properties, # dict of np.ndarrays
  1620. connectivity_matrix_MN_to_RC,
  1621. connectivity_matrix_RC_to_MN,
  1622. connectivity_matrix_MN_to_MN,
  1623. spike_trains_MN, # list of 1D np.ndarrays
  1624. spike_trains_RC, # list of 1D np.ndarrays
  1625. common_input_MN, # dict of 1D np.arrays (one per pool)
  1626. common_input_power_total, common_input_power_spectrum, # floats
  1627. independent_input_power_total, independent_input_power_spectrum): # dict with keys "frequencies" and "power"
  1628. output_file = os.path.join(directory_name, "simulation_output.h5")
  1629. with h5py.File(output_file, "w") as f:
  1630. # 1) simulation parameters
  1631. sim_grp = f.create_group("simulation_parameters")
  1632. param_dict = asdict(params)
  1633. param_dict.pop('RC_pair_indices', None) # Remove this value from the parameter list (it's a tuple and not readily writable)
  1634. for k, v in param_dict.items():
  1635. # numpy arrays → datasets
  1636. if isinstance(v, np.ndarray):
  1637. sim_grp.create_dataset(k, data=v)
  1638. # Brian2 quantities → two entries: numeric + unit
  1639. elif isinstance(v, Quantity):
  1640. sim_grp.create_dataset(k + "_value", data=v.magnitude)
  1641. sim_grp.attrs[k + "_unit"] = str(v.units)
  1642. # “simple” scalars OK as attrs
  1643. elif isinstance(v, (int, float, bool, str)):
  1644. sim_grp.attrs[k] = v
  1645. # anything else (lists, tuples, dicts, etc.) → JSON‐dumped string
  1646. else:
  1647. sim_grp.attrs[k] = json.dumps(v)
  1648. # 2) motoneuron and pool indices
  1649. mn_pool_grp = f.create_group("motoneurons_and_pools_indices")
  1650. # (a) pool_list_by_MN -- a string array
  1651. pool_list = motoneurons_and_pools_idx["pool_list_by_MN"]
  1652. # make a numpy array of dtype "variable‐length UTF‐8 string"
  1653. str_dt = string_dtype(encoding="utf-8")
  1654. mn_pool_grp.create_dataset(
  1655. "pool_list_by_MN",
  1656. data=np.array(pool_list, dtype=str_dt),
  1657. dtype=str_dt)
  1658. # (b) idx_of_MN_by_pool -- subgroup of integer arrays
  1659. by_pool_grp = mn_pool_grp.create_group("idx_of_MN_by_pool")
  1660. for poolname, idx_array in motoneurons_and_pools_idx["idx_of_MN_by_pool"].items():
  1661. # poolname is something like "pool_0", "pool_1", ...
  1662. by_pool_grp.create_dataset(poolname, data=idx_array.astype(int))
  1663. # # Example to read back:
  1664. # with h5py.File(..., "r") as f:
  1665. # grp = f["motoneurons_and_pools_indices"]
  1666. # pool_list = grp["pool_list_by_MN"][()] # array of bytes→ decode to str if you like
  1667. # by_pool = grp["idx_of_MN_by_pool"]
  1668. # idx0 = by_pool["pool_0"][()] # MN indices in pool_0
  1669. # 3) motoneuron properties
  1670. mn_grp = f.create_group("motoneurons_properties")
  1671. for k, arr in motoneurons_properties.items():
  1672. mn_grp.create_dataset(k, data=arr)
  1673. # 2) connectivity
  1674. conn = f.create_group("connectivity")
  1675. conn.create_dataset("MN_to_RC", data=connectivity_matrix_MN_to_RC)
  1676. conn.create_dataset("RC_to_MN", data=connectivity_matrix_RC_to_MN)
  1677. conn.create_dataset("MN_to_MN", data=connectivity_matrix_MN_to_MN)
  1678. # 5) spike trains
  1679. spikes = f.create_group("spike_trains")
  1680. mn_spikes = spikes.create_group("MN")
  1681. for i, tr in enumerate(spike_trains_MN):
  1682. mn_spikes.create_dataset(f"MN_{i}", data=tr)
  1683. rc_spikes = spikes.create_group("RC")
  1684. for i, tr in enumerate(spike_trains_RC):
  1685. rc_spikes.create_dataset(f"RC_{i}", data=tr)
  1686. # 6) Synaptic input
  1687. input_grp = f.create_group("input")
  1688. # Common input
  1689. common_input_list = []
  1690. for pooli in sorted(common_input_MN):
  1691. common_input_list.append(common_input_MN[pooli]) # now this is a 2D numeric array: (n_pools, n_timepoints)
  1692. common_input_array = np.stack(common_input_list, axis=0)
  1693. input_grp.create_dataset("common_input", data=common_input_array)
  1694. # Power spectrums
  1695. # # Total
  1696. input_grp.attrs["total_power_common_input"] = common_input_power_total[0] # only from first pool
  1697. input_grp.attrs["total_power_independent_input"] = independent_input_power_total
  1698. # # Power spectrum (per frequency)
  1699. for input_type_i in ["common_input", "independent_input"]:
  1700. if input_type_i == "common_input":
  1701. power_spectrum_to_use = common_input_power_spectrum
  1702. input_grp.create_dataset("frequencies", data=power_spectrum_to_use["frequencies"])
  1703. else: # if input_type_i == "independent_input"
  1704. power_spectrum_to_use = independent_input_power_spectrum
  1705. input_grp.create_dataset(f"power_spectrum_{input_type_i}", data=power_spectrum_to_use["power"]) # only the power spectrum of the first common input is saved (for the independent input, it is the average over all MNs)
  1706. return output_file
  1707. ######################################
  1708. ### RUN SIMULATION
  1709. ######################################
  1710. def run_simulation(params=None):
  1711. """
  1712. Runs one simulation. `params` may be:
  1713. • None → use every default
  1714. • a SimulationParameters object → used directly
  1715. """
  1716. # Create new folder and get simulation index number
  1717. if params.make_unique_output_folder:
  1718. directory_name, sim_index = make_unique_output_dir(parent_folder=params.output_folder_name)
  1719. else:
  1720. directory_name = params.output_folder_name
  1721. sim_index = os.path.basename(os.path.normpath(directory_name))
  1722. # Initialize
  1723. _ensure_logging() # ensure logging is configured for _this_ process
  1724. logger = logging.getLogger(__name__)
  1725. logger.info(f"Initializing simulation {sim_index}...")
  1726. _set_filter_params(params)
  1727. start_scope() # Re-initialize Brian
  1728. start_time = time.time()
  1729. # # # SET PARAMETERS AND CREATE OUTPUT FOLDER
  1730. # Get parameters
  1731. if params is None:
  1732. params = SimulationParameters()
  1733. elif isinstance(params, SimulationParameters):
  1734. params = params
  1735. else:
  1736. raise ValueError("run_simulation() expects a SimulationParameters object, or None")
  1737. np.random.seed(params.random_seed)
  1738. # write JSON with all parameters
  1739. param_dict = asdict(params)
  1740. param_dict.pop('RC_pair_indices', None) # Remove this value from the parameter list (it's a tuple and not readily writable in a json file)
  1741. with open(f"{directory_name}/sim_parameters.json","w") as fp:
  1742. json.dump(param_dict, fp, indent=2, default=str)
  1743. # # # SET EQUATIONS
  1744. (MN_equations, RC_equations) = set_brian2_equations()
  1745. # # # GENERATE MOTOR NEURONS
  1746. (motoneuron_soma_diameters, motoneuron_normalized_soma_diameters,
  1747. pool_list_by_MN, idx_of_MN_by_pool) = generate_motor_neurons(full_pool_sim=params.full_pool_sim,
  1748. min_soma_diameter=params.min_soma_diameter, max_soma_diameter=params.max_soma_diameter,
  1749. nb_pools=params.nb_pools, nb_motoneurons_per_pool=params.nb_motoneurons_per_pool, total_nb_motoneurons=params.total_nb_motoneurons,
  1750. size_distribution_exponent=params.size_distribution_exponent,
  1751. mean_soma_diameter=params.mean_soma_diameter, sd_prct_soma_diameter=params.sd_prct_soma_diameter,
  1752. generate_figure=params.output_plots, savepath=directory_name)
  1753. # # # GENERATE MOTOR NEURONS ELECTROPHYSIOLOGICAL PROPERTIES
  1754. (motoneurons_resistance, motoneurons_input_weight,
  1755. motoneurons_capacitance, motoneurons_membrane_conductance,
  1756. motoneurons_membrane_time_constant,
  1757. motoneurons_AHP_conductance_decay_time_constant, motoneurons_refractory_periods,
  1758. motoneurons_rheobases, motoneurons_spike_transmission_delays,
  1759. motoneurons_properties_dict) = generate_motor_neuron_electrophysiological_properties(
  1760. total_nb_motoneurons=params.total_nb_motoneurons, motoneuron_soma_diameters=motoneuron_soma_diameters,
  1761. resistance_constant=params.resistance_constant, resistance_exponent=params.resistance_exponent,
  1762. capacitance_constant=params.capacitance_constant, capacitance_exponent=params.capacitance_exponent,
  1763. AHP_duration_constant=params.AHP_duration_constant, AHP_duration_exponent=params.AHP_duration_exponent,
  1764. rheobase_constant=params.rheobase_constant, rheobase_exponent=params.rheobase_exponent, rheobase_scaling=params.rheobase_scaling,
  1765. refractory_period_absolute=params.refractory_period_absolute,
  1766. axonal_conduction_velocity_constant=params.axonal_conduction_velocity_constant, axonal_conduction_velocity_exponent=params.axonal_conduction_velocity_exponent,
  1767. generate_figure=params.output_plots, savepath=directory_name)
  1768. # # # GENERATE AND DISTRIBUTE COMMON INPUT
  1769. (MN_excit_input, corr_mat, inputs_to_mn_weight_matrix, total_power, power_within_band_of_interest,
  1770. power_per_frequency_band) = generate_and_distribute_common_inputs(
  1771. fsamp=params.fsamp, nb_of_common_inputs=params.nb_of_common_inputs,
  1772. nb_pools=params.nb_pools, nb_motoneurons_per_pool=params.nb_motoneurons_per_pool,
  1773. duration_with_ignored_window=params.duration_with_ignored_window, edges_ignore_duration=params.edges_ignore_duration,
  1774. frequency_range_of_common_input=params.frequency_range_of_common_input,
  1775. common_input_std=params.common_input_std,
  1776. frequency_range_to_set_input_power=params.frequency_range_to_set_input_power,
  1777. set_arbitrary_correlation_between_excitatory_inputs=params.set_arbitrary_correlation_between_excitatory_inputs,
  1778. between_pool_excitatory_input_correlation=params.between_pool_excitatory_input_correlation,
  1779. set_same_excitatory_input_for_all_pools=params.set_same_excitatory_input_for_all_pools,
  1780. generate_figure=params.output_plots, savepath=directory_name)
  1781. # # # ADD INDEPENDENT INPUTS AND CREATE Brian2 TIMED ARRAY OBJECTS
  1782. (input_MN_timedarray_amp, input_RC_timedarray_volt,
  1783. total_power_independent, power_per_frequency_band_independent) = generate_independent_inputs(
  1784. fsamp=params.fsamp,
  1785. nb_pools=params.nb_pools, total_nb_motoneurons=params.total_nb_motoneurons, total_nb_renshaw_cells=params.total_nb_renshaw_cells,
  1786. low_pass_filter_of_MN_independent_input=params.low_pass_filter_of_MN_independent_input,
  1787. low_pass_filter_of_RC_independent_input=params.low_pass_filter_of_RC_independent_input,
  1788. ref_common_input_power = power_within_band_of_interest,
  1789. MN_independent_input_absolute_or_ratio=params.independent_input_absolute_or_ratio,
  1790. MN_independent_input_power=params.independent_input_power,
  1791. RC_independent_input_std=params.RC_independent_input_std,
  1792. MN_excit_input=MN_excit_input,
  1793. inputs_to_mn_weight_matrix=inputs_to_mn_weight_matrix,
  1794. excitatory_input_baseline=params.excitatory_input_baseline,
  1795. duration_with_ignored_window=params.duration_with_ignored_window,
  1796. edges_ignore_duration=params.edges_ignore_duration,
  1797. motoneurons_rheobases=motoneurons_rheobases,
  1798. generate_figure=params.output_plots, savepath=directory_name)
  1799. # # # IF DESIRED, DISPLAY POWER OF COMMON AND INDEPENDENT INPUTS
  1800. if params.output_plots:
  1801. plt.figure(figsize=(10,6))
  1802. scaling_power = 1/1e6
  1803. plt.fill_between(
  1804. power_per_frequency_band_independent['frequencies'],
  1805. power_per_frequency_band_independent['power']*scaling_power,
  1806. y2=0, color='blue', alpha=0.3)
  1807. plt.fill_between(
  1808. power_per_frequency_band['frequencies'],
  1809. power_per_frequency_band['power']*scaling_power,
  1810. y2=0, color='red', alpha=0.3)
  1811. plt.plot(power_per_frequency_band['frequencies'],
  1812. power_per_frequency_band['power']*scaling_power,
  1813. color='red', label=f'Common input power\ntotal={total_power[0]*scaling_power:.2f} a.u.')
  1814. plt.plot(power_per_frequency_band_independent['frequencies'],
  1815. power_per_frequency_band_independent['power']*scaling_power,
  1816. color='blue', label=f'Independent input power\ntotal={total_power_independent*scaling_power:.2f} a.u.')
  1817. plt.xlabel("Frequency (Hz)")
  1818. plt.ylabel("Power (a.u.)")
  1819. plt.legend(loc="upper right")
  1820. plt.savefig(f"{directory_name}/Power_of_common_and_independent_inputs.png")
  1821. # # # CREATE CONNECTIVITY BETWEEN MOTOR NEURONS AND RENSHAW CELLS
  1822. (MN_to_MN_connectivity_matrix,
  1823. MN_to_Renshaw_connectivity_matrix,
  1824. Renshaw_to_MNs_connectivity_matrix) = create_connectivity(
  1825. total_nb_motoneurons = params.total_nb_motoneurons, total_nb_renshaw_cells = params.total_nb_renshaw_cells,
  1826. nb_pools = params.nb_pools, nb_motoneurons_per_pool = params.nb_motoneurons_per_pool, RC_pair_indices = params.RC_pair_indices, nb_RCs_per_pool_pair = params.nb_RCs_per_pool_pair,
  1827. disynpatic_inhib_connections_desired_MN_MN = params.disynpatic_inhib_connections_desired_MN_MN,
  1828. split_MN_RC_ratio = params.split_MN_RC_ratio, motoneuron_normalized_soma_diameters = motoneuron_normalized_soma_diameters,
  1829. distribution_type = params.distribution_type, distribution_params = params.distribution_params,
  1830. disynaptic_inhib_received_arbitrary_adjustment = params.disynaptic_inhib_received_arbitrary_adjustment,
  1831. distribution_binary_weights = params.binary_connectivity,
  1832. generate_figure=params.output_plots, savepath=directory_name)
  1833. # logger = logging.getLogger(__name__)
  1834. # # # CREATE Brian2 NEURONGROUPS AND SYNAPSES OBJECTS
  1835. # set synaptic_IPSP_decay_time_constant_per_MN
  1836. if params.synaptic_IPSP_membrane_or_user_defined_time_constant == "user_defined":
  1837. synaptic_IPSP_decay_time_constant_per_MN = []
  1838. for mni in range(params.total_nb_motoneurons):
  1839. synaptic_IPSP_decay_time_constant_per_MN.append(params.synaptic_IPSP_decay_time_constant) # params.synaptic_IPSP_decay_time_constant is already in milliseconds
  1840. elif params.synaptic_IPSP_membrane_or_user_defined_time_constant == "membrane":
  1841. synaptic_IPSP_decay_time_constant_per_MN = motoneurons_membrane_time_constant * ms # motoneurons_membrane_time_constant is a list of unitless values
  1842. else:
  1843. raise ValueError("Please enter a valid value for synaptic_IPSP_decay_time_constant_per_MN: 'user_defined' or 'membrane'")
  1844. # logger.info(f"Synaptic tau: {synaptic_IPSP_decay_time_constant_per_MN}")
  1845. # actually create the objects
  1846. (motoneurons, renshaw_cells, synapses_MN_to_Renshaw, synapses_Renshaw_to_MN) = create_neurongroups_and_synapses_objects(
  1847. total_nb_motoneurons=params.total_nb_motoneurons, total_nb_renshaw_cells=params.total_nb_renshaw_cells,
  1848. MN_equations=MN_equations, RC_equations=RC_equations, voltage_rest=params.voltage_rest, voltage_thresh=params.voltage_thresh,
  1849. motoneurons_membrane_conductance=motoneurons_membrane_conductance, motoneurons_capacitance=motoneurons_capacitance,
  1850. motoneurons_input_weight=motoneurons_input_weight, synaptic_IPSP_decay_time_constant_per_MN=synaptic_IPSP_decay_time_constant_per_MN,
  1851. motoneurons_AHP_conductance_decay_time_constant=motoneurons_AHP_conductance_decay_time_constant, motoneurons_refractory_periods=motoneurons_refractory_periods,
  1852. AHP_conductance_delta_after_spiking=params.AHP_conductance_delta_after_spiking, MN_to_Renshaw_excit=params.MN_to_Renshaw_EPSP, Renshaw_to_MN_inhib=params.Renshaw_to_MN_IPSP,
  1853. scale_initial_IPSP_to_be_same_integral_regardless_of_synaptic_tau = params.scale_initial_IPSP_to_be_same_integral_regardless_of_synaptic_tau,
  1854. tau_Renshaw=params.tau_Renshaw, MN_RC_synpatic_delay=params.MN_RC_synpatic_delay, refractory_period_RC=params.refractory_period_RC,
  1855. MN_to_Renshaw_connectivity_matrix=MN_to_Renshaw_connectivity_matrix, Renshaw_to_MNs_connectivity_matrix=Renshaw_to_MNs_connectivity_matrix)
  1856. # # # RUN SIMULATION
  1857. # Initialize values
  1858. motoneurons.v = params.voltage_rest # in mV
  1859. renshaw_cells.v = params.voltage_rest # Initialize membrane potential
  1860. # Set monitors
  1861. monitor_spikes_motoneurons = SpikeMonitor(motoneurons, record=True)
  1862. # monitor_voltage_motoneurons = StateMonitor(motoneurons, 'v', record=True) # get the voltage trace
  1863. monitor_spikes_renshaw_cells = SpikeMonitor(renshaw_cells, record=True)
  1864. end_time = time.time()
  1865. elapsed_time = end_time - start_time
  1866. # logger.info(f"...Initialization of sim {sim_index} finished ({elapsed_time:.2f} seconds)")
  1867. logger.info(f"Starting simulation run for sim {sim_index} (total simulated time = {params.duration_with_ignored_window:.1f} seconds)...")
  1868. start_time = time.time()
  1869. # RUN THE SIMULATION
  1870. try:
  1871. run(params.duration_with_ignored_window * second)
  1872. except Exception as e:
  1873. # log the full Python traceback to your simulations_progress_log.log
  1874. logger.error("Exception during Brian2 run():\n" + traceback.format_exc())
  1875. # re-raise so Joblib will propagate (or at least you’ll see something)
  1876. raise
  1877. end_time = time.time()
  1878. elapsed_time = end_time - start_time
  1879. logger.info(f"...simulation {sim_index} finished! ({elapsed_time:.2f} seconds)")
  1880. spike_trains_MN, spike_trains_RC = get_spike_trains(
  1881. spike_monitor_MN=monitor_spikes_motoneurons, spike_monitor_RC=monitor_spikes_renshaw_cells,
  1882. spike_transmission_delay=motoneurons_spike_transmission_delays,
  1883. total_nb_motoneurons=params.total_nb_motoneurons, total_nb_renshaws=params.total_nb_renshaw_cells,
  1884. edges_ignore_duration=params.edges_ignore_duration, duration_with_ignored_window=params.duration_with_ignored_window,
  1885. motoneurons_soma_diameters=motoneuron_soma_diameters,
  1886. generate_figure=params.output_plots, savepath=directory_name
  1887. )
  1888. # Saving output
  1889. output_savefile = save_output_hdf5(directory_name=directory_name,
  1890. params=params,
  1891. motoneurons_and_pools_idx={"pool_list_by_MN": pool_list_by_MN,
  1892. "idx_of_MN_by_pool": idx_of_MN_by_pool},
  1893. motoneurons_properties=motoneurons_properties_dict,
  1894. connectivity_matrix_MN_to_RC=MN_to_Renshaw_connectivity_matrix, connectivity_matrix_RC_to_MN=Renshaw_to_MNs_connectivity_matrix,
  1895. connectivity_matrix_MN_to_MN=MN_to_MN_connectivity_matrix,
  1896. spike_trains_MN=spike_trains_MN, spike_trains_RC=spike_trains_RC,
  1897. common_input_MN=MN_excit_input,
  1898. common_input_power_total=total_power, common_input_power_spectrum=power_per_frequency_band,
  1899. independent_input_power_total=total_power_independent, independent_input_power_spectrum=power_per_frequency_band_independent)
  1900. logger.info(f"Data of simulation {sim_index} saved successfully to '{output_savefile}'.")
  1901. return output_savefile

simulator.py at commit 31e52ca, under MIT · at the source

Overview

  1. Université Côte d'Azur, LAMHESS, Nice, France
  2. Department of Bioengineering, Faculty of Engineering, Imperial College London, London, UK
  3. Nantes Université, Movement-Interactions-Performance, MIP, UR 4334, F-44000 Nantes, France
  4. The University of Queensland, School of Biomedical Sciences, Brisbane, QLD, Australia
  5. Institut Universitaire de France (IUF), Paris, France
Journal: Science advances, volume 12, issue 37, article eaee9425
Dates: received 21 December 2025; accepted 30 July 2026; published online 9 September 2026; in print September 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1126/sciadv.aee9425 · PMID 42715300 · PMCID PMC13557072 · OpenAlex W7212003845
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: human (organism)
Methods: Spectral & time-frequency, Smoothing, state filtering, decompositions, Connectivity, Preprocessing, Single-unit activity, calcium imaging, Physiology & signal measures, Machine learning
MeSH: Motor Neurons*, Muscle, Skeletal*, Action Potentials, Humans, Muscle Contraction (* major topic)
Topic: Muscle activation and electromyography studies (Biomedical Engineering, Engineering), according to OpenAlex
Funding: European Research Council (810346)
Citations: not cited yet (Europe PMC); 74 references in the paper

Abstract

Understanding how spinal circuits shape motor neuron behavior during muscle contractions remains a major challenge. Here, we combined large-scale motor unit recordings with simulation-based inference to generate probabilistic estimates of homonymous and heteronymous recurrent inhibition, a key spinal circuit that has remained largely inaccessible during natural voluntary contractions. We constructed synchronization cross-histograms from motor neuron spike trains and extracted features representative of recurrent inhibition. Because these features are also influenced by higher-frequency components of common synaptic input, we developed a simulation-based inference framework to disentangle these effects. Following validation, we applied this framework to experimental data from six muscles at two contraction intensities, revealing previously uncharacterized muscle- and intensity-dependent patterns: Recurrent inhibition decreased with contraction intensity in most muscles but increased in the vastus lateralis and medialis. The pipeline is openly available and designed for reuse on comparable datasets and for adaptation to diverse experimental contexts, including other spinal circuits.

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

Repository

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

FrancoisDernoncourt/Mapping_Recurrent_Inhibition

License: MIT
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: 31e52ca73d83a61dd2b5e1604fc26cf42aa6a589, 7 September 2026
Languages: Jupyter (9), R (2), Python (2)
Size: 2,068 files, 13 scripts
Software Heritage: not archived
Found in: “Data, code, and materials availability:”
Holds: README, license file, 11 notebooks
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: pandas (8 files), NumPy (7 files), Matplotlib (6 files), Brian 2 (5 files), h5py (4 files), SciPy (4 files), seaborn (4 files), PyTorch (3 files), ggplot2 (2 files), scikit-learn (2 files), tidyverse (2 files), broom (1 file), NetworkX (1 file), PyWavelets (1 file)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
15 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:

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

All data and code needed to evaluate and reproduce the results in the paper are present in the paper and/or the Supplementary Materials. This study did not generate new materials. The full dataset (raw electromyography data and edited motor unit spike trains), together with the analysis code, is available at https://doi.org/10.57745/4CKBPO. Code is also available as an additional resource at https://github.com/FrancoisDernoncourt/Mapping_Recurrent_Inhibition.

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, 5 authors, 5 MeSH terms, 1 funder, 69 references.

Cite

This paper

Dernoncourt, F., Avrillon, S., Cattagni, T., Farina, D., & Hug, F. (2026). Probabilistic inference of homonymous and heteronymous recurrent inhibition in human muscles from large-scale motor neuron recordings. Science advances, 12(37), eaee9425. https://doi.org/10.1126/sciadv.aee9425

BibTeX

@article{dernoncourt2026probabilistic,
author = {Dernoncourt, François and Avrillon, Simon and Cattagni, Thomas and Farina, Dario and Hug, François},
title = {{Probabilistic inference of homonymous and heteronymous recurrent inhibition in human muscles from large-scale motor neuron recordings}},
journal = {Science advances},
year = {2026},
month = sep,
volume = {12},
number = {37},
pages = {eaee9425},
publisher = {American Association for the Advancement of Science},
issn = {2375-2548},
doi = {10.1126/sciadv.aee9425},
url = {https://doi.org/10.1126/sciadv.aee9425},
pmid = {42715300},
pmcid = {PMC13557072}
}

RIS

TY - JOUR
AU - Dernoncourt, François
AU - Avrillon, Simon
AU - Cattagni, Thomas
AU - Farina, Dario
AU - Hug, François
TI - Probabilistic inference of homonymous and heteronymous recurrent inhibition in human muscles from large-scale motor neuron recordings
T2 - Science advances
J2 - Sci Adv
PY - 2026
DA - 2026/09/09
VL - 12
IS - 37
SP - eaee9425
SN - 2375-2548
PB - American Association for the Advancement of Science
DO - 10.1126/sciadv.aee9425
UR - https://doi.org/10.1126/sciadv.aee9425
LA - en
ER -

CSL-JSON

{
"id": "10.1126/sciadv.aee9425",
"type": "article-journal",
"title": "Probabilistic inference of homonymous and heteronymous recurrent inhibition in human muscles from large-scale motor neuron recordings",
"container-title": "Science advances",
"author": [
{
"family": "Dernoncourt",
"given": "François"
},
{
"family": "Avrillon",
"given": "Simon"
},
{
"family": "Cattagni",
"given": "Thomas"
},
{
"family": "Farina",
"given": "Dario"
},
{
"family": "Hug",
"given": "François"
}
],
"container-title-short": "Sci Adv",
"volume": "12",
"issue": "37",
"page": "eaee9425",
"DOI": "10.1126/sciadv.aee9425",
"PMID": "42715300",
"PMCID": "PMC13557072",
"ISSN": "2375-2548",
"publisher": "American Association for the Advancement of Science",
"URL": "https://doi.org/10.1126/sciadv.aee9425",
"language": "en",
"issued": {
"date-parts": [
[
2026,
9,
9
]
]
}
}

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.1186/s12984-026-02041-3 [code]
Mental tasks induce common modulations of oscillations in cortex and spinal cord.
Journal: Journal of neuroengineering and rehabilitation
In common: Matplotlib, 7 references, author Dario Farina
[2] doi:10.1038/s41467-026-74243-1 [code]
Spiking neural network decoders of finger forces from high-density intramuscular microelectrode arrays.
Journal: Nature communications
In common: PyTorch, seaborn, scikit-learn, 4 other tools, 2 references, author Dario Farina
[3] doi:10.1016/j.isci.2026.115488 [code]
An integrated &lt;i&gt;i&lt;/i&gt; &lt;i&gt;n vitro&lt;/i&gt; platform and biophysical modeling approach for studying synaptic transmission in isolated neuronal pairs.
Journal: iScience
In common: NetworkX, h5py, PyTorch, 6 other tools, 2 references
[4] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: broom, NetworkX, PyTorch, 8 other tools
[5] doi:10.1016/j.isci.2026.116055 [code]
Mapping the transcriptional diversity of calcium signaling in the mouse and human brain.
Journal: iScience
In common: NetworkX, h5py, PyTorch, 8 other tools
[6] doi:10.1038/s41467-026-75705-2 [code]
Redundant prefrontal hemispheres adapt storage strategy to working memory demands.
Journal: Nature communications
In common: Brian 2, h5py, seaborn, 5 other tools, 1 reference
[7] doi:10.1002/hipo.70089 [code]
The Role of Plasticity in Replay: Stability Through Anti-Hebbian Rules.
Journal: Hippocampus
In common: Brian 2, PyWavelets, seaborn, 4 other tools, 1 reference
[8] doi:10.1371/journal.pcbi.1014337 [code]
Fast reconstruction of degenerate populations of conductance-based neuron models from spike times.
Journal: PLoS computational biology
In common: PyTorch, seaborn, scikit-learn, 4 other tools, 3 references
[9] doi:10.1016/j.nicl.2026.104052 [code]
Beyond one-to-one mappings: Modelling distributed lesion-symptom relationships with multilayer networks.
Journal: NeuroImage. Clinical
In common: broom, NetworkX, ggplot2, 7 other tools
[10] doi:10.1093/bioinformatics/btag592 [code]
Network-based stratification of allele-specific expression reveals patient subgroups in Huntington's disease.
Journal: Bioinformatics (Oxford, England)
In common: broom, NetworkX, ggplot2, 7 other tools

Contribute

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

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

Request its removal

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

Discussion, reproductions, activity

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

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

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