1
2use ruda_model::{module::AutodiffModule, record::Record};
3
4use super::{
5 SimpleOptimizer,
6 adaptor::OptimizerAdaptor,
7 decay::{WeightDecay, WeightDecayConfig},
8};
9use crate::{LearningRate, grad_clipping::GradientClippingConfig};
10
11use ruda_model::config::Config;
12use ruda_model::tensor::backend::Backend;
13use ruda_model::tensor::{Tensor, backend::AutodiffBackend, ops::Device};
14
15#[derive(Config, Debug)]
17pub struct RmsPropConfig {
18 #[config(default = 0.99)]
20 alpha: f32,
21 #[config(default = 0.9)]
23 momentum: f32,
24 #[config(default = 1e-5)]
26 epsilon: f32,
27 #[config(default = false)]
29 centered: bool,
30 weight_decay: Option<WeightDecayConfig>,
32 grad_clipping: Option<GradientClippingConfig>,
34}
35
36impl RmsPropConfig {
37 pub fn build(&self) -> RmsProp {
39 let weight_decay = self.weight_decay.as_ref().map(WeightDecay::new);
40 RmsProp {
41 alpha: self.alpha,
42 centered: self.centered,
43 weight_decay,
44 momentum: RmsPropMomentum {
45 momentum: self.momentum,
46 epsilon: self.epsilon,
47 },
48 }
49 }
50
51 pub fn init<B: AutodiffBackend, M: AutodiffModule<B>>(
57 &self,
58 ) -> OptimizerAdaptor<RmsProp, M, B> {
59 let mut optim = OptimizerAdaptor::from(self.build());
60 if let Some(config) = &self.grad_clipping {
61 optim = optim.with_grad_clipping(config.init());
62 }
63
64 optim
65 }
66}
67
68#[derive(Clone)]
71pub struct RmsProp {
72 alpha: f32,
73 centered: bool,
75 momentum: RmsPropMomentum,
77 weight_decay: Option<WeightDecay>,
78}
79
80impl<B: Backend> SimpleOptimizer<B> for RmsProp {
81 type State<const D: usize> = RmsPropState<B, D>;
82
83 fn step<const D: usize>(
84 &self,
85 lr: LearningRate,
86 tensor: Tensor<B, D>,
87 mut grad: Tensor<B, D>,
88 state: Option<Self::State<D>>,
89 ) -> (Tensor<B, D>, Option<Self::State<D>>) {
90 let mut state_square_avg = None;
92 let mut state_centered = None;
93 let mut state_momentum = None;
94 if let Some(state) = state {
95 state_square_avg = Some(state.square_avg);
96 state_centered = Some(state.centered);
97 state_momentum = state.momentum;
98 }
99
100 if let Some(weight_decay) = &self.weight_decay {
102 grad = weight_decay.transform(grad, tensor.clone());
103 }
104
105 let (grad, state_square_avg) =
107 SquareAvgState::transform(self.alpha, grad, state_square_avg);
108
109 let (grad, state_square_avg, state_centered) = CenteredState::transform(
111 self.alpha,
112 self.centered,
113 grad,
114 state_square_avg,
115 state_centered,
116 );
117
118 let (grad, state_centered, state_momentum) =
120 self.momentum
121 .transform(grad, state_centered, state_momentum);
122
123 let state = RmsPropState::new(state_square_avg, state_centered, state_momentum);
125
126 let delta = grad.mul_scalar(lr);
128 (tensor - delta, Some(state))
129 }
130
131 fn to_device<const D: usize>(mut state: Self::State<D>, device: &Device<B>) -> Self::State<D> {
132 state.square_avg = state.square_avg.to_device(device);
133 state.centered = state.centered.to_device(device);
134 state.momentum = state.momentum.map(|momentum| momentum.to_device(device));
135 state
136 }
137}
138
139#[derive(Record, Clone, new)]
141pub struct RmsPropState<B: Backend, const D: usize> {
142 pub square_avg: SquareAvgState<B, D>,
144 pub centered: CenteredState<B, D>,
146 pub momentum: Option<RmsPropMomentumState<B, D>>,
148}
149
150#[derive(Record, Clone, new)]
152pub struct SquareAvgState<B: Backend, const D: usize> {
153 pub square_avg: Tensor<B, D>,
155}
156
157impl<B: Backend, const D: usize> SquareAvgState<B, D> {
158 fn transform(alpha: f32, grad: Tensor<B, D>, state: Option<Self>) -> (Tensor<B, D>, Self) {
160 match state {
161 Some(state) => {
162 let square_avg = state
163 .square_avg
164 .mul_scalar(alpha)
165 .add(grad.clone().square().mul_scalar(1. - alpha));
166 (grad, Self { square_avg })
167 }
168 _ => {
169 let square_avg = grad.clone().square().mul_scalar(1. - alpha);
170 (grad, Self { square_avg })
171 }
172 }
173 }
174
175 pub fn to_device(mut self, device: &B::Device) -> Self {
185 self.square_avg = self.square_avg.to_device(device);
186 self
187 }
188}
189
190#[derive(Record, Clone, new)]
192pub struct CenteredState<B: Backend, const D: usize> {
193 pub grad_avg: Option<Tensor<B, D>>,
195 pub avg: Tensor<B, D>,
197}
198
199impl<B: Backend, const D: usize> CenteredState<B, D> {
200 fn transform(
202 alpha: f32,
203 centered: bool,
204 grad: Tensor<B, D>,
205 square_avg_state: SquareAvgState<B, D>,
206 centered_state: Option<Self>,
207 ) -> (Tensor<B, D>, SquareAvgState<B, D>, Self) {
208 if centered {
209 let grad_avg_constant = grad.clone().mul_scalar(1. - alpha);
210 let grad_avg = match centered_state {
211 Some(state) => state
212 .grad_avg
213 .map_or(grad_avg_constant.clone(), move |grad_avg| {
214 grad_avg.mul_scalar(alpha).add(grad_avg_constant)
215 }),
216 _ => grad_avg_constant,
217 };
218 let avg = square_avg_state
219 .square_avg
220 .clone()
221 .sub(grad_avg.clone().square());
222
223 (
224 grad,
225 square_avg_state,
226 Self {
227 grad_avg: Some(grad_avg),
228 avg,
229 },
230 )
231 } else {
232 (
233 grad,
234 square_avg_state.clone(),
235 Self {
236 grad_avg: None,
237 avg: square_avg_state.square_avg,
238 },
239 )
240 }
241 }
242
243 pub fn to_device(mut self, device: &B::Device) -> Self {
253 self.grad_avg = self.grad_avg.map(|grad_avg| grad_avg.to_device(device));
254 self.avg = self.avg.to_device(device);
255 self
256 }
257}
258
259#[derive(Clone)]
262pub struct RmsPropMomentum {
263 momentum: f32,
264 epsilon: f32,
265}
266
267impl RmsPropMomentum {
268 fn transform<B: Backend, const D: usize>(
270 &self,
271 grad: Tensor<B, D>,
272 centered_state: CenteredState<B, D>,
273 momentum_state: Option<RmsPropMomentumState<B, D>>,
274 ) -> (
275 Tensor<B, D>,
276 CenteredState<B, D>,
277 Option<RmsPropMomentumState<B, D>>,
278 ) {
279 let grad = grad.div(centered_state.avg.clone().sqrt().add_scalar(self.epsilon));
280
281 if self.momentum > 0. {
282 let buf = match momentum_state {
283 Some(state) => state.buf.mul_scalar(self.momentum).add(grad),
284 _ => grad,
285 };
286 (
287 buf.clone(),
288 centered_state,
289 Some(RmsPropMomentumState { buf }),
290 )
291 } else {
292 (grad, centered_state, None)
293 }
294 }
295}
296
297#[derive(Record, Clone, new)]
299pub struct RmsPropMomentumState<B: Backend, const D: usize> {
300 buf: Tensor<B, D>,
301}
302
303impl<B: Backend, const D: usize> RmsPropMomentumState<B, D> {
304 pub fn to_device(mut self, device: &B::Device) -> Self {
314 self.buf = self.buf.to_device(device);
315 self
316 }
317}
318
319#[cfg(test)]
320mod tests {
321 use ruda_model::tensor::ops::FloatElem;
322 use ruda_model::tensor::{Shape, Tolerance};
323
324 use super::*;
325 use crate::TestAutodiffBackend;
326 use crate::optim::{GradientsParams, Optimizer};
327 use ruda_model::module::{Module, Param};
328 use ruda_model::tensor::{Distribution, Tensor, TensorData};
329 use ruda_nn::{Linear, LinearConfig, LinearRecord};
330
331 type FT = FloatElem<TestAutodiffBackend>;
332
333 const LEARNING_RATE: LearningRate = 0.01;
334
335 #[test]
336 fn test_rmsprop_optimizer_save_load_state() {
337 let device = Default::default();
338 let linear = LinearConfig::new(6, 6).init(&device);
339 let x = Tensor::<TestAutodiffBackend, 2>::random([2, 6], Distribution::Default, &device);
340 let mut optimizer = create_rmsprop();
341 let grads = linear.forward(x).backward();
342 let grads = GradientsParams::from_grads(grads, &linear);
343 let _linear = optimizer.step(LEARNING_RATE, linear, grads);
344
345 #[cfg(feature = "std")]
346 {
347 use ruda_model::record::{BinFileRecorder, FullPrecisionSettings, Recorder};
348
349 BinFileRecorder::<FullPrecisionSettings>::default()
350 .record(
351 optimizer.to_record(),
352 std::env::temp_dir().as_path().join("test_optim_rmsprop"),
353 )
354 .unwrap();
355 }
356 #[cfg(not(feature = "std"))]
357 {
358 use ruda_model::record::{BinBytesRecorder, FullPrecisionSettings, Recorder};
359
360 let result = BinBytesRecorder::<FullPrecisionSettings>::default()
361 .record(optimizer.to_record(), ())
362 .unwrap();
363 assert!(!result.is_empty());
364 }
365
366 let state_optim_before = optimizer.to_record();
367 let state_optim_before_copy = optimizer.to_record();
368 let optimizer = create_rmsprop();
369 let optimizer = optimizer.load_record(state_optim_before_copy);
370 let state_optim_after = optimizer.to_record();
371
372 assert_eq!(state_optim_before.len(), state_optim_after.len());
373 }
374
375 #[test]
377 fn test_rmsprop_optimizer_with_numbers_basic() {
378 let linear = given_linear_layer(
379 TensorData::from([
380 [1., 1., 1., 1., 1., 1.],
381 [1., 1., 1., 1., 1., 1.],
382 [1., 1., 1., 1., 1., 1.],
383 [1., 1., 1., 1., 1., 1.],
384 [1., 1., 1., 1., 1., 1.],
385 [1., 1., 1., 1., 1., 1.],
386 ]),
387 TensorData::from([0.5, 0.5, 0.5, 0.5, 0.5, 0.5]),
388 );
389 let device = Default::default();
390 let x_1 = Tensor::<TestAutodiffBackend, 2>::from_floats(
391 [
392 [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
393 [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
394 ],
395 &device,
396 )
397 .require_grad();
398 let x_2 = Tensor::<TestAutodiffBackend, 2>::from_floats(
399 [
400 [0.8491, 0.2108, 0.8939, 0.4433, 0.5527, 0.2528],
401 [0.3270, 0.0412, 0.5538, 0.9605, 0.3195, 0.9085],
402 ],
403 &device,
404 )
405 .require_grad();
406
407 let mut optimizer = RmsPropConfig::new()
408 .with_alpha(0.99)
409 .with_epsilon(1e-8)
410 .with_weight_decay(WeightDecayConfig::new(0.05).into())
411 .with_momentum(0.9)
412 .with_centered(false)
413 .init();
414
415 let grads = linear.forward(x_1).backward();
417 let grads = GradientsParams::from_grads(grads, &linear);
418 let linear = optimizer.step(LEARNING_RATE, linear, grads);
419
420 let grads = linear.forward(x_2).backward();
422 let grads = GradientsParams::from_grads(grads, &linear);
423 let linear = optimizer.step(LEARNING_RATE, linear, grads);
424
425 let state_updated = linear.into_record();
427
428 let (weight_updated, bias_updated) = (
429 state_updated.weight.to_data(),
430 state_updated.bias.unwrap().to_data(),
431 );
432
433 let weights_expected = TensorData::from([
437 [0.743937, 0.743937, 0.743937, 0.743937, 0.743937, 0.743937],
438 [0.783809, 0.783809, 0.783809, 0.783809, 0.783809, 0.783809],
439 [0.742881, 0.742881, 0.742881, 0.742881, 0.742881, 0.742881],
440 [0.740366, 0.740366, 0.740366, 0.740366, 0.740366, 0.740366],
441 [0.748005, 0.748005, 0.748005, 0.748005, 0.748005, 0.748005],
442 [0.743710, 0.743710, 0.743710, 0.743710, 0.743710, 0.743710],
443 ]);
444 let bias_expected =
445 TensorData::from([0.239199, 0.239199, 0.239199, 0.239199, 0.239199, 0.239199]);
446
447 let tolerance = Tolerance::absolute(1e-6);
448 bias_updated.assert_approx_eq::<FT>(&bias_expected, tolerance);
449 weight_updated.assert_approx_eq::<FT>(&weights_expected, tolerance);
450 }
451
452 #[test]
453 fn test_rmsprop_optimizer_with_numbers() {
454 let linear = given_linear_layer(
455 TensorData::from([
456 [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
457 [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
458 [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
459 [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
460 [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
461 [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
462 ]),
463 TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
464 );
465 let device = Default::default();
466 let x_1 = Tensor::<TestAutodiffBackend, 2>::from_floats(
467 [
468 [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
469 [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
470 ],
471 &device,
472 )
473 .require_grad();
474 let x_2 = Tensor::<TestAutodiffBackend, 2>::from_floats(
475 [
476 [0.8491, 0.2108, 0.8939, 0.4433, 0.5527, 0.2528],
477 [0.3270, 0.0412, 0.5538, 0.9605, 0.3195, 0.9085],
478 ],
479 &device,
480 )
481 .require_grad();
482
483 let mut optimizer = RmsPropConfig::new()
484 .with_alpha(0.99)
485 .with_epsilon(1e-8)
486 .with_weight_decay(WeightDecayConfig::new(0.05).into())
487 .with_momentum(0.9)
488 .with_centered(false)
489 .init();
490
491 let grads = linear.forward(x_1).backward();
492 let grads = GradientsParams::from_grads(grads, &linear);
493 let linear = optimizer.step(LEARNING_RATE, linear, grads);
494
495 let grads = linear.forward(x_2).backward();
496 let grads = GradientsParams::from_grads(grads, &linear);
497 let linear = optimizer.step(LEARNING_RATE, linear, grads);
498
499 let state_updated = linear.into_record();
500 let weights_expected = TensorData::from([
501 [
502 -0.576399, -0.118494, 0.148353, 0.064070, -0.169983, -0.188779,
503 ],
504 [
505 -0.135571, -0.231448, -0.578445, 0.041143, -0.018162, -0.504207,
506 ],
507 [
508 -0.275990, -0.222397, -0.553153, -0.008625, -0.534956, 0.055967,
509 ],
510 [
511 -0.557575, -0.480979, -0.631072, -0.557675, -0.335686, -0.096997,
512 ],
513 [
514 0.078313, -0.469618, 0.119993, -0.424341, 0.127890, -0.281912,
515 ],
516 [
517 -0.271996, -0.268097, -0.130324, -0.064037, -0.226805, 0.127126,
518 ],
519 ]);
520 let bias_expected = TensorData::from([
521 -0.651299, -0.172400, -0.357800, -0.143200, -0.124200, -0.247800,
522 ]);
523
524 let (weight_updated, bias_updated) = (
525 state_updated.weight.to_data(),
526 state_updated.bias.unwrap().to_data(),
527 );
528
529 let tolerance = Tolerance::absolute(1e-6);
533 bias_updated.assert_approx_eq::<FT>(&bias_expected, tolerance);
534 weight_updated.assert_approx_eq::<FT>(&weights_expected, tolerance);
535 }
536
537 fn given_linear_layer(weight: TensorData, bias: TensorData) -> Linear<TestAutodiffBackend> {
538 let device = Default::default();
539 let record = LinearRecord {
540 weight: Param::from_data(weight, &device),
541 bias: Some(Param::from_data(bias, &device)),
542 };
543
544 LinearConfig::new(6, 6).init(&device).load_record(record)
545 }
546
547 #[allow(dead_code)]
548 fn create_random_tensor() -> Tensor<TestAutodiffBackend, 2> {
549 Tensor::<TestAutodiffBackend, 2>::random(
550 Shape::new([2, 20]),
551 Distribution::Default,
552 &Default::default(),
553 )
554 }
555
556 fn create_rmsprop()
557 -> OptimizerAdaptor<RmsProp, Linear<TestAutodiffBackend>, TestAutodiffBackend> {
558 RmsPropConfig {
559 alpha: 0.99,
560 epsilon: 1e-9,
561 centered: false,
562 weight_decay: Some(WeightDecayConfig { penalty: 0.05 }),
563 momentum: 0.9,
564 grad_clipping: None,
565 }
566 .init()
567 }
568}