// analysis_weight.mbt — average synaptic weight between two neuron sets.
//
// Port of `refs/SNNUtils.jl/src/analysis/weights.jl`:
// - `average_weight(pre, post, synapse)` — mean of all synapse weights
// whose pre-neuron ∈ `pre` and post-neuron ∈ `post`.
// - `average_weight_dynamics(pre, post, synapse, record)` — same but
// averaged across a per-timestep weight record matrix.
//
// Note: Julia's `average_weight` uses `@unpack rowptr, colptr, I, J, index, W = synapse`
// where `rowptr` indexes POSTSYNAPTIC neurons (transposed CSR layout).
// Our `SpikingSynapse.matrix` is in standard row-major CSR (rows = pre),
// so we iterate pre → outgoing edges → filter post-neuron ∈ post_pop.
//
// Returns Float (mean weight). Empty intersection returns 0.0F.
///|
/// Check if `arr` (Array[Int]) contains `v`. Linear scan.
fn contains_int_local(arr : Array[Int], v : Int) -> Bool {
let mut i = 0
let n = arr.length()
let mut found = false
while i < n {
if arr[i] == v {
found = true
}
i = i + 1
}
found
}
///|
/// Mean of all synaptic weights whose pre-neuron ∈ `pre_pop_neurons`
/// AND post-neuron ∈ `post_pop_neurons`. Returns 0.0F if the set is
/// empty. Walks the CSR by pre-neuron (row-major).
pub fn average_weight(
pre_pop_neurons : Array[Int],
post_pop_neurons : Array[Int],
synapse : SpikingSynapse,
) -> Float {
let mut sum : Float = 0.0F
let mut count : Int = 0
let pre_n = synapse.pre.n
let post_n = synapse.post.n
let mat = synapse.matrix
// Walk each pre neuron j. If j ∈ pre_pop_neurons, scan its outgoing edges.
for j in 0.. Array[Float] {
let pre_n = synapse.pre.n
let post_n = synapse.post.n
let mat = synapse.matrix
// Collect the indices of matching edges first.
let matched : Array[Int] = []
for j in 0.. Array[Int] {
let pre_n = synapse.pre.n
let post_n = synapse.post.n
let mat = synapse.matrix
let result : Array[Int] = []
for j in 0.. Unit {
let matched = weights_indices(pre_pop_neurons, post_pop_neurons, synapse)
let mut k = 0
while k < matched.length() {
let s = matched[k]
synapse.matrix.vals[s] = synapse.matrix.vals[s] * factor
k = k + 1
}
}