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#[derive(Config, Debug, PartialEq)]
15pub enum Initializer {
16 Constant {
18 value: f64,
20 },
21 Ones,
23 Zeros,
25 Uniform {
27 min: f64,
29
30 max: f64,
32 },
33 Normal {
35 mean: f64,
37
38 std: f64,
40 },
41 KaimingUniform {
43 gain: f64,
45
46 fan_out_only: bool,
48 },
49 KaimingNormal {
51 gain: f64,
53
54 fan_out_only: bool,
56 },
57 XavierUniform {
61 gain: f64,
63 },
64 XavierNormal {
68 gain: f64,
70 },
71 Orthogonal {
75 gain: f64,
77 },
78}
79
80impl Initializer {
81 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 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 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 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 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 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 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 let q_matmul_r = qr.0.clone().matmul(qr.1.clone());
555
556 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 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.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 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 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 let shape = [3];
624 let _: Tensor<TB, 1> = Initializer::Orthogonal { gain }
625 .init(shape, &Default::default())
626 .into_value();
627 }
628}