Skip to main content

minidx_core/layers/
residual.rs

1use crate::Dtype;
2
3/// A residual connection around some module.
4#[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}