#![allow(clippy::module_name_repetitions)]
#![allow(
clippy::cast_possible_truncation,
clippy::cast_possible_wrap,
clippy::cast_sign_loss
)]
use crate::Error;
pub const SCALE: i32 = 1000;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum SynapseType {
#[default]
Excitatory,
Inhibitory,
Modulatory,
}
#[derive(Debug, Clone)]
pub struct Synapse {
pub pre_neuron_id: u16,
pub post_neuron_id: u16,
pub synapse_type: SynapseType,
pub weight: i16,
pub max_weight: i16,
pub min_weight: i16,
pub tau_rise_us: u16,
pub tau_decay_us: u16,
pub raw_stdp_delta: i64,
pub absorbed_delta: i64,
}
impl Synapse {
pub fn new(pre_id: u16, post_id: u16, weight: i16) -> Result<Self, Error> {
if pre_id == post_id {
return Err(Error::InvalidParameter);
}
let synapse_type = if weight >= 0 {
SynapseType::Excitatory
} else {
SynapseType::Inhibitory
};
let (tau_rise_us, tau_decay_us, max_weight, min_weight) = biological_params(synapse_type);
Ok(Self {
pre_neuron_id: pre_id,
post_neuron_id: post_id,
synapse_type,
weight,
max_weight,
min_weight,
tau_rise_us,
tau_decay_us,
raw_stdp_delta: 0,
absorbed_delta: 0,
})
}
pub fn update_weight(&mut self, delta_weight: i16) {
let target = self
.weight
.saturating_add(delta_weight)
.clamp(self.min_weight, self.max_weight);
let applied = target - self.weight;
self.raw_stdp_delta += i64::from(delta_weight);
self.absorbed_delta += i64::from(delta_weight) - i64::from(applied);
self.weight = target;
}
}
#[cfg(test)]
impl Synapse {
fn normalized_weight(&self) -> i16 {
let abs_max = self.max_weight.unsigned_abs();
if abs_max > 0 {
(i32::from(self.weight) * 100 / i32::from(abs_max)) as i16
} else {
0
}
}
}
impl Default for Synapse {
fn default() -> Self {
Self::new(0, 1, 100).unwrap_or(Self {
pre_neuron_id: 0,
post_neuron_id: 1,
synapse_type: SynapseType::Excitatory,
weight: 100,
max_weight: 2000,
min_weight: 0,
tau_rise_us: 500,
tau_decay_us: 5_000,
raw_stdp_delta: 0,
absorbed_delta: 0,
})
}
}
fn biological_params(t: SynapseType) -> (u16, u16, i16, i16) {
match t {
SynapseType::Excitatory => (500, 5_000, 2000, 0), SynapseType::Inhibitory => (300, 10_000, 0, -2000), SynapseType::Modulatory => (1000, 50_000, 1000, -1000), }
}
#[derive(Debug, Clone)]
pub struct STDPRule {
pub tau_plus_us: u32,
pub tau_minus_us: u32,
pub a_plus: i16,
pub a_minus: i16,
pub learning_rate: u16,
}
impl STDPRule {
#[must_use]
pub fn new() -> Self {
Self {
tau_plus_us: 20_000,
tau_minus_us: 20_000,
a_plus: 50,
a_minus: -53, learning_rate: 100,
}
}
#[must_use]
pub fn calculate_weight_change(&self, dt_us: i32) -> i16 {
if dt_us > 0 {
let decay = (i64::from(dt_us).abs() * i64::from(SCALE)) / i64::from(self.tau_minus_us);
if decay < 10_000 {
let factor = (i64::from(SCALE) - decay).max(0); ((i64::from(self.a_minus) * factor * i64::from(self.learning_rate))
/ (i64::from(SCALE) * i64::from(SCALE))) as i16
} else {
0
}
} else {
let decay = (i64::from(dt_us).abs() * i64::from(SCALE)) / i64::from(self.tau_plus_us);
if decay < 10_000 {
let factor = (i64::from(SCALE) - decay).max(0); ((i64::from(self.a_plus) * factor * i64::from(self.learning_rate))
/ (i64::from(SCALE) * i64::from(SCALE))) as i16
} else {
0
}
}
}
}
impl Default for STDPRule {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::shadow_unrelated)]
use super::*;
use proptest::prelude::*;
#[test]
fn new_excitatory_synapse_has_default_params() {
let s = Synapse::new(1, 2, 100).expect("valid ids");
assert_eq!(s.pre_neuron_id, 1);
assert_eq!(s.post_neuron_id, 2);
assert_eq!(s.weight, 100);
assert_eq!(s.synapse_type, SynapseType::Excitatory);
assert_eq!(s.tau_rise_us, 500); assert_eq!(s.tau_decay_us, 5_000);
assert_eq!(s.max_weight, 2000);
assert_eq!(s.min_weight, 0);
}
#[test]
fn negative_weight_makes_inhibitory() {
let s = Synapse::new(1, 2, -100).expect("valid ids");
assert_eq!(s.synapse_type, SynapseType::Inhibitory);
assert_eq!(s.tau_rise_us, 300); assert_eq!(s.tau_decay_us, 10_000);
assert_eq!(s.max_weight, 0);
assert_eq!(s.min_weight, -2000);
}
#[test]
fn self_connection_rejected() {
let err = Synapse::new(5, 5, 100).unwrap_err();
assert_eq!(err, Error::InvalidParameter);
}
#[test]
fn weight_bounds_directly_configurable() {
let mut s = Synapse::new(1, 2, 150).expect("valid ids");
s.min_weight = -500;
s.max_weight = 500;
assert_eq!(s.max_weight, 500);
assert_eq!(s.min_weight, -500);
}
#[test]
fn biological_params_map_modulatory_type() {
let (tau_rise, tau_decay, _max_w, _min_w) = biological_params(SynapseType::Modulatory);
assert_eq!(tau_decay, 50_000);
assert!(tau_rise > 0);
}
#[test]
fn normalized_weight_is_percentage_of_max() {
let mut s = Synapse::new(1, 2, 100).expect("valid ids");
assert_eq!(s.max_weight, 2000);
assert_eq!(s.normalized_weight(), 5);
s.weight = s.max_weight;
assert_eq!(s.normalized_weight(), 100);
s.weight = -s.max_weight;
assert_eq!(s.normalized_weight(), -100);
s.max_weight = 0;
assert_eq!(s.normalized_weight(), 0, "zero max must not divide by zero");
}
#[test]
fn weight_clamped_at_bounds() {
let mut s = Synapse::new(1, 2, 100).expect("valid ids");
s.update_weight(10_000); assert_eq!(s.weight, s.max_weight);
s.update_weight(-10_000); assert_eq!(s.weight, s.min_weight);
}
#[test]
fn stdp_ltp_when_pre_before_post() {
let rule = STDPRule::new();
let dt_us: i32 = -5_000; let delta = rule.calculate_weight_change(dt_us);
assert!(
delta >= 0,
"pre-before-post must produce LTP (>=0), got {delta}"
);
}
#[test]
fn stdp_ltd_when_post_before_pre() {
let rule = STDPRule::new();
let dt_us: i32 = 5_000;
let delta = rule.calculate_weight_change(dt_us);
assert!(
delta <= 0,
"post-before-pre must produce LTD (<=0), got {delta}"
);
}
#[test]
fn stdp_zero_outside_window() {
let rule = STDPRule::new();
let far_dt: i32 = -200_000;
assert_eq!(rule.calculate_weight_change(far_dt), 0);
let far_dt_pos: i32 = 200_000;
assert_eq!(rule.calculate_weight_change(far_dt_pos), 0);
}
#[test]
fn stdp_zero_at_zero_dt() {
let rule = STDPRule::new();
let delta = rule.calculate_weight_change(0);
let expected =
(i32::from(rule.a_plus) * SCALE * i32::from(rule.learning_rate)) / (SCALE * SCALE);
assert_eq!(delta, expected as i16);
}
#[test]
fn stdp_decay_monotonic_with_abs_dt() {
let rule = STDPRule::new();
let small = rule.calculate_weight_change(-1_000).abs();
let large = rule.calculate_weight_change(-10_000).abs();
assert!(
large <= small,
"|delta| must decay with |dt|: small={small}, large={large}"
);
}
#[test]
fn stdp_zero_just_past_the_old_overflow_boundary() {
let rule = STDPRule::new();
assert_eq!(rule.calculate_weight_change(2_200_000), 0, "LTD branch");
assert_eq!(rule.calculate_weight_change(-2_200_000), 0, "LTP branch");
assert_eq!(rule.calculate_weight_change(i32::MAX), 0);
assert_eq!(rule.calculate_weight_change(i32::MIN), 0);
}
#[test]
fn stdp_extreme_learning_rate_does_not_overflow_the_product() {
let mut rule = STDPRule::new();
rule.learning_rate = u16::MAX;
let delta = rule.calculate_weight_change(0);
let expected = (i64::from(rule.a_plus) * i64::from(SCALE) * i64::from(u16::MAX))
/ (i64::from(SCALE) * i64::from(SCALE));
assert_eq!(
i64::from(delta),
expected,
"no wrap at lr = u16::MAX, dt = 0"
);
}
proptest! {
#[test]
fn prop_self_connection_always_rejected(id in 0u16..=1000, weight in -3000i16..=3000) {
let result = Synapse::new(id, id, weight);
prop_assert!(result.is_err());
}
#[test]
fn prop_weight_clamped(
weight in -2000i16..=2000,
delta in -5000i16..=5000,
) {
let mut s = Synapse::new(1, 2, weight).unwrap_or_else(|_| {
Synapse::new(1, 2, 0).expect("fallback synapse")
});
s.update_weight(delta);
prop_assert!(s.weight >= s.min_weight);
prop_assert!(s.weight <= s.max_weight);
}
#[test]
fn prop_stdp_sign_convention(dt_us in -200_000i32..=200_000) {
let rule = STDPRule::new();
let delta = rule.calculate_weight_change(dt_us);
if dt_us > 0 {
prop_assert!(delta <= 0, "post-before-pre must produce LTD");
} else if dt_us < 0 {
prop_assert!(delta >= 0, "pre-before-post must produce LTP");
}
}
#[test]
fn prop_stdp_zero_outside_window(multiplier in 11u32..=100) {
let rule = STDPRule::new();
let far_dt_pos = (rule.tau_minus_us * multiplier) as i32;
let far_dt_neg = -((rule.tau_plus_us * multiplier) as i32);
prop_assert_eq!(rule.calculate_weight_change(far_dt_pos), 0);
prop_assert_eq!(rule.calculate_weight_change(far_dt_neg), 0);
}
#[test]
fn prop_stdp_dt_full_i32_range(dt_us in i32::MIN..=i32::MAX) {
let rule = STDPRule::new();
let delta = rule.calculate_weight_change(dt_us);
let outside =
i64::from(dt_us).abs() >= 10 * i64::from(rule.tau_plus_us.max(rule.tau_minus_us));
if outside {
prop_assert_eq!(delta, 0, "outside the window at dt={}", dt_us);
} else if dt_us > 0 {
prop_assert!(delta <= 0, "in-window post-before-pre must be LTD at dt={}", dt_us);
} else {
prop_assert!(delta >= 0, "in-window pre-before-post must be LTP at dt={}", dt_us);
}
}
}
}