1use crate::error::{RillError, checked_increment, ensure_finite};
4#[cfg(feature = "serde")]
5use crate::persistence::ValidateState;
6
7#[derive(Debug, Clone)]
9#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
10#[non_exhaustive]
11pub struct SgdConfig {
12 pub learning_rate: f64,
14 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#[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 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 pub const fn learning_rate(&self) -> f64 {
68 self.config.learning_rate
69 }
70
71 pub const fn l2(&self) -> f64 {
73 self.config.l2
74 }
75
76 pub const fn param_count(&self) -> usize {
78 self.feature_count + 1
79 }
80
81 pub const fn samples_seen(&self) -> u64 {
83 self.samples_seen
84 }
85
86 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 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 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 assert!((w[0] - 9.0).abs() < 1e-12);
198 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}