Skip to main content

rill_ml/optim/
sgd.rs

1//! Stochastic gradient descent optimizer.
2
3use crate::error::{RillError, checked_increment, ensure_finite};
4#[cfg(feature = "serde")]
5use crate::persistence::ValidateState;
6
7/// Configuration for [`Sgd`].
8#[derive(Debug, Clone)]
9#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
10#[non_exhaustive]
11pub struct SgdConfig {
12    /// Learning rate. Must be finite and strictly positive.
13    pub learning_rate: f64,
14    /// L2 regularization strength. Must be finite and non-negative.
15    pub l2: f64,
16}
17
18impl Default for SgdConfig {
19    fn default() -> Self {
20        Self {
21            learning_rate: 0.01,
22            l2: 0.0,
23        }
24    }
25}
26
27/// SGD optimizer with optional L2 regularization.
28///
29/// The update rule for each weight `w_i` is:
30/// ```text
31/// w_i -= lr * (grad_i + l2 * w_i)
32/// ```
33/// The intercept is not regularized.
34#[derive(Debug, Clone)]
35#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
36pub struct Sgd {
37    feature_count: usize,
38    config: SgdConfig,
39    samples_seen: u64,
40}
41
42impl Sgd {
43    /// Create a new SGD optimizer.
44    pub fn new(feature_count: usize, config: SgdConfig) -> Result<Self, RillError> {
45        if feature_count == 0 {
46            return Err(RillError::EmptyFeatures);
47        }
48        ensure_finite("learning_rate", config.learning_rate)?;
49        ensure_finite("l2", config.l2)?;
50        if config.learning_rate <= 0.0 {
51            return Err(RillError::InvalidLearningRate(config.learning_rate));
52        }
53        if config.l2 < 0.0 {
54            return Err(RillError::InvalidParameter {
55                name: "l2",
56                value: config.l2,
57            });
58        }
59        Ok(Self {
60            feature_count,
61            config,
62            samples_seen: 0,
63        })
64    }
65
66    /// The configured learning rate.
67    pub const fn learning_rate(&self) -> f64 {
68        self.config.learning_rate
69    }
70
71    /// The configured L2 regularization.
72    pub const fn l2(&self) -> f64 {
73        self.config.l2
74    }
75
76    /// Number of parameters (features + intercept).
77    pub const fn param_count(&self) -> usize {
78        self.feature_count + 1
79    }
80
81    /// Number of samples processed.
82    pub const fn samples_seen(&self) -> u64 {
83        self.samples_seen
84    }
85
86    /// Apply one gradient step.
87    pub fn step(
88        &mut self,
89        weights: &mut [f64],
90        intercept: &mut f64,
91        grad_weights: &[f64],
92        grad_intercept: f64,
93    ) -> Result<(), RillError> {
94        if weights.len() != self.feature_count {
95            return Err(RillError::DimensionMismatch {
96                expected: self.feature_count,
97                actual: weights.len(),
98            });
99        }
100        if grad_weights.len() != self.feature_count {
101            return Err(RillError::DimensionMismatch {
102                expected: self.feature_count,
103                actual: grad_weights.len(),
104            });
105        }
106        for &gradient in grad_weights {
107            ensure_finite("grad_weight", gradient)?;
108        }
109        ensure_finite("grad_intercept", grad_intercept)?;
110        let next_samples = checked_increment(self.samples_seen, "SGD sample")?;
111        let lr = self.config.learning_rate;
112        let l2 = self.config.l2;
113        let next_weights = weights
114            .iter()
115            .zip(grad_weights)
116            .map(|(&weight, &gradient)| {
117                let regularized_gradient = gradient + l2 * weight;
118                ensure_finite("regularized gradient", regularized_gradient)?;
119                let next_weight = weight - lr * regularized_gradient;
120                ensure_finite("weight", next_weight)?;
121                Ok(next_weight)
122            })
123            .collect::<Result<Vec<_>, RillError>>()?;
124        let next_intercept = *intercept - lr * grad_intercept;
125        ensure_finite("intercept", next_intercept)?;
126
127        weights.copy_from_slice(&next_weights);
128        *intercept = next_intercept;
129        self.samples_seen = next_samples;
130        Ok(())
131    }
132
133    /// Reset the sample counter.
134    pub fn reset(&mut self) {
135        self.samples_seen = 0;
136    }
137}
138
139#[cfg(feature = "serde")]
140impl ValidateState for Sgd {
141    fn validate_state(&self) -> Result<(), RillError> {
142        if self.feature_count == 0 {
143            return Err(RillError::EmptyFeatures);
144        }
145        ensure_finite("learning_rate", self.config.learning_rate)?;
146        ensure_finite("l2", self.config.l2)?;
147        if self.config.learning_rate <= 0.0 {
148            return Err(RillError::InvalidLearningRate(self.config.learning_rate));
149        }
150        if self.config.l2 < 0.0 {
151            return Err(RillError::InvalidParameter {
152                name: "l2",
153                value: self.config.l2,
154            });
155        }
156        Ok(())
157    }
158}
159
160#[cfg(test)]
161mod tests {
162    use super::*;
163
164    #[test]
165    fn sgd_updates_weights() {
166        let mut opt = Sgd::new(
167            2,
168            SgdConfig {
169                learning_rate: 0.1,
170                l2: 0.0,
171            },
172        )
173        .unwrap();
174        let mut w = vec![0.0, 0.0];
175        let mut b = 0.0;
176        opt.step(&mut w, &mut b, &[1.0, 2.0], 0.5).unwrap();
177        // w -= 0.1 * grad -> [-0.1, -0.2]
178        assert!((w[0] + 0.1).abs() < 1e-12);
179        assert!((w[1] + 0.2).abs() < 1e-12);
180        assert!((b + 0.05).abs() < 1e-12);
181    }
182
183    #[test]
184    fn sgd_l2_regularization() {
185        let mut opt = Sgd::new(
186            1,
187            SgdConfig {
188                learning_rate: 0.1,
189                l2: 1.0,
190            },
191        )
192        .unwrap();
193        let mut w = vec![10.0];
194        let mut b = 0.0;
195        opt.step(&mut w, &mut b, &[0.0], 0.0).unwrap();
196        // w -= 0.1 * (0 + 1*10) = -1.0 -> w = 9.0
197        assert!((w[0] - 9.0).abs() < 1e-12);
198        // intercept not regularized -> b unchanged
199        assert!((b - 0.0).abs() < 1e-12);
200    }
201
202    #[test]
203    fn invalid_learning_rate_rejected() {
204        assert!(
205            Sgd::new(
206                1,
207                SgdConfig {
208                    learning_rate: 0.0,
209                    l2: 0.0
210                }
211            )
212            .is_err()
213        );
214        assert!(
215            Sgd::new(
216                1,
217                SgdConfig {
218                    learning_rate: -1.0,
219                    l2: 0.0
220                }
221            )
222            .is_err()
223        );
224    }
225
226    #[test]
227    fn invalid_l2_rejected() {
228        assert!(
229            Sgd::new(
230                1,
231                SgdConfig {
232                    learning_rate: 0.1,
233                    l2: -1.0
234                }
235            )
236            .is_err()
237        );
238    }
239
240    #[test]
241    fn failed_step_is_atomic() {
242        let mut opt = Sgd::new(2, SgdConfig::default()).unwrap();
243        let mut weights = vec![1.0, 2.0];
244        let mut intercept = 3.0;
245        let result = opt.step(&mut weights, &mut intercept, &[1.0, f64::INFINITY], 1.0);
246        assert!(result.is_err());
247        assert_eq!(weights, vec![1.0, 2.0]);
248        assert_eq!(intercept, 3.0);
249        assert_eq!(opt.samples_seen(), 0);
250    }
251}