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 self.grad_sq_weights.fill(0.0);
165 self.grad_sq_intercept = 0.0;
166 self.samples_seen = 0;
167 }
168}
169
170#[cfg(feature = "serde")]
171impl ValidateState for AdaGrad {
172 fn validate_state(&self) -> Result<(), RillError> {
173 if self.feature_count == 0 {
174 return Err(RillError::EmptyFeatures);
175 }
176 ensure_finite("learning_rate", self.config.learning_rate)?;
177 ensure_finite("l2", self.config.l2)?;
178 ensure_finite("epsilon", self.config.epsilon)?;
179 if self.config.learning_rate <= 0.0 {
180 return Err(RillError::InvalidLearningRate(self.config.learning_rate));
181 }
182 if self.config.l2 < 0.0 {
183 return Err(RillError::InvalidParameter {
184 name: "l2",
185 value: self.config.l2,
186 });
187 }
188 if self.config.epsilon <= 0.0 {
189 return Err(RillError::InvalidParameter {
190 name: "epsilon",
191 value: self.config.epsilon,
192 });
193 }
194 if self.grad_sq_weights.len() != self.feature_count {
195 return Err(RillError::InvalidState(format!(
196 "adagrad grad_sq_weights length {} does not match feature_count {}",
197 self.grad_sq_weights.len(),
198 self.feature_count
199 )));
200 }
201 ensure_finite("grad_sq_intercept", self.grad_sq_intercept)?;
202 for &g in &self.grad_sq_weights {
203 ensure_finite("grad_sq_weights", g)?;
204 if g < 0.0 {
205 return Err(RillError::InvalidState(format!(
206 "adagrad grad_sq_weights must be non-negative, got {g}"
207 )));
208 }
209 }
210 if self.grad_sq_intercept < 0.0 {
211 return Err(RillError::InvalidState(format!(
212 "adagrad grad_sq_intercept must be non-negative, got {}",
213 self.grad_sq_intercept
214 )));
215 }
216 Ok(())
217 }
218}
219
220#[cfg(test)]
221mod tests {
222 use super::*;
223
224 #[test]
225 fn adagrad_decreases_weights() {
226 let mut opt = AdaGrad::new(2, AdaGradConfig::default()).unwrap();
227 let mut w = vec![0.0, 0.0];
228 let mut b = 0.0;
229 opt.step(&mut w, &mut b, &[1.0, 2.0], 1.0).unwrap();
230 assert!(w[0] < 0.0);
233 assert!(w[1] < 0.0);
234 }
235
236 #[test]
237 fn adagrad_learning_rate_decreases() {
238 let mut opt = AdaGrad::new(
239 1,
240 AdaGradConfig {
241 learning_rate: 1.0,
242 l2: 0.0,
243 epsilon: 1e-12,
244 },
245 )
246 .unwrap();
247 let mut w = vec![0.0];
248 let mut b = 0.0;
249 opt.step(&mut w, &mut b, &[1.0], 0.0).unwrap();
250 let step1 = w[0].abs();
251 opt.step(&mut w, &mut b, &[1.0], 0.0).unwrap();
252 let step2 = w[0].abs() - step1;
253 assert!(step2 < step1);
255 }
256
257 #[test]
258 fn invalid_config_rejected() {
259 assert!(
260 AdaGrad::new(
261 1,
262 AdaGradConfig {
263 learning_rate: 0.0,
264 l2: 0.0,
265 epsilon: 1e-8
266 }
267 )
268 .is_err()
269 );
270 assert!(
271 AdaGrad::new(
272 1,
273 AdaGradConfig {
274 learning_rate: 0.1,
275 l2: -1.0,
276 epsilon: 1e-8
277 }
278 )
279 .is_err()
280 );
281 assert!(
282 AdaGrad::new(
283 1,
284 AdaGradConfig {
285 learning_rate: 0.1,
286 l2: 0.0,
287 epsilon: 0.0
288 }
289 )
290 .is_err()
291 );
292 }
293
294 #[test]
295 fn failed_step_is_atomic() {
296 let mut opt = AdaGrad::new(2, AdaGradConfig::default()).unwrap();
297 let mut weights = vec![1.0, 2.0];
298 let mut intercept = 3.0;
299 let result = opt.step(&mut weights, &mut intercept, &[1.0, f64::MAX], 1.0);
300 assert!(result.is_err());
301 assert_eq!(weights, vec![1.0, 2.0]);
302 assert_eq!(intercept, 3.0);
303 assert_eq!(opt.samples_seen(), 0);
304 assert_eq!(opt.grad_sq_weights, vec![0.0, 0.0]);
305 }
306}