#!/usr/bin/env python3 """ sim_worker.py — Standalone circuit simulator for subprocess-based parallelism. Called by batch_generate_v2.py via subprocess.Popen: python3 sim_worker.py Reads circuit config from JSON, builds HH network in Jaxley, sweeps 11 ACh levels, computes 11 summary statistics, writes output JSON. ACh models: v14a: Synaptic suppression only (same as v13 production data) v14b: Mild depolarization only (0.0002 nA max, sigmoid onset at ach=0.3) v14c: Both mechanisms combined Output: /circuit_XXXXX.json (list of 11 dicts, one per ACh level) On error: /circuit_XXXXX.error (error message) Exit codes: 0 = success, 1 = failure """ import sys import json import math import os import time import traceback # Suppress Jaxley's verbose per-synapse/per-recording print() noise # (Jaxley uses print(), not logging, so we redirect stdout temporarily) import logging logging.getLogger("jaxley").setLevel(logging.ERROR) logging.getLogger("jax").setLevel(logging.WARNING) # --------------------------------------------------------------------------- # Constants (must match batch_generate.py / training pipeline) # --------------------------------------------------------------------------- ACH_LEVELS = [0.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0] # OU noise parameters (Destexhe 2001 background synaptic bombardment) OU_MU = -0.001 # nA — mean current (below rheobase ~0.003) OU_SIGMA = 0.001 # nA — noise amplitude OU_TAU = 5.0 # ms — correlation time # Synaptic conductances — very weak coupling (AI regime) EXC_GS = 0.000005 # μS (0.05× default) INH_GS = 0.00003 # μS (inhibition-dominant) E_SYN_EXC = 0.0 # mV (excitatory reversal) E_SYN_INH = -80.0 # mV (inhibitory reversal — GABA-A like) # --------------------------------------------------------------------------- # ACh Model Functions # --------------------------------------------------------------------------- def ach_to_mu_shift_v14a(ach: float) -> float: """v14a: No depolarization (synaptic suppression only).""" return 0.0 def ach_to_syn_scale_v14a(ach: float) -> float: """v14a: E2 exponential synaptic suppression (Ramaswamy 2018 Table 2).""" return max(0.05, math.exp(-2.3 * ach)) def ach_to_mu_shift_v14b(ach: float) -> float: """v14b: Mild depolarization only. +0.0002 nA max (6× smaller than v10's 0.0012 which caused synchrony). Sigmoid onset at ach=0.3 to avoid low-ACh correlation spikes. """ return 0.0002 / (1.0 + math.exp(-10.0 * (ach - 0.3))) def ach_to_syn_scale_v14b(ach: float) -> float: """v14b: No synaptic suppression (depolarization only).""" return 1.0 def ach_to_mu_shift_v14c(ach: float) -> float: """v14c: Same depolarization as v14b.""" return 0.0002 / (1.0 + math.exp(-10.0 * (ach - 0.3))) def ach_to_syn_scale_v14c(ach: float) -> float: """v14c: Same synaptic suppression as v14a.""" return max(0.05, math.exp(-2.3 * ach)) ACH_MODELS = { "v14a": (ach_to_mu_shift_v14a, ach_to_syn_scale_v14a), "v14b": (ach_to_mu_shift_v14b, ach_to_syn_scale_v14b), "v14c": (ach_to_mu_shift_v14c, ach_to_syn_scale_v14c), } # --------------------------------------------------------------------------- # OU Noise Generator # --------------------------------------------------------------------------- def generate_ou_current(dt_ms, n_steps, mu, sigma, tau_ms, rng): """Ornstein-Uhlenbeck process: colored noise for background synaptic input.""" import numpy as np x = np.zeros(n_steps) x[0] = mu dt_s = dt_ms / 1000.0 tau_s = tau_ms / 1000.0 noise_coeff = sigma * math.sqrt(2.0 * dt_s / tau_s) for step in range(1, n_steps): x[step] = x[step-1] + dt_s * (mu - x[step-1]) / tau_s + noise_coeff * rng.randn() return x # --------------------------------------------------------------------------- # Statistics (character-for-character match with batch_generate.py) # --------------------------------------------------------------------------- def compute_statistics(spike_trains, n_exc, n_inh, duration_ms): """Compute population-level summary statistics from spike trains.""" import numpy as np n_total = n_exc + n_inh duration_s = duration_ms / 1000.0 # Firing Rates rates = [len(st) / duration_s for st in spike_trains] exc_rates = rates[:n_exc] inh_rates = rates[n_exc:] mean_rate = float(np.mean(rates)) if rates else 0.0 mean_exc_rate = float(np.mean(exc_rates)) if exc_rates else 0.0 mean_inh_rate = float(np.mean(inh_rates)) if inh_rates else 0.0 # CV of ISI cv_isis = [] for st in spike_trains: if len(st) > 2: isis = np.diff(st) if np.mean(isis) > 0: cv_isis.append(float(np.std(isis) / np.mean(isis))) mean_cv_isi = float(np.mean(cv_isis)) if cv_isis else 0.0 # Fano Factor bin_size_ms = 50.0 n_bins = int(duration_ms / bin_size_ms) if n_bins > 0: bin_counts = np.zeros((n_total, n_bins)) for i, st in enumerate(spike_trains): for t in st: b = min(int(t / bin_size_ms), n_bins - 1) bin_counts[i, b] += 1 fano_factors = [] for i in range(n_total): m = np.mean(bin_counts[i]) if m > 0: fano_factors.append(float(np.var(bin_counts[i]) / m)) mean_fano = float(np.mean(fano_factors)) if fano_factors else 1.0 else: mean_fano = 1.0 # Population Synchrony pop_bin_ms = 5.0 n_pop_bins = int(duration_ms / pop_bin_ms) if n_pop_bins > 0: pop_rate = np.zeros(n_pop_bins) for st in spike_trains: for t in st: b = min(int(t / pop_bin_ms), n_pop_bins - 1) pop_rate[b] += 1 pop_rate /= n_total pop_mean = np.mean(pop_rate) synchrony_index = float(np.var(pop_rate) / pop_mean) if pop_mean > 0 else 0.0 else: synchrony_index = 0.0 # Pairwise Correlations (sample) n_sample = min(100, n_exc * (n_exc - 1) // 2) if n_sample > 0 and n_pop_bins > 0: corr_bin_ms = 10.0 n_corr_bins = int(duration_ms / corr_bin_ms) binned = np.zeros((n_exc, n_corr_bins)) for i in range(n_exc): for t in spike_trains[i]: b = min(int(t / corr_bin_ms), n_corr_bins - 1) binned[i, b] += 1 pair_rng = np.random.RandomState(0) correlations = [] for _ in range(n_sample): i, j = pair_rng.choice(n_exc, 2, replace=False) x, y = binned[i], binned[j] if np.std(x) > 0 and np.std(y) > 0: r = float(np.corrcoef(x, y)[0, 1]) if not np.isnan(r): correlations.append(r) mean_pairwise_corr = float(np.mean(correlations)) if correlations else 0.0 else: mean_pairwise_corr = 0.0 # Power Spectrum if n_pop_bins > 10: from scipy import signal freqs, psd = signal.welch(pop_rate, fs=1000.0 / pop_bin_ms, nperseg=min(256, n_pop_bins)) peak_freq = float(freqs[np.argmax(psd)]) total_power = float(np.sum(psd)) else: peak_freq = 0.0 total_power = 0.0 return { "mean_firing_rate": round(mean_rate, 3), "mean_exc_rate": round(mean_exc_rate, 3), "mean_inh_rate": round(mean_inh_rate, 3), "mean_cv_isi": round(mean_cv_isi, 3), "mean_fano_factor": round(mean_fano, 3), "synchrony_index": round(synchrony_index, 4), "mean_pairwise_corr": round(mean_pairwise_corr, 4), "peak_frequency_hz": round(peak_freq, 3), "total_spectral_power": round(total_power, 6), "n_active_neurons": sum(1 for r in rates if r > 0.5), "total_spikes": sum(len(st) for st in spike_trains), } # --------------------------------------------------------------------------- # Main Simulation # --------------------------------------------------------------------------- def simulate_circuit(config: dict, output_dir: str, ach_model: str): """Build one circuit, sweep all 11 ACh levels, save results.""" import numpy as np import copy # Lazy imports — each subprocess gets its own JAX runtime import jaxley as jx from jaxley.channels import HH from jaxley.synapses import IonotropicSynapse # Get ACh functions for this model mu_shift_fn, syn_scale_fn = ACH_MODELS[ach_model] circuit_id = config["circuit_id"] n_exc = config["n_exc"] n_inh = config["n_inh"] conn_prob = config["conn_prob"] sim_duration_ms = config["sim_duration_ms"] seed = config["seed"] n_total = n_exc + n_inh rng = np.random.RandomState(seed) dt = 0.025 n_steps = int(sim_duration_ms / dt) t_start = time.time() # === Phase 1: Build network ONCE === cells = [] for i in range(n_total): comp = jx.Compartment() comp.insert(HH()) cells.append(comp) net = jx.Network(cells) # Connect with pre-computed adjacency adjacency = rng.random((n_total, n_total)) < conn_prob np.fill_diagonal(adjacency, False) syn_count = 0 for pre_idx in range(n_total): for post_idx in range(n_total): if adjacency[pre_idx, post_idx]: jx.connect(net.cell(pre_idx), net.cell(post_idx), IonotropicSynapse()) syn_count += 1 t_build = time.time() - t_start # Set per-synapse gS and e_syn based on E/I identity if syn_count > 0: gs_arr = np.empty(syn_count) esyn_arr = np.empty(syn_count) syn_idx = 0 for pre_idx in range(n_total): for post_idx in range(n_total): if adjacency[pre_idx, post_idx]: if pre_idx < n_exc: gs_arr[syn_idx] = EXC_GS esyn_arr[syn_idx] = E_SYN_EXC else: gs_arr[syn_idx] = INH_GS esyn_arr[syn_idx] = E_SYN_INH syn_idx += 1 net.edges["IonotropicSynapse_gS"] = gs_arr net.edges["IonotropicSynapse_e_syn"] = esyn_arr net.edges["IonotropicSynapse_s"] = 0.0 # CRITICAL: no phantom current mean_in_degree = max(1.0, syn_count / n_total) # Pre-generate per-neuron noise seeds (shared across ACh levels) noise_seeds = rng.randint(0, 1000000, n_total) stim_len = int(sim_duration_ms / dt) + 1 # === Phase 2: Sequential ACh sweep using deepcopy === results = [] for ach_idx, ach in enumerate(ACH_LEVELS): t0 = time.time() net_copy = copy.deepcopy(net) # ACh modulates EXCITATORY synapses (muscarinic suppression) syn_scale = syn_scale_fn(ach) if syn_count > 0: exc_mask = net_copy.edges["IonotropicSynapse_e_syn"] == E_SYN_EXC net_copy.edges.loc[exc_mask, "IonotropicSynapse_gS"] = EXC_GS * syn_scale # ACh shifts mean current toward threshold mu_shift = mu_shift_fn(ach) mu_eff = OU_MU + mu_shift sigma_eff = OU_SIGMA # ACh does NOT change noise amplitude # Stimulate each neuron with independent OU noise for i in range(n_total): neuron_rng = np.random.RandomState(noise_seeds[i]) ou_trace = generate_ou_current(dt, n_steps, mu_eff, sigma_eff, OU_TAU, neuron_rng) i_stim = np.zeros(stim_len) ou_len = min(len(ou_trace), stim_len) i_stim[:ou_len] = ou_trace[:ou_len] net_copy.cell(i).stimulate(i_stim) net_copy.cell(i).record("v") voltages = jx.integrate(net_copy, delta_t=dt) t_sim = time.time() - t0 # Extract spikes (only from stable period: after 200ms warmup) v_np = np.array(voltages) warmup_steps = int(200.0 / dt) spike_trains = [] for i in range(n_total): v_stable = v_np[i, warmup_steps:] crossings = np.where((v_stable[:-1] < 0.0) & (v_stable[1:] >= 0.0))[0] spike_times_ms = ((crossings + warmup_steps) * dt).tolist() spike_trains.append(spike_times_ms) # Compute statistics on stable period only stable_duration_ms = sim_duration_ms - 200.0 stats = compute_statistics(spike_trains, n_exc, n_inh, stable_duration_ms) results.append({ "circuit_id": circuit_id, "ach_level": ach, "ach_model": ach_model, "n_exc": n_exc, "n_inh": n_inh, "n_total": n_total, "conn_prob": conn_prob, "n_synapses": syn_count, "mean_in_degree": round(mean_in_degree, 1), "gS_exc_effective": round(EXC_GS * syn_scale, 8), "ou_mu_effective": round(mu_eff, 6), "ou_sigma_effective": round(sigma_eff, 6), "ou_tau": OU_TAU, "sim_duration_ms": sim_duration_ms, "seed": seed, "sim_time_s": round(t_sim, 2), "statistics": stats, }) total_time = time.time() - t_start # === Save results === os.makedirs(output_dir, exist_ok=True) outpath = os.path.join(output_dir, f"circuit_{circuit_id:05d}.json") with open(outpath, "w") as f: json.dump(results, f) # Print summary to stderr (stdout may be suppressed) print( f"OK circuit={circuit_id} model={ach_model} " f"build={t_build:.1f}s total={total_time:.1f}s " f"rate@0={results[0]['statistics']['mean_firing_rate']:.1f}Hz " f"corr@0={results[0]['statistics']['mean_pairwise_corr']:.4f} " f"corr@1={results[-1]['statistics']['mean_pairwise_corr']:.4f}", file=sys.stderr, ) return results # --------------------------------------------------------------------------- # Entry point # --------------------------------------------------------------------------- def main(): if len(sys.argv) != 4: print(f"Usage: python3 {sys.argv[0]} ", file=sys.stderr) print(f" ach_model: v14a | v14b | v14c", file=sys.stderr) sys.exit(1) config_path = sys.argv[1] output_dir = sys.argv[2] ach_model = sys.argv[3] if ach_model not in ACH_MODELS: print(f"ERROR: Unknown ACh model '{ach_model}'. Must be one of: {list(ACH_MODELS.keys())}", file=sys.stderr) sys.exit(1) # Read config with open(config_path, "r") as f: config = json.load(f) circuit_id = config["circuit_id"] try: simulate_circuit(config, output_dir, ach_model) except Exception as e: # Write error marker os.makedirs(output_dir, exist_ok=True) error_path = os.path.join(output_dir, f"circuit_{circuit_id:05d}.error") with open(error_path, "w") as f: f.write(f"{type(e).__name__}: {e}\n") f.write(traceback.format_exc()) print(f"FAIL circuit={circuit_id}: {e}", file=sys.stderr) sys.exit(1) if __name__ == "__main__": main()