optirs_core/optimizers/
adagrad.rs1use scirs2_core::ndarray::{Array, Dimension, IxDyn, ScalarOperand, Zip};
4use scirs2_core::numeric::Float;
5use std::fmt::Debug;
6
7use crate::error::{OptimError, Result};
8use crate::optimizers::Optimizer;
9
10#[derive(Debug, Clone)]
39pub struct Adagrad<A: Float + ScalarOperand + Debug> {
40 learning_rate: A,
42 epsilon: A,
44 weight_decay: A,
46 sum_squared_grads: Option<Vec<Array<A, IxDyn>>>,
48}
49
50impl<A: Float + ScalarOperand + Debug + Send + Sync> Adagrad<A> {
51 pub fn new(learning_rate: A) -> Self {
57 Self {
58 learning_rate,
59 epsilon: A::from(1e-10).expect("Adagrad: default epsilon (1e-10) must fit in A"),
60 weight_decay: A::zero(),
61 sum_squared_grads: None,
62 }
63 }
64
65 pub fn new_with_config(learning_rate: A, epsilon: A, weight_decay: A) -> Self {
73 Self {
74 learning_rate,
75 epsilon,
76 weight_decay,
77 sum_squared_grads: None,
78 }
79 }
80
81 pub fn set_epsilon(&mut self, epsilon: A) -> &mut Self {
83 self.epsilon = epsilon;
84 self
85 }
86
87 pub fn get_epsilon(&self) -> A {
89 self.epsilon
90 }
91
92 pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
94 self.weight_decay = weight_decay;
95 self
96 }
97
98 pub fn get_weight_decay(&self) -> A {
100 self.weight_decay
101 }
102
103 pub fn reset(&mut self) {
105 self.sum_squared_grads = None;
106 }
107
108 fn ensure_state(&mut self, index: usize, dim: &IxDyn) {
110 let accumulators = self.sum_squared_grads.get_or_insert_with(Vec::new);
111 while accumulators.len() <= index {
112 accumulators.push(Array::zeros(dim.clone()));
113 }
114 if accumulators[index].raw_dim() != *dim {
115 accumulators[index] = Array::zeros(dim.clone());
116 }
117 }
118
119 pub fn step_indexed<D: Dimension>(
124 &mut self,
125 index: usize,
126 params: &Array<A, D>,
127 gradients: &Array<A, D>,
128 ) -> Result<Array<A, D>> {
129 if params.shape() != gradients.shape() {
130 return Err(OptimError::DimensionMismatch(format!(
131 "Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
132 params.shape(),
133 gradients.shape()
134 )));
135 }
136
137 let dim = params.raw_dim().into_dyn();
138 self.ensure_state(index, &dim);
139
140 let lr = self.learning_rate;
141 let eps = self.epsilon;
142 let weight_decay = self.weight_decay;
143 let use_weight_decay = weight_decay > A::zero();
144
145 let accumulators = self.sum_squared_grads.as_mut().ok_or_else(|| {
146 OptimError::InvalidConfig("Adagrad state not initialized".to_string())
147 })?;
148
149 let mut updated = params.to_owned();
150 let mut params_view = updated.view_mut().into_dyn();
151 let gradients_view = gradients.view().into_dyn();
152
153 Zip::from(&mut params_view)
154 .and(&gradients_view)
155 .and(&mut accumulators[index])
156 .for_each(|p, &g, acc| {
157 let grad = if use_weight_decay {
158 g + weight_decay * *p
159 } else {
160 g
161 };
162 *acc = *acc + grad * grad;
163 *p = *p - lr * grad / (acc.sqrt() + eps);
164 });
165 drop(params_view);
166
167 Ok(updated)
168 }
169}
170
171impl<A, D> Optimizer<A, D> for Adagrad<A>
172where
173 A: Float + ScalarOperand + Debug + Send + Sync,
174 D: Dimension,
175{
176 fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
177 self.step_indexed(0, params, gradients)
178 }
179
180 fn step_list(
181 &mut self,
182 params_list: &[&Array<A, D>],
183 gradients_list: &[&Array<A, D>],
184 ) -> Result<Vec<Array<A, D>>> {
185 if params_list.len() != gradients_list.len() {
186 return Err(OptimError::InvalidConfig(format!(
187 "Number of parameter arrays ({}) does not match number of gradient arrays ({})",
188 params_list.len(),
189 gradients_list.len()
190 )));
191 }
192
193 let mut results = Vec::with_capacity(params_list.len());
194 for (index, (params, grads)) in params_list.iter().zip(gradients_list.iter()).enumerate() {
195 results.push(self.step_indexed(index, params, grads)?);
196 }
197 Ok(results)
198 }
199
200 fn get_learning_rate(&self) -> A {
201 self.learning_rate
202 }
203
204 fn set_learning_rate(&mut self, learning_rate: A) {
205 self.learning_rate = learning_rate;
206 }
207}