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        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        // first step: scale = sqrt(g^2 + eps) ≈ |g|
231        // w0 -= 0.1 / 1.0 * 1.0 = -0.1
232        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        // second step should be smaller than first
254        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}