1use crate::conversions::{array2_to_mat, mat_ref_to_array2, mat_to_array2};
7use ndarray::{Array1, Array2};
8use oxiblas_core::scalar::Field;
9use oxiblas_lapack::{cholesky, evd, lu, qr, solve, svd};
10use oxiblas_matrix::Mat;
11
12pub use evd::Eigenvalue;
14
15#[derive(Debug, Clone)]
21pub enum LapackError {
22 Singular(String),
24 NotPositiveDefinite(String),
26 DimensionMismatch(String),
28 NotConverged(String),
30 Other(String),
32}
33
34impl std::fmt::Display for LapackError {
35 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
36 match self {
37 Self::Singular(msg) => write!(f, "Singular matrix: {msg}"),
38 Self::NotPositiveDefinite(msg) => write!(f, "Not positive definite: {msg}"),
39 Self::DimensionMismatch(msg) => write!(f, "Dimension mismatch: {msg}"),
40 Self::NotConverged(msg) => write!(f, "Did not converge: {msg}"),
41 Self::Other(msg) => write!(f, "LAPACK error: {msg}"),
42 }
43 }
44}
45
46impl std::error::Error for LapackError {}
47
48pub type LapackResult<T> = Result<T, LapackError>;
50
51#[derive(Debug, Clone)]
57pub struct LuResult<T> {
58 pub l: Array2<T>,
60 pub u: Array2<T>,
62 pub perm: Vec<usize>,
64}
65
66impl<T: Field + Clone> LuResult<T>
67where
68 T: bytemuck::Zeroable,
69{
70 pub fn solve(&self, b: &Array1<T>) -> Array1<T> {
72 let n = self.l.dim().0;
73 assert_eq!(b.len(), n, "b length must match matrix dimension");
74
75 let mut pb: Vec<T> = b.iter().cloned().collect();
85 for k in 0..n {
86 let pk = self.perm[k];
87 if k != pk {
88 pb.swap(k, pk);
89 }
90 }
91
92 let mut y: Vec<T> = vec![T::zero(); n];
94 for i in 0..n {
95 let mut sum = pb[i];
96 for j in 0..i {
97 sum -= self.l[[i, j]] * y[j];
98 }
99 y[i] = sum;
100 }
101
102 let mut x: Vec<T> = vec![T::zero(); n];
104 for i in (0..n).rev() {
105 let mut sum = y[i];
106 for j in (i + 1)..n {
107 sum -= self.u[[i, j]] * x[j];
108 }
109 x[i] = sum / self.u[[i, i]];
110 }
111
112 Array1::from_vec(x)
113 }
114
115 pub fn det(&self) -> T {
117 let n = self.l.dim().0;
118 let mut det = T::one();
119
120 for i in 0..n {
122 det *= self.u[[i, i]];
123 }
124
125 let num_swaps = self
135 .perm
136 .iter()
137 .enumerate()
138 .filter(|&(k, &pk)| k != pk)
139 .count();
140
141 if num_swaps % 2 == 1 {
142 det = T::zero() - det;
143 }
144
145 det
146 }
147}
148
149pub fn lu_ndarray<T: Field + Clone>(a: &Array2<T>) -> LapackResult<LuResult<T>>
159where
160 T: bytemuck::Zeroable,
161{
162 let mat = array2_to_mat(a);
163
164 match lu::Lu::compute(mat.as_ref()) {
165 Ok(lu_decomp) => {
166 let l = mat_to_array2(&lu_decomp.l_factor());
168 let u = mat_to_array2(&lu_decomp.u_factor());
169
170 let perm = lu_decomp.pivot().to_vec();
172
173 Ok(LuResult { l, u, perm })
174 }
175 Err(e) => Err(LapackError::Singular(format!("{e:?}"))),
176 }
177}
178
179#[derive(Debug, Clone)]
185pub struct QrResult<T> {
186 pub q: Array2<T>,
188 pub r: Array2<T>,
190}
191
192impl<T: Field + Clone> QrResult<T> {
193 pub fn solve_least_squares(&self, b: &Array1<T>) -> Array1<T> {
195 let (m, n) = (self.q.dim().0, self.r.dim().1);
196 assert_eq!(b.len(), m, "b length must match matrix rows");
197
198 let mut qtb: Array1<T> = Array1::from_vec(vec![T::zero(); n]);
200 for j in 0..n {
201 let mut sum = T::zero();
202 for i in 0..m {
203 sum += self.q[[i, j]].conj() * b[i];
204 }
205 qtb[j] = sum;
206 }
207
208 let mut x: Array1<T> = Array1::from_vec(vec![T::zero(); n]);
210 for i in (0..n).rev() {
211 let mut sum = qtb[i];
212 for j in (i + 1)..n {
213 sum -= self.r[[i, j]] * x[j];
214 }
215 x[i] = sum / self.r[[i, i]];
216 }
217
218 x
219 }
220}
221
222pub fn qr_ndarray<T: Field + Clone>(a: &Array2<T>) -> LapackResult<QrResult<T>>
232where
233 T: bytemuck::Zeroable + oxiblas_core::scalar::Real,
234{
235 let mat = array2_to_mat(a);
236
237 match qr::Qr::compute(mat.as_ref()) {
238 Ok(qr_decomp) => {
239 let q = mat_to_array2(&qr_decomp.q());
240 let r = mat_to_array2(&qr_decomp.r());
241
242 Ok(QrResult { q, r })
243 }
244 Err(e) => Err(LapackError::Other(format!("{e:?}"))),
245 }
246}
247
248#[derive(Debug, Clone)]
254pub struct SvdResult<T> {
255 pub u: Array2<T>,
257 pub s: Array1<T>,
259 pub vt: Array2<T>,
261}
262
263impl<T: Field + Clone> SvdResult<T> {
264 pub fn rank(&self, tol: T) -> usize {
266 self.s.iter().filter(|&s| s.abs() > tol.abs()).count()
267 }
268}
269
270pub fn svd_ndarray<T>(a: &Array2<T>) -> LapackResult<SvdResult<T>>
280where
281 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
282{
283 let mat = array2_to_mat(a);
284
285 match svd::Svd::compute(mat.as_ref()) {
286 Ok(svd_decomp) => {
287 let u = mat_ref_to_array2(svd_decomp.u());
288 let s = Array1::from_vec(svd_decomp.singular_values().to_vec());
289 let vt = mat_ref_to_array2(svd_decomp.vt());
290
291 Ok(SvdResult { u, s, vt })
292 }
293 Err(e) => Err(LapackError::NotConverged(format!("{e:?}"))),
294 }
295}
296
297pub fn svd_truncated<T>(a: &Array2<T>, k: usize) -> LapackResult<SvdResult<T>>
301where
302 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
303{
304 let svd_result = svd_ndarray(a)?;
305
306 let actual_k = k.min(svd_result.s.len());
307
308 let u = svd_result.u.slice(ndarray::s![.., ..actual_k]).to_owned();
310 let s = svd_result.s.slice(ndarray::s![..actual_k]).to_owned();
311 let vt = svd_result.vt.slice(ndarray::s![..actual_k, ..]).to_owned();
312
313 Ok(SvdResult { u, s, vt })
314}
315
316#[derive(Debug, Clone)]
322pub struct SymEvdResult<T> {
323 pub eigenvalues: Array1<T>,
325 pub eigenvectors: Array2<T>,
327}
328
329pub fn eig_symmetric<T>(a: &Array2<T>) -> LapackResult<SymEvdResult<T>>
339where
340 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
341{
342 let (m, n) = a.dim();
343 if m != n {
344 return Err(LapackError::DimensionMismatch(
345 "Matrix must be square".to_string(),
346 ));
347 }
348
349 let mat = array2_to_mat(a);
350
351 match evd::SymmetricEvd::compute(mat.as_ref()) {
352 Ok(evd_result) => {
353 let eigenvalues = Array1::from_vec(evd_result.eigenvalues().to_vec());
354 let evec_ref = evd_result.eigenvectors();
356 let (rows, cols) = (evec_ref.nrows(), evec_ref.ncols());
357 let eigenvectors = Array2::from_shape_fn((rows, cols), |(i, j)| evec_ref[(i, j)]);
358
359 Ok(SymEvdResult {
360 eigenvalues,
361 eigenvectors,
362 })
363 }
364 Err(e) => Err(LapackError::NotConverged(format!("{e:?}"))),
365 }
366}
367
368pub fn eigvals_symmetric<T>(a: &Array2<T>) -> LapackResult<Array1<T>>
370where
371 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
372{
373 eig_symmetric(a).map(|result| result.eigenvalues)
374}
375
376#[derive(Debug, Clone)]
382pub struct ComplexSvdResult<T>
383where
384 T: oxiblas_core::Scalar,
385{
386 pub u: Array2<T>,
388 pub s: Array1<T::Real>,
390 pub vt: Array2<T>,
392}
393
394pub fn svd_complex_ndarray<T>(a: &Array2<T>) -> LapackResult<ComplexSvdResult<T>>
405where
406 T: Field + oxiblas_core::scalar::ComplexScalar + Clone + bytemuck::Zeroable,
407 T::Real: oxiblas_core::scalar::Real,
408{
409 let mat = array2_to_mat(a);
410
411 match svd::ComplexSvd::compute(mat.as_ref()) {
412 Ok(svd_decomp) => {
413 let u = mat_ref_to_array2(svd_decomp.u().as_ref());
414 let s = Array1::from_vec(svd_decomp.singular_values().to_vec());
415 let vh = mat_ref_to_array2(svd_decomp.vh().as_ref());
416
417 Ok(ComplexSvdResult { u, s, vt: vh })
418 }
419 Err(e) => Err(LapackError::NotConverged(format!("{e:?}"))),
420 }
421}
422
423pub fn qr_complex_ndarray<T>(a: &Array2<T>) -> LapackResult<QrResult<T>>
431where
432 T: Field + oxiblas_core::scalar::ComplexScalar + Clone + bytemuck::Zeroable,
433 T::Real: oxiblas_core::scalar::Real,
434{
435 let mat = array2_to_mat(a);
436
437 match qr::UnitaryQr::compute(mat.as_ref()) {
438 Ok(qr_decomp) => {
439 let q_mat = qr_decomp.q();
440 let r_mat = qr_decomp.r();
441 let q = mat_to_array2(&q_mat);
442 let r = mat_to_array2(&r_mat);
443 Ok(QrResult { q, r })
444 }
445 Err(e) => Err(LapackError::NotConverged(format!("{e:?}"))),
446 }
447}
448
449pub fn cholesky_hermitian_ndarray<T>(a: &Array2<T>) -> LapackResult<CholeskyResult<T>>
457where
458 T: Field + oxiblas_core::scalar::ComplexScalar + Clone + bytemuck::Zeroable,
459 T::Real: oxiblas_core::scalar::Real,
460{
461 let (m, n) = a.dim();
462 if m != n {
463 return Err(LapackError::DimensionMismatch(
464 "Matrix must be square".to_string(),
465 ));
466 }
467
468 let mat = array2_to_mat(a);
469
470 match cholesky::HermitianCholesky::compute(mat.as_ref()) {
471 Ok(chol) => {
472 let l_mat = chol.l_factor();
473 let l = mat_to_array2(&l_mat);
474 Ok(CholeskyResult { l })
475 }
476 Err(e) => Err(LapackError::NotPositiveDefinite(format!("{e:?}"))),
477 }
478}
479
480#[derive(Debug, Clone)]
482pub struct HermitianEvdResult<T>
483where
484 T: oxiblas_core::Scalar,
485{
486 pub eigenvalues: Array1<T>,
488 pub eigenvectors: Array2<T>,
490}
491
492pub fn eig_hermitian_ndarray<T>(a: &Array2<T>) -> LapackResult<(Array1<T::Real>, Array2<T>)>
502where
503 T: Field + oxiblas_core::scalar::ComplexScalar + Clone + bytemuck::Zeroable,
504 T::Real: oxiblas_core::scalar::Real + Clone + bytemuck::Zeroable,
505{
506 let (m, n) = a.dim();
507 if m != n {
508 return Err(LapackError::DimensionMismatch(
509 "Matrix must be square".to_string(),
510 ));
511 }
512
513 let mat = array2_to_mat(a);
514
515 match evd::HermitianEvd::compute(mat.as_ref()) {
516 Ok(evd_result) => {
517 let eigenvalues = Array1::from_vec(evd_result.eigenvalues().to_vec());
518
519 let evec_ref = evd_result.eigenvectors();
521 let eigenvectors = mat_ref_to_array2(evec_ref);
522
523 Ok((eigenvalues, eigenvectors))
524 }
525 Err(e) => Err(LapackError::NotConverged(format!("{e:?}"))),
526 }
527}
528
529#[derive(Debug, Clone)]
535pub struct CholeskyResult<T> {
536 pub l: Array2<T>,
538}
539
540impl<T: Field + Clone> CholeskyResult<T> {
541 pub fn solve(&self, b: &Array1<T>) -> Array1<T> {
543 let n = self.l.dim().0;
544 assert_eq!(b.len(), n, "b length must match matrix dimension");
545
546 let mut y: Array1<T> = Array1::from_vec(vec![T::zero(); n]);
548 for i in 0..n {
549 let mut sum = b[i];
550 for j in 0..i {
551 sum -= self.l[[i, j]] * y[j];
552 }
553 y[i] = sum / self.l[[i, i]];
554 }
555
556 let mut x: Array1<T> = Array1::from_vec(vec![T::zero(); n]);
558 for i in (0..n).rev() {
559 let mut sum = y[i];
560 for j in (i + 1)..n {
561 sum -= self.l[[j, i]].conj() * x[j];
562 }
563 x[i] = sum / self.l[[i, i]].conj();
564 }
565
566 x
567 }
568
569 pub fn det(&self) -> T {
571 let n = self.l.dim().0;
572 let mut det = T::one();
573 for i in 0..n {
574 let diag = self.l[[i, i]];
575 det = det * diag * diag;
576 }
577 det
578 }
579}
580
581pub fn cholesky_ndarray<T>(a: &Array2<T>) -> LapackResult<CholeskyResult<T>>
591where
592 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
593{
594 let (m, n) = a.dim();
595 if m != n {
596 return Err(LapackError::DimensionMismatch(
597 "Matrix must be square".to_string(),
598 ));
599 }
600
601 let mat = array2_to_mat(a);
602
603 match cholesky::Cholesky::compute(mat.as_ref()) {
604 Ok(chol) => {
605 let l = mat_to_array2(&chol.l_factor());
606 Ok(CholeskyResult { l })
607 }
608 Err(e) => Err(LapackError::NotPositiveDefinite(format!("{e:?}"))),
609 }
610}
611
612pub fn solve_ndarray<T>(a: &Array2<T>, b: &Array1<T>) -> LapackResult<Array1<T>>
625where
626 T: Field + Clone + bytemuck::Zeroable,
627{
628 let (m, n) = a.dim();
629 if m != n {
630 return Err(LapackError::DimensionMismatch(
631 "Matrix must be square".to_string(),
632 ));
633 }
634 if b.len() != n {
635 return Err(LapackError::DimensionMismatch(
636 "b length must match matrix dimension".to_string(),
637 ));
638 }
639
640 let a_mat = array2_to_mat(a);
641 let mut b_mat: Mat<T> = Mat::zeros(n, 1);
643 for i in 0..n {
644 b_mat[(i, 0)] = b[i];
645 }
646
647 match solve::solve(a_mat.as_ref(), b_mat.as_ref()) {
648 Ok(x_mat) => {
649 let x: Vec<T> = (0..n).map(|i| x_mat[(i, 0)]).collect();
651 Ok(Array1::from_vec(x))
652 }
653 Err(e) => Err(LapackError::Singular(format!("{e:?}"))),
654 }
655}
656
657pub fn solve_multiple_ndarray<T>(a: &Array2<T>, b: &Array2<T>) -> LapackResult<Array2<T>>
666where
667 T: Field + Clone + bytemuck::Zeroable,
668{
669 let (m, n) = a.dim();
670 let (b_rows, _b_cols) = b.dim();
671
672 if m != n {
673 return Err(LapackError::DimensionMismatch(
674 "Matrix must be square".to_string(),
675 ));
676 }
677 if b_rows != n {
678 return Err(LapackError::DimensionMismatch(
679 "b rows must match matrix dimension".to_string(),
680 ));
681 }
682
683 let a_mat = array2_to_mat(a);
684 let b_mat = array2_to_mat(b);
685
686 match solve::solve_multiple(a_mat.as_ref(), b_mat.as_ref()) {
687 Ok(x_mat) => Ok(mat_to_array2(&x_mat)),
688 Err(e) => Err(LapackError::Singular(format!("{e:?}"))),
689 }
690}
691
692pub fn lstsq_ndarray<T>(a: &Array2<T>, b: &Array1<T>) -> LapackResult<Array1<T>>
694where
695 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
696{
697 let m = a.dim().0;
698 let a_mat = array2_to_mat(a);
699 let mut b_mat: Mat<T> = Mat::zeros(m, 1);
701 for i in 0..m {
702 b_mat[(i, 0)] = b[i];
703 }
704
705 match solve::lstsq(a_mat.as_ref(), b_mat.as_ref()) {
706 Ok(result) => {
707 let n = result.solution.nrows();
709 let x: Vec<T> = (0..n).map(|i| result.solution[(i, 0)]).collect();
710 Ok(Array1::from_vec(x))
711 }
712 Err(e) => Err(LapackError::NotConverged(format!("{e:?}"))),
713 }
714}
715
716pub fn inv_ndarray<T>(a: &Array2<T>) -> LapackResult<Array2<T>>
722where
723 T: Field + Clone + bytemuck::Zeroable,
724{
725 let (m, n) = a.dim();
726 if m != n {
727 return Err(LapackError::DimensionMismatch(
728 "Matrix must be square".to_string(),
729 ));
730 }
731
732 let a_mat = array2_to_mat(a);
733
734 match oxiblas_lapack::utils::inv(a_mat.as_ref()) {
735 Ok(inv_mat) => Ok(mat_to_array2(&inv_mat)),
736 Err(e) => Err(LapackError::Singular(format!("{e:?}"))),
737 }
738}
739
740pub fn pinv_ndarray<T>(a: &Array2<T>) -> LapackResult<Array2<T>>
742where
743 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
744{
745 let a_mat = array2_to_mat(a);
746
747 match oxiblas_lapack::utils::pinv_default(a_mat.as_ref()) {
748 Ok(result) => Ok(mat_to_array2(&result.pinv)),
749 Err(e) => Err(LapackError::NotConverged(format!("{e:?}"))),
750 }
751}
752
753pub fn det_ndarray<T>(a: &Array2<T>) -> LapackResult<T>
759where
760 T: Field + Clone + bytemuck::Zeroable,
761{
762 let (m, n) = a.dim();
763 if m != n {
764 return Err(LapackError::DimensionMismatch(
765 "Matrix must be square".to_string(),
766 ));
767 }
768
769 let a_mat = array2_to_mat(a);
770
771 match oxiblas_lapack::utils::det(a_mat.as_ref()) {
772 Ok(d) => Ok(d),
773 Err(e) => Err(LapackError::Other(format!("{e:?}"))),
774 }
775}
776
777pub fn cond_ndarray<T>(a: &Array2<T>) -> LapackResult<T>
783where
784 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
785{
786 let svd_result = svd_ndarray(a)?;
787 let n = svd_result.s.len();
788
789 if n == 0 {
790 return Ok(T::one());
791 }
792
793 let sigma_max = svd_result.s[0];
794 let sigma_min = svd_result.s[n - 1];
795
796 if sigma_min == T::zero() {
798 Ok(T::from_f64(1e15).unwrap_or(T::one()))
800 } else {
801 Ok(sigma_max / sigma_min)
802 }
803}
804
805pub fn rank_ndarray<T>(a: &Array2<T>) -> LapackResult<usize>
811where
812 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
813{
814 let (m, n) = a.dim();
815 let svd_result = svd_ndarray(a)?;
816
817 if svd_result.s.is_empty() {
818 return Ok(0);
819 }
820
821 let sigma_max = svd_result.s[0];
823 let eps = T::from_f64(1e-14).unwrap_or(T::zero());
824 let dim_scale = T::from_f64(m.max(n) as f64).unwrap_or(T::one());
825 let tol = dim_scale * eps * sigma_max;
826
827 Ok(svd_result.rank(tol))
828}
829
830#[derive(Debug, Clone)]
836pub struct RandomizedSvdResult<T> {
837 pub u: Array2<T>,
839 pub s: Array1<T>,
841 pub v: Array2<T>,
843}
844
845pub fn rsvd_ndarray<T>(a: &Array2<T>, k: usize) -> LapackResult<RandomizedSvdResult<T>>
867where
868 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
869{
870 let mat = array2_to_mat(a);
871
872 match svd::RandomizedSvd::compute(mat.as_ref(), k) {
873 Ok(rsvd) => {
874 let u = mat_ref_to_array2(rsvd.u());
875 let s = Array1::from_vec(rsvd.singular_values().to_vec());
876 let v = mat_ref_to_array2(rsvd.v());
877
878 Ok(RandomizedSvdResult { u, s, v })
879 }
880 Err(e) => Err(LapackError::Other(format!("{e:?}"))),
881 }
882}
883
884pub fn rsvd_power_ndarray<T>(
894 a: &Array2<T>,
895 k: usize,
896 power_iterations: usize,
897) -> LapackResult<RandomizedSvdResult<T>>
898where
899 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
900{
901 let mat = array2_to_mat(a);
902
903 let config = svd::RandomizedSvdConfig::new(k).with_power_iterations(power_iterations);
904
905 match svd::RandomizedSvd::compute_with_config(mat.as_ref(), config) {
906 Ok(rsvd) => {
907 let u = mat_ref_to_array2(rsvd.u());
908 let s = Array1::from_vec(rsvd.singular_values().to_vec());
909 let v = mat_ref_to_array2(rsvd.v());
910
911 Ok(RandomizedSvdResult { u, s, v })
912 }
913 Err(e) => Err(LapackError::Other(format!("{e:?}"))),
914 }
915}
916
917#[derive(Debug, Clone)]
923pub struct SchurResult<T> {
924 pub q: Array2<T>,
926 pub t: Array2<T>,
928 pub eigenvalues: Vec<Eigenvalue<T>>,
930}
931
932pub fn schur_ndarray<T>(a: &Array2<T>) -> LapackResult<SchurResult<T>>
945where
946 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
947{
948 let (m, n) = a.dim();
949 if m != n {
950 return Err(LapackError::DimensionMismatch(
951 "Matrix must be square".to_string(),
952 ));
953 }
954
955 let mat = array2_to_mat(a);
956
957 match evd::Schur::compute(mat.as_ref()) {
958 Ok(schur) => {
959 let q = mat_ref_to_array2(schur.q());
960 let t = mat_ref_to_array2(schur.t());
961 let eigenvalues = schur.eigenvalues().to_vec();
962
963 Ok(SchurResult { q, t, eigenvalues })
964 }
965 Err(e) => Err(LapackError::NotConverged(format!("{e:?}"))),
966 }
967}
968
969#[derive(Debug, Clone)]
975pub struct GeneralEvdResult<T> {
976 pub eigenvalues: Vec<Eigenvalue<T>>,
978 pub eigenvectors_real: Option<Array2<T>>,
980 pub eigenvectors_imag: Option<Array2<T>>,
982 pub left_eigenvectors_real: Option<Array2<T>>,
984 pub left_eigenvectors_imag: Option<Array2<T>>,
986}
987
988pub fn eig_ndarray<T>(a: &Array2<T>) -> LapackResult<GeneralEvdResult<T>>
1000where
1001 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
1002{
1003 let (m, n) = a.dim();
1004 if m != n {
1005 return Err(LapackError::DimensionMismatch(
1006 "Matrix must be square".to_string(),
1007 ));
1008 }
1009
1010 let mat = array2_to_mat(a);
1011
1012 match evd::GeneralEvd::compute(mat.as_ref()) {
1013 Ok(evd_result) => {
1014 let eigenvalues = evd_result.eigenvalues().to_vec();
1015
1016 let eigenvectors_real = evd_result
1018 .eigenvectors_real()
1019 .map(|vr| mat_ref_to_array2(vr));
1020
1021 let eigenvectors_imag = evd_result
1022 .eigenvectors_imag()
1023 .map(|vi| mat_ref_to_array2(vi));
1024
1025 let left_eigenvectors_real = evd_result
1027 .left_eigenvectors_real()
1028 .map(|vl| mat_ref_to_array2(vl));
1029
1030 let left_eigenvectors_imag = evd_result
1031 .left_eigenvectors_imag()
1032 .map(|vl| mat_ref_to_array2(vl));
1033
1034 Ok(GeneralEvdResult {
1035 eigenvalues,
1036 eigenvectors_real,
1037 eigenvectors_imag,
1038 left_eigenvectors_real,
1039 left_eigenvectors_imag,
1040 })
1041 }
1042 Err(e) => Err(LapackError::NotConverged(format!("{e:?}"))),
1043 }
1044}
1045
1046pub fn eigvals_ndarray<T>(a: &Array2<T>) -> LapackResult<Vec<Eigenvalue<T>>>
1048where
1049 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
1050{
1051 let (m, n) = a.dim();
1052 if m != n {
1053 return Err(LapackError::DimensionMismatch(
1054 "Matrix must be square".to_string(),
1055 ));
1056 }
1057
1058 let mat = array2_to_mat(a);
1059
1060 match evd::GeneralEvd::eigenvalues_only(mat.as_ref()) {
1061 Ok(evd_result) => Ok(evd_result.eigenvalues().to_vec()),
1062 Err(e) => Err(LapackError::NotConverged(format!("{e:?}"))),
1063 }
1064}
1065
1066pub fn tridiag_solve_ndarray<T>(
1083 dl: &Array1<T>,
1084 d: &Array1<T>,
1085 du: &Array1<T>,
1086 b: &Array1<T>,
1087) -> LapackResult<Array1<T>>
1088where
1089 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
1090{
1091 let n = d.len();
1092
1093 if n == 0 {
1098 if !dl.is_empty() || !du.is_empty() || !b.is_empty() {
1099 return Err(LapackError::DimensionMismatch(
1100 "Tridiagonal dimensions must be consistent".to_string(),
1101 ));
1102 }
1103 return Ok(Array1::from_vec(Vec::new()));
1104 }
1105
1106 if dl.len() != n - 1 || du.len() != n - 1 || b.len() != n {
1107 return Err(LapackError::DimensionMismatch(
1108 "Tridiagonal dimensions must be consistent".to_string(),
1109 ));
1110 }
1111
1112 let dl_vec: Vec<T> = dl.iter().cloned().collect();
1113 let d_vec: Vec<T> = d.iter().cloned().collect();
1114 let du_vec: Vec<T> = du.iter().cloned().collect();
1115 let b_vec: Vec<T> = b.iter().cloned().collect();
1116
1117 match solve::tridiag_solve(&dl_vec, &d_vec, &du_vec, &b_vec) {
1118 Ok(x) => Ok(Array1::from_vec(x)),
1119 Err(e) => Err(LapackError::Singular(format!("{e:?}"))),
1120 }
1121}
1122
1123pub fn tridiag_solve_spd_ndarray<T>(
1136 d: &Array1<T>,
1137 e: &Array1<T>,
1138 b: &Array1<T>,
1139) -> LapackResult<Array1<T>>
1140where
1141 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
1142{
1143 let n = d.len();
1144
1145 if n == 0 {
1150 if !e.is_empty() || !b.is_empty() {
1151 return Err(LapackError::DimensionMismatch(
1152 "Tridiagonal dimensions must be consistent".to_string(),
1153 ));
1154 }
1155 return Ok(Array1::from_vec(Vec::new()));
1156 }
1157
1158 if e.len() != n - 1 || b.len() != n {
1159 return Err(LapackError::DimensionMismatch(
1160 "Tridiagonal dimensions must be consistent".to_string(),
1161 ));
1162 }
1163
1164 let d_vec: Vec<T> = d.iter().cloned().collect();
1165 let e_vec: Vec<T> = e.iter().cloned().collect();
1166 let b_vec: Vec<T> = b.iter().cloned().collect();
1167
1168 match solve::tridiag_solve_spd(&d_vec, &e_vec, &b_vec) {
1169 Ok(x) => Ok(Array1::from_vec(x)),
1170 Err(e) => Err(LapackError::NotPositiveDefinite(format!("{e:?}"))),
1171 }
1172}
1173
1174pub fn tridiag_solve_multiple_ndarray<T>(
1185 dl: &Array1<T>,
1186 d: &Array1<T>,
1187 du: &Array1<T>,
1188 b: &Array2<T>,
1189) -> LapackResult<Array2<T>>
1190where
1191 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
1192{
1193 let n = d.len();
1194 let (b_rows, b_cols) = b.dim();
1195
1196 if n == 0 {
1200 if !dl.is_empty() || !du.is_empty() || b_rows != 0 {
1201 return Err(LapackError::DimensionMismatch(
1202 "Tridiagonal dimensions must be consistent".to_string(),
1203 ));
1204 }
1205 return Ok(Array2::zeros((0, b_cols)));
1206 }
1207
1208 if dl.len() != n - 1 || du.len() != n - 1 || b_rows != n {
1209 return Err(LapackError::DimensionMismatch(
1210 "Tridiagonal dimensions must be consistent".to_string(),
1211 ));
1212 }
1213
1214 let dl_vec: Vec<T> = dl.iter().cloned().collect();
1215 let d_vec: Vec<T> = d.iter().cloned().collect();
1216 let du_vec: Vec<T> = du.iter().cloned().collect();
1217 let b_mat = array2_to_mat(b);
1218
1219 match solve::tridiag_solve_multiple(&dl_vec, &d_vec, &du_vec, b_mat.as_ref()) {
1220 Ok(x_mat) => Ok(mat_to_array2(&x_mat)),
1221 Err(e) => Err(LapackError::Singular(format!("{e:?}"))),
1222 }
1223}
1224
1225pub fn low_rank_approx_ndarray<T>(a: &Array2<T>, k: usize) -> LapackResult<Array2<T>>
1240where
1241 T: Field + Clone + bytemuck::Zeroable + oxiblas_core::scalar::Real,
1242{
1243 let mat = array2_to_mat(a);
1244
1245 match svd::low_rank_approximation(mat.as_ref(), k) {
1246 Ok(approx) => Ok(mat_to_array2(&approx)),
1247 Err(e) => Err(LapackError::Other(format!("{e:?}"))),
1248 }
1249}
1250
1251#[cfg(test)]
1252mod tests {
1253 use super::*;
1254 use ndarray::array;
1255
1256 #[test]
1257 fn test_lu_decomposition() {
1258 let a = array![[2.0f64, 1.0], [1.0, 3.0]];
1259 let lu = lu_ndarray(&a).unwrap();
1260
1261 let n = a.dim().0;
1263 for i in 0..n {
1264 for j in 0..n {
1265 let mut sum = 0.0f64;
1266 for k in 0..n {
1267 sum += lu.l[[i, k]] * lu.u[[k, j]];
1268 }
1269 let perm_i = lu.perm.iter().position(|&p| p == i).unwrap();
1270 assert!((sum - a[[perm_i, j]]).abs() < 1e-10);
1271 }
1272 }
1273 }
1274
1275 #[test]
1276 fn test_lu_determinant() {
1277 let a = array![[2.0f64, 1.0], [1.0, 3.0]];
1278 let lu = lu_ndarray(&a).unwrap();
1279 let det = lu.det();
1280 assert!((det - 5.0).abs() < 1e-10);
1282 }
1283
1284 #[test]
1285 fn test_lu_solve() {
1286 let a = array![[2.0f64, 1.0], [1.0, 3.0]];
1287 let b = array![5.0f64, 7.0];
1288 let lu = lu_ndarray(&a).unwrap();
1289 let x = lu.solve(&b);
1290
1291 let ax0 = a[[0, 0]] * x[0] + a[[0, 1]] * x[1];
1293 let ax1 = a[[1, 0]] * x[0] + a[[1, 1]] * x[1];
1294 assert!((ax0 - b[0]).abs() < 1e-10);
1295 assert!((ax1 - b[1]).abs() < 1e-10);
1296 }
1297
1298 #[test]
1316 fn test_lu_solve_and_det_with_two_pivot_swaps() {
1317 let a = array![[2.0f64, 1.0, 1.0], [4.0, 3.0, 3.0], [8.0, 7.0, 9.0]];
1318
1319 let lu = lu_ndarray(&a).unwrap();
1320
1321 let num_swaps = lu
1324 .perm
1325 .iter()
1326 .enumerate()
1327 .filter(|&(k, &pk)| k != pk)
1328 .count();
1329 assert!(
1330 num_swaps >= 2,
1331 "test matrix must force at least two row swaps, got perm = {:?}",
1332 lu.perm
1333 );
1334
1335 let det = lu.det();
1337 assert!(
1338 (det - 4.0).abs() < 1e-10,
1339 "det = {det}, expected +4 (sign must be positive for an even swap count)"
1340 );
1341
1342 let x_true = [1.0f64, 2.0, 3.0];
1344 let b = array![7.0f64, 19.0, 49.0];
1345 let x = lu.solve(&b);
1346 for i in 0..3 {
1347 assert!(
1348 (x[i] - x_true[i]).abs() < 1e-10,
1349 "x[{i}] = {}, expected {}",
1350 x[i],
1351 x_true[i]
1352 );
1353 }
1354
1355 for i in 0..3 {
1357 let axi = a[[i, 0]] * x[0] + a[[i, 1]] * x[1] + a[[i, 2]] * x[2];
1358 assert!(
1359 (axi - b[i]).abs() < 1e-10,
1360 "residual row {i}: {axi} != {}",
1361 b[i]
1362 );
1363 }
1364 }
1365
1366 #[test]
1369 fn test_tridiag_empty_inputs_no_overflow() {
1370 let empty1: Array1<f64> = Array1::from_vec(Vec::new());
1371
1372 let x = tridiag_solve_ndarray(&empty1, &empty1, &empty1, &empty1).unwrap();
1373 assert_eq!(x.len(), 0);
1374
1375 let x_spd = tridiag_solve_spd_ndarray(&empty1, &empty1, &empty1).unwrap();
1376 assert_eq!(x_spd.len(), 0);
1377
1378 let b_empty: Array2<f64> = Array2::zeros((0, 3));
1379 let x_multi = tridiag_solve_multiple_ndarray(&empty1, &empty1, &empty1, &b_empty).unwrap();
1380 assert_eq!(x_multi.dim(), (0, 3));
1381 }
1382
1383 #[test]
1384 fn test_qr_decomposition() {
1385 let a = array![[1.0f64, 2.0], [3.0, 4.0], [5.0, 6.0]];
1386 let qr = qr_ndarray(&a).unwrap();
1387
1388 let qt = qr.q.t();
1390 let qtq = crate::blas::matmul(&qt.to_owned(), &qr.q);
1391 for i in 0..qtq.dim().0 {
1392 for j in 0..qtq.dim().1 {
1393 let expected = if i == j { 1.0 } else { 0.0 };
1394 assert!(
1395 (qtq[[i, j]] - expected).abs() < 1e-10,
1396 "Q^T Q[{},{}] = {}, expected {}",
1397 i,
1398 j,
1399 qtq[[i, j]],
1400 expected
1401 );
1402 }
1403 }
1404
1405 let qr_product = crate::blas::matmul(&qr.q, &qr.r);
1407 for i in 0..a.dim().0 {
1408 for j in 0..a.dim().1 {
1409 assert!(
1410 (qr_product[[i, j]] - a[[i, j]]).abs() < 1e-10,
1411 "QR[{},{}] = {}, A = {}",
1412 i,
1413 j,
1414 qr_product[[i, j]],
1415 a[[i, j]]
1416 );
1417 }
1418 }
1419 }
1420
1421 #[test]
1422 fn test_svd() {
1423 let a = array![[1.0f64, 2.0], [3.0, 4.0], [5.0, 6.0]];
1424 let svd = svd_ndarray(&a).unwrap();
1425
1426 let (m, n) = a.dim();
1428 let k = svd.s.len();
1429
1430 for i in 0..m {
1431 for j in 0..n {
1432 let mut sum = 0.0f64;
1433 for l in 0..k {
1434 sum += svd.u[[i, l]] * svd.s[l] * svd.vt[[l, j]];
1435 }
1436 assert!(
1437 (sum - a[[i, j]]).abs() < 1e-10,
1438 "Reconstructed[{},{}] = {}, A = {}",
1439 i,
1440 j,
1441 sum,
1442 a[[i, j]]
1443 );
1444 }
1445 }
1446 }
1447
1448 #[test]
1449 fn test_symmetric_evd() {
1450 let a = array![[4.0f64, 1.0], [1.0, 3.0]];
1452 let evd = eig_symmetric(&a).unwrap();
1453
1454 assert!(evd.eigenvalues.len() == 2);
1456
1457 for (idx, &lambda) in evd.eigenvalues.iter().enumerate() {
1459 let v = evd.eigenvectors.column(idx);
1460 let av = crate::blas::matvec(&a, &v.to_owned());
1461 let lambda_v: Array1<f64> = v.iter().map(|&x| lambda * x).collect();
1462
1463 for i in 0..2 {
1464 assert!(
1465 (av[i] - lambda_v[i]).abs() < 1e-10,
1466 "Av[{}] = {}, λv[{}] = {}",
1467 i,
1468 av[i],
1469 i,
1470 lambda_v[i]
1471 );
1472 }
1473 }
1474 }
1475
1476 #[test]
1477 fn test_cholesky() {
1478 let a = array![[4.0f64, 2.0], [2.0, 5.0]];
1480 let chol = cholesky_ndarray(&a).unwrap();
1481
1482 let lt = chol.l.t();
1484 let llt = crate::blas::matmul(&chol.l, <.to_owned());
1485
1486 for i in 0..2 {
1487 for j in 0..2 {
1488 assert!(
1489 (llt[[i, j]] - a[[i, j]]).abs() < 1e-10,
1490 "LLT[{},{}] = {}, A = {}",
1491 i,
1492 j,
1493 llt[[i, j]],
1494 a[[i, j]]
1495 );
1496 }
1497 }
1498 }
1499
1500 #[test]
1501 fn test_solve() {
1502 let a = array![[2.0f64, 1.0], [1.0, 3.0]];
1503 let b = array![5.0f64, 7.0];
1504 let x = solve_ndarray(&a, &b).unwrap();
1505
1506 let ax = crate::blas::matvec(&a, &x);
1508 assert!((ax[0] - b[0]).abs() < 1e-10);
1509 assert!((ax[1] - b[1]).abs() < 1e-10);
1510 }
1511
1512 #[test]
1513 fn test_inverse() {
1514 let a = array![[4.0f64, 7.0], [2.0, 6.0]];
1515 let a_inv = inv_ndarray(&a).unwrap();
1516
1517 let product = crate::blas::matmul(&a, &a_inv);
1519 for i in 0..2 {
1520 for j in 0..2 {
1521 let expected = if i == j { 1.0 } else { 0.0 };
1522 assert!(
1523 (product[[i, j]] - expected).abs() < 1e-10,
1524 "A*A^-1[{},{}] = {}, expected {}",
1525 i,
1526 j,
1527 product[[i, j]],
1528 expected
1529 );
1530 }
1531 }
1532 }
1533
1534 #[test]
1535 fn test_determinant() {
1536 let a = array![[2.0f64, 1.0], [1.0, 3.0]];
1537 let det = det_ndarray(&a).unwrap();
1538 assert!((det - 5.0).abs() < 1e-10);
1540 }
1541
1542 #[test]
1543 fn test_condition_number() {
1544 let a = array![[1.0f64, 0.0], [0.0, 1.0]];
1545 let cond = cond_ndarray(&a).unwrap();
1546 assert!((cond - 1.0).abs() < 1e-10);
1548 }
1549
1550 #[test]
1551 fn test_rank() {
1552 let a = array![[1.0f64, 2.0], [3.0, 4.0]];
1554 let r = rank_ndarray(&a).unwrap();
1555 assert_eq!(r, 2);
1556
1557 let b = array![[1.0f64, 2.0], [2.0, 4.0]];
1559 let r2 = rank_ndarray(&b).unwrap();
1560 assert_eq!(r2, 1);
1561 }
1562
1563 #[test]
1568 fn test_rsvd_basic() {
1569 let a = array![
1571 [1.0f64, 2.0, 3.0, 4.0],
1572 [5.0, 6.0, 7.0, 8.0],
1573 [9.0, 10.0, 11.0, 12.0]
1574 ];
1575
1576 let rsvd = rsvd_ndarray(&a, 2).unwrap();
1577
1578 assert_eq!(rsvd.s.len(), 2);
1580
1581 assert!(rsvd.s[0] > rsvd.s[1]);
1583 assert!(rsvd.s[1] >= 0.0);
1584
1585 assert_eq!(rsvd.u.dim(), (3, 2));
1587
1588 assert_eq!(rsvd.v.dim(), (4, 2));
1590 }
1591
1592 #[test]
1593 fn test_rsvd_approximation_quality() {
1594 let a = Array2::from_shape_fn((10, 8), |(i, j)| (i as f64) * 0.1 + (j as f64) * 0.2);
1596
1597 let rsvd = rsvd_ndarray(&a, 2).unwrap();
1598
1599 let (m, n) = a.dim();
1601 let k = rsvd.s.len();
1602
1603 let mut approx: Array2<f64> = Array2::zeros((m, n));
1604 for i in 0..m {
1605 for j in 0..n {
1606 for l in 0..k {
1607 approx[[i, j]] += rsvd.u[[i, l]] * rsvd.s[l] * rsvd.v[[j, l]];
1608 }
1609 }
1610 }
1611
1612 let mut diff_norm = 0.0f64;
1614 for i in 0..m {
1615 for j in 0..n {
1616 let diff = a[[i, j]] - approx[[i, j]];
1617 diff_norm += diff.powi(2);
1618 }
1619 }
1620 diff_norm = diff_norm.sqrt();
1621
1622 assert!(diff_norm < 1e-10, "Reconstruction error = {}", diff_norm);
1624 }
1625
1626 #[test]
1627 fn test_rsvd_power_iteration() {
1628 let a = Array2::from_shape_fn((20, 15), |(i, j)| ((i * j) as f64).sin() + 0.1 * (i as f64));
1629
1630 let rsvd = rsvd_power_ndarray(&a, 3, 2).unwrap();
1631
1632 assert_eq!(rsvd.s.len(), 3);
1633 assert!(rsvd.s[0] >= rsvd.s[1]);
1634 assert!(rsvd.s[1] >= rsvd.s[2]);
1635 }
1636
1637 #[test]
1642 fn test_schur_triangular() {
1643 let a = array![[1.0f64, 2.0], [0.0, 3.0]];
1645
1646 let schur = schur_ndarray(&a).unwrap();
1647
1648 assert_eq!(schur.eigenvalues.len(), 2);
1650
1651 let evs: Vec<f64> = schur.eigenvalues.iter().map(|e| e.real).collect();
1652 assert!(evs.contains(&1.0) || evs.iter().any(|&x| (x - 1.0).abs() < 1e-10));
1653 assert!(evs.contains(&3.0) || evs.iter().any(|&x| (x - 3.0).abs() < 1e-10));
1654 }
1655
1656 #[test]
1657 fn test_schur_reconstruction() {
1658 let a = array![[4.0f64, 1.0], [2.0, 3.0]];
1659
1660 let schur = schur_ndarray(&a).unwrap();
1661
1662 let qt = schur.q.t();
1664 let qt_owned = qt.to_owned();
1665 let qr_temp = crate::blas::matmul(&schur.q, &schur.t);
1666 let reconstructed = crate::blas::matmul(&qr_temp, &qt_owned);
1667
1668 for i in 0..2 {
1669 for j in 0..2 {
1670 assert!(
1671 (reconstructed[[i, j]] - a[[i, j]]).abs() < 1e-10,
1672 "Reconstruction failed at [{},{}]: {} vs {}",
1673 i,
1674 j,
1675 reconstructed[[i, j]],
1676 a[[i, j]]
1677 );
1678 }
1679 }
1680 }
1681
1682 #[test]
1683 fn test_schur_orthogonality() {
1684 let a = array![[1.0f64, 2.0, 3.0], [4.0, 5.0, 6.0], [7.0, 8.0, 10.0]];
1685
1686 let schur = schur_ndarray(&a).unwrap();
1687
1688 let qt = schur.q.t();
1690 let qtq = crate::blas::matmul(&qt.to_owned(), &schur.q);
1691
1692 for i in 0..3 {
1693 for j in 0..3 {
1694 let expected = if i == j { 1.0 } else { 0.0 };
1695 assert!(
1696 (qtq[[i, j]] - expected).abs() < 1e-10,
1697 "Q^T Q[{},{}] = {}, expected {}",
1698 i,
1699 j,
1700 qtq[[i, j]],
1701 expected
1702 );
1703 }
1704 }
1705 }
1706
1707 #[test]
1712 fn test_eig_real_eigenvalues() {
1713 let a = array![[4.0f64, 1.0], [1.0, 3.0]];
1715
1716 let evd = eig_ndarray(&a).unwrap();
1717
1718 assert_eq!(evd.eigenvalues.len(), 2);
1719
1720 for ev in &evd.eigenvalues {
1722 assert!(
1723 ev.imag.abs() < 1e-10,
1724 "Expected real eigenvalue, got imag = {}",
1725 ev.imag
1726 );
1727 }
1728 }
1729
1730 #[test]
1731 fn test_eig_complex_eigenvalues() {
1732 let a = array![[0.0f64, -1.0], [1.0, 0.0]];
1734
1735 let evd = eig_ndarray(&a).unwrap();
1736
1737 assert_eq!(evd.eigenvalues.len(), 2);
1738
1739 let has_complex = evd.eigenvalues.iter().any(|e| e.imag.abs() > 0.5);
1741 assert!(has_complex, "Expected complex eigenvalues");
1742
1743 for ev in &evd.eigenvalues {
1745 assert!(ev.real.abs() < 1e-10, "Expected real part ≈ 0");
1746 }
1747 }
1748
1749 #[test]
1750 fn test_eigvals_only() {
1751 let a = array![[1.0f64, 2.0], [0.0, 3.0]];
1752
1753 let evs = eigvals_ndarray(&a).unwrap();
1754
1755 assert_eq!(evs.len(), 2);
1756
1757 let reals: Vec<f64> = evs.iter().map(|e| e.real).collect();
1759 assert!(reals.iter().any(|&x| (x - 1.0).abs() < 1e-10));
1760 assert!(reals.iter().any(|&x| (x - 3.0).abs() < 1e-10));
1761 }
1762
1763 #[test]
1768 fn test_tridiag_solve() {
1769 let dl = array![-1.0f64, -1.0];
1774 let d = array![2.0f64, 2.0, 2.0];
1775 let du = array![-1.0f64, -1.0];
1776 let b = array![1.0f64, 0.0, 1.0];
1777
1778 let x = tridiag_solve_ndarray(&dl, &d, &du, &b).unwrap();
1779
1780 assert_eq!(x.len(), 3);
1781
1782 let tx0 = d[0] * x[0] + du[0] * x[1];
1784 let tx1 = dl[0] * x[0] + d[1] * x[1] + du[1] * x[2];
1785 let tx2 = dl[1] * x[1] + d[2] * x[2];
1786
1787 assert!((tx0 - b[0]).abs() < 1e-10);
1788 assert!((tx1 - b[1]).abs() < 1e-10);
1789 assert!((tx2 - b[2]).abs() < 1e-10);
1790 }
1791
1792 #[test]
1793 fn test_tridiag_solve_spd() {
1794 let d = array![4.0f64, 4.0, 4.0];
1800 let e = array![1.0f64, 1.0]; let b = array![5.0f64, 6.0, 5.0];
1802
1803 let x = tridiag_solve_spd_ndarray(&d, &e, &b).unwrap();
1804
1805 assert_eq!(x.len(), 3);
1806
1807 let tx0 = d[0] * x[0] + e[0] * x[1];
1809 let tx1 = e[0] * x[0] + d[1] * x[1] + e[1] * x[2];
1810 let tx2 = e[1] * x[1] + d[2] * x[2];
1811
1812 assert!((tx0 - b[0]).abs() < 1e-10, "tx0 = {}, b[0] = {}", tx0, b[0]);
1813 assert!((tx1 - b[1]).abs() < 1e-10, "tx1 = {}, b[1] = {}", tx1, b[1]);
1814 assert!((tx2 - b[2]).abs() < 1e-10, "tx2 = {}, b[2] = {}", tx2, b[2]);
1815 }
1816
1817 #[test]
1818 fn test_tridiag_solve_multiple() {
1819 let dl = array![-1.0f64, -1.0];
1820 let d = array![2.0f64, 2.0, 2.0];
1821 let du = array![-1.0f64, -1.0];
1822 let b = array![[1.0f64, 0.0], [0.0, 1.0], [1.0, 0.0]];
1823
1824 let x = tridiag_solve_multiple_ndarray(&dl, &d, &du, &b).unwrap();
1825
1826 assert_eq!(x.dim(), (3, 2));
1827
1828 for j in 0..2 {
1830 let tx0 = d[0] * x[[0, j]] + du[0] * x[[1, j]];
1831 let tx1 = dl[0] * x[[0, j]] + d[1] * x[[1, j]] + du[1] * x[[2, j]];
1832 let tx2 = dl[1] * x[[1, j]] + d[2] * x[[2, j]];
1833
1834 assert!((tx0 - b[[0, j]]).abs() < 1e-10);
1835 assert!((tx1 - b[[1, j]]).abs() < 1e-10);
1836 assert!((tx2 - b[[2, j]]).abs() < 1e-10);
1837 }
1838 }
1839
1840 #[test]
1845 fn test_low_rank_approx() {
1846 let u = array![1.0f64, 2.0, 3.0];
1848 let v = array![4.0, 5.0, 6.0, 7.0];
1849
1850 let mut a = Array2::zeros((3, 4));
1851 for i in 0..3 {
1852 for j in 0..4 {
1853 a[[i, j]] = u[i] * v[j];
1854 }
1855 }
1856
1857 let approx = low_rank_approx_ndarray(&a, 1).unwrap();
1859
1860 assert_eq!(approx.dim(), a.dim());
1861
1862 for i in 0..3 {
1863 for j in 0..4 {
1864 assert!(
1865 (approx[[i, j]] - a[[i, j]]).abs() < 1e-10,
1866 "Approximation failed at [{},{}]",
1867 i,
1868 j
1869 );
1870 }
1871 }
1872 }
1873
1874 #[test]
1875 fn test_low_rank_approx_truncation() {
1876 let a = array![
1877 [1.0f64, 2.0, 3.0],
1878 [4.0, 5.0, 6.0],
1879 [7.0, 8.0, 9.0],
1880 [10.0, 11.0, 12.0]
1881 ];
1882
1883 let approx = low_rank_approx_ndarray(&a, 2).unwrap();
1884
1885 assert_eq!(approx.dim(), (4, 3));
1886
1887 let mut diff_norm = 0.0f64;
1890 let mut orig_norm = 0.0f64;
1891 for i in 0..4 {
1892 for j in 0..3 {
1893 diff_norm += (a[[i, j]] - approx[[i, j]]).powi(2);
1894 orig_norm += a[[i, j]].powi(2);
1895 }
1896 }
1897
1898 let rel_error = diff_norm.sqrt() / orig_norm.sqrt();
1900 assert!(rel_error < 0.1, "Relative error = {}", rel_error);
1901 }
1902}