Skip to main content

ruda_optim/optim/
decay.rs

1
2use ruda_model::config::Config;
3use ruda_model::record::Record;
4use ruda_model::tensor::Tensor;
5use ruda_model::tensor::backend::Backend;
6
7/// Configuration to create [weight decay](WeightDecay).
8#[derive(Config, Debug)]
9pub struct WeightDecayConfig {
10    /// L2 penalty.
11    pub penalty: f32,
12}
13
14/// State of [weight decay](WeightDecay).
15#[derive(Record, Clone, new)]
16pub struct WeightDecayState<B: Backend, const D: usize> {
17    pub(crate) grad_last_step: Tensor<B, D>,
18}
19
20/// Weight decay implementation that transforms gradients.
21#[derive(Clone)]
22pub struct WeightDecay {
23    penalty: f32,
24}
25
26impl WeightDecay {
27    /// Creates a new [weight decay](WeightDecay) from a [config](WeightDecayConfig).
28    pub fn new(config: &WeightDecayConfig) -> Self {
29        Self {
30            penalty: config.penalty,
31        }
32    }
33
34    /// Transforms a gradient.
35    ///
36    /// # Arguments
37    ///
38    /// * `grad` - Gradient to transform.
39    /// * `tensor` - Tensor param of the last iteration.
40    ///
41    /// # Returns
42    ///
43    /// * `grad` - Transformed gradient.
44    pub fn transform<B: Backend, const D: usize>(
45        &self,
46        grad: Tensor<B, D>,
47        tensor: Tensor<B, D>,
48    ) -> Tensor<B, D> {
49        tensor.mul_scalar(self.penalty).add(grad)
50    }
51}
52
53impl<B: Backend, const D: usize> WeightDecayState<B, D> {
54    /// Moves the state to a device.
55    ///
56    /// # Arguments
57    ///
58    /// * `device` - Device to move the state to.
59    ///
60    /// # Returns
61    ///
62    /// * `self` - Moved state.
63    pub fn to_device(mut self, device: &B::Device) -> Self {
64        self.grad_last_step = self.grad_last_step.to_device(device);
65        self
66    }
67}