Skip to main content

entrenar/optim/
sgd.rs

1//! Stochastic Gradient Descent optimizer
2
3use super::Optimizer;
4use crate::Tensor;
5use ndarray::Array1;
6
7/// SGD optimizer with optional momentum
8pub struct SGD {
9    lr: f32,
10    momentum: f32,
11    velocities: Vec<Option<Array1<f32>>>,
12}
13
14impl SGD {
15    /// Create a new SGD optimizer
16    pub fn new(lr: f32, momentum: f32) -> Self {
17        Self { lr, momentum, velocities: Vec::new() }
18    }
19
20    /// Initialize velocities if needed
21    fn ensure_velocities(&mut self, params: &[Tensor]) {
22        if self.velocities.is_empty() {
23            self.velocities = params.iter().map(|_| None).collect();
24        }
25    }
26}
27
28impl Optimizer for SGD {
29    fn step(&mut self, params: &mut [Tensor]) {
30        self.ensure_velocities(params);
31
32        for (i, param) in params.iter_mut().enumerate() {
33            if let Some(grad) = param.grad() {
34                // Use SIMD for large tensors (>= 16 elements for meaningful speedup)
35                if grad.len() >= 16 {
36                    let grad_slice = grad.as_slice().expect("grad array is contiguous");
37                    let param_slice =
38                        param.data_mut().as_slice_mut().expect("param array is contiguous");
39
40                    if self.momentum > 0.0 {
41                        // Initialize velocity if needed
42                        if self.velocities[i].is_none() {
43                            self.velocities[i] = Some(Array1::zeros(grad.len()));
44                        }
45
46                        let velocity =
47                            self.velocities[i].as_mut().expect("velocity buffer initialized above");
48                        let velocity_slice =
49                            velocity.as_slice_mut().expect("velocity array is contiguous");
50
51                        // PyTorch SGD+momentum (F-SGD-MOMENTUM-LRSCHED-001):
52                        // buffer stays UNSCALED so a mid-training lr change
53                        // (LR schedule) applies the FRESH lr each step.
54                        //   b = momentum * b + grad
55                        //   param -= lr * b   (lr read fresh, not baked into b)
56                        // First scale the buffer by momentum.
57                        for v in velocity_slice.iter_mut() {
58                            *v *= self.momentum;
59                        }
60
61                        // b += grad (a=1.0) using SIMD axpy
62                        super::simd::simd_axpy(1.0, grad_slice, velocity_slice);
63
64                        // param += -lr * b (lr applied fresh at update time)
65                        super::simd::simd_axpy(-self.lr, velocity_slice, param_slice);
66                    } else {
67                        // Simple SGD: param -= lr * grad (using SIMD axpy)
68                        super::simd::simd_axpy(-self.lr, grad_slice, param_slice);
69                    }
70                } else {
71                    // Fallback to scalar implementation for small tensors
72                    if self.momentum > 0.0 {
73                        // PyTorch SGD+momentum (F-SGD-MOMENTUM-LRSCHED-001):
74                        // UNSCALED buffer b = momentum * b + grad, then
75                        // param -= lr * b with lr read FRESH each step (so an
76                        // LR schedule never carries a stale lr in the buffer).
77                        let velocity = if let Some(v) = &self.velocities[i] {
78                            v * self.momentum + &grad
79                        } else {
80                            grad.clone()
81                        };
82
83                        *param.data_mut() = param.data() - &(&velocity * self.lr);
84                        self.velocities[i] = Some(velocity);
85                    } else {
86                        // Simple SGD: param -= lr * grad
87                        *param.data_mut() = param.data() - &(&grad * self.lr);
88                    }
89                }
90            }
91        }
92    }
93
94    fn lr(&self) -> f32 {
95        self.lr
96    }
97
98    fn set_lr(&mut self, lr: f32) {
99        self.lr = lr;
100    }
101}
102
103#[cfg(test)]
104mod tests {
105    use super::*;
106
107    #[test]
108    fn test_sgd_small_tensor_no_momentum() {
109        let param = Tensor::from_vec(vec![1.0, 2.0, 3.0], true);
110        param.set_grad(Array1::from_vec(vec![0.1, 0.2, 0.3]));
111
112        let mut opt = SGD::new(0.1, 0.0);
113        opt.step(&mut [param.clone()]);
114        // Small tensor path, no momentum
115    }
116
117    #[test]
118    fn test_sgd_small_tensor_with_momentum() {
119        let param = Tensor::from_vec(vec![1.0, 2.0, 3.0], true);
120        param.set_grad(Array1::from_vec(vec![0.1, 0.2, 0.3]));
121
122        let mut opt = SGD::new(0.1, 0.9);
123        // First step initializes velocity from scratch
124        opt.step(&mut [param.clone()]);
125
126        // Second step uses existing velocity
127        param.set_grad(Array1::from_vec(vec![0.1, 0.2, 0.3]));
128        opt.step(&mut [param.clone()]);
129    }
130
131    #[test]
132    fn test_sgd_large_tensor_with_momentum() {
133        // >= 16 elements to trigger SIMD path
134        let data: Vec<f32> = (0..20).map(|i| i as f32).collect();
135        let grad: Vec<f32> = vec![0.1; 20];
136
137        let param = Tensor::from_vec(data, true);
138        param.set_grad(Array1::from_vec(grad.clone()));
139
140        let mut opt = SGD::new(0.1, 0.9);
141        opt.step(&mut [param.clone()]);
142
143        // Second step with existing velocity
144        param.set_grad(Array1::from_vec(grad));
145        opt.step(&mut [param.clone()]);
146    }
147
148    #[test]
149    fn test_sgd_lr_getter_setter() {
150        let mut opt = SGD::new(0.1, 0.0);
151        assert!((opt.lr() - 0.1).abs() < 1e-6);
152        opt.set_lr(0.01);
153        assert!((opt.lr() - 0.01).abs() < 1e-6);
154    }
155
156    /// FALSIFY F-SGD-MOMENTUM-LRSCHED-001 (scalar path):
157    /// SGD-with-momentum must match PyTorch under an LR schedule.
158    ///
159    /// PyTorch SGD+momentum: `b = mu*b + g` (UNSCALED buffer), then
160    /// `theta -= lr*b` (lr read FRESH each step). aprender previously baked
161    /// `lr` into the velocity buffer, so after `set_lr` the momentum term
162    /// carried a STALE lr → divergence.
163    ///
164    /// Closed-form (g=1.0, mu=0.9, theta0=0.0, lr 0.1 → set_lr(0.01)):
165    ///   b1 = 0.9*0 + 1   = 1.0   ; theta1 = 0    - 0.1 *1.0 = -0.1
166    ///   b2 = 0.9*1 + 1   = 1.9   ; theta2 = -0.1 - 0.01*1.9 = -0.119
167    /// On the buggy (lr-baked) path theta2 = -0.200 (~40% off).
168    #[test]
169    fn falsify_sgd_momentum_lrsched_scalar() {
170        // 1 element → scalar fallback path (< 16 elements).
171        let param = Tensor::from_vec(vec![0.0], true);
172        let mut opt = SGD::new(0.1, 0.9);
173
174        param.set_grad(Array1::from_vec(vec![1.0]));
175        let mut params = [param];
176        opt.step(&mut params);
177        // theta1 = -0.1
178        assert!(
179            (params[0].data()[0] - (-0.1)).abs() < 1e-6,
180            "FALSIFIED: theta1 = {} != -0.1 (PyTorch step 1)",
181            params[0].data()[0]
182        );
183
184        opt.set_lr(0.01);
185        params[0].set_grad(Array1::from_vec(vec![1.0]));
186        opt.step(&mut params);
187        // theta2 = -0.119 (PyTorch). Buggy lr-baked path gives -0.200.
188        assert!(
189            (params[0].data()[0] - (-0.119)).abs() < 1e-6,
190            "FALSIFIED F-SGD-MOMENTUM-LRSCHED-001 (scalar): theta2 = {} != -0.119 \
191             (PyTorch rule b=mu*b+g, theta-=lr*b). lr baked into velocity buffer?",
192            params[0].data()[0]
193        );
194    }
195
196    /// FALSIFY F-SGD-MOMENTUM-LRSCHED-001 (SIMD path): identical assertion as
197    /// the scalar test but with 16 elements (>= 16 triggers the SIMD path).
198    /// Every element is independent and shares the same closed-form, so each
199    /// must equal -0.119 after the lr-scheduled second step.
200    #[test]
201    fn falsify_sgd_momentum_lrsched_simd() {
202        let n = 16;
203        let param = Tensor::from_vec(vec![0.0; n], true);
204        let mut opt = SGD::new(0.1, 0.9);
205
206        param.set_grad(Array1::from_vec(vec![1.0; n]));
207        let mut params = [param];
208        opt.step(&mut params);
209        for &x in params[0].data().iter() {
210            assert!(
211                (x - (-0.1)).abs() < 1e-6,
212                "FALSIFIED: theta1 = {x} != -0.1 (SIMD, PyTorch step 1)"
213            );
214        }
215
216        opt.set_lr(0.01);
217        params[0].set_grad(Array1::from_vec(vec![1.0; n]));
218        opt.step(&mut params);
219        for &x in params[0].data().iter() {
220            assert!(
221                (x - (-0.119)).abs() < 1e-6,
222                "FALSIFIED F-SGD-MOMENTUM-LRSCHED-001 (SIMD): theta2 = {x} != -0.119 \
223                 (PyTorch rule b=mu*b+g, theta-=lr*b)."
224            );
225        }
226    }
227
228    /// Control: constant lr must NOT regress. With lr fixed at 0.1, two steps
229    /// of (g=1.0, mu=0.9, theta0=0.0):
230    ///   b1=1.0; theta1 = -0.1
231    ///   b2=1.9; theta2 = -0.1 - 0.1*1.9 = -0.29
232    /// This holds identically for the old and new rules (lr never changes).
233    #[test]
234    fn test_sgd_momentum_constant_lr_no_regression() {
235        // Scalar path.
236        let param = Tensor::from_vec(vec![0.0], true);
237        let mut opt = SGD::new(0.1, 0.9);
238        param.set_grad(Array1::from_vec(vec![1.0]));
239        let mut params = [param];
240        opt.step(&mut params);
241        assert!((params[0].data()[0] - (-0.1)).abs() < 1e-6);
242        params[0].set_grad(Array1::from_vec(vec![1.0]));
243        opt.step(&mut params);
244        assert!(
245            (params[0].data()[0] - (-0.29)).abs() < 1e-6,
246            "constant-lr scalar regression: theta2 = {} != -0.29",
247            params[0].data()[0]
248        );
249
250        // SIMD path (16 elements), same closed-form.
251        let n = 16;
252        let param = Tensor::from_vec(vec![0.0; n], true);
253        let mut opt = SGD::new(0.1, 0.9);
254        param.set_grad(Array1::from_vec(vec![1.0; n]));
255        let mut params = [param];
256        opt.step(&mut params);
257        params[0].set_grad(Array1::from_vec(vec![1.0; n]));
258        opt.step(&mut params);
259        for &x in params[0].data().iter() {
260            assert!(
261                (x - (-0.29)).abs() < 1e-6,
262                "constant-lr SIMD regression: theta2 = {x} != -0.29"
263            );
264        }
265    }
266
267    #[test]
268    fn test_sgd_no_grad_skips() {
269        let param = Tensor::from_vec(vec![1.0, 2.0, 3.0], false);
270        // No gradient set
271
272        let mut opt = SGD::new(0.1, 0.0);
273        opt.step(&mut [param.clone()]); // Should not panic
274    }
275}