// Sparse matrix (CSR format) — bit-exact port of SNNModels.jl/src/utils/sparse_matrix.jl.
//
// The Julia source uses SparseMatrixCSC (compressed sparse column)
// but we use CSR (compressed sparse row) here for simpler pre-synaptic
// iteration: for each pre-synaptic neuron, iterate over its
// outgoing connections.
//
// Julia reference (src/utils/sparse_matrix.jl):
// sparse_matrix(post, pre, μ, σ, p) = ... → W, I, J, index, rowptr
// dsparse(W, I, J) → SparseMatrixCSC{Float32, Int}
// update_sparse_matrix!(c, W) → refresh from dense matrix
// connect!(c, j, i, μ) → set (i,j) weight then rebuild
//
// For our CSR layout we store (rowptr, colptr, vals) per pre-syn row.
// `connect!(c, post, pre, w)` mutates the dense weight buffer and
// then rebuilds the CSR.
///|
/// Sparse matrix in compressed-sparse-row format.
pub struct SparseMatrixCSR {
rows : Int
cols : Int
// rowptr[i] = start index in colptr/vals of row i's first non-zero.
// Length = rows + 1.
rowptr : Array[Int]
// Column index of each non-zero. Length = nnz.
colptr : Array[Int]
// Value of each non-zero. Length = nnz.
vals : Array[Float]
}
///|
/// Number of stored non-zeros.
pub fn SparseMatrixCSR::nnz(m : SparseMatrixCSR) -> Int {
m.vals.length()
}
///|
/// Build an empty N×M sparse matrix.
pub fn SparseMatrixCSR::empty(rows : Int, cols : Int) -> SparseMatrixCSR {
let rowptr : Array[Int] = Array::make(rows + 1, 0)
{ rows, cols, rowptr, colptr: [], vals: [] }
}
///|
/// Build a CSR sparse matrix from a dense Float32 matrix.
/// Density = fraction of non-zero entries preserved.
pub fn SparseMatrixCSR::from_dense(
dense : Array[Array[Float]],
threshold : Float,
) -> SparseMatrixCSR {
let rows = dense.length()
let cols = if rows > 0 { dense[0].length() } else { 0 }
let rowptr : Array[Int] = Array::make(rows + 1, 0)
let colptr : Array[Int] = []
let vals : Array[Float] = []
for i in 0.. threshold {
colptr.push(j)
vals.push(v)
}
}
}
rowptr[rows] = vals.length()
{ rows, cols, rowptr, colptr, vals }
}
///|
/// Connection rules for random sparse matrices. Mirrors Julia SNN's
/// `rule` keyword on `sparse_matrix(post, pre, μ, σ, p; rule=...)`:
/// - `Bernoulli`: each edge independently kept with probability `p`
/// - `FixedIn` : each post-synaptic neuron receives from exactly
/// `p * Npre` pre-synaptic neurons (chosen uniformly
/// without replacement). Same in-degree for all posts.
/// - `FixedOut` : each pre-synaptic neuron projects to exactly
/// `p * Npost` post-synaptic neurons. Same out-degree
/// for all pres.
pub(all) enum ConnectRule {
Bernoulli
FixedIn
FixedOut
}
///|
/// Build a sparse matrix with random Normal(μ, σ) weights and
/// connection probability `p`. Equivalent to Julia's
/// `sparse_matrix(post, pre, μ, σ, p)`.
///
/// `rng` is the Xoshiro stream. The Float32 draw path is bit-exact
/// with Julia's `rand(rng, Float32)` for the same seed.
pub fn SparseMatrixCSR::random(
rows : Int,
cols : Int,
mu : Float,
sigma : Float,
p : Float,
rng : Xoshiro,
) -> SparseMatrixCSR {
// Default rule is Bernoulli for backward compatibility.
SparseMatrixCSR::random_with_rule(rows, cols, mu, sigma, p, Bernoulli, rng)
}
///|
/// Build a sparse matrix with the given connection rule. See
/// `ConnectRule` for the semantics of each rule.
pub fn SparseMatrixCSR::random_with_rule(
rows : Int,
cols : Int,
mu : Float,
sigma : Float,
p : Float,
rule : ConnectRule,
rng : Xoshiro,
) -> SparseMatrixCSR {
// Build the dense weight matrix first. We use Float for weights
// (matches Julia's Float32 default).
let dense : Array[Array[Float]] = Array::make(rows, [])
for i in 0.. {
// Each weight independently zeroed with prob (1 - p).
for i in 0..= p {
dense[i][j] = 0.0F
}
}
}
}
FixedIn => {
// For each post j, keep exactly `round(p*rows)` pre-synaptic
// neurons (chosen uniformly without replacement). The rest go
// to 0.
let n_keep = (Float::from_int(rows) * p).to_int()
if n_keep > 0 && n_keep <= rows {
for j in 0..= rows { rows - 1 } else { r_idx }
// Swap pre_idx[k] and pre_idx[r_clamped]
let tmp = pre_idx[k]
pre_idx[k] = pre_idx[r_clamped]
pre_idx[r_clamped] = tmp
}
for k in 0.. {
// For each pre i, keep exactly `round(p*cols)` post-synaptic
// neurons.
let n_keep = (Float::from_int(cols) * p).to_int()
if n_keep > 0 && n_keep <= cols {
for i in 0..= cols { cols - 1 } else { r_idx }
let tmp = post_idx[k]
post_idx[k] = post_idx[r_clamped]
post_idx[r_clamped] = tmp
}
for k in 0.. Float {
if i < 0 || i >= m.rows {
return 0.0F
}
let start = m.rowptr[i]
let end = m.rowptr[i + 1]
// Linear scan within row i (rows are typically short)
for k in start.. Unit {
if i < 0 || i >= m.rows {
return
}
let start = m.rowptr[i]
let end = m.rowptr[i + 1]
// Look for existing entry
let mut found = false
for k in start.. Unit {
let rows = m.rows
for i in 0.. Unit {
let rows = m.rows
for i in 0.. Int {
if i < 0 || i >= m.rows {
return 0
}
m.rowptr[i + 1] - m.rowptr[i]
}