// stdp_gerstner.mbt — STDP Gerstner 1996 pair-based exponential (v0.40.0).
//
// Gerstner, W., Kempter, R., van Hemmen, J. L., & Wagner, H. (1996).
// A neuronal learning rule for sub-millisecond temporal coding.
// Nature, 383(6595), 76–78.
//
// Each synapse carries a pre-synaptic trace tpre[j] and post-synaptic
// trace tpost[i]. On a pre-spike, the trace increases by A_pre; on a
// post-spike, the trace increases by A_post. Both traces decay
// exponentially with their own time constants τpre and τpost.
//
// On a post-spike: W[s] += A_pre · tpre[J[s]] (LTP via pre trace).
// On a pre-spike : W[s] += A_post · tpost[I[s]] (LTD via post trace).
//
// All updates clamped to [Wmin, Wmax]. Float32 throughout.

///|
/// STDP Gerstner parameters.
pub(all) struct StdpGerstnerParam {
  a_pre : Float    // LTP learning rate (post-spike reads pre-trace)
  a_post : Float   // LTD learning rate (pre-spike reads post-trace)
  tau_pre : Float  // pre-synaptic trace time constant (ms)
  tau_post : Float // post-synaptic trace time constant (ms)
  w_min : Float
  w_max : Float
}

///|
pub fn StdpGerstnerParam::new() -> StdpGerstnerParam {
  {
    a_pre: 0.01F,
    a_post: 0.01F,
    tau_pre: 20.0F,
    tau_post: 20.0F,
    w_min: 0.0F,
    w_max: 1.0F,
  }
}

///|
/// STDP Gerstner state (per-network traces).
pub struct StdpGerstnerState {
  n_pre : Int
  n_post : Int
  tpre : Array[Float]
  tpost : Array[Float]
  last_pre : Array[Float]
  last_post : Array[Float]
}

///|
pub fn StdpGerstnerState::new(n_pre : Int, n_post : Int) -> StdpGerstnerState {
  {
    n_pre,
    n_post,
    tpre: Array::make(n_pre, 0.0F),
    tpost: Array::make(n_post, 0.0F),
    last_pre: Array::make(n_pre, 0.0F),
    last_post: Array::make(n_post, 0.0F),
  }
}

///|
/// One step of STDP Gerstner plasticity. `now` is the current
/// simulation time (ms). `fire_pre[j]` / `fire_post[i]` are boolean
/// spike indicators. `w` is the dense weight matrix indexed by
/// (i*n_pre + j) for post i, pre j; updated in-place.
pub fn stdp_gerstner_step(
  state : StdpGerstnerState,
  param : StdpGerstnerParam,
  fire_pre : Array[Bool],
  fire_post : Array[Bool],
  w : Array[Float],
  now : Float,
) -> Unit {
  let n_pre = state.n_pre
  let n_post = state.n_post
  let inv_tau_pre = 1.0F / param.tau_pre
  let inv_tau_post = 1.0F / param.tau_post
  // Update pre traces.
  for j in 0.. state.last_pre[j] {
      now - state.last_pre[j]
    } else {
      0.0F
    }
    let decay = expf(-dt_pre * inv_tau_pre)
    if fire_pre[j] {
      state.tpre[j] = state.tpre[j] * decay + param.a_pre
      state.last_pre[j] = now
    } else {
      state.tpre[j] = state.tpre[j] * decay
    }
  }
  // Update post traces.
  for i in 0.. state.last_post[i] {
      now - state.last_post[i]
    } else {
      0.0F
    }
    let decay = expf(-dt_post * inv_tau_post)
    if fire_post[i] {
      state.tpost[i] = state.tpost[i] * decay + param.a_post
      state.last_post[i] = now
    } else {
      state.tpost[i] = state.tpost[i] * decay
    }
  }
  // Weight updates.
  for i in 0.. param.w_max {
        w[s] = param.w_max
      } else {
        w[s] = w_new
      }
    }
  }
}