// neuron_tripod_het.mbt — Tripod with heterogeneous per-neuron
// AdEx parameters. Mirrors AdExHet (v0.10.9) and Julia's
// `Tripod(..., NMDA=EyalNMDA, param=AdExParameter{Vector{Float32}})`.
//
// Per-neuron fields: vt/vr/el/tm/r/dt_slope/tw/a/b/c/gl. Dendrite
// geometry (d1/d2) stays shared — Julia's TripodHet uses a single
// Dendrite struct across the population. Heun predictor-corrector
// integration reads per-ne parameters inside the per-neuron loop.
//
// Bit-exact note: this matches Julia's `Tripod{Vector{Float32}}`
// behavior. The soma params vary per neuron but dendrite params +
// AdExPostSpike + Heun correction order are shared across neurons
// (same as Julia).
///|
/// TripodHet — Tripod with heterogeneous per-neuron AdEx soma params.
/// Same compartments (:soma + :d1 + :d2) and same Heun integration as
/// `Tripod`, but each neuron reads its own AdExParameterHet fields.
pub struct TripodHet {
// Soma (AdEx) parameters — per-neuron arrays.
soma_param : AdExParameterHet
soma_spike : AdExPostSpike
// Soma-membrane derived scalars (per-neuron). Stored explicitly so
// we don't recompute c = tm/(1000*r) inside the inner loop.
c : Array[Float]
gl : Array[Float]
// Dendrite passive parameters (d1 + d2) — shared across neurons.
d1 : Dendrite
d2 : Dendrite
// External input currents.
i_s : Array[Float]
i_d1 : Array[Float]
i_d2 : Array[Float]
// Synaptic conductances (single-exp on soma + 2 dendrites).
ge_s : Array[Float]
gi_s : Array[Float]
ge_d1 : Array[Float]
gi_d1 : Array[Float]
ge_d2 : Array[Float]
gi_d2 : Array[Float]
glu_s : Array[Float]
gaba_s : Array[Float]
glu_d1 : Array[Float]
gaba_d1 : Array[Float]
glu_d2 : Array[Float]
gaba_d2 : Array[Float]
// Synapse parameters (scalar).
e_e : Float
e_i : Float
tau_e : Float
tau_i : Float
gsyn_e : Float
gsyn_i : Float
// State.
n : Int
v_s : Array[Float]
w_s : Array[Float]
v_d1 : Array[Float]
v_d2 : Array[Float]
fire : Array[Bool]
threshold : Array[Float]
tabs : Array[Int]
// Heun temp arrays (4 * n).
dv : Array[Float]
dv_temp : Array[Float]
syn_curr_s : Array[Float]
syn_curr_d1 : Array[Float]
syn_curr_d2 : Array[Float]
}
///|
/// Construct a TripodHet population. `soma_param` must already be
/// filled with per-neuron arrays of length N. Dendrite parameters
/// are derived from `d1`/`d2` (Julia's `human_dend` defaults).
pub fn TripodHet::new(
n : Int,
soma_param : AdExParameterHet,
rng : Xoshiro,
) -> TripodHet {
// Per-neuron c[i] = tm[i] / r[i] (matches Julia: C = tm / R, with
// tm in ms and R in MΩ giving C in pF when multiplied by 1e-3
// inside the dv formula).
let c : Array[Float] = Array::make(n, 0.0F)
let gl : Array[Float] = Array::make(n, 0.0F)
for k in 0.. Unit {
let n = p.n
for i in 0.. Unit {
let n = p.n
for i in 0.. Unit {
let n = p.n
for i in 0.. Unit {
let n = p.n
for i in 0.. Unit {
let n = p.n
let mut k = 0
while k < n {
let vt = p.soma_param.vt[k]
let el = p.soma_param.el[k]
let tm = p.soma_param.tm[k]
let r_val = p.soma_param.r[k]
let dt_slope = p.soma_param.dt_slope[k]
let tw = p.soma_param.tw[k]
let a = p.soma_param.a[k]
let c_val = p.c[k]
let gl_val = p.gl[k]
let ds : Float = if store_temp { p.dv_temp[k * 4 + 0] * dt } else { 0.0F }
let dd1 : Float = if store_temp { p.dv_temp[k * 4 + 1] * dt } else { 0.0F }
let dd2 : Float = if store_temp { p.dv_temp[k * 4 + 2] * dt } else { 0.0F }
let dw : Float = if store_temp { p.dv_temp[k * 4 + 3] } else { 0.0F }
let ic1 = -(p.v_d1[k] + dd1 - p.v_s[k] - ds) * p.d1.gax[k]
let ic2 = -(p.v_d2[k] + dd2 - p.v_s[k] - ds) * p.d2.gax[k]
let exp_term = if dt_slope < 0.0F {
0.0F
} else {
dt_slope * expf((p.v_s[k] + ds - p.threshold[k]) / dt_slope)
}
let dv_s_val = (gl_val * (el - p.v_s[k] - ds) + exp_term -
p.w_s[k] - dw - p.syn_curr_s[k] - (ic1 + ic2) +
p.i_s[k]) / c_val
let dv_d1_val = ((el - p.v_d1[k] - dd1) * p.d1.gm[k] -
p.syn_curr_d1[k] + ic1 + p.i_d1[k]) / p.d1.c[k]
let dv_d2_val = ((el - p.v_d2[k] - dd2) * p.d2.gm[k] -
p.syn_curr_d2[k] + ic2 + p.i_d2[k]) / p.d2.c[k]
let dw_val = (a * (p.v_s[k] + ds - el) - (p.w_s[k] + dw)) / tw
if store_temp {
p.dv_temp[k * 4 + 0] = dv_s_val
p.dv_temp[k * 4 + 1] = dv_d1_val
p.dv_temp[k * 4 + 2] = dv_d2_val
p.dv_temp[k * 4 + 3] = dw_val
} else {
p.dv[k * 4 + 0] = dv_s_val
p.dv[k * 4 + 1] = dv_d1_val
p.dv[k * 4 + 2] = dv_d2_val
p.dv[k * 4 + 3] = dw_val
}
let _ = vt
let _ = tm
let _ = r_val
k = k + 1
}
}
///|
/// Update TripodHet for one timestep. Same Julia `integrate!` order
/// as `step_tripod`, but reads per-neuron AdEx params.
pub fn step_tripod_het(p : TripodHet, dt : Float) -> Unit {
let n = p.n
let at = p.soma_spike.at
let tau_a = p.soma_spike.tau_a
let tabs_const = p.soma_spike.tabs_const
let up = p.soma_spike.up
let ap_membrane = p.soma_spike.ap_membrane
let tabs_steps : Int = ((up + tabs_const) / dt).to_int()
// 1. Synapses.
tripod_het_soma_step_synapses(p, dt)
tripod_het_dend_step_synapses(p, dt)
// 2. Synaptic currents.
tripod_het_syn_curr_soma(p)
tripod_het_syn_curr_dends(p)
// 3. Heun integration.
tripod_het_heun_step(p, dt, false)
for i in 0..<(n * 4) {
p.dv_temp[i] = p.dv[i]
}
tripod_het_heun_step(p, dt, true)
// 4 + 5. Apply + spike detection (Julia's exact per-neuron order).
for k in 0.. 0 {
// Refractory: v_s = Vr, v_d += axial coupling.
p.v_s[k] = vr_k
p.v_d1[k] = p.v_d1[k] + dt * (p.v_s[k] - p.v_d1[k]) * p.d1.gax[k] / p.d1.c[k]
p.v_d2[k] = p.v_d2[k] + dt * (p.v_s[k] - p.v_d2[k]) * p.d2.gax[k] / p.d2.c[k]
continue
}
// Active phase. Detect fire BEFORE Heun apply.
let v_s_pred = p.v_s[k] + p.dv[k * 4 + 0] * dt
p.fire[k] = v_s_pred >= -10.0F
if p.fire[k] {
p.dv[k * 4 + 0] = ap_membrane - p.v_s[k]
p.v_s[k] = ap_membrane
p.w_s[k] = p.w_s[k] + b_k
p.threshold[k] = p.threshold[k] + at
p.tabs[k] = tabs_steps
continue
}
// No spike: apply Heun correction.
p.v_s[k] = p.v_s[k] + 0.5F * dt * (p.dv[k * 4 + 0] + p.dv_temp[k * 4 + 0])
p.v_d1[k] = p.v_d1[k] + 0.5F * dt * (p.dv[k * 4 + 1] + p.dv_temp[k * 4 + 1])
p.v_d2[k] = p.v_d2[k] + 0.5F * dt * (p.dv[k * 4 + 2] + p.dv_temp[k * 4 + 2])
p.w_s[k] = p.w_s[k] + 0.5F * dt * (p.dv[k * 4 + 3] + p.dv_temp[k * 4 + 3])
}
}