1use 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#[derive(Debug, Clone)]
55pub struct STDPConfig {
56 pub a_plus: f64,
58 pub a_minus: f64,
60 pub tau_plus: f64,
62 pub tau_minus: f64,
64 pub w_max: f64,
66 pub w_min: f64,
68 pub spike_threshold: f64,
70 pub tau_membrane: f64,
72 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#[derive(Debug, Clone)]
94struct SpikeState {
95 last_spike_time: Option<f64>,
97 membrane_potential: f64,
99 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
113pub struct STDPOptimizer {
119 lr: f32,
121 config: STDPConfig,
123 current_time: f64,
125 param_groups: Vec<Arc<RwLock<Tensor>>>,
127 spike_states: HashMap<String, SpikeState>,
129 eligibility_traces: HashMap<String, Tensor>,
131}
132
133impl STDPOptimizer {
134 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 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 pub fn with_defaults(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> OptimizerResult<Self> {
174 Self::new(params, lr, STDPConfig::default())
175 }
176
177 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 let grad_norm = gradient.norm()?.item()?;
186 let grad_norm_f64 = grad_norm as f64;
187
188 state.membrane_potential =
190 state.membrane_potential * (1.0 - 1.0 / self.config.tau_membrane) + grad_norm_f64;
191
192 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 if state.spike_history.len() > 100 {
201 state.spike_history.pop_front();
202 }
203
204 state.membrane_potential = self.config.v_reset;
206 Ok(true)
207 } else {
208 Ok(false)
209 }
210 }
211
212 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 self.config.a_plus * (-dt / self.config.tau_plus).exp()
219 } else {
220 -self.config.a_minus * (dt.abs() / self.config.tau_minus).exp()
222 }
223 }
224
225 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 let decay = 0.95; *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 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 self.update_eligibility_trace(¶m_key, &grad)?;
262
263 let spiked = self.detect_spike(¶m_key, &grad)?;
265
266 if spiked {
267 let state = &self.spike_states[¶m_key];
269
270 let mut total_stdp_change = 0.0;
272
273 if let Some(current_spike_time) = state.last_spike_time {
274 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 let trace = self.eligibility_traces[¶m_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 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 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 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 Ok(())
364 }
365}
366
367#[derive(Debug, Clone)]
373pub struct EventDrivenConfig {
374 pub spike_threshold: f64,
376 pub refractory_period: usize,
378 pub min_update_interval: usize,
380 pub adaptive_threshold: bool,
382 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
398pub struct EventDrivenOptimizer {
403 lr: f32,
405 config: EventDrivenConfig,
407 param_groups: Vec<Arc<RwLock<Tensor>>>,
409 steps_since_spike: HashMap<String, usize>,
411 adaptive_thresholds: HashMap<String, f64>,
413 momentum_buffers: HashMap<String, Tensor>,
415 momentum: f32,
417}
418
419impl EventDrivenOptimizer {
420 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 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 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 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 if *steps < self.config.refractory_period {
476 return Ok(false);
477 }
478
479 if *steps < self.config.min_update_interval {
481 return Ok(false);
482 }
483
484 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 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 for (_key, steps) in self.steps_since_spike.iter_mut() {
510 *steps += 1;
511 }
512
513 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 if self.should_spike(¶m_key, &grad)? {
526 self.steps_since_spike.insert(param_key.clone(), 0);
528
529 let buffer = self
531 .momentum_buffers
532 .get_mut(¶m_key)
533 .expect("momentum_buffers should exist for param_key");
534 *buffer = buffer.mul_scalar(self.momentum)?;
535 *buffer = buffer.add(&grad)?;
536
537 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#[derive(Debug, Clone)]
611pub struct TemporalCreditConfig {
612 pub trace_decay: f64,
614 pub discount_factor: f64,
616 pub max_trace_length: usize,
618 pub use_dopamine_modulation: bool,
620 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
636pub struct TemporalCreditOptimizer {
641 lr: f32,
643 config: TemporalCreditConfig,
645 param_groups: Vec<Arc<RwLock<Tensor>>>,
647 eligibility_traces: HashMap<String, Tensor>,
649 reward_history: VecDeque<f64>,
651 dopamine_level: f64,
653}
654
655impl TemporalCreditOptimizer {
656 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 pub fn with_defaults(params: Vec<Arc<RwLock<Tensor>>>, lr: f32) -> OptimizerResult<Self> {
695 Self::new(params, lr, TemporalCreditConfig::default())
696 }
697
698 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 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 pub fn update_dopamine(&mut self, reward: f64) {
716 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 self.dopamine_level = reward - avg_reward + self.config.baseline_dopamine;
727 }
728
729 pub fn step_with_reward(&mut self, reward: f64) -> OptimizerResult<()> {
731 self.update_dopamine(reward);
733
734 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 self.update_traces(&gradients)?;
746
747 for (i, param) in self.param_groups.iter().enumerate() {
749 let param_key = format!("param_{}", i);
750 let trace = &self.eligibility_traces[¶m_key];
751
752 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 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 }
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#[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 {
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 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 {
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 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 {
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 optimizer.step_with_reward(1.0)?;
945
946 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 {
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 optimizer.zero_grad();
967
968 let p = param.read();
970 assert!(p.grad().is_none());
971
972 Ok(())
973 }
974}