1use crate::TorshResult;
30use torsh_core::{DeviceType, TorshError};
31use torsh_tensor::Tensor;
32
33#[derive(Debug, Clone)]
35pub struct RandomizedConfig {
36 pub target_rank: usize,
38 pub oversampling: usize,
40 pub n_power_iter: usize,
42 pub random_seed: Option<u64>,
44 pub tolerance: f32,
46}
47
48impl Default for RandomizedConfig {
49 fn default() -> Self {
50 Self {
51 target_rank: 10,
52 oversampling: 10,
53 n_power_iter: 2,
54 random_seed: None,
55 tolerance: 1e-6,
56 }
57 }
58}
59
60impl RandomizedConfig {
61 pub fn fast(target_rank: usize) -> Self {
63 Self {
64 target_rank,
65 oversampling: 5,
66 n_power_iter: 0,
67 random_seed: None,
68 tolerance: 1e-4,
69 }
70 }
71
72 pub fn accurate(target_rank: usize) -> Self {
74 Self {
75 target_rank,
76 oversampling: 20,
77 n_power_iter: 4,
78 random_seed: None,
79 tolerance: 1e-8,
80 }
81 }
82
83 pub fn with_seed(mut self, seed: u64) -> Self {
85 self.random_seed = Some(seed);
86 self
87 }
88}
89
90fn generate_random_matrix(
95 rows: usize,
96 cols: usize,
97 device: DeviceType,
98 _seed: Option<u64>,
99) -> TorshResult<Tensor> {
100 let mut data = Vec::with_capacity(rows * cols);
103 let scale = 1.0 / (rows as f32).sqrt();
104
105 for i in 0..(rows * cols) {
107 let x = ((i as f32 * 12.9898 + 78.233).sin() * 43758.5453).fract();
109 let y = ((i as f32 * 93.9898 + 47.233).sin() * 29341.5453).fract();
110
111 let r = (-2.0 * x.ln()).sqrt();
113 let theta = 2.0 * std::f32::consts::PI * y;
114 let val = r * theta.cos() * scale;
115
116 data.push(val);
117 }
118
119 Tensor::from_data(data, vec![rows, cols], device)
120}
121
122pub fn randomized_range_finder(matrix: &Tensor, config: &RandomizedConfig) -> TorshResult<Tensor> {
143 if matrix.shape().ndim() != 2 {
144 return Err(TorshError::InvalidArgument(
145 "Randomized range finder requires 2D matrix".to_string(),
146 ));
147 }
148
149 let (m, n) = (matrix.shape().dims()[0], matrix.shape().dims()[1]);
150 let ell = config.target_rank + config.oversampling;
151
152 if ell > n.min(m) {
153 return Err(TorshError::InvalidArgument(format!(
154 "Target rank + oversampling ({}) exceeds matrix dimensions ({})",
155 ell,
156 n.min(m)
157 )));
158 }
159
160 let omega = generate_random_matrix(n, ell, matrix.device(), config.random_seed)?;
162
163 let mut y = matrix.matmul(&omega)?;
165
166 for _ in 0..config.n_power_iter {
169 let at = matrix.t()?;
170 let z = at.matmul(&y)?;
171 y = matrix.matmul(&z)?;
172 }
173
174 let (q, _) = crate::decomposition::qr(&y)?;
176
177 Ok(q)
179}
180
181pub fn randomized_qb(matrix: &Tensor, config: &RandomizedConfig) -> TorshResult<(Tensor, Tensor)> {
196 let q = randomized_range_finder(matrix, config)?;
198
199 let qt = q.t()?;
201 let b = qt.matmul(matrix)?;
202
203 Ok((q, b))
204}
205
206pub fn randomized_svd(
237 matrix: &Tensor,
238 config: &RandomizedConfig,
239) -> TorshResult<(Tensor, Tensor, Tensor)> {
240 if matrix.shape().ndim() != 2 {
241 return Err(TorshError::InvalidArgument(
242 "Randomized SVD requires 2D matrix".to_string(),
243 ));
244 }
245
246 let (q, b) = randomized_qb(matrix, config)?;
248
249 let (u_b, s, vt) = crate::decomposition::svd(&b, false)?;
251
252 let u = q.matmul(&u_b)?;
254
255 let k = config.target_rank;
257 let (m, _) = (u.shape().dims()[0], u.shape().dims()[1]);
258 let n_vt = vt.shape().dims()[0];
259
260 let mut u_k_data = vec![0.0f32; m * k];
262 for i in 0..m {
263 for j in 0..k {
264 u_k_data[i * k + j] = u.get(&[i, j])?;
265 }
266 }
267 let u_k = Tensor::from_data(u_k_data, vec![m, k], matrix.device())?;
268
269 let s_len = s.shape().dims()[0].min(k);
271 let mut s_k_data = vec![0.0f32; k];
272 for i in 0..s_len {
273 s_k_data[i] = s.get(&[i])?;
274 }
275 let s_k = Tensor::from_data(s_k_data, vec![k], matrix.device())?;
276
277 let n = vt.shape().dims()[1];
279 let mut vt_k_data = vec![0.0f32; k * n];
280 for i in 0..k.min(n_vt) {
281 for j in 0..n {
282 vt_k_data[i * n + j] = vt.get(&[i, j])?;
283 }
284 }
285 let vt_k = Tensor::from_data(vt_k_data, vec![k, n], matrix.device())?;
286
287 Ok((u_k, s_k, vt_k))
288}
289
290pub fn low_rank_approximation(
305 matrix: &Tensor,
306 rank: usize,
307 config: Option<&RandomizedConfig>,
308) -> TorshResult<Tensor> {
309 let default_config = RandomizedConfig::default();
310 let cfg = config.unwrap_or(&default_config);
311
312 let mut cfg_modified = cfg.clone();
313 cfg_modified.target_rank = rank;
314
315 let (u, s, vt) = randomized_svd(matrix, &cfg_modified)?;
316
317 let k = s.shape().dims()[0];
320 let m = u.shape().dims()[0];
321 let mut u_s_data = vec![0.0f32; m * k];
322
323 for i in 0..m {
324 for j in 0..k {
325 let u_val = u.get(&[i, j])?;
326 let s_val = s.get(&[j])?;
327 u_s_data[i * k + j] = u_val * s_val;
328 }
329 }
330
331 let u_s = Tensor::from_data(u_s_data, vec![m, k], matrix.device())?;
332
333 u_s.matmul(&vt)
335}
336
337pub fn estimate_rank(matrix: &Tensor, config: &RandomizedConfig) -> TorshResult<usize> {
351 let (_, s, _) = randomized_svd(matrix, config)?;
352
353 let s_len = s.shape().dims()[0];
354 let mut rank = 0;
355
356 for i in 0..s_len {
357 let sv = s.get(&[i])?;
358 if sv.abs() > config.tolerance {
359 rank += 1;
360 }
361 }
362
363 Ok(rank)
364}
365
366pub fn approximation_error(matrix: &Tensor, approximation: &Tensor) -> TorshResult<f32> {
380 let diff = matrix.sub(approximation)?;
381 crate::matrix_functions::matrix_norm(&diff, Some("fro"))
382}
383
384pub fn randomized_trace(matrix: &Tensor, num_samples: usize) -> TorshResult<f32> {
403 if matrix.shape().ndim() != 2 {
404 return Err(TorshError::InvalidArgument(
405 "Trace estimation requires 2D matrix".to_string(),
406 ));
407 }
408
409 let (m, n) = (matrix.shape().dims()[0], matrix.shape().dims()[1]);
410 if m != n {
411 return Err(TorshError::InvalidArgument(
412 "Trace estimation requires square matrix".to_string(),
413 ));
414 }
415
416 let mut trace_sum = 0.0f32;
417
418 for i in 0..num_samples {
419 let mut v_data = vec![0.0f32; n];
421 for j in 0..n {
422 let hash = ((i * n + j) as f32 * 12.9898).sin() * 43758.5453;
423 v_data[j] = if hash.fract() > 0.5 { 1.0 } else { -1.0 };
424 }
425 let v = Tensor::from_data(v_data, vec![n], matrix.device())?;
426
427 let av = matrix.matmul(&v.unsqueeze(1)?)?;
429 let av = av.squeeze(1)?;
430
431 let mut vt_av = 0.0f32;
433 for j in 0..n {
434 vt_av += v.get(&[j])? * av.get(&[j])?;
435 }
436
437 trace_sum += vt_av;
438 }
439
440 Ok(trace_sum / num_samples as f32)
441}
442
443#[cfg(test)]
444mod tests {
445 use super::*;
446 use approx::assert_relative_eq;
447
448 fn create_low_rank_matrix() -> TorshResult<Tensor> {
449 let u = vec![1.0f32, 2.0, 3.0, 4.0];
452 let v = vec![1.0f32, 2.0, 3.0];
453
454 let mut data = vec![0.0f32; 12]; for i in 0..4 {
456 for j in 0..3 {
457 data[i * 3 + j] = u[i] * v[j];
458 }
459 }
460
461 Tensor::from_data(data, vec![4, 3], DeviceType::Cpu)
462 }
463
464 #[test]
465 fn test_generate_random_matrix() -> TorshResult<()> {
466 let random_mat = generate_random_matrix(10, 5, DeviceType::Cpu, Some(42))?;
467
468 assert_eq!(random_mat.shape().dims(), &[10, 5]);
469
470 let mut has_nonzero = false;
472 let mut max_abs = 0.0f32;
473
474 for i in 0..10 {
475 for j in 0..5 {
476 let val = random_mat.get(&[i, j])?;
477 if val.abs() > 0.001 {
478 has_nonzero = true;
479 }
480 max_abs = max_abs.max(val.abs());
481 }
482 }
483
484 assert!(has_nonzero);
485 assert!(max_abs < 10.0); Ok(())
488 }
489
490 #[test]
491 #[ignore] fn test_randomized_range_finder() -> TorshResult<()> {
493 let matrix = create_low_rank_matrix()?;
494 let config = RandomizedConfig {
495 target_rank: 2,
496 oversampling: 1,
497 n_power_iter: 0, random_seed: Some(42),
499 tolerance: 1e-6,
500 };
501
502 let q = randomized_range_finder(&matrix, &config)?;
503
504 assert_eq!(q.shape().dims()[0], 4); let k = q.shape().dims()[1];
509 assert!(k > 0);
510 assert!(k <= 3);
511
512 for i in 0..4 {
514 for j in 0..k {
515 let val = q.get(&[i, j])?;
516 assert!(
517 val.is_finite(),
518 "Q contains non-finite value at ({}, {})",
519 i,
520 j
521 );
522 }
523 }
524
525 Ok(())
526 }
527
528 #[test]
529 #[ignore] fn test_randomized_qb() -> TorshResult<()> {
531 let matrix = create_low_rank_matrix()?;
532 let config = RandomizedConfig {
533 target_rank: 2,
534 oversampling: 1,
535 n_power_iter: 0, random_seed: Some(42),
537 tolerance: 1e-6,
538 };
539
540 let (q, b) = randomized_qb(&matrix, &config)?;
541
542 assert_eq!(q.shape().dims()[0], 4); assert_eq!(b.shape().dims()[1], 3); let approx = q.matmul(&b)?;
548 assert_eq!(approx.shape().dims(), matrix.shape().dims());
549
550 for i in 0..4 {
552 for j in 0..3 {
553 let val = approx.get(&[i, j])?;
554 assert!(val.is_finite());
555 }
556 }
557
558 Ok(())
559 }
560
561 #[test]
562 #[ignore] fn test_randomized_svd() -> TorshResult<()> {
564 let matrix = create_low_rank_matrix()?;
565 let config = RandomizedConfig {
566 target_rank: 2,
567 oversampling: 1,
568 n_power_iter: 0, random_seed: Some(42),
570 tolerance: 1e-6,
571 };
572
573 let (u, s, vt) = randomized_svd(&matrix, &config)?;
574
575 assert_eq!(u.shape().dims()[0], 4); assert_eq!(u.shape().dims()[1], 2); assert_eq!(s.shape().dims()[0], 2); assert_eq!(vt.shape().dims()[0], 2); assert_eq!(vt.shape().dims()[1], 3); for i in 0..2 {
584 let sv = s.get(&[i])?;
585 assert!(sv.is_finite());
586 }
587
588 Ok(())
589 }
590
591 #[test]
592 #[ignore] fn test_low_rank_approximation() -> TorshResult<()> {
594 let matrix = create_low_rank_matrix()?;
595 let config = RandomizedConfig {
596 target_rank: 2,
597 oversampling: 1,
598 n_power_iter: 0, random_seed: Some(42),
600 tolerance: 1e-6,
601 };
602 let approx = low_rank_approximation(&matrix, 2, Some(&config))?;
603
604 assert_eq!(approx.shape().dims(), matrix.shape().dims());
605
606 for i in 0..4 {
608 for j in 0..3 {
609 let val = approx.get(&[i, j])?;
610 assert!(val.is_finite());
611 }
612 }
613
614 Ok(())
615 }
616
617 #[test]
618 #[ignore] fn test_estimate_rank() -> TorshResult<()> {
620 let matrix = create_low_rank_matrix()?;
621 let config = RandomizedConfig {
622 target_rank: 2,
623 oversampling: 1,
624 n_power_iter: 0, random_seed: Some(42),
626 tolerance: 0.1, };
628
629 let estimated_rank = estimate_rank(&matrix, &config)?;
630
631 assert!(estimated_rank > 0);
633 assert!(estimated_rank <= 2);
634
635 Ok(())
636 }
637
638 #[test]
639 fn test_randomized_trace() -> TorshResult<()> {
640 let data = vec![1.0f32, 0.0, 0.0, 0.0, 2.0, 0.0, 0.0, 0.0, 3.0];
642 let matrix = Tensor::from_data(data, vec![3, 3], DeviceType::Cpu)?;
643
644 let estimated_trace = randomized_trace(&matrix, 100)?;
645
646 assert_relative_eq!(estimated_trace, 6.0, epsilon = 1.0);
648
649 Ok(())
650 }
651
652 #[test]
653 fn test_config_builders() -> TorshResult<()> {
654 let fast = RandomizedConfig::fast(5);
655 assert_eq!(fast.target_rank, 5);
656 assert_eq!(fast.n_power_iter, 0);
657
658 let accurate = RandomizedConfig::accurate(10);
659 assert_eq!(accurate.target_rank, 10);
660 assert_eq!(accurate.n_power_iter, 4);
661
662 let with_seed = RandomizedConfig::default().with_seed(123);
663 assert_eq!(with_seed.random_seed, Some(123));
664
665 Ok(())
666 }
667
668 #[test]
669 fn test_approximation_error() -> TorshResult<()> {
670 let matrix = Tensor::from_data(vec![1.0f32, 2.0, 3.0, 4.0], vec![2, 2], DeviceType::Cpu)?;
671
672 let approx = Tensor::from_data(vec![1.1f32, 2.1, 3.1, 4.1], vec![2, 2], DeviceType::Cpu)?;
673
674 let error = approximation_error(&matrix, &approx)?;
675
676 assert_relative_eq!(error, 0.2, epsilon = 1e-5);
678
679 Ok(())
680 }
681
682 #[test]
683 fn test_error_cases() -> TorshResult<()> {
684 let vec1d = Tensor::from_data(vec![1.0f32, 2.0, 3.0], vec![3], DeviceType::Cpu)?;
686 let config = RandomizedConfig::default();
687
688 assert!(randomized_range_finder(&vec1d, &config).is_err());
689 assert!(randomized_svd(&vec1d, &config).is_err());
690
691 let matrix = Tensor::from_data(vec![1.0f32, 2.0, 3.0, 4.0], vec![2, 2], DeviceType::Cpu)?;
693
694 let bad_config = RandomizedConfig {
695 target_rank: 10,
696 oversampling: 10,
697 n_power_iter: 1,
698 random_seed: None,
699 tolerance: 1e-6,
700 };
701
702 assert!(randomized_range_finder(&matrix, &bad_config).is_err());
703
704 Ok(())
705 }
706}