#[derive(Debug, Clone)]
pub struct SparseSynapseMatrix {
weights: Vec<i16>,
col_indices: Vec<u16>,
synapse_indices: Vec<usize>,
pre_ids: Vec<u16>,
row_ptrs: Vec<u32>,
weight_index_of: Vec<usize>,
rv_row_ptrs: Vec<u32>,
rv_pre_ids: Vec<u16>,
rv_syn_indices: Vec<usize>,
neuron_count: u16,
}
impl SparseSynapseMatrix {
#[must_use]
pub fn new(neuron_count: u16, estimated_synapses: usize) -> Self {
Self {
weights: Vec::with_capacity(estimated_synapses),
col_indices: Vec::with_capacity(estimated_synapses),
synapse_indices: Vec::with_capacity(estimated_synapses),
pre_ids: Vec::with_capacity(estimated_synapses),
row_ptrs: vec![0; neuron_count as usize + 1],
weight_index_of: Vec::with_capacity(estimated_synapses),
rv_row_ptrs: Vec::new(),
rv_pre_ids: Vec::with_capacity(estimated_synapses),
rv_syn_indices: Vec::with_capacity(estimated_synapses),
neuron_count,
}
}
pub fn add(&mut self, pre_id: u16, post_id: u16, weight: i16, synapse_index: usize) {
debug_assert!(
pre_id < self.neuron_count,
"pre_id {pre_id} ≥ neuron_count {}",
self.neuron_count
);
self.weights.push(weight);
self.col_indices.push(post_id);
self.synapse_indices.push(synapse_index);
self.pre_ids.push(pre_id);
self.weight_index_of.push(synapse_index);
for row in (pre_id as usize + 1)..self.row_ptrs.len() {
self.row_ptrs[row] += 1;
}
}
pub fn finalize(&mut self) {
let n = self.neuron_count as usize;
let m = self.weights.len();
let ins_post: Vec<u16> = self.col_indices.clone();
let ins_syn: Vec<usize> = self.synapse_indices.clone();
let mut degree = vec![0u32; n];
for &pre in &self.pre_ids {
degree[pre as usize] += 1;
}
self.row_ptrs = vec![0; n + 1];
for (k, °) in degree.iter().enumerate() {
self.row_ptrs[k + 1] = self.row_ptrs[k] + deg;
}
let mut new_weights = vec![0i16; m];
let mut new_col = vec![0u16; m];
let mut new_syn = vec![0usize; m];
let mut cursor = self.row_ptrs.clone();
for (i, &pre) in self.pre_ids.iter().enumerate() {
let pos = cursor[pre as usize] as usize;
new_weights[pos] = self.weights[i];
new_col[pos] = self.col_indices[i];
new_syn[pos] = self.synapse_indices[i];
cursor[pre as usize] += 1;
}
self.weights = new_weights;
self.col_indices = new_col;
self.synapse_indices = new_syn;
let mut inv = vec![0usize; m];
for (pos, &syn_idx) in self.synapse_indices.iter().enumerate() {
if syn_idx < m {
inv[syn_idx] = pos;
}
}
self.weight_index_of = inv;
let mut post_degree = vec![0u32; n];
for &post in &ins_post {
let p = post as usize;
if p < n {
post_degree[p] += 1;
}
}
let mut rv_row_ptrs = vec![0u32; n + 1];
for (k, °) in post_degree.iter().enumerate() {
rv_row_ptrs[k + 1] = rv_row_ptrs[k] + deg;
}
let mut rv_pre_ids = vec![0u16; m];
let mut rv_syn_indices = vec![0usize; m];
let mut rcursor = rv_row_ptrs.clone();
for i in 0..m {
let post = ins_post[i] as usize;
if post < n {
let pos = rcursor[post] as usize;
rv_pre_ids[pos] = self.pre_ids[i];
rv_syn_indices[pos] = ins_syn[i];
rcursor[post] += 1;
}
}
self.rv_row_ptrs = rv_row_ptrs;
self.rv_pre_ids = rv_pre_ids;
self.rv_syn_indices = rv_syn_indices;
}
#[must_use]
pub fn connections(&self, pre_id: u16) -> SynapseIter<'_> {
debug_assert!(pre_id < self.neuron_count);
let start = self.row_ptrs[pre_id as usize] as usize;
let end = self.row_ptrs[pre_id as usize + 1] as usize;
SynapseIter {
weights: &self.weights[start..end],
col_indices: &self.col_indices[start..end],
synapse_indices: &self.synapse_indices[start..end],
pos: 0,
}
}
#[must_use]
pub fn incoming(&self, post_id: u16) -> IncomingIter<'_> {
if self.rv_row_ptrs.is_empty() {
return IncomingIter {
pre_ids: &[],
syn_indices: &[],
pos: 0,
};
}
debug_assert!(post_id < self.neuron_count);
let start = self.rv_row_ptrs[post_id as usize] as usize;
let end = self.rv_row_ptrs[post_id as usize + 1] as usize;
IncomingIter {
pre_ids: &self.rv_pre_ids[start..end],
syn_indices: &self.rv_syn_indices[start..end],
pos: 0,
}
}
pub fn set_weight(&mut self, synapse_index: usize, weight: i16) {
if let Some(&pos) = self.weight_index_of.get(synapse_index) {
if let Some(slot) = self.weights.get_mut(pos) {
*slot = weight;
}
}
}
pub fn clear(&mut self) {
self.weights.clear();
self.col_indices.clear();
self.synapse_indices.clear();
self.pre_ids.clear();
self.weight_index_of.clear();
self.rv_row_ptrs.clear();
self.rv_pre_ids.clear();
self.rv_syn_indices.clear();
self.row_ptrs.iter_mut().for_each(|p| *p = 0);
}
#[must_use]
pub fn len(&self) -> usize {
self.weights.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.weights.is_empty()
}
}
#[derive(Debug, Clone)]
pub struct SynapseIter<'a> {
weights: &'a [i16],
col_indices: &'a [u16],
synapse_indices: &'a [usize],
pos: usize,
}
impl Iterator for SynapseIter<'_> {
type Item = (u16, i16, usize);
fn next(&mut self) -> Option<Self::Item> {
if self.pos < self.weights.len() {
let item = (
self.col_indices[self.pos],
self.weights[self.pos],
self.synapse_indices[self.pos],
);
self.pos += 1;
Some(item)
} else {
None
}
}
}
#[derive(Debug, Clone)]
pub struct IncomingIter<'a> {
pre_ids: &'a [u16],
syn_indices: &'a [usize],
pos: usize,
}
impl Iterator for IncomingIter<'_> {
type Item = (u16, usize);
fn next(&mut self) -> Option<Self::Item> {
if self.pos < self.pre_ids.len() {
let item = (self.pre_ids[self.pos], self.syn_indices[self.pos]);
self.pos += 1;
Some(item)
} else {
None
}
}
}