1use ndarray::{Array1, Array2};
31use num_traits::Float;
32use oxiblas_core::scalar::{Field, Scalar};
33use oxiblas_sparse::csc::CscMatrix;
34use oxiblas_sparse::csr::CsrMatrix;
35
36#[inline]
57fn retain_entry<T: Scalar>(val: T, tolerance: Option<<T as Scalar>::Real>) -> bool {
58 if val.real().is_nan() || val.imag().is_nan() {
61 return true;
62 }
63
64 match tolerance {
65 Some(tol) => Scalar::abs(val) > tol,
66 None => val != T::zero(),
67 }
68}
69
70pub fn array2_to_csr<T: Scalar + Clone + Field>(arr: &Array2<T>) -> CsrMatrix<T> {
106 array2_to_csr_with_tolerance(arr, None)
107}
108
109pub fn array2_to_csr_with_tolerance<T: Scalar + Clone + Field>(
127 arr: &Array2<T>,
128 tolerance: Option<<T as Scalar>::Real>,
129) -> CsrMatrix<T> {
130 let (nrows, ncols) = arr.dim();
131
132 let mut row_ptrs = Vec::with_capacity(nrows + 1);
133 let mut col_indices = Vec::new();
134 let mut values = Vec::new();
135
136 row_ptrs.push(0);
137
138 for i in 0..nrows {
139 for j in 0..ncols {
140 let val = arr[[i, j]];
141 if retain_entry(val, tolerance) {
142 col_indices.push(j);
143 values.push(val);
144 }
145 }
146 row_ptrs.push(values.len());
147 }
148
149 unsafe { CsrMatrix::new_unchecked(nrows, ncols, row_ptrs, col_indices, values) }
155}
156
157pub fn csr_to_array2<T: Scalar + Clone + Field>(csr: &CsrMatrix<T>) -> Array2<T> {
165 let (nrows, ncols) = csr.shape();
166 let mut result = Array2::zeros((nrows, ncols));
167
168 for i in 0..nrows {
169 for (col, val) in csr.row_iter(i) {
170 result[[i, col]] = *val;
171 }
172 }
173
174 result
175}
176
177pub fn array2_to_csc<T: Scalar + Clone + Field>(arr: &Array2<T>) -> CscMatrix<T> {
201 array2_to_csc_with_tolerance(arr, None)
202}
203
204pub fn array2_to_csc_with_tolerance<T: Scalar + Clone + Field>(
222 arr: &Array2<T>,
223 tolerance: Option<<T as Scalar>::Real>,
224) -> CscMatrix<T> {
225 let (nrows, ncols) = arr.dim();
226
227 let mut col_ptrs = Vec::with_capacity(ncols + 1);
228 let mut row_indices = Vec::new();
229 let mut values = Vec::new();
230
231 col_ptrs.push(0);
232
233 for j in 0..ncols {
234 for i in 0..nrows {
235 let val = arr[[i, j]];
236 if retain_entry(val, tolerance) {
237 row_indices.push(i);
238 values.push(val);
239 }
240 }
241 col_ptrs.push(values.len());
242 }
243
244 unsafe { CscMatrix::new_unchecked(nrows, ncols, col_ptrs, row_indices, values) }
246}
247
248pub fn csc_to_array2<T: Scalar + Clone + Field>(csc: &CscMatrix<T>) -> Array2<T> {
256 let (nrows, ncols) = csc.shape();
257 let mut result = Array2::zeros((nrows, ncols));
258
259 for j in 0..ncols {
260 for (row, val) in csc.col_iter(j) {
261 result[[row, j]] = *val;
262 }
263 }
264
265 result
266}
267
268pub fn spmv_ndarray<T: Scalar + Clone + Field>(a: &CsrMatrix<T>, x: &Array1<T>) -> Array1<T> {
287 assert_eq!(
288 x.len(),
289 a.ncols(),
290 "Vector length {} must match matrix columns {}",
291 x.len(),
292 a.ncols()
293 );
294
295 let x_vec: Vec<T> = x.iter().cloned().collect();
296 let mut y_vec = vec![T::zero(); a.nrows()];
297
298 oxiblas_sparse::ops::spmv(T::one(), a, &x_vec, T::zero(), &mut y_vec);
299
300 Array1::from_vec(y_vec)
301}
302
303pub fn spmv_full_ndarray<T: Scalar + Clone + Field>(
317 alpha: T,
318 a: &CsrMatrix<T>,
319 x: &Array1<T>,
320 beta: T,
321 y: &mut Array1<T>,
322) {
323 assert_eq!(x.len(), a.ncols(), "x length must match matrix columns");
324 assert_eq!(y.len(), a.nrows(), "y length must match matrix rows");
325
326 let x_vec: Vec<T> = x.iter().cloned().collect();
327
328 if let Some(y_slice) = y.as_slice_mut() {
329 oxiblas_sparse::ops::spmv(alpha, a, &x_vec, beta, y_slice);
330 } else {
331 let mut y_vec: Vec<T> = y.iter().cloned().collect();
332 oxiblas_sparse::ops::spmv(alpha, a, &x_vec, beta, &mut y_vec);
333 for (yi, val) in y.iter_mut().zip(y_vec) {
334 *yi = val;
335 }
336 }
337}
338
339#[derive(Debug, Clone)]
345pub enum SparseNdarrayError {
346 NotSquare {
348 nrows: usize,
350 ncols: usize,
352 },
353 DimensionMismatch {
355 matrix_dim: usize,
357 vector_len: usize,
359 },
360 NotConverged {
362 iterations: usize,
364 residual_norm: f64,
366 },
367 SolverError(String),
369}
370
371impl core::fmt::Display for SparseNdarrayError {
372 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
373 match self {
374 Self::NotSquare { nrows, ncols } => {
375 write!(f, "Matrix must be square: got {nrows}x{ncols}")
376 }
377 Self::DimensionMismatch {
378 matrix_dim,
379 vector_len,
380 } => {
381 write!(
382 f,
383 "Dimension mismatch: matrix dim={matrix_dim}, vector len={vector_len}"
384 )
385 }
386 Self::NotConverged {
387 iterations,
388 residual_norm,
389 } => {
390 write!(
391 f,
392 "CG did not converge after {iterations} iterations (residual={residual_norm})"
393 )
394 }
395 Self::SolverError(msg) => write!(f, "Solver error: {msg}"),
396 }
397 }
398}
399
400impl std::error::Error for SparseNdarrayError {}
401
402pub fn sparse_solve_ndarray(
421 a: &CsrMatrix<f64>,
422 b: &Array1<f64>,
423) -> Result<Array1<f64>, SparseNdarrayError> {
424 let (nrows, ncols) = a.shape();
425
426 if nrows != ncols {
427 return Err(SparseNdarrayError::NotSquare { nrows, ncols });
428 }
429
430 if b.len() != nrows {
431 return Err(SparseNdarrayError::DimensionMismatch {
432 matrix_dim: nrows,
433 vector_len: b.len(),
434 });
435 }
436
437 let b_vec: Vec<f64> = b.iter().copied().collect();
438 let x0 = vec![0.0f64; nrows];
439
440 let tol = 1e-10;
441 let max_iter = nrows * 2 + 100;
442
443 match oxiblas_sparse::linalg::cg(a, &b_vec, &x0, tol, max_iter) {
444 Ok(result) => {
445 if result.converged {
446 Ok(Array1::from_vec(result.x))
447 } else {
448 Err(SparseNdarrayError::NotConverged {
449 iterations: result.iterations,
450 residual_norm: result.residual_norm,
451 })
452 }
453 }
454 Err(e) => Err(SparseNdarrayError::SolverError(e.to_string())),
455 }
456}
457
458pub fn sparse_solve_ndarray_with_options(
475 a: &CsrMatrix<f64>,
476 b: &Array1<f64>,
477 tol: f64,
478 max_iter: usize,
479) -> Result<Array1<f64>, SparseNdarrayError> {
480 let (nrows, ncols) = a.shape();
481
482 if nrows != ncols {
483 return Err(SparseNdarrayError::NotSquare { nrows, ncols });
484 }
485
486 if b.len() != nrows {
487 return Err(SparseNdarrayError::DimensionMismatch {
488 matrix_dim: nrows,
489 vector_len: b.len(),
490 });
491 }
492
493 let b_vec: Vec<f64> = b.iter().copied().collect();
494 let x0 = vec![0.0f64; nrows];
495
496 match oxiblas_sparse::linalg::cg(a, &b_vec, &x0, tol, max_iter) {
497 Ok(result) => {
498 if result.converged {
499 Ok(Array1::from_vec(result.x))
500 } else {
501 Err(SparseNdarrayError::NotConverged {
502 iterations: result.iterations,
503 residual_norm: result.residual_norm,
504 })
505 }
506 }
507 Err(e) => Err(SparseNdarrayError::SolverError(e.to_string())),
508 }
509}
510
511#[cfg(test)]
512mod tests {
513 use super::*;
514 use ndarray::array;
515
516 #[test]
521 fn test_array2_to_csr_basic() {
522 let a = array![[1.0f64, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
523 let csr = array2_to_csr(&a);
524
525 assert_eq!(csr.nrows(), 3);
526 assert_eq!(csr.ncols(), 3);
527 assert_eq!(csr.nnz(), 5);
528
529 assert_eq!(csr.get(0, 0), Some(&1.0));
530 assert_eq!(csr.get(0, 2), Some(&2.0));
531 assert_eq!(csr.get(1, 1), Some(&3.0));
532 assert_eq!(csr.get(2, 0), Some(&4.0));
533 assert_eq!(csr.get(2, 2), Some(&5.0));
534
535 assert_eq!(csr.get(0, 1), None);
537 assert_eq!(csr.get(1, 0), None);
538 }
539
540 #[test]
541 fn test_csr_to_array2_basic() {
542 let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
543 let col_indices = vec![0, 2, 1, 0, 2];
544 let row_ptrs = vec![0, 2, 3, 5];
545 let csr = CsrMatrix::new(3, 3, row_ptrs, col_indices, values)
546 .expect("Failed to create CSR matrix");
547
548 let arr = csr_to_array2(&csr);
549 assert_eq!(arr.dim(), (3, 3));
550 assert!((arr[[0, 0]] - 1.0).abs() < 1e-15);
551 assert!((arr[[0, 1]]).abs() < 1e-15);
552 assert!((arr[[0, 2]] - 2.0).abs() < 1e-15);
553 assert!((arr[[1, 1]] - 3.0).abs() < 1e-15);
554 assert!((arr[[2, 0]] - 4.0).abs() < 1e-15);
555 assert!((arr[[2, 2]] - 5.0).abs() < 1e-15);
556 }
557
558 #[test]
559 fn test_roundtrip_csr() {
560 let original = array![
561 [1.0f64, 0.0, 3.0, 0.0],
562 [0.0, 5.0, 0.0, 7.0],
563 [9.0, 0.0, 11.0, 0.0]
564 ];
565
566 let csr = array2_to_csr(&original);
567 let recovered = csr_to_array2(&csr);
568
569 assert_eq!(original.dim(), recovered.dim());
570 for i in 0..3 {
571 for j in 0..4 {
572 assert!(
573 (original[[i, j]] - recovered[[i, j]]).abs() < 1e-15,
574 "Mismatch at ({}, {})",
575 i,
576 j
577 );
578 }
579 }
580 }
581
582 #[test]
583 fn test_array2_to_csc_basic() {
584 let a = array![[1.0f64, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
585 let csc = array2_to_csc(&a);
586
587 assert_eq!(csc.nrows(), 3);
588 assert_eq!(csc.ncols(), 3);
589 assert_eq!(csc.nnz(), 5);
590
591 assert_eq!(csc.get(0, 0), Some(&1.0));
592 assert_eq!(csc.get(0, 2), Some(&2.0));
593 assert_eq!(csc.get(1, 1), Some(&3.0));
594 assert_eq!(csc.get(2, 0), Some(&4.0));
595 assert_eq!(csc.get(2, 2), Some(&5.0));
596 }
597
598 #[test]
599 fn test_csc_to_array2_basic() {
600 let values = vec![1.0f64, 4.0, 3.0, 2.0, 5.0];
601 let row_indices = vec![0, 2, 1, 0, 2];
602 let col_ptrs = vec![0, 2, 3, 5];
603 let csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values)
604 .expect("Failed to create CSC matrix");
605
606 let arr = csc_to_array2(&csc);
607 assert_eq!(arr.dim(), (3, 3));
608 assert!((arr[[0, 0]] - 1.0).abs() < 1e-15);
609 assert!((arr[[2, 0]] - 4.0).abs() < 1e-15);
610 assert!((arr[[1, 1]] - 3.0).abs() < 1e-15);
611 assert!((arr[[0, 2]] - 2.0).abs() < 1e-15);
612 assert!((arr[[2, 2]] - 5.0).abs() < 1e-15);
613 }
614
615 #[test]
616 fn test_roundtrip_csc() {
617 let original = array![
618 [0.0f64, 2.0, 0.0],
619 [4.0, 0.0, 6.0],
620 [0.0, 8.0, 0.0],
621 [10.0, 0.0, 12.0]
622 ];
623
624 let csc = array2_to_csc(&original);
625 let recovered = csc_to_array2(&csc);
626
627 assert_eq!(original.dim(), recovered.dim());
628 for i in 0..4 {
629 for j in 0..3 {
630 assert!(
631 (original[[i, j]] - recovered[[i, j]]).abs() < 1e-15,
632 "Mismatch at ({}, {})",
633 i,
634 j
635 );
636 }
637 }
638 }
639
640 #[test]
641 fn test_empty_matrix_csr() {
642 let a: Array2<f64> = Array2::zeros((3, 4));
643 let csr = array2_to_csr(&a);
644 assert_eq!(csr.nnz(), 0);
645 assert_eq!(csr.shape(), (3, 4));
646
647 let recovered = csr_to_array2(&csr);
648 for i in 0..3 {
649 for j in 0..4 {
650 assert!(recovered[[i, j]].abs() < 1e-15);
651 }
652 }
653 }
654
655 #[test]
656 fn test_dense_matrix_csr() {
657 let a = array![[1.0f64, 2.0], [3.0, 4.0]];
658 let csr = array2_to_csr(&a);
659 assert_eq!(csr.nnz(), 4);
660 }
661
662 #[test]
663 fn test_array2_to_csr_exact_zero_sparsification_retains_tiny_values() {
664 let tiny = 1e-300f64;
668 let a = array![[tiny, 0.0], [0.0, 1.0]];
669 let csr = array2_to_csr(&a);
670 assert_eq!(csr.nnz(), 2);
671 assert_eq!(csr.get(0, 0), Some(&tiny));
672 }
673
674 #[test]
675 fn test_array2_to_csc_exact_zero_sparsification_retains_tiny_values() {
676 let tiny = 1e-300f64;
677 let a = array![[tiny, 0.0], [0.0, 1.0]];
678 let csc = array2_to_csc(&a);
679 assert_eq!(csc.nnz(), 2);
680 assert_eq!(csc.get(0, 0), Some(&tiny));
681 }
682
683 #[test]
684 fn test_array2_to_csr_never_drops_nan() {
685 let a = array![[f64::NAN, 0.0], [0.0, 1.0]];
686 let csr = array2_to_csr(&a);
687 assert_eq!(csr.nnz(), 2);
689 let stored = csr.get(0, 0).copied().expect("NaN entry must be stored");
690 assert!(stored.is_nan());
691 }
692
693 #[test]
694 fn test_array2_to_csc_never_drops_nan() {
695 let a = array![[f64::NAN, 0.0], [0.0, 1.0]];
696 let csc = array2_to_csc(&a);
697 assert_eq!(csc.nnz(), 2);
698 let stored = csc.get(0, 0).copied().expect("NaN entry must be stored");
699 assert!(stored.is_nan());
700 }
701
702 #[test]
703 fn test_array2_to_csr_exact_zero_drops_exact_zero_only() {
704 let a = array![[0.0f64, -0.0], [1.0, 0.0]];
706 let csr = array2_to_csr(&a);
707 assert_eq!(csr.nnz(), 1);
708 }
709
710 #[test]
711 fn test_array2_to_csr_with_tolerance_is_opt_in() {
712 let small = 1e-10f64;
714 let a = array![[small, 1.0]];
715 let default_csr = array2_to_csr(&a);
716 assert_eq!(default_csr.nnz(), 2);
717
718 let tol_csr = array2_to_csr_with_tolerance(&a, Some(1e-6));
720 assert_eq!(tol_csr.nnz(), 1);
721 assert_eq!(tol_csr.get(0, 1), Some(&1.0));
722 }
723
724 #[test]
725 fn test_array2_to_csc_with_tolerance_is_opt_in() {
726 let small = 1e-10f64;
727 let a = array![[small, 1.0]];
728 let default_csc = array2_to_csc(&a);
729 assert_eq!(default_csc.nnz(), 2);
730
731 let tol_csc = array2_to_csc_with_tolerance(&a, Some(1e-6));
732 assert_eq!(tol_csc.nnz(), 1);
733 assert_eq!(tol_csc.get(0, 1), Some(&1.0));
734 }
735
736 #[test]
737 fn test_array2_to_csr_with_tolerance_still_never_drops_nan() {
738 let a = array![[f64::NAN, 1.0]];
741 let tol_csr = array2_to_csr_with_tolerance(&a, Some(1e-6));
742 assert_eq!(tol_csr.nnz(), 2);
743 let stored = tol_csr
744 .get(0, 0)
745 .copied()
746 .expect("NaN entry must be stored even with tolerance");
747 assert!(stored.is_nan());
748 }
749
750 #[test]
751 fn test_identity_csr() {
752 let n = 5;
753 let mut a = Array2::<f64>::zeros((n, n));
754 for i in 0..n {
755 a[[i, i]] = 1.0;
756 }
757
758 let csr = array2_to_csr(&a);
759 assert_eq!(csr.nnz(), n);
760
761 for i in 0..n {
762 assert_eq!(csr.get(i, i), Some(&1.0));
763 }
764 }
765
766 #[test]
767 fn test_f32_conversions() {
768 let a = array![[1.0f32, 0.0, 2.0], [0.0, 3.0, 0.0]];
769 let csr = array2_to_csr(&a);
770 assert_eq!(csr.nnz(), 3);
771
772 let recovered = csr_to_array2(&csr);
773 assert!((recovered[[0, 0]] - 1.0f32).abs() < 1e-6);
774 assert!((recovered[[0, 2]] - 2.0f32).abs() < 1e-6);
775 assert!((recovered[[1, 1]] - 3.0f32).abs() < 1e-6);
776 }
777
778 #[test]
783 fn test_spmv_ndarray_basic() {
784 let a = array![[1.0f64, 0.0, 2.0], [0.0, 3.0, 0.0], [4.0, 0.0, 5.0]];
785 let csr = array2_to_csr(&a);
786 let x = array![1.0f64, 1.0, 1.0];
787
788 let y = spmv_ndarray(&csr, &x);
789
790 assert!((y[0] - 3.0).abs() < 1e-10);
794 assert!((y[1] - 3.0).abs() < 1e-10);
795 assert!((y[2] - 9.0).abs() < 1e-10);
796 }
797
798 #[test]
799 fn test_spmv_ndarray_identity() {
800 let n = 10;
801 let csr: CsrMatrix<f64> = CsrMatrix::eye(n);
802 let x = Array1::from_shape_fn(n, |i| (i + 1) as f64);
803
804 let y = spmv_ndarray(&csr, &x);
805
806 for i in 0..n {
807 assert!((y[i] - x[i]).abs() < 1e-15);
808 }
809 }
810
811 #[test]
812 fn test_spmv_full_ndarray() {
813 let a = array![[2.0f64, 0.0], [0.0, 3.0]];
814 let csr = array2_to_csr(&a);
815 let x = array![1.0f64, 2.0];
816 let mut y = array![10.0f64, 20.0];
817
818 spmv_full_ndarray(2.0, &csr, &x, 0.5, &mut y);
822
823 assert!((y[0] - 9.0).abs() < 1e-10);
824 assert!((y[1] - 22.0).abs() < 1e-10);
825 }
826
827 #[test]
832 fn test_sparse_solve_identity() {
833 let n = 5;
834 let csr: CsrMatrix<f64> = CsrMatrix::eye(n);
835 let b = Array1::from_shape_fn(n, |i| (i + 1) as f64);
836
837 let x = sparse_solve_ndarray(&csr, &b).expect("Solve should succeed for identity");
838
839 for i in 0..n {
840 assert!(
841 (x[i] - b[i]).abs() < 1e-8,
842 "Mismatch at {}: got {}, expected {}",
843 i,
844 x[i],
845 b[i]
846 );
847 }
848 }
849
850 #[test]
851 fn test_sparse_solve_spd() {
852 let values = vec![4.0, -1.0, -1.0, 4.0, -1.0, -1.0, 4.0];
854 let col_indices = vec![0, 1, 0, 1, 2, 1, 2];
855 let row_ptrs = vec![0, 2, 5, 7];
856 let csr = CsrMatrix::new(3, 3, row_ptrs, col_indices, values)
857 .expect("Failed to create CSR matrix");
858
859 let b = array![3.0f64, 2.0, 3.0];
860 let x = sparse_solve_ndarray(&csr, &b).expect("Solve should succeed for SPD matrix");
861
862 let residual = spmv_ndarray(&csr, &x);
864 for i in 0..3 {
865 assert!(
866 (residual[i] - b[i]).abs() < 1e-8,
867 "Residual mismatch at {}: got {}, expected {}",
868 i,
869 residual[i],
870 b[i]
871 );
872 }
873 }
874
875 #[test]
876 fn test_sparse_solve_larger_spd() {
877 let n = 20;
878 let mut values = Vec::new();
879 let mut col_indices = Vec::new();
880 let mut row_ptrs = vec![0usize];
881
882 for i in 0..n {
883 if i > 0 {
884 values.push(-1.0f64);
885 col_indices.push(i - 1);
886 }
887 values.push(4.0f64);
888 col_indices.push(i);
889 if i < n - 1 {
890 values.push(-1.0f64);
891 col_indices.push(i + 1);
892 }
893 row_ptrs.push(values.len());
894 }
895
896 let csr = CsrMatrix::new(n, n, row_ptrs, col_indices, values)
897 .expect("Failed to create CSR matrix");
898
899 let b = Array1::from_shape_fn(n, |i| (i + 1) as f64);
900 let x = sparse_solve_ndarray(&csr, &b).expect("Solve should succeed for larger SPD");
901
902 let residual = spmv_ndarray(&csr, &x);
903 for i in 0..n {
904 assert!(
905 (residual[i] - b[i]).abs() < 1e-6,
906 "Residual mismatch at {}: got {}, expected {}",
907 i,
908 residual[i],
909 b[i]
910 );
911 }
912 }
913
914 #[test]
915 fn test_sparse_solve_not_square() {
916 let csr: CsrMatrix<f64> = CsrMatrix::zeros(3, 4);
917 let b = array![1.0f64, 2.0, 3.0];
918 let result = sparse_solve_ndarray(&csr, &b);
919 assert!(result.is_err());
920 }
921
922 #[test]
923 fn test_sparse_solve_dimension_mismatch() {
924 let csr: CsrMatrix<f64> = CsrMatrix::eye(3);
925 let b = array![1.0f64, 2.0]; let result = sparse_solve_ndarray(&csr, &b);
927 assert!(result.is_err());
928 }
929
930 #[test]
931 fn test_sparse_solve_with_options() {
932 let csr: CsrMatrix<f64> = CsrMatrix::eye(3);
933 let b = array![1.0f64, 2.0, 3.0];
934
935 let x = sparse_solve_ndarray_with_options(&csr, &b, 1e-12, 100)
936 .expect("Solve with options should succeed");
937
938 for i in 0..3 {
939 assert!((x[i] - b[i]).abs() < 1e-10);
940 }
941 }
942}