Skip to main content

ruda_optim/optim/
adagrad.rs

1
2use ruda_model::{module::AutodiffModule, record::Record};
3
4use ruda_model::config::Config;
5use ruda_model::tensor::{Tensor, backend::AutodiffBackend};
6use ruda_model::tensor::{backend::Backend, ops::Device};
7
8use super::{
9    SimpleOptimizer,
10    adaptor::OptimizerAdaptor,
11    decay::{WeightDecay, WeightDecayConfig},
12};
13use crate::{LearningRate, grad_clipping::GradientClippingConfig};
14
15/// AdaGrad configuration.
16#[derive(Config, Debug)]
17pub struct AdaGradConfig {
18    #[config(default = 0.)]
19    lr_decay: f64,
20    #[config(default = 1e-5)]
21    epsilon: f32,
22    /// [Weight decay](WeightDecayConfig) config.
23    weight_decay: Option<WeightDecayConfig>,
24    /// [Gradient Clipping](GradientClippingConfig) config.
25    grad_clipping: Option<GradientClippingConfig>,
26}
27
28/// AdaGrad optimizer
29#[derive(Clone)]
30pub struct AdaGrad {
31    lr_decay: LrDecay,
32    weight_decay: Option<WeightDecay>,
33}
34
35/// AdaGrad state.
36#[derive(Record, Clone, new)]
37pub struct AdaGradState<B: Backend, const D: usize> {
38    lr_decay: LrDecayState<B, D>,
39}
40
41impl<B: Backend> SimpleOptimizer<B> for AdaGrad {
42    type State<const D: usize> = AdaGradState<B, D>;
43
44    fn step<const D: usize>(
45        &self,
46        lr: LearningRate,
47        tensor: Tensor<B, D>,
48        mut grad: Tensor<B, D>,
49        state: Option<Self::State<D>>,
50    ) -> (Tensor<B, D>, Option<Self::State<D>>) {
51        let mut state_lr_decay = None;
52
53        if let Some(state) = state {
54            state_lr_decay = Some(state.lr_decay);
55        }
56
57        if let Some(weight_decay) = &self.weight_decay {
58            grad = weight_decay.transform(grad, tensor.clone());
59        }
60
61        let (grad, state_lr_decay) = self.lr_decay.transform(grad, lr, state_lr_decay);
62
63        let state = AdaGradState::new(state_lr_decay);
64
65        (tensor - grad, Some(state))
66    }
67
68    fn to_device<const D: usize>(mut state: Self::State<D>, device: &Device<B>) -> Self::State<D> {
69        state.lr_decay = state.lr_decay.to_device(device);
70        state
71    }
72}
73
74impl AdaGradConfig {
75    /// Build an [`AdaGrad`] from the config.
76    pub fn build(&self) -> AdaGrad {
77        AdaGrad {
78            lr_decay: LrDecay {
79                lr_decay: self.lr_decay,
80                epsilon: self.epsilon,
81            },
82            weight_decay: self.weight_decay.as_ref().map(WeightDecay::new),
83        }
84    }
85
86    /// Initialize AdaGrad optimizer.
87    ///
88    /// # Returns
89    ///
90    /// Returns an optimizer that can be used to optimize a module.
91    pub fn init<B: AutodiffBackend, M: AutodiffModule<B>>(
92        &self,
93    ) -> OptimizerAdaptor<AdaGrad, M, B> {
94        let mut optim = OptimizerAdaptor::from(self.build());
95        if let Some(config) = &self.grad_clipping {
96            optim = optim.with_grad_clipping(config.init());
97        }
98        optim
99    }
100}
101
102/// Learning rate decay state (also includes sum state).
103#[derive(Record, new, Clone)]
104pub struct LrDecayState<B: Backend, const D: usize> {
105    time: usize,
106    sum: Tensor<B, D>,
107}
108
109#[derive(Clone)]
110struct LrDecay {
111    lr_decay: f64,
112    epsilon: f32,
113}
114
115impl LrDecay {
116    pub fn transform<B: Backend, const D: usize>(
117        &self,
118        grad: Tensor<B, D>,
119        lr: LearningRate,
120        lr_decay_state: Option<LrDecayState<B, D>>,
121    ) -> (Tensor<B, D>, LrDecayState<B, D>) {
122        let state = if let Some(mut state) = lr_decay_state {
123            state.sum = state.sum.add(grad.clone().square());
124            state.time += 1;
125            state
126        } else {
127            LrDecayState::new(1, grad.clone().square())
128        };
129
130        let new_lr = lr / (1. + (state.time as f64 - 1.) * self.lr_decay);
131
132        let grad = grad
133            .div(state.sum.clone().sqrt().add_scalar(self.epsilon))
134            .mul_scalar(new_lr);
135
136        (grad, state)
137    }
138}
139
140impl<B: Backend, const D: usize> LrDecayState<B, D> {
141    /// Move state to device.
142    ///
143    /// # Arguments
144    ///
145    /// * `device` - Device to move state to.
146    ///
147    /// # Returns
148    ///
149    /// Returns state moved to device.
150    pub fn to_device(mut self, device: &B::Device) -> Self {
151        self.sum = self.sum.to_device(device);
152        self
153    }
154}
155
156#[cfg(test)]
157mod tests {
158    use ruda_model::tensor::Tolerance;
159    use ruda_model::tensor::ops::FloatElem;
160
161    use super::*;
162    use crate::TestAutodiffBackend;
163    use crate::{GradientsParams, Optimizer};
164    use ruda_model::module::{Module, Param};
165    use ruda_model::tensor::{Distribution, Tensor, TensorData};
166    use ruda_nn::{Linear, LinearConfig, LinearRecord};
167
168    const LEARNING_RATE: LearningRate = 0.01;
169
170    #[test]
171    fn test_adagrad_optimizer_save_load_state() {
172        let device = Default::default();
173        let linear = LinearConfig::new(6, 6).init(&device);
174        let x = Tensor::<TestAutodiffBackend, 2>::random([2, 6], Distribution::Default, &device);
175        let mut optimizer = create_adagrad();
176        let grads = linear.forward(x).backward();
177        let grads = GradientsParams::from_grads(grads, &linear);
178        let _linear = optimizer.step(LEARNING_RATE, linear, grads);
179
180        #[cfg(feature = "std")]
181        {
182            use ruda_model::record::{BinFileRecorder, FullPrecisionSettings, Recorder};
183
184            BinFileRecorder::<FullPrecisionSettings>::default()
185                .record(
186                    optimizer.to_record(),
187                    std::env::temp_dir().as_path().join("test_optim_adagrad"),
188                )
189                .unwrap();
190        }
191        #[cfg(not(feature = "std"))]
192        {
193            use ruda_model::record::{BinBytesRecorder, FullPrecisionSettings, Recorder};
194
195            let result = BinBytesRecorder::<FullPrecisionSettings>::default()
196                .record(optimizer.to_record(), ())
197                .unwrap();
198            assert!(!result.is_empty());
199        }
200
201        let state_optim_before = optimizer.to_record();
202        let state_optim_before_copy = optimizer.to_record();
203        let optimizer = create_adagrad();
204        let optimizer = optimizer.load_record(state_optim_before_copy);
205        let state_optim_after = optimizer.to_record();
206
207        assert_eq!(state_optim_before.len(), state_optim_after.len());
208    }
209
210    #[test]
211    fn test_adagrad_optimizer_with_numbers() {
212        let device = Default::default();
213        let linear = given_linear_layer(
214            TensorData::from([
215                [-0.3206, 0.1374, 0.4043, 0.3200, 0.0859, 0.0671],
216                [0.0777, -0.0185, -0.3667, 0.2550, 0.1955, -0.2922],
217                [-0.0190, 0.0346, -0.2962, 0.2484, -0.2780, 0.3130],
218                [-0.2980, -0.2214, -0.3715, -0.2981, -0.0761, 0.1626],
219                [0.3300, -0.2182, 0.3717, -0.1729, 0.3796, -0.0304],
220                [-0.0159, -0.0120, 0.1258, 0.1921, 0.0293, 0.3833],
221            ]),
222            TensorData::from([-0.3905, 0.0884, -0.0970, 0.1176, 0.1366, 0.0130]),
223        );
224        let x_1 = Tensor::<TestAutodiffBackend, 2>::from_floats(
225            [
226                [0.6294, 0.0940, 0.8176, 0.8824, 0.5228, 0.4310],
227                [0.7152, 0.9559, 0.7893, 0.5684, 0.5939, 0.8883],
228            ],
229            &device,
230        )
231        .require_grad();
232        let x_2 = Tensor::<TestAutodiffBackend, 2>::from_floats(
233            [
234                [0.8491, 0.2108, 0.8939, 0.4433, 0.5527, 0.2528],
235                [0.3270, 0.0412, 0.5538, 0.9605, 0.3195, 0.9085],
236            ],
237            &device,
238        )
239        .require_grad();
240
241        let mut optimizer = AdaGradConfig::new()
242            .with_epsilon(1e-8)
243            .with_lr_decay(0.5)
244            .init();
245
246        let grads = linear.forward(x_1).backward();
247        let grads = GradientsParams::from_grads(grads, &linear);
248        let linear = optimizer.step(LEARNING_RATE, linear, grads);
249
250        let grads = linear.forward(x_2).backward();
251        let grads = GradientsParams::from_grads(grads, &linear);
252        let linear = optimizer.step(LEARNING_RATE, linear, grads);
253
254        let state_updated = linear.into_record();
255        let weights_expected = TensorData::from([
256            [-0.334989, 0.123011, 0.389911, 0.305611, 0.071511, 0.052711],
257            [
258                0.066144, -0.030056, -0.378256, 0.243444, 0.183944, -0.303756,
259            ],
260            [
261                -0.033462, 0.020138, -0.310662, 0.233938, -0.292462, 0.298538,
262            ],
263            [
264                -0.312636, -0.236036, -0.386136, -0.312736, -0.090736, 0.147964,
265            ],
266            [
267                0.315896, -0.232304, 0.357596, -0.187004, 0.365496, -0.044504,
268            ],
269            [-0.030305, -0.026405, 0.111395, 0.177695, 0.014895, 0.368895],
270        ]);
271        let bias_expected = TensorData::from([
272            -0.405214, 0.073686, -0.111714, 0.102886, 0.121886, -0.001714,
273        ]);
274
275        let (weight_updated, bias_updated) = (
276            state_updated.weight.val().into_data(),
277            state_updated.bias.unwrap().val().into_data(),
278        );
279
280        type FT = FloatElem<TestAutodiffBackend>;
281        let tolerance = Tolerance::absolute(1e-6);
282        bias_updated.assert_approx_eq::<FT>(&bias_expected, tolerance);
283        weight_updated.assert_approx_eq::<FT>(&weights_expected, tolerance);
284    }
285
286    fn given_linear_layer(weight: TensorData, bias: TensorData) -> Linear<TestAutodiffBackend> {
287        let device = Default::default();
288        let record = LinearRecord {
289            weight: Param::from_data(weight, &device),
290            bias: Some(Param::from_data(bias, &device)),
291        };
292
293        LinearConfig::new(6, 6).init(&device).load_record(record)
294    }
295
296    fn create_adagrad()
297    -> OptimizerAdaptor<AdaGrad, Linear<TestAutodiffBackend>, TestAutodiffBackend> {
298        let config = AdaGradConfig::new();
299        AdaGrad {
300            lr_decay: LrDecay {
301                lr_decay: config.lr_decay,
302                epsilon: config.epsilon,
303            },
304            weight_decay: config.weight_decay.as_ref().map(WeightDecay::new),
305        }
306        .into()
307    }
308}