Skip to main content

rill_ml/optim/
adagrad.rs

1//! AdaGrad optimizer.
2//!
3//! Maintains a per-parameter sum of squared gradients and scales the learning
4//! rate accordingly. Time complexity per step: `O(d)`. Space: `O(d)`.
5
6use crate::error::{RillError, checked_finite_add, checked_increment, ensure_finite};
7#[cfg(feature = "serde")]
8use crate::persistence::ValidateState;
9
10/// Configuration for [`AdaGrad`].
11#[derive(Debug, Clone)]
12#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
13#[non_exhaustive]
14pub struct AdaGradConfig {
15    /// Global learning rate. Must be finite and strictly positive.
16    pub learning_rate: f64,
17    /// L2 regularization strength.
18    pub l2: f64,
19    /// Small constant added to the denominator for numerical stability.
20    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/// AdaGrad optimizer.
34///
35/// Update rule:
36/// ```text
37/// g2_i += grad_i^2
38/// w_i -= lr / sqrt(g2_i + epsilon) * (grad_i + l2 * w_i)
39/// ```
40#[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    /// Create a new AdaGrad optimizer.
52    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    /// Number of parameters (features + intercept).
84    pub const fn param_count(&self) -> usize {
85        self.feature_count + 1
86    }
87
88    /// Number of samples processed.
89    pub const fn samples_seen(&self) -> u64 {
90        self.samples_seen
91    }
92
93    /// Apply one gradient step.
94    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    /// Reset to initial state.
163    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        // first step: scale = sqrt(g^2 + eps) ≈ |g|
233        // w0 -= 0.1 / 1.0 * 1.0 = -0.1
234        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        // second step should be smaller than first
256        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}