pub const RM_STDP_TAU_PLUS: f32 = 20.0;
pub const RM_STDP_TAU_MINUS: f32 = 20.0;
pub const RM_STDP_A_PLUS: f32 = 0.01;
pub const RM_STDP_A_MINUS: f32 = 0.012;
pub const RM_STDP_W_MIN: f32 = 0.0;
pub const RM_STDP_W_MAX: f32 = 2.0;
const _: () = assert!(RM_STDP_W_MIN < RM_STDP_W_MAX);
const _: () = assert!(RM_STDP_A_MINUS >= RM_STDP_A_PLUS);
pub struct EligibilityTrace {
pub value: f32,
pub tau: f32,
}
pub struct RmStdpConfig {
pub tau_eligibility: f32,
pub reward_lr: f32,
pub w_min: f32,
pub w_max: f32,
}
impl EligibilityTrace {
pub fn decay(&mut self) {
let tau = self.tau.max(f32::EPSILON);
self.value *= (-1.0 / tau).exp();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decay_scales_value_by_exp_neg_inv_tau() {
let mut trace = EligibilityTrace {
value: 1.0,
tau: 50.0,
};
let expected_factor = (-1.0_f32 / 50.0).exp();
trace.decay();
assert!((trace.value - expected_factor).abs() < 1e-6);
}
#[test]
fn decay_applied_repeatedly_compounds_toward_zero() {
let mut trace = EligibilityTrace {
value: 1.0,
tau: 50.0,
};
let factor = (-1.0_f32 / 50.0).exp();
for _ in 0..5 {
trace.decay();
}
let expected = factor.powi(5);
assert!((trace.value - expected).abs() < 1e-5);
assert!(trace.value < 1.0);
}
#[test]
fn decay_preserves_sign_for_negative_values() {
let mut trace = EligibilityTrace {
value: -1.0,
tau: 50.0,
};
trace.decay();
assert!(trace.value < 0.0);
}
#[test]
fn decay_with_zero_tau_does_not_panic_or_diverge() {
let mut trace = EligibilityTrace {
value: 1.0,
tau: 0.0,
};
trace.decay();
assert!(trace.value.is_finite());
assert!(trace.value >= 0.0);
}
#[test]
fn decay_with_negative_tau_does_not_panic_or_diverge() {
let mut trace = EligibilityTrace {
value: 1.0,
tau: -10.0,
};
trace.decay();
assert!(trace.value.is_finite());
assert!(trace.value >= 0.0);
}
}