Skip to main content

ruda_optim/optim/
sgd.rs

1
2use super::SimpleOptimizer;
3use super::adaptor::OptimizerAdaptor;
4use super::decay::{WeightDecay, WeightDecayConfig};
5use super::momentum::{Momentum, MomentumConfig, MomentumState};
6use crate::LearningRate;
7use crate::grad_clipping::GradientClippingConfig;
8use ruda_model::config::Config;
9use ruda_model::module::AutodiffModule;
10use ruda_model::record::Record;
11use ruda_model::tensor::Tensor;
12use ruda_model::tensor::backend::{AutodiffBackend, Backend};
13
14/// Configuration to create the [Sgd](Sgd) optimizer.
15#[derive(Config, Debug)]
16pub struct SgdConfig {
17    /// [Weight decay](WeightDecayConfig) config.
18    weight_decay: Option<WeightDecayConfig>,
19    /// [Momentum](MomentumConfig) config.
20    momentum: Option<MomentumConfig>,
21    /// [Gradient Clipping](GradientClippingConfig) config.
22    gradient_clipping: Option<GradientClippingConfig>,
23}
24
25/// Optimizer that implements stochastic gradient descent with momentum.
26///
27/// The optimizer can be configured with [SgdConfig](SgdConfig).
28#[derive(Clone)]
29pub struct Sgd<B: Backend> {
30    momentum: Option<Momentum<B>>,
31    weight_decay: Option<WeightDecay>,
32}
33
34/// State of [Sgd](Sgd).
35#[derive(Record, Clone, new)]
36pub struct SgdState<B: Backend, const D: usize> {
37    /// The current state of the momentum (if any).
38    pub momentum: Option<MomentumState<B, D>>,
39}
40
41impl SgdConfig {
42    /// Build a [`Sgd`] from the config.
43    pub fn build<B: Backend>(&self) -> Sgd<B> {
44        Sgd {
45            momentum: self.momentum.as_ref().map(Momentum::new),
46            weight_decay: self.weight_decay.as_ref().map(WeightDecay::new),
47        }
48    }
49
50    /// Creates a new [SgdConfig](SgdConfig) with default values.
51    pub fn init<B: AutodiffBackend, M: AutodiffModule<B>>(
52        &self,
53    ) -> OptimizerAdaptor<Sgd<B::InnerBackend>, M, B> {
54        let mut optim = OptimizerAdaptor::from(self.build());
55        if let Some(config) = &self.gradient_clipping {
56            optim = optim.with_grad_clipping(config.init());
57        }
58        optim
59    }
60}
61
62impl<B: Backend> SimpleOptimizer<B> for Sgd<B> {
63    type State<const D: usize> = SgdState<B, D>;
64
65    fn step<const D: usize>(
66        &self,
67        lr: LearningRate,
68        tensor: Tensor<B, D>,
69        mut grad: Tensor<B, D>,
70        state: Option<Self::State<D>>,
71    ) -> (Tensor<B, D>, Option<Self::State<D>>) {
72        let mut state_momentum = None;
73
74        if let Some(state) = state {
75            state_momentum = state.momentum;
76        }
77
78        if let Some(weight_decay) = &self.weight_decay {
79            grad = weight_decay.transform(grad, tensor.clone());
80        }
81
82        if let Some(momentum) = &self.momentum {
83            let (grad_out, state) = momentum.transform(grad, state_momentum);
84            state_momentum = Some(state);
85            grad = grad_out;
86        }
87
88        let state = SgdState::new(state_momentum);
89        let delta = grad.mul_scalar(lr);
90
91        (tensor - delta, Some(state))
92    }
93
94    fn to_device<const D: usize>(mut state: Self::State<D>, device: &B::Device) -> Self::State<D> {
95        state.momentum = state.momentum.map(|state| state.to_device(device));
96        state
97    }
98}
99
100#[cfg(test)]
101mod tests {
102    use super::*;
103    use crate::{
104        TestAutodiffBackend, TestBackend,
105        grad_clipping::GradientClipping,
106        optim::{GradientsParams, Optimizer},
107    };
108    use ruda_model::tensor::{Distribution, Shape};
109    use ruda_nn::{Linear, LinearConfig};
110
111    const LEARNING_RATE: LearningRate = 0.02;
112
113    #[test]
114    fn with_updated_params_should_have_state() {
115        let device = Default::default();
116        let layer = layer::<TestAutodiffBackend>(&device);
117        let mut optim = sgd_with_all();
118        let loss = layer.forward(random_tensor::<TestAutodiffBackend>(&device));
119        let grads = loss.backward();
120        let grads = GradientsParams::from_grads(grads, &layer);
121        let _layer = optim.step(LEARNING_RATE, layer, grads);
122
123        let record = optim.to_record();
124
125        assert!(!record.is_empty());
126    }
127
128    #[test]
129    fn without_updated_params_should_not_have_state() {
130        let optim = sgd_with_all();
131        let record = optim.to_record();
132        assert!(record.is_empty());
133    }
134
135    #[test]
136    fn can_attach_gradient_clipping() {
137        let optim = sgd_with_all().with_grad_clipping(GradientClipping::Value(0.5));
138        assert!(optim.has_gradient_clipping());
139    }
140
141    #[test]
142    fn should_load_state() {
143        let device = Default::default();
144        let layer = layer::<TestAutodiffBackend>(&device);
145        let mut optim = sgd_with_all();
146        let loss = layer.forward(random_tensor(&device));
147        let grads = loss.backward();
148        let grads = GradientsParams::from_grads(grads, &layer);
149        let _layer = optim.step(LEARNING_RATE, layer, grads);
150
151        let record = optim.to_record();
152        let optim_new = sgd_with_all();
153        let record_new = optim_new.to_record();
154        let optim_new = optim_new.load_record(record.clone());
155        let state_restored = optim_new.to_record();
156
157        assert_ne!(record.len(), record_new.len());
158        assert_eq!(record.len(), state_restored.len());
159    }
160
161    fn random_tensor<B: Backend>(device: &B::Device) -> Tensor<B, 2> {
162        Tensor::<B, 2>::random(Shape::new([2, 20]), Distribution::Default, device)
163    }
164
165    fn layer<B: Backend>(device: &B::Device) -> Linear<B> {
166        LinearConfig::new(20, 20).init(device)
167    }
168
169    fn sgd_with_all()
170    -> OptimizerAdaptor<Sgd<TestBackend>, Linear<TestAutodiffBackend>, TestAutodiffBackend> {
171        SgdConfig {
172            weight_decay: Some(WeightDecayConfig { penalty: 0.05 }),
173            momentum: Some(MomentumConfig {
174                momentum: 0.9,
175                dampening: 0.1,
176                nesterov: true,
177            }),
178            gradient_clipping: None,
179        }
180        .init()
181    }
182}