minidx_core/layers/
residual.rs1use crate::Dtype;
2
3#[derive(Clone, Debug, Default)]
5pub struct Residual<E: Dtype, const I: usize, M: Default + crate::Module<[E; I]>> {
6 pub module: M,
7 pub dt: std::marker::PhantomData<E>,
8}
9
10impl<E: Dtype, const I: usize, M: Default + crate::Module<[E; I], Output = [E; I]>>
11 crate::Module<[E; I]> for Residual<E, I, M>
12{
13 type Output = M::Output;
14
15 fn forward(&self, x: &[E; I]) -> Result<Self::Output, crate::Error> {
16 let mut out = self.module.forward(x)?;
17 out.iter_mut().zip(x).for_each(|(o, &x)| *o += x);
18 Ok(out)
19 }
20}
21
22impl<
23 E: Dtype,
24 const I: usize,
25 M: Default + crate::Module<[E; I], Output = [E; I]> + crate::TracedModule<[E; I]>,
26 > crate::TracedModule<[E; I]> for Residual<E, I, M>
27{
28 type Trace = M::Trace;
29
30 fn traced_forward(
31 &self,
32 x: [E; I],
33 ) -> Result<(<Self as crate::Module<[E; I]>>::Output, Self::Trace), crate::Error> {
34 let (mut out, trace) = self.module.traced_forward(x)?;
35 out.iter_mut().zip(x).for_each(|(o, x)| *o += x);
36 Ok((out, trace))
37 }
38}
39
40impl<
41 E: Dtype,
42 const I: usize,
43 M: Default
44 + crate::Module<[E; I], Output = [E; I]>
45 + crate::TracedModule<[E; I]>
46 + crate::BackpropModule<[E; I]>,
47 > crate::BackpropModule<[E; I]> for Residual<E, I, M>
48{
49 type SelfGrads = M::SelfGrads;
50
51 fn backprop(
52 &self,
53 trace: &<M as crate::TracedModule<[E; I]>>::Trace,
54 grads_wrt_output: <M as crate::Module<[E; I]>>::Output,
55 ) -> ([E; I], Self::SelfGrads) {
56 let (mut out, mod_grads) = self.module.backprop(trace, grads_wrt_output.clone());
57 out.iter_mut()
58 .zip(grads_wrt_output.into_iter())
59 .for_each(|(o, x)| *o += x);
60 (out, mod_grads)
61 }
62
63 fn update(
64 &mut self,
65 applyer: &mut impl crate::optimizers::GradApplyer,
66 updates: Self::SelfGrads,
67 ) -> Result<(), crate::Error> {
68 self.module.update(applyer, updates)
69 }
70}
71
72impl<E: Dtype, const I: usize, M: Default + crate::Module<[E; I]> + crate::LoadableModule>
73 crate::LoadableModule for Residual<E, I, M>
74{
75 fn save(
76 &self,
77 path: String,
78 dict: &mut std::collections::HashMap<String, Vec<f64>>,
79 ) -> Result<(), crate::LoadSaveError> {
80 self.module.save(path + ".inner", dict)
81 }
82
83 fn load(
84 &mut self,
85 path: String,
86 dict: &std::collections::HashMap<String, Vec<f64>>,
87 ) -> Result<(), crate::LoadSaveError> {
88 self.module.load(path + ".inner", dict)
89 }
90}
91
92impl<E: Dtype, const I: usize, M: Default + crate::Module<[E; I]> + crate::ResetParams>
93 crate::ResetParams for Residual<E, I, M>
94{
95 fn rand_params<RNG: rand::Rng>(
96 &mut self,
97 rng: &mut RNG,
98 scale: f32,
99 ) -> Result<(), crate::Error> {
100 self.module.rand_params(rng, scale)?;
101 Ok(())
102 }
103}