Skip to main content

ruda_optim/optim/
rmsprop.rs

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/// Configuration to create the [RmsProp](RmsProp) optimizer.
16#[derive(Config, Debug)]
17pub struct RmsPropConfig {
18    /// Smoothing constant.
19    #[config(default = 0.99)]
20    alpha: f32,
21    /// momentum for RmsProp.
22    #[config(default = 0.9)]
23    momentum: f32,
24    /// A value required for numerical stability.
25    #[config(default = 1e-5)]
26    epsilon: f32,
27    /// if True, compute the centered RmsProp, the gradient is normalized by an estimation of its variance
28    #[config(default = false)]
29    centered: bool,
30    /// [Weight decay](WeightDecayConfig) config.
31    weight_decay: Option<WeightDecayConfig>,
32    /// [Gradient Clipping](GradientClippingConfig) config.
33    grad_clipping: Option<GradientClippingConfig>,
34}
35
36impl RmsPropConfig {
37    /// Build a [`RmsProp`] from the config.
38    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    /// Initialize RmsProp optimizer.
52    ///
53    /// # Returns
54    ///
55    /// Returns an optimizer that can be used to optimize a module.
56    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/// Optimizer that implements stochastic gradient descent with momentum.
69/// The optimizer can be configured with [RmsPropConfig](RmsPropConfig).
70#[derive(Clone)]
71pub struct RmsProp {
72    alpha: f32,
73    // epsilon: f32,
74    centered: bool,
75    // momentum: Option<Momentum<B>>,
76    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        // fetch state for params
91        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        // weight_decay transform
101        if let Some(weight_decay) = &self.weight_decay {
102            grad = weight_decay.transform(grad, tensor.clone());
103        }
104
105        // square_avg transform
106        let (grad, state_square_avg) =
107            SquareAvgState::transform(self.alpha, grad, state_square_avg);
108
109        // centered transform
110        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        // momentum transform
119        let (grad, state_centered, state_momentum) =
120            self.momentum
121                .transform(grad, state_centered, state_momentum);
122
123        // transition state
124        let state = RmsPropState::new(state_square_avg, state_centered, state_momentum);
125
126        // tensor param transform
127        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/// State of [RmsProp](RmsProp)
140#[derive(Record, Clone, new)]
141pub struct RmsPropState<B: Backend, const D: usize> {
142    /// Current squared average state.
143    pub square_avg: SquareAvgState<B, D>,
144    /// Current centered state
145    pub centered: CenteredState<B, D>,
146    /// Current gradient momentum, if any.
147    pub momentum: Option<RmsPropMomentumState<B, D>>,
148}
149
150/// [SquareAvgState](SquareAvgState) is to store and pass optimizer step params.
151#[derive(Record, Clone, new)]
152pub struct SquareAvgState<B: Backend, const D: usize> {
153    /// Current squared average.
154    pub square_avg: Tensor<B, D>,
155}
156
157impl<B: Backend, const D: usize> SquareAvgState<B, D> {
158    /// transform [SquareAvgState] to the next step
159    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    /// Moves the state to a device.
176    ///
177    /// # Arguments
178    ///
179    /// * `device` - Device to move the state to.
180    ///
181    /// # Returns
182    ///
183    /// * `self` - Moved state.
184    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/// [CenteredState](CenteredState) is to store and pass optimizer step params.
191#[derive(Record, Clone, new)]
192pub struct CenteredState<B: Backend, const D: usize> {
193    /// The averaged gradient to calculate the centered gradient, if available.
194    pub grad_avg: Option<Tensor<B, D>>,
195    /// The current average value.
196    pub avg: Tensor<B, D>,
197}
198
199impl<B: Backend, const D: usize> CenteredState<B, D> {
200    /// transform [CenteredState] to the next step
201    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    /// Moves the state to a device.
244    ///
245    /// # Arguments
246    ///
247    /// * `device` - Device to move the state to.
248    ///
249    /// # Returns
250    ///
251    /// * `self` - Moved state.
252    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/// [RmsPropMomentum](RmsPropMomentum) is to store config status for optimizer.
260/// (, which is stored in [optimizer](RmsProp) itself and not passed in during `step()` calculation)
261#[derive(Clone)]
262pub struct RmsPropMomentum {
263    momentum: f32,
264    epsilon: f32,
265}
266
267impl RmsPropMomentum {
268    /// transform [grad](Tensor) and [RmsPropMomentumState] to the next step
269    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/// [RmsPropMomentumState](RmsPropMomentumState) is to store and pass optimizer step params.
298#[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    /// Moves the state to a device.
305    ///
306    /// # Arguments
307    ///
308    /// * `device` - Device to move the state to.
309    ///
310    /// # Returns
311    ///
312    /// * `self` - Moved state.
313    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    /// used for test differences and debug
376    #[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        // println!("linear is {:?}", linear);
416        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        // println!("linear is {:?}", linear);
421        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        // println!("linear is {:?}", linear);
426        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        // println!("\nweight_updated\n{:?}", weight_updated);
434        // println!("\nbias_updated\n{:?}", bias_updated);
435
436        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        // println!("\nweight_updated\n{:?}", weight_updated);
530        // println!("\nbias_updated\n{:?}", bias_updated);
531
532        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}