1use crate::matmul::MatMulImpl;
2use crate::{Float, Shape};
3
4#[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}