minidx_core/layers/
lr_modifier.rs1use crate::{Dtype, Gradients};
2
3#[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}