Skip to main content

minidx_core/layers/
rmsdiv.rs

1use crate::matmul::MatMulImpl;
2use crate::{Float, Shape};
3
4/// A layer which divides each feature by the RMS of all features.
5#[derive(Clone, Debug, Default)]
6pub struct RMSDiv<E: Float + MatMulImpl, const I: usize> {
7    pub(crate) marker: std::marker::PhantomData<[E; I]>,
8}
9
10impl<E: Float + MatMulImpl, const I: usize> RMSDiv<E, I> {
11    #[inline]
12    fn forward(&self, input: &[E; I]) -> [E; I] {
13        let i = E::from_usize(I).unwrap();
14        let rms: E = (input.iter().fold(E::default(), |acc, x| acc + (*x * *x)) / i).sqrt();
15
16        let mut out: [E; I] = input.clone();
17        for o in out.iter_mut() {
18            *o /= rms;
19        }
20        out
21    }
22
23    #[inline]
24    fn jacobian(&self, input: &[E; I]) -> [[E; I]; I] {
25        let n = E::from_usize(I).unwrap();
26        let rms: E = (input.iter().fold(E::default(), |acc, x| acc + (*x * *x))
27            / E::from_usize(I).unwrap())
28        .sqrt();
29
30        let mut jacobian: [[E; I]; I] = [[E::default(); I]; I];
31        for (i, iv) in input.iter().enumerate() {
32            for (j, jv) in input.iter().enumerate() {
33                jacobian[i][j] = if i == j {
34                    (n * (rms * rms) - (*iv * *iv)) / (n * rms * rms * rms)
35                } else {
36                    -*jv * *iv / (n * rms * rms * rms)
37                };
38            }
39        }
40        jacobian
41    }
42
43    fn backward(&self, output_gradients: &[E; I], jacobian: &[[E; I]; I]) -> [E; I] {
44        let mut out: [E; I] = [E::default(); I];
45
46        let (m, k) = (1, I);
47        let n = I;
48        E::matmul(
49            (m, k, n),
50            true,
51            output_gradients.as_ptr(),
52            Shape::strides(&(1, I)),
53            jacobian.as_ptr() as *const E,
54            Shape::strides(&(I, I)),
55            out.as_mut_ptr(),
56            Shape::strides(&(1, I)),
57        );
58
59        out
60    }
61}
62
63impl<E: Float + MatMulImpl, const I: usize> crate::BaseModule for RMSDiv<E, I> {}
64
65impl<E: Float + MatMulImpl, const I: usize> crate::Module<[E; I]> for RMSDiv<E, I> {
66    type Output = [E; I];
67
68    fn forward(&self, x: &[E; I]) -> Result<Self::Output, crate::Error> {
69        Ok(RMSDiv::forward(self, &x))
70    }
71}
72
73impl<E: Float + MatMulImpl, const I: usize> crate::RevModule<[E; I]> for RMSDiv<E, I> {
74    type SelfGrads = ();
75
76    fn reverse(&self, inputs: &[E; I], grads_wrt_output: &[E; I]) -> ([E; I], Self::SelfGrads) {
77        (
78            RMSDiv::backward(self, grads_wrt_output, &RMSDiv::jacobian(self, inputs)),
79            (),
80        )
81    }
82
83    fn apply(
84        &mut self,
85        _applyer: &mut impl crate::optimizers::GradApplyer,
86        _updates: Self::SelfGrads,
87    ) -> Result<(), crate::Error> {
88        Ok(())
89    }
90}
91
92impl<E: Float + MatMulImpl, const I: usize> crate::ResetParams for RMSDiv<E, I> {
93    fn rand_params<RNG: rand::Rng>(
94        &mut self,
95        _rng: &mut RNG,
96        _scale: f32,
97    ) -> Result<(), crate::Error> {
98        Ok(())
99    }
100}
101
102impl<E: Float + MatMulImpl, const I: usize> crate::VisualizableUnit for RMSDiv<E, I> {
103    const KIND: &'static str = "rmsdiv";
104    type Params = ();
105    fn params(&self) -> &Self::Params {
106        &()
107    }
108}
109
110impl<E: Float + MatMulImpl, const I: usize> crate::LoadableModule for RMSDiv<E, I> {
111    fn save(
112        &self,
113        _path: String,
114        _dict: &mut std::collections::HashMap<String, Vec<f64>>,
115    ) -> Result<(), crate::LoadSaveError> {
116        Ok(())
117    }
118
119    fn load(
120        &mut self,
121        _path: String,
122        _dict: &std::collections::HashMap<String, Vec<f64>>,
123    ) -> Result<(), crate::LoadSaveError> {
124        Ok(())
125    }
126}
127
128#[cfg(test)]
129mod tests {
130    use super::*;
131
132    #[test]
133    fn test_forward() {
134        let layer = RMSDiv::<f32, 2>::default();
135        assert_eq!(layer.forward(&[2.0, 2.0]), [1.0, 1.0],);
136    }
137
138    #[test]
139    fn test_backward() {
140        let layer = RMSDiv::<f32, 2>::default();
141        assert_eq!(
142            layer.backward(&[1.0, 1.0], &layer.jacobian(&[2.0, 2.0])),
143            [0.0, 0.0],
144        );
145        let grads = layer.backward(&[1.0, 1.0], &layer.jacobian(&[4.0, 8.0]));
146        assert!(grads[0] > 0.062, "[0]: {:?}", grads[0]);
147        assert!(grads[0] < 0.0634, "[0]: {:?}", grads[0]);
148        assert!(grads[1] > -0.0317, "[1]: {:?}", grads[1]);
149        assert!(grads[1] < -0.0315, "[1]: {:?}", grads[1]);
150    }
151}