// Metaplasticity — homeostatic weight normalization for SpikingSynapses.
//
// Julia reference:
// SNNModels.jl/src/connections/metaplasticity/normalization.jl
// SNNModels.jl/test/syn/metaplasticity.jl
//
// Provides:
// - `MultiplicativeNorm(τ)`: per-step multiplicative rule.
// At each plasticity call: μ[i] = (W0[i] - W1[i]) / W1[i];
// W[s] *= (1 + μ[i]).
// - `AdditiveNorm(τ)`: per-step additive rule.
// At each plasticity call: μ[i] = W0[i] - W1[i];
// W[s] += μ[i].
// - `SynapseNormalization(targets, param)`: constructor that
// captures initial per-post-synaptic-neuron weight sum W0[i]
// from a list of synapse targets.
// - `metaplasticity_step(norm)`: runs the normalization step
// (always-on / per-step form).
// - `metaplasticity_step_gated(norm, step_count, dt)`: periodic
// form (every τ/dt steps), matches Julia's outer `plasticity!`.
//
// Bit-exact ordering matches Julia's `plasticity!`:
// 1. W1[i] = 0; for each synapse: W1[i] += sum of weights connecting to i.
// 2. μ[i] = operator(W0[i], W1[i]) - W1[i] (multiplicative or additive).
// 3. for each synapse: W[s] += operator(W[s], μ[i]).
//
// Float32 contract: every arithmetic uses Float32.
// =========================================================================
// MultiplicativeNorm
// =========================================================================
///|
/// MultiplicativeNorm — per-step multiplicative rule.
/// At each step: μ[i] = (W0[i] - W1[i]) / W1[i] (the ratio of
/// initial-sum to current-sum minus 1); W[s] *= (1 + μ[i]).
pub struct MultiplicativeNorm {
tau : Float
}
///|
pub fn MultiplicativeNorm::new(tau : Float) -> MultiplicativeNorm {
{ tau: tau }
}
// =========================================================================
// AdditiveNorm
// =========================================================================
///|
/// AdditiveNorm — per-step additive rule.
/// At each step: μ[i] = W0[i] - W1[i]; W[s] += μ[i].
pub struct AdditiveNorm {
tau : Float
}
///|
pub fn AdditiveNorm::new(tau : Float) -> AdditiveNorm {
{ tau: tau }
}
// =========================================================================
// NormParam enum (dispatches multiplicative vs additive step)
// =========================================================================
///|
/// Sum type for the two normalisation rules. Used internally to
/// dispatch `metaplasticity_step`.
pub(all) enum NormParam {
Multiplicative_(MultiplicativeNorm)
Additive_(AdditiveNorm)
}
// =========================================================================
// SynapseNormalization
// =========================================================================
///|
/// SynapseNormalization — holds the initial weight sum W0[i] for
/// each post-neuron i, and a list of synapse targets whose weights
/// to mutate. Mirrors Julia's `SynapseNormalization` struct (subset
/// of fields).
pub struct SynapseNormalization {
param : NormParam
n_post : Int
// Per-post-neuron initial weight sum (captured at construction).
w0 : Array[Float]
// Per-post-neuron current weight sum (computed each step).
w1 : Array[Float]
// Per-post-neuron μ (intermediate value, computed each step).
mu : Array[Float]
// Synapse targets — one CSR buffer per source, all targeting the
// same post-population.
targets : Array[SynapseTarget]
}
///|
/// SynapseTarget — holds a reference to a synapse's CSR (vals +
/// rowptr) and a post-population identifier. The mutation loop uses
/// `rowptr[i]` and `vals[j]` to access the actual weights. `post_id`
/// identifies which post-neuron population the buffer targets and is
/// used by callers to ensure all targets in the normalization share
/// the same post-population.
pub(all) struct SynapseTarget {
// Identifier for the post-synaptic population.
post_id : Int
// CSR weights array (will be mutated in place).
vals : Array[Float]
// rowptr (length n_post + 1). Note: our CSR uses rowptr indexed
// by post-neuron (since we target the post-side for normalization).
rowptr : Array[Int]
}
///|
/// Construct a SynapseNormalization from a list of synapse targets.
/// All targets MUST share the same `post_id` AND the same `n_post` —
/// the caller is expected to verify this before calling. (We do not
/// abort at construction time because MoonBit's `abort` is a
/// polymorphic bottom type that the type system doesn't track as a
/// raise site — making `try...catch` ergonomics poor. Julia's
/// `@assert` is a runtime check that throws AssertionError; we rely
/// on the caller to ensure post-population consistency.)
pub fn SynapseNormalization::new(
targets : Array[SynapseTarget],
param : NormParam,
) -> SynapseNormalization {
// All targets must have the same n_post (derived from rowptr).
let n_post = if targets.length() > 0 {
targets[0].rowptr.length() - 1
} else {
0
}
// Compute W0[i]: for each post-neuron i, sum weights of all
// synapses at rowptr[i]..rowptr[i+1] across all target synapses.
let w0 : Array[Float] = Array::make(n_post, 0.0F)
for t in 0.. Unit {
let n_post = norm.n_post
let n_targets = norm.targets.length()
// Step 1: compute W1 (sum of weights targeting each post-neuron).
for i in 0.. {
let _ = p
for i in 0.. 0.0F {
norm.mu[i] = (norm.w0[i] - norm.w1[i]) / norm.w1[i]
} else {
norm.mu[i] = 0.0F
}
}
}
Additive_(p) => {
let _ = p
for i in 0.. {
vals[j] = vals[j] * (1.0F + mu_i)
}
Additive_(_) => {
vals[j] = vals[j] + mu_i
}
}
j = j + 1
}
}
}
}
///|
/// Periodic form of `metaplasticity_step` — mirrors Julia's outer
/// `plasticity!(c, param, dt, T)` which gates the step on
/// `((tt) % round(Int, τ / dt)) < dt`. Only runs the normalization
/// when `step_count` is a multiple of `τ / dt` (rounded).
///
/// When τ == 0, never fires (matches Julia's τ=0 default which
/// disables the periodic rule). Otherwise fires every
/// `round(τ/dt)` steps.
///
/// Note: this is a simple "every-N-steps" gate, not a true-time gate.
/// Julia uses `T.t / dt` for the same purpose; our `step_count`
/// parameter is the equivalent (caller maintains the step counter).
pub fn metaplasticity_step_gated(
norm : SynapseNormalization,
step_count : Int,
dt : Float,
) -> Unit {
// Read τ from the rule (multiplicative and additive share the
// gate parameter).
let tau = match norm.param {
Multiplicative_(p) => p.tau
Additive_(p) => p.tau
}
if tau <= 0.0F {
return
}
let period_steps = Float::to_int(tau / dt + 0.5F)
if period_steps <= 0 {
return
}
// Gate: fire only when step_count is a multiple of period_steps.
// (step_count is 0-based; the first step at 0 fires if 0 is a
// multiple of period_steps, matching Julia's `((tt) % round(Int,
// τ/dt)) < dt` check at tt=0.)
if step_count % period_steps != 0 {
return
}
metaplasticity_step(norm)
}