1use crate::error::{RillError, checked_finite_add, checked_increment, ensure_finite};
7#[cfg(feature = "serde")]
8use crate::persistence::ValidateState;
9
10#[derive(Debug, Clone)]
12#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
13#[non_exhaustive]
14pub struct AdaGradConfig {
15 pub learning_rate: f64,
17 pub l2: f64,
19 pub epsilon: f64,
21}
22
23impl Default for AdaGradConfig {
24 fn default() -> Self {
25 Self {
26 learning_rate: 0.1,
27 l2: 0.0,
28 epsilon: 1e-8,
29 }
30 }
31}
32
33#[derive(Debug, Clone)]
41#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
42pub struct AdaGrad {
43 feature_count: usize,
44 config: AdaGradConfig,
45 grad_sq_weights: Vec<f64>,
46 grad_sq_intercept: f64,
47 samples_seen: u64,
48}
49
50impl AdaGrad {
51 pub fn new(feature_count: usize, config: AdaGradConfig) -> Result<Self, RillError> {
53 if feature_count == 0 {
54 return Err(RillError::EmptyFeatures);
55 }
56 ensure_finite("learning_rate", config.learning_rate)?;
57 ensure_finite("l2", config.l2)?;
58 ensure_finite("epsilon", config.epsilon)?;
59 if config.learning_rate <= 0.0 {
60 return Err(RillError::InvalidLearningRate(config.learning_rate));
61 }
62 if config.l2 < 0.0 {
63 return Err(RillError::InvalidParameter {
64 name: "l2",
65 value: config.l2,
66 });
67 }
68 if config.epsilon <= 0.0 {
69 return Err(RillError::InvalidParameter {
70 name: "epsilon",
71 value: config.epsilon,
72 });
73 }
74 Ok(Self {
75 feature_count,
76 config,
77 grad_sq_weights: vec![0.0; feature_count],
78 grad_sq_intercept: 0.0,
79 samples_seen: 0,
80 })
81 }
82
83 pub const fn param_count(&self) -> usize {
85 self.feature_count + 1
86 }
87
88 pub const fn samples_seen(&self) -> u64 {
90 self.samples_seen
91 }
92
93 pub fn step(
95 &mut self,
96 weights: &mut [f64],
97 intercept: &mut f64,
98 grad_weights: &[f64],
99 grad_intercept: f64,
100 ) -> Result<(), RillError> {
101 if weights.len() != self.feature_count {
102 return Err(RillError::DimensionMismatch {
103 expected: self.feature_count,
104 actual: weights.len(),
105 });
106 }
107 if grad_weights.len() != self.feature_count {
108 return Err(RillError::DimensionMismatch {
109 expected: self.feature_count,
110 actual: grad_weights.len(),
111 });
112 }
113 for &gradient in grad_weights {
114 ensure_finite("grad_weight", gradient)?;
115 }
116 ensure_finite("grad_intercept", grad_intercept)?;
117 let next_samples = checked_increment(self.samples_seen, "AdaGrad sample")?;
118 let lr = self.config.learning_rate;
119 let l2 = self.config.l2;
120 let eps = self.config.epsilon;
121
122 let mut next_grad_sq_weights = Vec::with_capacity(self.feature_count);
123 let mut next_weights = Vec::with_capacity(self.feature_count);
124 for (i, (&weight, &gradient)) in weights.iter().zip(grad_weights).enumerate() {
125 let squared_gradient = gradient * gradient;
126 ensure_finite("squared gradient", squared_gradient)?;
127 let accumulator = checked_finite_add(
128 self.grad_sq_weights[i],
129 squared_gradient,
130 "AdaGrad accumulator",
131 )?;
132 let scale = (accumulator + eps).sqrt();
133 ensure_finite("AdaGrad scale", scale)?;
134 let regularized_gradient = gradient + l2 * weight;
135 ensure_finite("regularized gradient", regularized_gradient)?;
136 let next_weight = weight - lr / scale * regularized_gradient;
137 ensure_finite("weight", next_weight)?;
138 next_grad_sq_weights.push(accumulator);
139 next_weights.push(next_weight);
140 }
141
142 let squared_intercept_gradient = grad_intercept * grad_intercept;
143 ensure_finite("squared intercept gradient", squared_intercept_gradient)?;
144 let next_grad_sq_intercept = checked_finite_add(
145 self.grad_sq_intercept,
146 squared_intercept_gradient,
147 "AdaGrad intercept accumulator",
148 )?;
149 let intercept_scale = (next_grad_sq_intercept + eps).sqrt();
150 ensure_finite("AdaGrad intercept scale", intercept_scale)?;
151 let next_intercept = *intercept - lr / intercept_scale * grad_intercept;
152 ensure_finite("intercept", next_intercept)?;
153
154 self.grad_sq_weights = next_grad_sq_weights;
155 self.grad_sq_intercept = next_grad_sq_intercept;
156 weights.copy_from_slice(&next_weights);
157 *intercept = next_intercept;
158 self.samples_seen = next_samples;
159 Ok(())
160 }
161
162 pub fn reset(&mut self) {
164 for g in &mut self.grad_sq_weights {
165 *g = 0.0;
166 }
167 self.grad_sq_intercept = 0.0;
168 self.samples_seen = 0;
169 }
170}
171
172#[cfg(feature = "serde")]
173impl ValidateState for AdaGrad {
174 fn validate_state(&self) -> Result<(), RillError> {
175 if self.feature_count == 0 {
176 return Err(RillError::EmptyFeatures);
177 }
178 ensure_finite("learning_rate", self.config.learning_rate)?;
179 ensure_finite("l2", self.config.l2)?;
180 ensure_finite("epsilon", self.config.epsilon)?;
181 if self.config.learning_rate <= 0.0 {
182 return Err(RillError::InvalidLearningRate(self.config.learning_rate));
183 }
184 if self.config.l2 < 0.0 {
185 return Err(RillError::InvalidParameter {
186 name: "l2",
187 value: self.config.l2,
188 });
189 }
190 if self.config.epsilon <= 0.0 {
191 return Err(RillError::InvalidParameter {
192 name: "epsilon",
193 value: self.config.epsilon,
194 });
195 }
196 if self.grad_sq_weights.len() != self.feature_count {
197 return Err(RillError::InvalidState(format!(
198 "adagrad grad_sq_weights length {} does not match feature_count {}",
199 self.grad_sq_weights.len(),
200 self.feature_count
201 )));
202 }
203 ensure_finite("grad_sq_intercept", self.grad_sq_intercept)?;
204 for &g in &self.grad_sq_weights {
205 ensure_finite("grad_sq_weights", g)?;
206 if g < 0.0 {
207 return Err(RillError::InvalidState(format!(
208 "adagrad grad_sq_weights must be non-negative, got {g}"
209 )));
210 }
211 }
212 if self.grad_sq_intercept < 0.0 {
213 return Err(RillError::InvalidState(format!(
214 "adagrad grad_sq_intercept must be non-negative, got {}",
215 self.grad_sq_intercept
216 )));
217 }
218 Ok(())
219 }
220}
221
222#[cfg(test)]
223mod tests {
224 use super::*;
225
226 #[test]
227 fn adagrad_decreases_weights() {
228 let mut opt = AdaGrad::new(2, AdaGradConfig::default()).unwrap();
229 let mut w = vec![0.0, 0.0];
230 let mut b = 0.0;
231 opt.step(&mut w, &mut b, &[1.0, 2.0], 1.0).unwrap();
232 assert!(w[0] < 0.0);
235 assert!(w[1] < 0.0);
236 }
237
238 #[test]
239 fn adagrad_learning_rate_decreases() {
240 let mut opt = AdaGrad::new(
241 1,
242 AdaGradConfig {
243 learning_rate: 1.0,
244 l2: 0.0,
245 epsilon: 1e-12,
246 },
247 )
248 .unwrap();
249 let mut w = vec![0.0];
250 let mut b = 0.0;
251 opt.step(&mut w, &mut b, &[1.0], 0.0).unwrap();
252 let step1 = w[0].abs();
253 opt.step(&mut w, &mut b, &[1.0], 0.0).unwrap();
254 let step2 = w[0].abs() - step1;
255 assert!(step2 < step1);
257 }
258
259 #[test]
260 fn invalid_config_rejected() {
261 assert!(
262 AdaGrad::new(
263 1,
264 AdaGradConfig {
265 learning_rate: 0.0,
266 l2: 0.0,
267 epsilon: 1e-8
268 }
269 )
270 .is_err()
271 );
272 assert!(
273 AdaGrad::new(
274 1,
275 AdaGradConfig {
276 learning_rate: 0.1,
277 l2: -1.0,
278 epsilon: 1e-8
279 }
280 )
281 .is_err()
282 );
283 assert!(
284 AdaGrad::new(
285 1,
286 AdaGradConfig {
287 learning_rate: 0.1,
288 l2: 0.0,
289 epsilon: 0.0
290 }
291 )
292 .is_err()
293 );
294 }
295
296 #[test]
297 fn failed_step_is_atomic() {
298 let mut opt = AdaGrad::new(2, AdaGradConfig::default()).unwrap();
299 let mut weights = vec![1.0, 2.0];
300 let mut intercept = 3.0;
301 let result = opt.step(&mut weights, &mut intercept, &[1.0, f64::MAX], 1.0);
302 assert!(result.is_err());
303 assert_eq!(weights, vec![1.0, 2.0]);
304 assert_eq!(intercept, 3.0);
305 assert_eq!(opt.samples_seen(), 0);
306 assert_eq!(opt.grad_sq_weights, vec![0.0, 0.0]);
307 }
308}