Skip to main content

minidx_core/layers/
lr_modifier.rs

1use crate::{Dtype, Gradients};
2
3/// A wrapper which locally adjusts the learning rate.
4#[derive(Clone, Debug, Default)]
5pub struct LR<E: Dtype, const I: usize, M: Default + crate::Module<[E; I]>> {
6    pub module: M,
7    pub update_multiplier: f32,
8    pub dt: std::marker::PhantomData<E>,
9}
10
11impl<E: Dtype, const I: usize, M: Default + crate::Module<[E; I]>> crate::Module<[E; I]>
12    for LR<E, I, M>
13{
14    type Output = M::Output;
15
16    fn forward(&self, x: &[E; I]) -> Result<Self::Output, crate::Error> {
17        self.module.forward(x)
18    }
19}
20
21impl<
22        E: Dtype,
23        const I: usize,
24        M: Default + crate::Module<[E; I]> + crate::TracedModule<[E; I]>,
25    > crate::TracedModule<[E; I]> for LR<E, I, M>
26{
27    type Trace = M::Trace;
28
29    fn traced_forward(
30        &self,
31        x: [E; I],
32    ) -> Result<(<Self as crate::Module<[E; I]>>::Output, Self::Trace), crate::Error> {
33        self.module.traced_forward(x)
34    }
35}
36
37impl<
38        E: Dtype,
39        const I: usize,
40        M: Default
41            + crate::Module<[E; I]>
42            + crate::TracedModule<[E; I]>
43            + crate::BackpropModule<[E; I]>,
44    > crate::BackpropModule<[E; I]> for LR<E, I, M>
45where
46    M::SelfGrads: Gradients,
47{
48    type SelfGrads = M::SelfGrads;
49
50    fn backprop(
51        &self,
52        trace: &<M as crate::TracedModule<[E; I]>>::Trace,
53        grads_wrt_output: <M as crate::Module<[E; I]>>::Output,
54    ) -> ([E; I], Self::SelfGrads) {
55        self.module.backprop(trace, grads_wrt_output)
56    }
57
58    fn update(
59        &mut self,
60        applyer: &mut impl crate::optimizers::GradApplyer,
61        mut updates: Self::SelfGrads,
62    ) -> Result<(), crate::Error> {
63        use num_traits::FromPrimitive;
64        let m = <Self::SelfGrads as Gradients>::Concrete::from_f32(self.update_multiplier).unwrap();
65        updates.grad_iter_mut().for_each(|u| *u *= m);
66
67        self.module.update(applyer, updates)
68    }
69}
70
71impl<E: Dtype, const I: usize, M: Default + crate::Module<[E; I]> + crate::ResetParams>
72    crate::ResetParams for LR<E, I, M>
73{
74    fn rand_params<RNG: rand::Rng>(
75        &mut self,
76        rng: &mut RNG,
77        scale: f32,
78    ) -> Result<(), crate::Error> {
79        self.module.rand_params(rng, scale)
80    }
81}
82
83impl<E: Dtype, const I: usize, M: Default + crate::Module<[E; I]>> crate::LoadableModule
84    for LR<E, I, M>
85{
86    fn save(
87        &self,
88        _path: String,
89        _dict: &mut std::collections::HashMap<String, Vec<f64>>,
90    ) -> Result<(), crate::LoadSaveError> {
91        Ok(())
92    }
93
94    fn load(
95        &mut self,
96        _path: String,
97        _dict: &std::collections::HashMap<String, Vec<f64>>,
98    ) -> Result<(), crate::LoadSaveError> {
99        Ok(())
100    }
101}
102
103impl<E: Dtype, const I: usize, M: Default + crate::Module<[E; I]> + crate::VisualizableUnit>
104    crate::VisualizableUnit for LR<E, I, M>
105{
106    const KIND: &'static str = M::KIND;
107    type Params = M::Params;
108    fn params(&self) -> &Self::Params {
109        self.module.params()
110    }
111}