1use super::Optimizer;
4use crate::Tensor;
5use ndarray::Array1;
6
7pub struct SGD {
9 lr: f32,
10 momentum: f32,
11 velocities: Vec<Option<Array1<f32>>>,
12}
13
14impl SGD {
15 pub fn new(lr: f32, momentum: f32) -> Self {
17 Self { lr, momentum, velocities: Vec::new() }
18 }
19
20 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 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 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 for v in velocity_slice.iter_mut() {
58 *v *= self.momentum;
59 }
60
61 super::simd::simd_axpy(1.0, grad_slice, velocity_slice);
63
64 super::simd::simd_axpy(-self.lr, velocity_slice, param_slice);
66 } else {
67 super::simd::simd_axpy(-self.lr, grad_slice, param_slice);
69 }
70 } else {
71 if self.momentum > 0.0 {
73 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 *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 }
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 opt.step(&mut [param.clone()]);
125
126 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 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 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 #[test]
169 fn falsify_sgd_momentum_lrsched_scalar() {
170 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 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 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 #[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 #[test]
234 fn test_sgd_momentum_constant_lr_no_regression() {
235 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 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 let mut opt = SGD::new(0.1, 0.0);
273 opt.step(&mut [param.clone()]); }
275}