use crate::{Optimizer, OptimizerError, OptimizerResult, OptimizerState};
use parking_lot::RwLock;
use scirs2_core::ndarray::{Array1, Array2};
use scirs2_core::random::{thread_rng, Uniform};
use std::collections::{HashMap, VecDeque};
use std::sync::Arc;
use torsh_tensor::Tensor;
#[derive(Debug, Clone)]
pub struct STDPConfig {
pub a_plus: f64,
pub a_minus: f64,
pub tau_plus: f64,
pub tau_minus: f64,
pub w_max: f64,
pub w_min: f64,
pub spike_threshold: f64,
pub tau_membrane: f64,
pub v_reset: f64,
}
impl Default for STDPConfig {
fn default() -> Self {
Self {
a_plus: 0.01,
a_minus: 0.01,
tau_plus: 20.0,
tau_minus: 20.0,
w_max: 1.0,
w_min: -1.0,
spike_threshold: 1.0,
tau_membrane: 10.0,
v_reset: 0.0,
}
}
}
#[derive(Debug, Clone)]
struct SpikeState {
last_spike_time: Option<f64>,
membrane_potential: f64,
spike_history: VecDeque<(f64, f64)>,
}
impl Default for SpikeState {
fn default() -> Self {
Self {
last_spike_time: None,
membrane_potential: 0.0,
spike_history: VecDeque::with_capacity(100),
}
}
}
pub struct STDPOptimizer {
lr: f32,
config: STDPConfig,
current_time: f64,
param_groups: Vec<Arc<RwLock<Tensor>>>,
spike_states: HashMap<String, SpikeState>,
eligibility_traces: HashMap<String, Tensor>,
}
impl STDPOptimizer {
pub fn new(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
config: STDPConfig,
) -> OptimizerResult<Self> {
if lr <= 0.0 {
return Err(OptimizerError::InvalidParameter(format!(
"Invalid learning rate: {}",
lr
)));
}
let mut spike_states = HashMap::new();
let mut eligibility_traces = HashMap::new();
for (i, param) in params.iter().enumerate() {
let param_key = format!("param_{}", i);
spike_states.insert(param_key.clone(), SpikeState::default());
let param_read = param.read();
let shape_owned = param_read.shape().dims().to_vec();
drop(param_read);
let trace = torsh_tensor::creation::zeros(&shape_owned)?;
eligibility_traces.insert(param_key, trace);
}
Ok(Self {
lr,
config,
current_time: 0.0,
param_groups: params,
spike_states,
eligibility_traces,
})
}
pub fn with_defaults(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> OptimizerResult<Self> {
Self::new(params, lr, STDPConfig::default())
}
fn detect_spike(&mut self, param_key: &str, gradient: &Tensor) -> OptimizerResult<bool> {
let state = self
.spike_states
.get_mut(param_key)
.expect("spike_states should exist for param_key");
let grad_norm = gradient.norm()?.item()?;
let grad_norm_f64 = grad_norm as f64;
state.membrane_potential =
state.membrane_potential * (1.0 - 1.0 / self.config.tau_membrane) + grad_norm_f64;
if state.membrane_potential > self.config.spike_threshold {
state.last_spike_time = Some(self.current_time);
state
.spike_history
.push_back((self.current_time, grad_norm_f64));
if state.spike_history.len() > 100 {
state.spike_history.pop_front();
}
state.membrane_potential = self.config.v_reset;
Ok(true)
} else {
Ok(false)
}
}
fn compute_stdp_change(&self, pre_time: f64, post_time: f64) -> f64 {
let dt = post_time - pre_time;
if dt > 0.0 {
self.config.a_plus * (-dt / self.config.tau_plus).exp()
} else {
-self.config.a_minus * (dt.abs() / self.config.tau_minus).exp()
}
}
fn update_eligibility_trace(
&mut self,
param_key: &str,
gradient: &Tensor,
) -> OptimizerResult<()> {
let trace = self
.eligibility_traces
.get_mut(param_key)
.expect("eligibility_traces should exist for param_key");
let decay = 0.95; *trace = trace.mul_scalar(decay)?;
*trace = trace.add(gradient)?;
Ok(())
}
}
impl Optimizer for STDPOptimizer {
fn step(&mut self) -> OptimizerResult<()> {
self.current_time += 1.0;
let mut gradients = Vec::new();
for param in self.param_groups.iter() {
let grad_opt = param.read().grad().clone();
gradients.push(grad_opt);
}
for (i, grad_opt) in gradients.into_iter().enumerate() {
let param_key = format!("param_{}", i);
if let Some(grad) = grad_opt {
self.update_eligibility_trace(¶m_key, &grad)?;
let spiked = self.detect_spike(¶m_key, &grad)?;
if spiked {
let state = &self.spike_states[¶m_key];
let mut total_stdp_change = 0.0;
if let Some(current_spike_time) = state.last_spike_time {
for (past_time, _magnitude) in state.spike_history.iter() {
if *past_time != current_spike_time {
let stdp_change =
self.compute_stdp_change(*past_time, current_spike_time);
total_stdp_change += stdp_change;
}
}
}
let trace = self.eligibility_traces[¶m_key].clone();
let update = trace.mul_scalar(self.lr * (1.0 + total_stdp_change as f32))?;
let param = &self.param_groups[i];
let mut param_write = param.write();
let new_param = param_write.sub(&update)?;
let clamped =
new_param.clamp(self.config.w_min as f32, self.config.w_max as f32)?;
*param_write = clamped;
}
}
}
Ok(())
}
fn zero_grad(&mut self) {
for param in &self.param_groups {
param.write().set_grad(None);
}
}
fn get_lr(&self) -> Vec<f32> {
vec![self.lr]
}
fn set_lr(&mut self, lr: f32) {
self.lr = lr;
}
fn add_param_group(
&mut self,
params: Vec<Arc<RwLock<Tensor>>>,
_options: HashMap<String, f32>,
) {
let start_idx = self.param_groups.len();
self.param_groups.extend(params.iter().cloned());
for (i, _param) in params.iter().enumerate() {
let param_key = format!("param_{}", start_idx + i);
self.spike_states
.insert(param_key.clone(), SpikeState::default());
}
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
let mut state = OptimizerState {
optimizer_type: "STDP".to_string(),
version: "1.0".to_string(),
param_groups: vec![],
state: HashMap::new(),
global_state: HashMap::new(),
};
state.global_state.insert("lr".to_string(), self.lr);
state
.global_state
.insert("current_time".to_string(), self.current_time as f32);
state
.global_state
.insert("a_plus".to_string(), self.config.a_plus as f32);
state
.global_state
.insert("a_minus".to_string(), self.config.a_minus as f32);
Ok(state)
}
fn load_state_dict(&mut self, _state: OptimizerState) -> OptimizerResult<()> {
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct EventDrivenConfig {
pub spike_threshold: f64,
pub refractory_period: usize,
pub min_update_interval: usize,
pub adaptive_threshold: bool,
pub threshold_adapt_rate: f64,
}
impl Default for EventDrivenConfig {
fn default() -> Self {
Self {
spike_threshold: 0.1,
refractory_period: 5,
min_update_interval: 1,
adaptive_threshold: true,
threshold_adapt_rate: 0.01,
}
}
}
pub struct EventDrivenOptimizer {
lr: f32,
config: EventDrivenConfig,
param_groups: Vec<Arc<RwLock<Tensor>>>,
steps_since_spike: HashMap<String, usize>,
adaptive_thresholds: HashMap<String, f64>,
momentum_buffers: HashMap<String, Tensor>,
momentum: f32,
}
impl EventDrivenOptimizer {
pub fn new(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
momentum: f32,
config: EventDrivenConfig,
) -> OptimizerResult<Self> {
if lr <= 0.0 {
return Err(OptimizerError::InvalidParameter(format!(
"Invalid learning rate: {}",
lr
)));
}
let mut steps_since_spike = HashMap::new();
let mut adaptive_thresholds = HashMap::new();
let mut momentum_buffers = HashMap::new();
for (i, param) in params.iter().enumerate() {
let param_key = format!("param_{}", i);
steps_since_spike.insert(param_key.clone(), 0);
adaptive_thresholds.insert(param_key.clone(), config.spike_threshold);
let param_read = param.read();
let shape_owned = param_read.shape().dims().to_vec();
drop(param_read);
let buffer = torsh_tensor::creation::zeros(&shape_owned)?;
momentum_buffers.insert(param_key, buffer);
}
Ok(Self {
lr,
config,
param_groups: params,
steps_since_spike,
adaptive_thresholds,
momentum_buffers,
momentum,
})
}
pub fn with_defaults(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> OptimizerResult<Self> {
Self::new(params, lr, 0.9, EventDrivenConfig::default())
}
fn should_spike(&mut self, param_key: &str, gradient: &Tensor) -> OptimizerResult<bool> {
let steps = self
.steps_since_spike
.get(param_key)
.expect("steps_since_spike should exist for param_key");
if *steps < self.config.refractory_period {
return Ok(false);
}
if *steps < self.config.min_update_interval {
return Ok(false);
}
let grad_norm = gradient.norm()?.item()?;
let grad_norm_f64 = grad_norm as f64;
let threshold = self.adaptive_thresholds[param_key];
let should_spike = grad_norm_f64 > threshold;
if self.config.adaptive_threshold {
let new_threshold = if should_spike {
threshold * (1.0 + self.config.threshold_adapt_rate)
} else {
threshold * (1.0 - self.config.threshold_adapt_rate)
};
self.adaptive_thresholds
.insert(param_key.to_string(), new_threshold.max(1e-6));
}
Ok(should_spike)
}
}
impl Optimizer for EventDrivenOptimizer {
fn step(&mut self) -> OptimizerResult<()> {
for (_key, steps) in self.steps_since_spike.iter_mut() {
*steps += 1;
}
let mut gradients = Vec::new();
for param in self.param_groups.iter() {
let grad_opt = param.read().grad().clone();
gradients.push(grad_opt);
}
for (i, grad_opt) in gradients.into_iter().enumerate() {
let param_key = format!("param_{}", i);
if let Some(grad) = grad_opt {
if self.should_spike(¶m_key, &grad)? {
self.steps_since_spike.insert(param_key.clone(), 0);
let buffer = self
.momentum_buffers
.get_mut(¶m_key)
.expect("momentum_buffers should exist for param_key");
*buffer = buffer.mul_scalar(self.momentum)?;
*buffer = buffer.add(&grad)?;
let update = buffer.mul_scalar(self.lr)?;
let param = &self.param_groups[i];
let mut param_write = param.write();
*param_write = param_write.sub(&update)?;
}
}
}
Ok(())
}
fn zero_grad(&mut self) {
for param in &self.param_groups {
param.write().set_grad(None);
}
}
fn get_lr(&self) -> Vec<f32> {
vec![self.lr]
}
fn set_lr(&mut self, lr: f32) {
self.lr = lr;
}
fn add_param_group(
&mut self,
params: Vec<Arc<RwLock<Tensor>>>,
_options: HashMap<String, f32>,
) {
let start_idx = self.param_groups.len();
self.param_groups.extend(params.iter().cloned());
for (i, _param) in params.iter().enumerate() {
let param_key = format!("param_{}", start_idx + i);
self.steps_since_spike.insert(param_key.clone(), 0);
self.adaptive_thresholds
.insert(param_key.clone(), self.config.spike_threshold);
}
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
let mut state = OptimizerState {
optimizer_type: "EventDriven".to_string(),
version: "1.0".to_string(),
param_groups: vec![],
state: HashMap::new(),
global_state: HashMap::new(),
};
state.global_state.insert("lr".to_string(), self.lr);
state
.global_state
.insert("momentum".to_string(), self.momentum);
Ok(state)
}
fn load_state_dict(&mut self, _state: OptimizerState) -> OptimizerResult<()> {
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct TemporalCreditConfig {
pub trace_decay: f64,
pub discount_factor: f64,
pub max_trace_length: usize,
pub use_dopamine_modulation: bool,
pub baseline_dopamine: f64,
}
impl Default for TemporalCreditConfig {
fn default() -> Self {
Self {
trace_decay: 0.95,
discount_factor: 0.99,
max_trace_length: 100,
use_dopamine_modulation: true,
baseline_dopamine: 1.0,
}
}
}
pub struct TemporalCreditOptimizer {
lr: f32,
config: TemporalCreditConfig,
param_groups: Vec<Arc<RwLock<Tensor>>>,
eligibility_traces: HashMap<String, Tensor>,
reward_history: VecDeque<f64>,
dopamine_level: f64,
}
impl TemporalCreditOptimizer {
pub fn new(
params: Vec<Arc<RwLock<Tensor>>>,
lr: f32,
config: TemporalCreditConfig,
) -> OptimizerResult<Self> {
if lr <= 0.0 {
return Err(OptimizerError::InvalidParameter(format!(
"Invalid learning rate: {}",
lr
)));
}
let mut eligibility_traces = HashMap::new();
for (i, param) in params.iter().enumerate() {
let param_key = format!("param_{}", i);
let param_read = param.read();
let shape_owned = param_read.shape().dims().to_vec();
drop(param_read);
let trace = torsh_tensor::creation::zeros(&shape_owned)?;
eligibility_traces.insert(param_key, trace);
}
let max_trace_length = config.max_trace_length;
let baseline_dopamine = config.baseline_dopamine;
Ok(Self {
lr,
config,
param_groups: params,
eligibility_traces,
reward_history: VecDeque::with_capacity(max_trace_length),
dopamine_level: baseline_dopamine,
})
}
pub fn with_defaults(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> OptimizerResult<Self> {
Self::new(params, lr, TemporalCreditConfig::default())
}
fn update_traces(&mut self, gradients: &HashMap<String, Tensor>) -> OptimizerResult<()> {
for (key, grad) in gradients {
let trace = self
.eligibility_traces
.get_mut(key)
.expect("eligibility_traces should exist for key");
let decay_factor = (self.config.trace_decay * self.config.discount_factor) as f32;
*trace = trace.mul_scalar(decay_factor)?;
*trace = trace.add(grad)?;
}
Ok(())
}
pub fn update_dopamine(&mut self, reward: f64) {
self.reward_history.push_back(reward);
if self.reward_history.len() > self.config.max_trace_length {
self.reward_history.pop_front();
}
let avg_reward: f64 =
self.reward_history.iter().sum::<f64>() / self.reward_history.len() as f64;
self.dopamine_level = reward - avg_reward + self.config.baseline_dopamine;
}
pub fn step_with_reward(&mut self, reward: f64) -> OptimizerResult<()> {
self.update_dopamine(reward);
let mut gradients = HashMap::new();
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
let param_read = param.read();
if let Some(grad) = param_read.grad() {
gradients.insert(param_key, grad.clone());
}
}
self.update_traces(&gradients)?;
for (i, param) in self.param_groups.iter().enumerate() {
let param_key = format!("param_{}", i);
let trace = &self.eligibility_traces[¶m_key];
let modulation = if self.config.use_dopamine_modulation {
self.dopamine_level as f32
} else {
1.0
};
let update = trace.mul_scalar(self.lr * modulation)?;
let mut param_write = param.write();
*param_write = param_write.sub(&update)?;
}
Ok(())
}
}
impl Optimizer for TemporalCreditOptimizer {
fn step(&mut self) -> OptimizerResult<()> {
self.step_with_reward(0.0)
}
fn zero_grad(&mut self) {
for param in &self.param_groups {
param.write().set_grad(None);
}
}
fn get_lr(&self) -> Vec<f32> {
vec![self.lr]
}
fn set_lr(&mut self, lr: f32) {
self.lr = lr;
}
fn add_param_group(
&mut self,
params: Vec<Arc<RwLock<Tensor>>>,
_options: HashMap<String, f32>,
) {
let start_idx = self.param_groups.len();
self.param_groups.extend(params.iter().cloned());
for (i, _param) in params.iter().enumerate() {
let param_key = format!("param_{}", start_idx + i);
}
}
fn state_dict(&self) -> OptimizerResult<OptimizerState> {
let mut state = OptimizerState {
optimizer_type: "TemporalCredit".to_string(),
version: "1.0".to_string(),
param_groups: vec![],
state: HashMap::new(),
global_state: HashMap::new(),
};
state.global_state.insert("lr".to_string(), self.lr);
state
.global_state
.insert("dopamine_level".to_string(), self.dopamine_level as f32);
Ok(state)
}
fn load_state_dict(&mut self, _state: OptimizerState) -> OptimizerResult<()> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use torsh_tensor::creation::randn;
#[test]
fn test_stdp_optimizer_creation() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
let optimizer = STDPOptimizer::with_defaults(vec![param], 0.01)?;
assert_eq!(optimizer.param_groups.len(), 1);
assert_eq!(optimizer.spike_states.len(), 1);
Ok(())
}
#[test]
fn test_stdp_config_default() {
let config = STDPConfig::default();
assert_eq!(config.a_plus, 0.01);
assert_eq!(config.a_minus, 0.01);
assert_eq!(config.tau_plus, 20.0);
assert!(config.w_max > config.w_min);
}
#[test]
fn test_event_driven_optimizer_creation() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[5, 5])?));
let optimizer = EventDrivenOptimizer::with_defaults(vec![param], 0.01)?;
assert_eq!(optimizer.param_groups.len(), 1);
assert_eq!(optimizer.steps_since_spike.len(), 1);
Ok(())
}
#[test]
fn test_event_driven_config_default() {
let config = EventDrivenConfig::default();
assert_eq!(config.spike_threshold, 0.1);
assert_eq!(config.refractory_period, 5);
assert!(config.adaptive_threshold);
}
#[test]
fn test_temporal_credit_optimizer_creation() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[3, 3])?));
let optimizer = TemporalCreditOptimizer::with_defaults(vec![param], 0.01)?;
assert_eq!(optimizer.param_groups.len(), 1);
assert_eq!(optimizer.eligibility_traces.len(), 1);
Ok(())
}
#[test]
fn test_temporal_credit_config_default() {
let config = TemporalCreditConfig::default();
assert_eq!(config.trace_decay, 0.95);
assert_eq!(config.discount_factor, 0.99);
assert!(config.use_dopamine_modulation);
}
#[test]
fn test_stdp_step() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
{
let mut p = param.write();
let grad = randn::<f32>(&[2, 2])?;
p.set_grad(Some(grad));
}
let mut optimizer = STDPOptimizer::with_defaults(vec![param.clone()], 0.01)?;
optimizer.step()?;
Ok(())
}
#[test]
fn test_event_driven_step() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
{
let mut p = param.write();
let grad = randn::<f32>(&[2, 2])?;
p.set_grad(Some(grad));
}
let mut optimizer = EventDrivenOptimizer::with_defaults(vec![param.clone()], 0.01)?;
optimizer.step()?;
Ok(())
}
#[test]
fn test_temporal_credit_step_with_reward() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
{
let mut p = param.write();
let grad = randn::<f32>(&[2, 2])?;
p.set_grad(Some(grad));
}
let mut optimizer = TemporalCreditOptimizer::with_defaults(vec![param.clone()], 0.01)?;
optimizer.step_with_reward(1.0)?;
assert!(optimizer.dopamine_level > 0.0);
Ok(())
}
#[test]
fn test_zero_grad() -> OptimizerResult<()> {
let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
{
let mut p = param.write();
let grad = randn::<f32>(&[2, 2])?;
p.set_grad(Some(grad));
}
let mut optimizer = STDPOptimizer::with_defaults(vec![param.clone()], 0.01)?;
optimizer.zero_grad();
let p = param.read();
assert!(p.grad().is_none());
Ok(())
}
}