Skip to main content

ruda_model/module/
initializer.rs

1use crate::tensor::Shape;
2
3use crate::config::Config;
4use crate::module::{Param, ParamId};
5use crate::tensor::backend::Backend;
6use crate::tensor::{Distribution, Tensor, s};
7
8
9#[cfg(not(feature = "std"))]
10#[allow(unused_imports)]
11use num_traits::Float as _;
12
13/// Enum specifying with what values a tensor should be initialized
14#[derive(Config, Debug, PartialEq)]
15pub enum Initializer {
16    /// Fills tensor with specified value everywhere
17    Constant {
18        /// The value to fill the tensor with
19        value: f64,
20    },
21    /// Fills tensor with 1s everywhere
22    Ones,
23    /// Fills tensor with 0s everywhere
24    Zeros,
25    /// Fills tensor with values drawn uniformly between specified values
26    Uniform {
27        /// The minimum value to draw from
28        min: f64,
29
30        /// The maximum value to draw from
31        max: f64,
32    },
33    /// Fills tensor with values drawn from normal distribution with specified mean and std
34    Normal {
35        /// The mean of the normal distribution
36        mean: f64,
37
38        /// The standard deviation of the normal distribution
39        std: f64,
40    },
41    /// Fills tensor with values according to the uniform version of Kaiming initialization
42    KaimingUniform {
43        /// The gain to use in initialization formula
44        gain: f64,
45
46        /// Whether to use fan out only in initialization formula
47        fan_out_only: bool,
48    },
49    /// Fills tensor with values according to the uniform version of Kaiming initialization
50    KaimingNormal {
51        /// The gain to use in initialization formula
52        gain: f64,
53
54        /// Whether to use fan out only in initialization formula
55        fan_out_only: bool,
56    },
57    /// Fills tensor with values according to the uniform version of Xavier Glorot initialization
58    /// described in [Understanding the difficulty of training deep feedforward neural networks
59    /// ](https://proceedings.mlr.press/v9/glorot10a/glorot10a.pdf)
60    XavierUniform {
61        /// The gain to use in initialization formula
62        gain: f64,
63    },
64    /// Fills tensor with values according to the normal version of Xavier Glorot initialization
65    /// described in [Understanding the difficulty of training deep feedforward neural networks
66    /// ](https://proceedings.mlr.press/v9/glorot10a/glorot10a.pdf)
67    XavierNormal {
68        /// The gain to use in initialization formula
69        gain: f64,
70    },
71    /// Fills tensor with values according to the (semi) orthogonal initialization
72    /// described in [Exact solutions to the nonlinear dynamics of learning in deep linear neural networks`
73    ///  - [Saxe, A. et al. (2013)](https://arxiv.org/abs/1312.6120)
74    Orthogonal {
75        /// The gain to use in initialization formula
76        gain: f64,
77    },
78}
79
80impl Initializer {
81    /// Inits a tensor parameter of given shape with values depending on initializer kind.
82    ///
83    /// # Params
84    ///
85    /// - shape: Shape of the initiated tensor.
86    pub fn init<B: Backend, const D: usize, S: Into<Shape>>(
87        &self,
88        shape: S,
89        device: &B::Device,
90    ) -> Param<Tensor<B, D>> {
91        self.init_with(shape, None, None, device)
92    }
93
94    /// Inits a tensor parameter of given shape with values depending on initializer kind.
95    ///
96    /// # Params
97    ///
98    /// - shape: Shape of the initiated tensor.
99    pub fn init_with<B: Backend, const D: usize, S: Into<Shape>>(
100        &self,
101        shape: S,
102        fan_in: Option<usize>,
103        fan_out: Option<usize>,
104        device: &B::Device,
105    ) -> Param<Tensor<B, D>> {
106        let device = device.clone();
107        let shape: Shape = shape.into();
108        let config = self.clone();
109        let shape_for_closure = shape.clone();
110
111        Param::uninitialized(
112            ParamId::new(),
113            move |device, require_grad| {
114                let config = config.clone();
115                let shape = shape.clone();
116                B::memory_persistent_allocations(device, (), move |_| {
117                    let mut tensor = config.init_tensor(shape.clone(), fan_in, fan_out, device);
118
119                    if require_grad {
120                        tensor = tensor.require_grad();
121                    }
122
123                    tensor
124                })
125            },
126            device,
127            true,
128            shape_for_closure,
129        )
130    }
131
132    fn init_tensor<B: Backend, const D: usize, S: Into<Shape>>(
133        &self,
134        shape: S,
135        fan_in: Option<usize>,
136        fan_out: Option<usize>,
137        device: &B::Device,
138    ) -> Tensor<B, D> {
139        let shape = shape.into();
140        match self {
141            Initializer::Constant { value } => Tensor::<B, D>::full(shape, *value, device),
142            Initializer::Ones => Tensor::<B, D>::ones(shape, device),
143            Initializer::Zeros => Tensor::<B, D>::zeros(shape, device),
144            Initializer::Uniform { min, max } => uniform_draw(shape, *min, *max, device),
145            Initializer::Normal { mean, std } => normal_draw(shape, *mean, *std, device),
146            Initializer::KaimingUniform { gain, fan_out_only } => {
147                let a = 3.0f64.sqrt() * *gain * self.kaiming_std(*fan_out_only, fan_in, fan_out);
148                uniform_draw(shape, -a, a, device)
149            }
150            Initializer::KaimingNormal { gain, fan_out_only } => {
151                let std = *gain * self.kaiming_std(*fan_out_only, fan_in, fan_out);
152                normal_draw(shape, 0.0, std, device)
153            }
154            Initializer::XavierUniform { gain } => {
155                let a = 3.0f64.sqrt() * *gain * self.xavier_std(fan_in, fan_out);
156                uniform_draw(shape, -a, a, device)
157            }
158            Initializer::XavierNormal { gain } => {
159                let std = *gain * self.xavier_std(fan_in, fan_out);
160                normal_draw(shape, 0.0, std, device)
161            }
162            Initializer::Orthogonal { gain } => {
163                // following the implementation in pytorch:
164                // https://github.com/pytorch/pytorch/blob/v2.7.0/torch/nn/init.py#L574
165
166                assert!(
167                    D >= 2,
168                    "Expected D (in Tensor<B, D>) to be greater or equal 2; (D >= 2)"
169                );
170
171                let rows: usize = shape.dims::<D>()[0];
172                let cols: usize = shape.num_elements() / rows;
173
174                let mut t: Tensor<B, 2> = normal_draw([rows, cols], 0.0, 1.0, device);
175
176                if rows < cols {
177                    t = t.transpose();
178                }
179
180                let (q, r) = qr_decomposition(t, device);
181                let [r_rows, r_cols] = r.clone().dims();
182
183                let diag_r = Tensor::<B, 2>::ones([1, r_rows], device)
184                    .matmul(Tensor::<B, 2>::eye(r_cols, device).mul(r.clone()));
185
186                let ph = diag_r.clone().sign();
187
188                let mut q = q.mul(ph);
189
190                if rows < cols {
191                    q = q.transpose();
192                }
193
194                q.reshape(shape).mul_scalar(*gain)
195            }
196        }
197    }
198
199    fn kaiming_std(
200        &self,
201        fan_out_only: bool,
202        fan_in: Option<usize>,
203        fan_out: Option<usize>,
204    ) -> f64 {
205        let fan = if fan_out_only { fan_out } else { fan_in };
206        let fan = fan.expect(
207            "Can't use Kaiming initialization without specifying fan. Use init_with method.",
208        );
209
210        1.0 / (fan as f64).sqrt()
211    }
212
213    fn xavier_std(&self, fan_in: Option<usize>, fan_out: Option<usize>) -> f64 {
214        let fan_in = fan_in.expect(
215            "Can't use Xavier initialization without specifying fan in. Use init_with method and \
216             provide fan_in.",
217        );
218        let fan_out = fan_out.expect(
219            "Can't use Xavier initialization without specifying fan out. Use init_with method and \
220             provide fan_out.",
221        );
222        (2.0 / (fan_in + fan_out) as f64).sqrt()
223    }
224}
225
226fn uniform_draw<B: Backend, const D: usize, S: Into<Shape>>(
227    shape: S,
228    low: f64,
229    high: f64,
230    device: &B::Device,
231) -> Tensor<B, D> {
232    let distribution = Distribution::Uniform(low, high);
233    Tensor::<B, D>::random(shape, distribution, device)
234}
235
236fn normal_draw<B: Backend, const D: usize, S: Into<Shape>>(
237    shape: S,
238    mean: f64,
239    std: f64,
240    device: &B::Device,
241) -> Tensor<B, D> {
242    let distribution = Distribution::Normal(mean, std);
243    Tensor::<B, D>::random(shape, distribution, device)
244}
245
246fn qr_decomposition<B: Backend>(
247    a: Tensor<B, 2>,
248    device: &B::Device,
249) -> (Tensor<B, 2>, Tensor<B, 2>) {
250    // Calculate the QR decomposition using Gram-Schmidt-process: https://en.wikipedia.org/wiki/Gram%E2%80%93Schmidt_process
251
252    let [m, n] = a.clone().dims();
253    let mut q = Tensor::<B, 2>::zeros([m, n], device);
254    let mut r = Tensor::<B, 2>::zeros([n, n], device);
255
256    for j in 0..n {
257        let mut v: Tensor<B, 1> = a.clone().slice(s![.., j..=j]).squeeze_dim(1);
258
259        for i in 0..j {
260            let q_i: Tensor<B, 1> = q.clone().slice(s![.., i..=i]).squeeze_dim(1);
261            let r_ij = q_i.clone().mul(v.clone()).sum();
262
263            r = r
264                .clone()
265                .slice_assign([i..i + 1, j..j + 1], r_ij.clone().unsqueeze());
266
267            v = v - q_i.mul(r_ij);
268        }
269
270        // norm of v
271        let r_jj = v
272            .clone()
273            .powf(Tensor::from_floats([2.0], device))
274            .sum()
275            .sqrt();
276
277        r = r
278            .clone()
279            .slice_assign([j..j + 1, j..j + 1], r_jj.clone().unsqueeze());
280
281        let q_j = v / r_jj;
282
283        q = q
284            .clone()
285            .slice_assign([0..m, j..j + 1], q_j.unsqueeze_dim(1));
286    }
287
288    (q, r)
289}
290
291#[cfg(test)]
292mod tests {
293    use super::*;
294
295    use ruda_tensor::api::{ElementConversion, TensorData};
296    use num_traits::Pow;
297
298    pub type TB = ruda_tensor_host::Host;
299    use ruda_tensor::api::{Tolerance, ops::FloatElem};
300    type FT = FloatElem<TB>;
301
302    fn assert_normal_init(expected_mean: f64, expected_var: f64, tensor: &Tensor<TB, 2>) {
303        let (actual_vars, actual_means) = tensor.clone().var_mean(0);
304        let actual_vars = actual_vars.to_data();
305        let actual_vars = actual_vars.as_slice::<FT>().unwrap();
306        let actual_means = actual_means.to_data();
307        let actual_means = actual_means.as_slice::<FT>().unwrap();
308
309        for i in 0..tensor.shape()[0] {
310            let actual_var = actual_vars[i] as f64;
311            let actual_mean = actual_means[i] as f64;
312
313            assert!(
314                (expected_var - actual_var).abs() <= 0.1,
315                "Expected variance to be between {expected_var} += 0.1, but got {actual_var}"
316            );
317            assert!(
318                (expected_mean - actual_mean).abs() <= 0.1,
319                "Expected mean to be between {expected_mean} += 0.1, but got {actual_mean}"
320            );
321        }
322    }
323
324    #[test]
325    fn initializer_uniform_init() {
326        let device = Default::default();
327        TB::seed(&device, 0);
328
329        let (min, max) = (0.0, 1.0);
330        let uniform = Initializer::Uniform { min, max };
331        let tensor: Tensor<TB, 4> = uniform.init([2, 2, 2, 2], &Default::default()).into_value();
332
333        tensor
334            .into_data()
335            .assert_within_range::<FT>(min.elem()..max.elem());
336    }
337
338    #[test]
339    fn initializer_normal_init() {
340        // seed random generator
341        let device = Default::default();
342        TB::seed(&device, 0);
343
344        let (mean, std) = (0.0, 1.0);
345        let normal: Tensor<TB, 1> = Initializer::Normal { mean, std }
346            .init([10000], &Default::default())
347            .into_value();
348        let (var_act, mean_act) = normal.var_mean(0);
349
350        let var_act: f32 = var_act.into_scalar().elem();
351        let mean_act: f32 = mean_act.into_scalar().elem();
352
353        assert!(
354            var_act > 0.9 && var_act < 1.1,
355            "Expected variance to be between 1.0 += 0.1, but got {var_act}"
356        );
357        assert!(
358            mean_act > -0.1 && mean_act < 0.1,
359            "Expected mean to be between 0.0 += 0.1, but got {mean_act}"
360        );
361    }
362
363    #[test]
364    fn initializer_constant_init() {
365        let value = 5.0;
366        let constants: Tensor<TB, 4> = Initializer::Constant { value }
367            .init([2, 2, 2, 2], &Default::default())
368            .into_value();
369        constants.sum().to_data().assert_approx_eq::<FT>(
370            &TensorData::from([value as f32 * 16.0]),
371            Tolerance::default(),
372        );
373    }
374
375    #[test]
376    fn initializer_zeros_init() {
377        let zeros: Tensor<TB, 4> = Initializer::Zeros
378            .init([2, 2, 2, 2], &Default::default())
379            .into_value();
380        zeros
381            .sum()
382            .to_data()
383            .assert_approx_eq::<FT>(&TensorData::from([0.0]), Tolerance::default());
384    }
385
386    #[test]
387    fn initializer_ones_init() {
388        let ones: Tensor<TB, 4> = Initializer::Ones
389            .init([2, 2, 2, 2], &Default::default())
390            .into_value();
391        ones.sum()
392            .to_data()
393            .assert_approx_eq::<FT>(&TensorData::from([16.0]), Tolerance::default());
394    }
395
396    #[test]
397    fn initializer_kaiming_uniform_init() {
398        let device = Default::default();
399        TB::seed(&device, 0);
400
401        let gain = 2_f64;
402        let (fan_in, fan_out) = (5, 6);
403        let k = (gain * (3.0 / fan_in as f64).sqrt()).elem::<FT>();
404
405        let tensor: Tensor<TB, 2> = Initializer::KaimingUniform {
406            gain,
407            fan_out_only: false,
408        }
409        .init_with([fan_out, fan_in], Some(fan_in), None, &Default::default())
410        .into_value();
411        tensor.into_data().assert_within_range(-k..k);
412    }
413
414    #[test]
415    fn initializer_kaiming_normal_init() {
416        let device = Default::default();
417        TB::seed(&device, 0);
418
419        let gain = 2.;
420        let (fan_in, fan_out) = (1000, 10);
421        let expected_mean = 0_f64;
422
423        let expected_var = (gain * (1. / (fan_in as f64)).sqrt()).pow(2.);
424        let tensor: Tensor<TB, 2> = Initializer::KaimingNormal {
425            gain,
426            fan_out_only: false,
427        }
428        .init_with([fan_out, fan_in], Some(fan_in), None, &Default::default())
429        .into_value();
430        assert_normal_init(expected_mean, expected_var, &tensor)
431    }
432
433    #[test]
434    fn initializer_kaiming_uniform_init_bias() {
435        let device = Default::default();
436        TB::seed(&device, 0);
437
438        let gain = 2_f64;
439        let shape = [3];
440        let fan_in = 5;
441        let k = (gain * (3.0 / fan_in as f64).sqrt()).elem::<FT>();
442
443        let tensor: Tensor<TB, 1> = Initializer::KaimingUniform {
444            gain,
445            fan_out_only: false,
446        }
447        .init_with(shape, Some(fan_in), None, &Default::default())
448        .into_value();
449        tensor.into_data().assert_within_range(-k..k);
450    }
451
452    #[test]
453    fn initializer_kaiming_uniform_init_fan_out() {
454        let device = Default::default();
455        TB::seed(&device, 0);
456
457        let gain = 2_f64;
458        let (fan_in, fan_out) = (5, 6);
459        let k = (gain * (3.0 / fan_out as f64).sqrt()).elem::<FT>();
460
461        let tensor: Tensor<TB, 2> = Initializer::KaimingUniform {
462            gain,
463            fan_out_only: true,
464        }
465        .init_with([fan_out, fan_in], None, Some(fan_out), &Default::default())
466        .into_value();
467        tensor.into_data().assert_within_range(-k..k);
468    }
469
470    #[test]
471    #[should_panic]
472    fn initializer_kaiming_uniform_no_fan() {
473        let device = Default::default();
474        TB::seed(&device, 0);
475
476        let gain = 2_f64;
477        let (fan_in, fan_out) = (5, 6);
478
479        let _: Tensor<TB, 2> = Initializer::KaimingUniform {
480            gain,
481            fan_out_only: false,
482        }
483        .init([fan_out, fan_in], &Default::default())
484        .into_value();
485    }
486
487    #[test]
488    fn initializer_xavier_uniform_init() {
489        let device = Default::default();
490        TB::seed(&device, 0);
491
492        let gain = 2.;
493        let (fan_in, fan_out) = (5, 6);
494        let bound = (gain * (6. / (fan_in + fan_out) as f64).sqrt()).elem::<FT>();
495        let tensor: Tensor<TB, 2> = Initializer::XavierUniform { gain }
496            .init_with(
497                [fan_out, fan_in],
498                Some(fan_in),
499                Some(fan_out),
500                &Default::default(),
501            )
502            .into_value();
503
504        tensor.into_data().assert_within_range(-bound..bound);
505    }
506
507    #[test]
508    fn initializer_xavier_normal_init() {
509        let device = Default::default();
510        TB::seed(&device, 0);
511
512        let gain = 2.;
513        let (fan_in, fan_out) = (1000, 10);
514        let expected_mean = 0_f64;
515
516        let expected_var = (gain * (2. / (fan_in as f64 + fan_out as f64)).sqrt()).powf(2.);
517        let tensor: Tensor<TB, 2> = Initializer::XavierNormal { gain }
518            .init_with(
519                [fan_out, fan_in],
520                Some(fan_in),
521                Some(fan_out),
522                &Default::default(),
523            )
524            .into_value();
525        assert_normal_init(expected_mean, expected_var, &tensor)
526    }
527
528    #[test]
529    #[should_panic]
530    fn initializer_xavier_uniform_no_fan() {
531        let device = Default::default();
532        TB::seed(&device, 0);
533
534        let gain = 2.;
535        let (fan_in, fan_out) = (5, 6);
536        let _: Tensor<TB, 2> = Initializer::XavierUniform { gain }
537            .init([fan_out, fan_in], &Default::default())
538            .into_value();
539    }
540
541    #[test]
542    fn test_qr_decomposition() {
543        let device = Default::default();
544        TB::seed(&device, 0);
545
546        // test values follow the example from https://pytorch.org/docs/stable/generated/torch.linalg.qr.html#torch.linalg.qr
547        let a = Tensor::<TB, 2>::from_floats(
548            [[12., -51., 4.], [6., 167., -68.], [-4., 24., -41.]],
549            &Default::default(),
550        );
551        let qr = qr_decomposition(a.clone(), &Default::default());
552
553        // Q @ R should reconstruct input `a`
554        let q_matmul_r = qr.0.clone().matmul(qr.1.clone());
555
556        // assert that the difference between input (`a`) and Q @ R is (almost) zero
557        q_matmul_r
558            .into_data()
559            .assert_approx_eq::<FT>(&a.into_data(), Tolerance::rel_abs(0.1, 0.1));
560    }
561
562    #[test]
563    fn initializer_orthogonal_correct() {
564        let device = Default::default();
565        TB::seed(&device, 0);
566
567        let gain = 1.;
568
569        // test 2D tensor
570        let size = 10;
571        let q: Tensor<TB, 2> = Initializer::Orthogonal { gain }
572            .init([size, size], &Default::default())
573            .into_value();
574        let eye = Tensor::<TB, 2>::eye(size, &Default::default());
575
576        // Q.T @ Q should be close to identity matrix
577        q.clone()
578            .transpose()
579            .matmul(q)
580            .into_data()
581            .assert_approx_eq::<FT>(&eye.into_data(), Tolerance::rel_abs(0.1, 0.1));
582    }
583
584    #[test]
585    fn initializer_orthogonal_init() {
586        let device = Default::default();
587        TB::seed(&device, 0);
588
589        let gain = 1.;
590
591        // test 2D tensor
592        let shape = [25, 30];
593        let t: Tensor<TB, 2> = Initializer::Orthogonal { gain }
594            .init(shape, &Default::default())
595            .into_value();
596        let dims = t.dims();
597        assert_eq!(
598            shape, dims,
599            "Expected the shape of the input tensor to match the shape of the output. ({shape:?}, {dims:?})"
600        );
601
602        // test 3D tensor
603        let shape = [24, 6, 85];
604        let t: Tensor<TB, 3> = Initializer::Orthogonal { gain }
605            .init(shape, &Default::default())
606            .into_value();
607        let dims = t.dims();
608        assert_eq!(
609            shape, dims,
610            "Expected the shape of the input tensor to match the shape of the output. ({shape:?}, {dims:?})"
611        );
612    }
613
614    #[test]
615    #[should_panic]
616    fn initializer_orthogonal_init_1d() {
617        let device = Default::default();
618        TB::seed(&device, 0);
619
620        let gain = 1.;
621
622        // test 1D tensor
623        let shape = [3];
624        let _: Tensor<TB, 1> = Initializer::Orthogonal { gain }
625            .init(shape, &Default::default())
626            .into_value();
627    }
628}