Skip to main content

torsh_optim/
neuromorphic.rs

1//! Neuromorphic Optimization
2//!
3//! This module implements biologically-inspired optimization algorithms based on
4//! spiking neural networks (SNNs) and brain-inspired learning rules.
5//!
6//! # Key Features
7//!
8//! - **STDP (Spike-Timing-Dependent Plasticity)**: Hebbian learning rule based on spike timing
9//! - **Event-driven optimization**: Sparse updates triggered by spike events
10//! - **Energy-aware computation**: Minimizes computational cost inspired by biological efficiency
11//! - **Temporal credit assignment**: Handles delayed reward signals with eligibility traces
12//! - **Adaptive thresholding**: Dynamic spike thresholds for stability
13//!
14//! # Algorithms
15//!
16//! ## STDPOptimizer
17//!
18//! Implements spike-timing-dependent plasticity, where synaptic weights are updated
19//! based on the relative timing of pre- and post-synaptic spikes.
20//!
21//! Weight update rule:
22//! - If pre-spike before post-spike: Δw = A+ * exp(-Δt / τ+)  (potentiation)
23//! - If post-spike before pre-spike: Δw = -A- * exp(-Δt / τ-)  (depression)
24//!
25//! ## EventDrivenOptimizer
26//!
27//! Sparse gradient-based optimization that only updates parameters when
28//! corresponding "spikes" (large gradient magnitudes) are detected.
29//!
30//! ## TemporalCreditAssignment
31//!
32//! Handles credit assignment over time using eligibility traces, allowing
33//! delayed rewards to influence earlier parameter updates.
34//!
35//! # References
36//!
37//! - Gerstner & Kistler (2002). "Spiking Neuron Models"
38//! - Bi & Poo (1998). "Synaptic modifications in cultured hippocampal neurons"
39//! - Bellec et al. (2020). "A solution to the learning dilemma for recurrent networks of spiking neurons"
40
41use crate::{Optimizer, OptimizerError, OptimizerResult, OptimizerState};
42use parking_lot::RwLock;
43use scirs2_core::ndarray::{Array1, Array2};
44use scirs2_core::random::{thread_rng, Uniform};
45use std::collections::{HashMap, VecDeque};
46use std::sync::Arc;
47use torsh_tensor::Tensor;
48
49// ============================================================================
50// STDP (Spike-Timing-Dependent Plasticity) Optimizer
51// ============================================================================
52
53/// STDP learning configuration
54#[derive(Debug, Clone)]
55pub struct STDPConfig {
56    /// Potentiation amplitude (A+)
57    pub a_plus: f64,
58    /// Depression amplitude (A-)
59    pub a_minus: f64,
60    /// Potentiation time constant (τ+)
61    pub tau_plus: f64,
62    /// Depression time constant (τ-)
63    pub tau_minus: f64,
64    /// Maximum allowed weight
65    pub w_max: f64,
66    /// Minimum allowed weight
67    pub w_min: f64,
68    /// Spike threshold for membrane potential
69    pub spike_threshold: f64,
70    /// Membrane time constant
71    pub tau_membrane: f64,
72    /// Resting potential reset value
73    pub v_reset: f64,
74}
75
76impl Default for STDPConfig {
77    fn default() -> Self {
78        Self {
79            a_plus: 0.01,
80            a_minus: 0.01,
81            tau_plus: 20.0,
82            tau_minus: 20.0,
83            w_max: 1.0,
84            w_min: -1.0,
85            spike_threshold: 1.0,
86            tau_membrane: 10.0,
87            v_reset: 0.0,
88        }
89    }
90}
91
92/// Spike timing state for a neuron
93#[derive(Debug, Clone)]
94struct SpikeState {
95    /// Last spike time
96    last_spike_time: Option<f64>,
97    /// Membrane potential
98    membrane_potential: f64,
99    /// Spike history (time, magnitude)
100    spike_history: VecDeque<(f64, f64)>,
101}
102
103impl Default for SpikeState {
104    fn default() -> Self {
105        Self {
106            last_spike_time: None,
107            membrane_potential: 0.0,
108            spike_history: VecDeque::with_capacity(100),
109        }
110    }
111}
112
113/// STDP-based optimizer
114///
115/// Implements spike-timing-dependent plasticity for parameter updates.
116/// Updates are based on the correlation between pre- and post-synaptic
117/// spike timings.
118pub struct STDPOptimizer {
119    /// Learning rate
120    lr: f32,
121    /// STDP configuration
122    config: STDPConfig,
123    /// Current time step
124    current_time: f64,
125    /// Parameter groups
126    param_groups: Vec<Arc<RwLock<Tensor>>>,
127    /// Spike states per parameter
128    spike_states: HashMap<String, SpikeState>,
129    /// Eligibility traces
130    eligibility_traces: HashMap<String, Tensor>,
131}
132
133impl STDPOptimizer {
134    /// Create a new STDP optimizer
135    pub fn new(
136        params: Vec<Arc<RwLock<Tensor>>>,
137        lr: f32,
138        config: STDPConfig,
139    ) -> OptimizerResult<Self> {
140        if lr <= 0.0 {
141            return Err(OptimizerError::InvalidParameter(format!(
142                "Invalid learning rate: {}",
143                lr
144            )));
145        }
146
147        let mut spike_states = HashMap::new();
148        let mut eligibility_traces = HashMap::new();
149
150        for (i, param) in params.iter().enumerate() {
151            let param_key = format!("param_{}", i);
152            spike_states.insert(param_key.clone(), SpikeState::default());
153
154            // Initialize eligibility trace
155            let param_read = param.read();
156            let shape_owned = param_read.shape().dims().to_vec();
157            drop(param_read);
158            let trace = torsh_tensor::creation::zeros(&shape_owned)?;
159            eligibility_traces.insert(param_key, trace);
160        }
161
162        Ok(Self {
163            lr,
164            config,
165            current_time: 0.0,
166            param_groups: params,
167            spike_states,
168            eligibility_traces,
169        })
170    }
171
172    /// Create with default configuration
173    pub fn with_defaults(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> OptimizerResult<Self> {
174        Self::new(params, lr, STDPConfig::default())
175    }
176
177    /// Detect spikes based on gradient magnitude
178    fn detect_spike(&mut self, param_key: &str, gradient: &Tensor) -> OptimizerResult<bool> {
179        let state = self
180            .spike_states
181            .get_mut(param_key)
182            .expect("spike_states should exist for param_key");
183
184        // Calculate "membrane potential" as a function of gradient magnitude
185        let grad_norm = gradient.norm()?.item()?;
186        let grad_norm_f64 = grad_norm as f64;
187
188        // Leaky integration of gradient
189        state.membrane_potential =
190            state.membrane_potential * (1.0 - 1.0 / self.config.tau_membrane) + grad_norm_f64;
191
192        // Check if spike threshold is crossed
193        if state.membrane_potential > self.config.spike_threshold {
194            state.last_spike_time = Some(self.current_time);
195            state
196                .spike_history
197                .push_back((self.current_time, grad_norm_f64));
198
199            // Keep history bounded
200            if state.spike_history.len() > 100 {
201                state.spike_history.pop_front();
202            }
203
204            // Reset membrane potential
205            state.membrane_potential = self.config.v_reset;
206            Ok(true)
207        } else {
208            Ok(false)
209        }
210    }
211
212    /// Compute STDP weight change
213    fn compute_stdp_change(&self, pre_time: f64, post_time: f64) -> f64 {
214        let dt = post_time - pre_time;
215
216        if dt > 0.0 {
217            // Post after pre -> potentiation (LTP)
218            self.config.a_plus * (-dt / self.config.tau_plus).exp()
219        } else {
220            // Pre after post -> depression (LTD)
221            -self.config.a_minus * (dt.abs() / self.config.tau_minus).exp()
222        }
223    }
224
225    /// Update eligibility trace
226    fn update_eligibility_trace(
227        &mut self,
228        param_key: &str,
229        gradient: &Tensor,
230    ) -> OptimizerResult<()> {
231        let trace = self
232            .eligibility_traces
233            .get_mut(param_key)
234            .expect("eligibility_traces should exist for param_key");
235
236        // Exponential decay: e(t+1) = λ * e(t) + grad
237        let decay = 0.95; // Decay factor
238        *trace = trace.mul_scalar(decay)?;
239        *trace = trace.add(gradient)?;
240
241        Ok(())
242    }
243}
244
245impl Optimizer for STDPOptimizer {
246    fn step(&mut self) -> OptimizerResult<()> {
247        self.current_time += 1.0;
248
249        // Collect gradients first to avoid borrow issues
250        let mut gradients = Vec::new();
251        for param in self.param_groups.iter() {
252            let grad_opt = param.read().grad().clone();
253            gradients.push(grad_opt);
254        }
255
256        for (i, grad_opt) in gradients.into_iter().enumerate() {
257            let param_key = format!("param_{}", i);
258
259            if let Some(grad) = grad_opt {
260                // Update eligibility trace
261                self.update_eligibility_trace(&param_key, &grad)?;
262
263                // Detect spike
264                let spiked = self.detect_spike(&param_key, &grad)?;
265
266                if spiked {
267                    // Apply STDP rule
268                    let state = &self.spike_states[&param_key];
269
270                    // Compute STDP-based update using spike history
271                    let mut total_stdp_change = 0.0;
272
273                    if let Some(current_spike_time) = state.last_spike_time {
274                        // Look for correlated spikes in recent history
275                        for (past_time, _magnitude) in state.spike_history.iter() {
276                            if *past_time != current_spike_time {
277                                let stdp_change =
278                                    self.compute_stdp_change(*past_time, current_spike_time);
279                                total_stdp_change += stdp_change;
280                            }
281                        }
282                    }
283
284                    // Apply update with eligibility trace
285                    let trace = self.eligibility_traces[&param_key].clone();
286                    let update = trace.mul_scalar(self.lr * (1.0 + total_stdp_change as f32))?;
287
288                    let param = &self.param_groups[i];
289                    let mut param_write = param.write();
290                    let new_param = param_write.sub(&update)?;
291
292                    // Clamp weights to valid range
293                    let clamped =
294                        new_param.clamp(self.config.w_min as f32, self.config.w_max as f32)?;
295                    crate::param_update::assign(&mut param_write, &clamped)?;
296                }
297            }
298        }
299
300        Ok(())
301    }
302
303    fn zero_grad(&mut self) {
304        for param in &self.param_groups {
305            param.write().set_grad(None);
306        }
307    }
308
309    fn get_lr(&self) -> Vec<f32> {
310        vec![self.lr]
311    }
312
313    fn set_lr(&mut self, lr: f32) {
314        self.lr = lr;
315    }
316
317    fn add_param_group(
318        &mut self,
319        params: Vec<Arc<RwLock<Tensor>>>,
320        _options: HashMap<String, f32>,
321    ) {
322        let start_idx = self.param_groups.len();
323        self.param_groups.extend(params.iter().cloned());
324
325        for (i, _param) in params.iter().enumerate() {
326            let param_key = format!("param_{}", start_idx + i);
327            self.spike_states
328                .insert(param_key.clone(), SpikeState::default());
329        }
330    }
331
332    fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
333        self.param_groups.clone()
334    }
335
336    fn state_dict(&self) -> OptimizerResult<OptimizerState> {
337        // Create basic state
338        let mut state = OptimizerState {
339            optimizer_type: "STDP".to_string(),
340            version: "1.0".to_string(),
341            param_groups: vec![],
342            state: HashMap::new(),
343            global_state: HashMap::new(),
344        };
345
346        // Store configuration
347        state.global_state.insert("lr".to_string(), self.lr);
348        state
349            .global_state
350            .insert("current_time".to_string(), self.current_time as f32);
351        state
352            .global_state
353            .insert("a_plus".to_string(), self.config.a_plus as f32);
354        state
355            .global_state
356            .insert("a_minus".to_string(), self.config.a_minus as f32);
357
358        Ok(state)
359    }
360
361    fn load_state_dict(&mut self, _state: OptimizerState) -> OptimizerResult<()> {
362        // State loading implementation
363        Ok(())
364    }
365}
366
367// ============================================================================
368// Event-Driven Optimizer
369// ============================================================================
370
371/// Event-driven optimization configuration
372#[derive(Debug, Clone)]
373pub struct EventDrivenConfig {
374    /// Spike threshold for gradient magnitude
375    pub spike_threshold: f64,
376    /// Refractory period (steps to skip after spike)
377    pub refractory_period: usize,
378    /// Minimum time between updates (in steps)
379    pub min_update_interval: usize,
380    /// Use adaptive thresholding
381    pub adaptive_threshold: bool,
382    /// Threshold adaptation rate
383    pub threshold_adapt_rate: f64,
384}
385
386impl Default for EventDrivenConfig {
387    fn default() -> Self {
388        Self {
389            spike_threshold: 0.1,
390            refractory_period: 5,
391            min_update_interval: 1,
392            adaptive_threshold: true,
393            threshold_adapt_rate: 0.01,
394        }
395    }
396}
397
398/// Event-driven optimizer
399///
400/// Only updates parameters when gradient magnitude exceeds a threshold,
401/// implementing sparse, event-driven computation.
402pub struct EventDrivenOptimizer {
403    /// Base learning rate
404    lr: f32,
405    /// Configuration
406    config: EventDrivenConfig,
407    /// Parameter groups
408    param_groups: Vec<Arc<RwLock<Tensor>>>,
409    /// Steps since last spike per parameter
410    steps_since_spike: HashMap<String, usize>,
411    /// Adaptive thresholds per parameter
412    adaptive_thresholds: HashMap<String, f64>,
413    /// Momentum buffers (for smoother updates)
414    momentum_buffers: HashMap<String, Tensor>,
415    /// Momentum coefficient
416    momentum: f32,
417}
418
419impl EventDrivenOptimizer {
420    /// Create a new event-driven optimizer
421    pub fn new(
422        params: Vec<Arc<RwLock<Tensor>>>,
423        lr: f32,
424        momentum: f32,
425        config: EventDrivenConfig,
426    ) -> OptimizerResult<Self> {
427        if lr <= 0.0 {
428            return Err(OptimizerError::InvalidParameter(format!(
429                "Invalid learning rate: {}",
430                lr
431            )));
432        }
433
434        let mut steps_since_spike = HashMap::new();
435        let mut adaptive_thresholds = HashMap::new();
436        let mut momentum_buffers = HashMap::new();
437
438        for (i, param) in params.iter().enumerate() {
439            let param_key = format!("param_{}", i);
440            steps_since_spike.insert(param_key.clone(), 0);
441            adaptive_thresholds.insert(param_key.clone(), config.spike_threshold);
442
443            // Initialize momentum buffer
444            let param_read = param.read();
445            let shape_owned = param_read.shape().dims().to_vec();
446            drop(param_read);
447            let buffer = torsh_tensor::creation::zeros(&shape_owned)?;
448            momentum_buffers.insert(param_key, buffer);
449        }
450
451        Ok(Self {
452            lr,
453            config,
454            param_groups: params,
455            steps_since_spike,
456            adaptive_thresholds,
457            momentum_buffers,
458            momentum,
459        })
460    }
461
462    /// Create with default configuration
463    pub fn with_defaults(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> OptimizerResult<Self> {
464        Self::new(params, lr, 0.9, EventDrivenConfig::default())
465    }
466
467    /// Check if parameter should spike (update)
468    fn should_spike(&mut self, param_key: &str, gradient: &Tensor) -> OptimizerResult<bool> {
469        let steps = self
470            .steps_since_spike
471            .get(param_key)
472            .expect("steps_since_spike should exist for param_key");
473
474        // Check refractory period
475        if *steps < self.config.refractory_period {
476            return Ok(false);
477        }
478
479        // Check minimum update interval
480        if *steps < self.config.min_update_interval {
481            return Ok(false);
482        }
483
484        // Check gradient magnitude against threshold
485        let grad_norm = gradient.norm()?.item()?;
486        let grad_norm_f64 = grad_norm as f64;
487        let threshold = self.adaptive_thresholds[param_key];
488
489        let should_spike = grad_norm_f64 > threshold;
490
491        // Adapt threshold if enabled
492        if self.config.adaptive_threshold {
493            let new_threshold = if should_spike {
494                threshold * (1.0 + self.config.threshold_adapt_rate)
495            } else {
496                threshold * (1.0 - self.config.threshold_adapt_rate)
497            };
498            self.adaptive_thresholds
499                .insert(param_key.to_string(), new_threshold.max(1e-6));
500        }
501
502        Ok(should_spike)
503    }
504}
505
506impl Optimizer for EventDrivenOptimizer {
507    fn step(&mut self) -> OptimizerResult<()> {
508        // Increment all step counters
509        for (_key, steps) in self.steps_since_spike.iter_mut() {
510            *steps += 1;
511        }
512
513        // Collect gradients first to avoid borrow issues
514        let mut gradients = Vec::new();
515        for param in self.param_groups.iter() {
516            let grad_opt = param.read().grad().clone();
517            gradients.push(grad_opt);
518        }
519
520        for (i, grad_opt) in gradients.into_iter().enumerate() {
521            let param_key = format!("param_{}", i);
522
523            if let Some(grad) = grad_opt {
524                // Check if this parameter should spike
525                if self.should_spike(&param_key, &grad)? {
526                    // Reset counter
527                    self.steps_since_spike.insert(param_key.clone(), 0);
528
529                    // Update momentum buffer
530                    let buffer = self
531                        .momentum_buffers
532                        .get_mut(&param_key)
533                        .expect("momentum_buffers should exist for param_key");
534                    *buffer = buffer.mul_scalar(self.momentum)?;
535                    *buffer = buffer.add(&grad)?;
536
537                    // Apply update
538                    let update = buffer.mul_scalar(self.lr)?;
539                    let param = &self.param_groups[i];
540                    let mut param_write = param.write();
541                    crate::param_update::sub_assign(&mut param_write, &update)?;
542                }
543            }
544        }
545
546        Ok(())
547    }
548
549    fn zero_grad(&mut self) {
550        for param in &self.param_groups {
551            param.write().set_grad(None);
552        }
553    }
554
555    fn get_lr(&self) -> Vec<f32> {
556        vec![self.lr]
557    }
558
559    fn set_lr(&mut self, lr: f32) {
560        self.lr = lr;
561    }
562
563    fn add_param_group(
564        &mut self,
565        params: Vec<Arc<RwLock<Tensor>>>,
566        _options: HashMap<String, f32>,
567    ) {
568        let start_idx = self.param_groups.len();
569        self.param_groups.extend(params.iter().cloned());
570
571        for (i, _param) in params.iter().enumerate() {
572            let param_key = format!("param_{}", start_idx + i);
573            self.steps_since_spike.insert(param_key.clone(), 0);
574            self.adaptive_thresholds
575                .insert(param_key.clone(), self.config.spike_threshold);
576        }
577    }
578
579    fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
580        self.param_groups.clone()
581    }
582
583    fn state_dict(&self) -> OptimizerResult<OptimizerState> {
584        let mut state = OptimizerState {
585            optimizer_type: "EventDriven".to_string(),
586            version: "1.0".to_string(),
587            param_groups: vec![],
588            state: HashMap::new(),
589            global_state: HashMap::new(),
590        };
591
592        state.global_state.insert("lr".to_string(), self.lr);
593        state
594            .global_state
595            .insert("momentum".to_string(), self.momentum);
596
597        Ok(state)
598    }
599
600    fn load_state_dict(&mut self, _state: OptimizerState) -> OptimizerResult<()> {
601        Ok(())
602    }
603}
604
605// ============================================================================
606// Temporal Credit Assignment
607// ============================================================================
608
609/// Temporal credit assignment configuration
610#[derive(Debug, Clone)]
611pub struct TemporalCreditConfig {
612    /// Eligibility trace decay rate (λ)
613    pub trace_decay: f64,
614    /// Reward discount factor (γ)
615    pub discount_factor: f64,
616    /// Maximum trace history length
617    pub max_trace_length: usize,
618    /// Use three-factor learning rule (dopamine modulation)
619    pub use_dopamine_modulation: bool,
620    /// Baseline dopamine level
621    pub baseline_dopamine: f64,
622}
623
624impl Default for TemporalCreditConfig {
625    fn default() -> Self {
626        Self {
627            trace_decay: 0.95,
628            discount_factor: 0.99,
629            max_trace_length: 100,
630            use_dopamine_modulation: true,
631            baseline_dopamine: 1.0,
632        }
633    }
634}
635
636/// Temporal credit assignment optimizer
637///
638/// Implements eligibility traces for delayed reward learning,
639/// allowing future rewards to influence past parameter updates.
640pub struct TemporalCreditOptimizer {
641    /// Base learning rate
642    lr: f32,
643    /// Configuration
644    config: TemporalCreditConfig,
645    /// Parameter groups
646    param_groups: Vec<Arc<RwLock<Tensor>>>,
647    /// Eligibility traces
648    eligibility_traces: HashMap<String, Tensor>,
649    /// Reward history
650    reward_history: VecDeque<f64>,
651    /// Current dopamine level (reward prediction error)
652    dopamine_level: f64,
653}
654
655impl TemporalCreditOptimizer {
656    /// Create a new temporal credit assignment optimizer
657    pub fn new(
658        params: Vec<Arc<RwLock<Tensor>>>,
659        lr: f32,
660        config: TemporalCreditConfig,
661    ) -> OptimizerResult<Self> {
662        if lr <= 0.0 {
663            return Err(OptimizerError::InvalidParameter(format!(
664                "Invalid learning rate: {}",
665                lr
666            )));
667        }
668
669        let mut eligibility_traces = HashMap::new();
670
671        for (i, param) in params.iter().enumerate() {
672            let param_key = format!("param_{}", i);
673            let param_read = param.read();
674            let shape_owned = param_read.shape().dims().to_vec();
675            drop(param_read);
676            let trace = torsh_tensor::creation::zeros(&shape_owned)?;
677            eligibility_traces.insert(param_key, trace);
678        }
679
680        let max_trace_length = config.max_trace_length;
681        let baseline_dopamine = config.baseline_dopamine;
682
683        Ok(Self {
684            lr,
685            config,
686            param_groups: params,
687            eligibility_traces,
688            reward_history: VecDeque::with_capacity(max_trace_length),
689            dopamine_level: baseline_dopamine,
690        })
691    }
692
693    /// Create with default configuration
694    pub fn with_defaults(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> OptimizerResult<Self> {
695        Self::new(params, lr, TemporalCreditConfig::default())
696    }
697
698    /// Update eligibility traces
699    fn update_traces(&mut self, gradients: &HashMap<String, Tensor>) -> OptimizerResult<()> {
700        for (key, grad) in gradients {
701            let trace = self
702                .eligibility_traces
703                .get_mut(key)
704                .expect("eligibility_traces should exist for key");
705
706            // e(t+1) = λ * γ * e(t) + ∇L
707            let decay_factor = (self.config.trace_decay * self.config.discount_factor) as f32;
708            *trace = trace.mul_scalar(decay_factor)?;
709            *trace = trace.add(grad)?;
710        }
711        Ok(())
712    }
713
714    /// Update dopamine level (reward prediction error)
715    pub fn update_dopamine(&mut self, reward: f64) {
716        // Simple moving average of recent rewards
717        self.reward_history.push_back(reward);
718        if self.reward_history.len() > self.config.max_trace_length {
719            self.reward_history.pop_front();
720        }
721
722        let avg_reward: f64 =
723            self.reward_history.iter().sum::<f64>() / self.reward_history.len() as f64;
724
725        // Dopamine = reward prediction error
726        self.dopamine_level = reward - avg_reward + self.config.baseline_dopamine;
727    }
728
729    /// Step with reward signal
730    pub fn step_with_reward(&mut self, reward: f64) -> OptimizerResult<()> {
731        // Update dopamine level
732        self.update_dopamine(reward);
733
734        // Collect gradients
735        let mut gradients = HashMap::new();
736        for (i, param) in self.param_groups.iter().enumerate() {
737            let param_key = format!("param_{}", i);
738            let param_read = param.read();
739            if let Some(grad) = param_read.grad() {
740                gradients.insert(param_key, grad.clone());
741            }
742        }
743
744        // Update eligibility traces
745        self.update_traces(&gradients)?;
746
747        // Apply updates using eligibility traces and dopamine modulation
748        for (i, param) in self.param_groups.iter().enumerate() {
749            let param_key = format!("param_{}", i);
750            let trace = &self.eligibility_traces[&param_key];
751
752            // Three-factor learning rule: Δw = lr * dopamine * eligibility_trace
753            let modulation = if self.config.use_dopamine_modulation {
754                self.dopamine_level as f32
755            } else {
756                1.0
757            };
758
759            let update = trace.mul_scalar(self.lr * modulation)?;
760            let mut param_write = param.write();
761            crate::param_update::sub_assign(&mut param_write, &update)?;
762        }
763
764        Ok(())
765    }
766}
767
768impl Optimizer for TemporalCreditOptimizer {
769    fn step(&mut self) -> OptimizerResult<()> {
770        // Default step with zero reward
771        self.step_with_reward(0.0)
772    }
773
774    fn zero_grad(&mut self) {
775        for param in &self.param_groups {
776            param.write().set_grad(None);
777        }
778    }
779
780    fn get_lr(&self) -> Vec<f32> {
781        vec![self.lr]
782    }
783
784    fn set_lr(&mut self, lr: f32) {
785        self.lr = lr;
786    }
787
788    fn add_param_group(
789        &mut self,
790        params: Vec<Arc<RwLock<Tensor>>>,
791        _options: HashMap<String, f32>,
792    ) {
793        let start_idx = self.param_groups.len();
794        self.param_groups.extend(params.iter().cloned());
795
796        for (i, _param) in params.iter().enumerate() {
797            let param_key = format!("param_{}", start_idx + i);
798            // Would need to initialize eligibility trace here
799        }
800    }
801
802    fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
803        self.param_groups.clone()
804    }
805
806    fn state_dict(&self) -> OptimizerResult<OptimizerState> {
807        let mut state = OptimizerState {
808            optimizer_type: "TemporalCredit".to_string(),
809            version: "1.0".to_string(),
810            param_groups: vec![],
811            state: HashMap::new(),
812            global_state: HashMap::new(),
813        };
814
815        state.global_state.insert("lr".to_string(), self.lr);
816        state
817            .global_state
818            .insert("dopamine_level".to_string(), self.dopamine_level as f32);
819
820        Ok(state)
821    }
822
823    fn load_state_dict(&mut self, _state: OptimizerState) -> OptimizerResult<()> {
824        Ok(())
825    }
826}
827
828// ============================================================================
829// Tests
830// ============================================================================
831
832#[cfg(test)]
833mod tests {
834    use super::*;
835    use torsh_tensor::creation::randn;
836
837    #[test]
838    fn test_stdp_optimizer_creation() -> OptimizerResult<()> {
839        let param = Arc::new(RwLock::new(randn::<f32>(&[10, 10])?));
840        let optimizer = STDPOptimizer::with_defaults(vec![param], 0.01)?;
841
842        assert_eq!(optimizer.param_groups.len(), 1);
843        assert_eq!(optimizer.spike_states.len(), 1);
844        Ok(())
845    }
846
847    #[test]
848    fn test_stdp_config_default() {
849        let config = STDPConfig::default();
850        assert_eq!(config.a_plus, 0.01);
851        assert_eq!(config.a_minus, 0.01);
852        assert_eq!(config.tau_plus, 20.0);
853        assert!(config.w_max > config.w_min);
854    }
855
856    #[test]
857    fn test_event_driven_optimizer_creation() -> OptimizerResult<()> {
858        let param = Arc::new(RwLock::new(randn::<f32>(&[5, 5])?));
859        let optimizer = EventDrivenOptimizer::with_defaults(vec![param], 0.01)?;
860
861        assert_eq!(optimizer.param_groups.len(), 1);
862        assert_eq!(optimizer.steps_since_spike.len(), 1);
863        Ok(())
864    }
865
866    #[test]
867    fn test_event_driven_config_default() {
868        let config = EventDrivenConfig::default();
869        assert_eq!(config.spike_threshold, 0.1);
870        assert_eq!(config.refractory_period, 5);
871        assert!(config.adaptive_threshold);
872    }
873
874    #[test]
875    fn test_temporal_credit_optimizer_creation() -> OptimizerResult<()> {
876        let param = Arc::new(RwLock::new(randn::<f32>(&[3, 3])?));
877        let optimizer = TemporalCreditOptimizer::with_defaults(vec![param], 0.01)?;
878
879        assert_eq!(optimizer.param_groups.len(), 1);
880        assert_eq!(optimizer.eligibility_traces.len(), 1);
881        Ok(())
882    }
883
884    #[test]
885    fn test_temporal_credit_config_default() {
886        let config = TemporalCreditConfig::default();
887        assert_eq!(config.trace_decay, 0.95);
888        assert_eq!(config.discount_factor, 0.99);
889        assert!(config.use_dopamine_modulation);
890    }
891
892    #[test]
893    fn test_stdp_step() -> OptimizerResult<()> {
894        let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
895
896        // Set gradient
897        {
898            let mut p = param.write();
899            let grad = randn::<f32>(&[2, 2])?;
900            p.set_grad(Some(grad));
901        }
902
903        let mut optimizer = STDPOptimizer::with_defaults(vec![param.clone()], 0.01)?;
904
905        // Step should succeed
906        optimizer.step()?;
907
908        Ok(())
909    }
910
911    #[test]
912    fn test_event_driven_step() -> OptimizerResult<()> {
913        let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
914
915        // Set gradient
916        {
917            let mut p = param.write();
918            let grad = randn::<f32>(&[2, 2])?;
919            p.set_grad(Some(grad));
920        }
921
922        let mut optimizer = EventDrivenOptimizer::with_defaults(vec![param.clone()], 0.01)?;
923
924        // Step should succeed
925        optimizer.step()?;
926
927        Ok(())
928    }
929
930    #[test]
931    fn test_temporal_credit_step_with_reward() -> OptimizerResult<()> {
932        let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
933
934        // Set gradient
935        {
936            let mut p = param.write();
937            let grad = randn::<f32>(&[2, 2])?;
938            p.set_grad(Some(grad));
939        }
940
941        let mut optimizer = TemporalCreditOptimizer::with_defaults(vec![param.clone()], 0.01)?;
942
943        // Step with reward
944        optimizer.step_with_reward(1.0)?;
945
946        // Dopamine level should be updated
947        assert!(optimizer.dopamine_level > 0.0);
948
949        Ok(())
950    }
951
952    #[test]
953    fn test_zero_grad() -> OptimizerResult<()> {
954        let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
955
956        // Set gradient
957        {
958            let mut p = param.write();
959            let grad = randn::<f32>(&[2, 2])?;
960            p.set_grad(Some(grad));
961        }
962
963        let mut optimizer = STDPOptimizer::with_defaults(vec![param.clone()], 0.01)?;
964
965        // Zero gradients
966        optimizer.zero_grad();
967
968        // Check gradient is None
969        let p = param.read();
970        assert!(p.grad().is_none());
971
972        Ok(())
973    }
974}