1use crate::bsc::BscMatrix;
19use crate::bsr::BsrMatrix;
20use crate::coo::CooMatrix;
21use crate::csc::CscMatrix;
22use crate::csr::CsrMatrix;
23use crate::dia::DiaMatrix;
24use crate::ell::EllMatrix;
25use crate::hyb::{HybMatrix, HybWidthStrategy};
26use crate::sell::{SellMatrix, SliceSize};
27use oxiblas_core::scalar::{Field, Real, Scalar};
28
29pub fn csr_to_csc<T: Scalar + Clone>(csr: &CsrMatrix<T>) -> CscMatrix<T> {
34 let nrows = csr.nrows();
35 let ncols = csr.ncols();
36 let nnz = csr.nnz();
37
38 if nnz == 0 {
39 return CscMatrix::zeros(nrows, ncols);
40 }
41
42 let mut col_counts = vec![0usize; ncols];
44 for &col in csr.col_indices() {
45 col_counts[col] += 1;
46 }
47
48 let mut col_ptrs = vec![0usize; ncols + 1];
50 for i in 0..ncols {
51 col_ptrs[i + 1] = col_ptrs[i] + col_counts[i];
52 }
53
54 let mut row_indices = vec![0usize; nnz];
56 let mut values = vec![T::zero(); nnz];
57 let mut write_pos = col_ptrs.clone();
58
59 for row in 0..nrows {
60 let start = csr.row_ptrs()[row];
61 let end = csr.row_ptrs()[row + 1];
62
63 for i in start..end {
64 let col = csr.col_indices()[i];
65 let pos = write_pos[col];
66
67 row_indices[pos] = row;
68 values[pos] = csr.values()[i].clone();
69
70 write_pos[col] += 1;
71 }
72 }
73
74 unsafe { CscMatrix::new_unchecked(nrows, ncols, col_ptrs, row_indices, values) }
76}
77
78pub fn csc_to_csr<T: Scalar + Clone>(csc: &CscMatrix<T>) -> CsrMatrix<T> {
83 let nrows = csc.nrows();
84 let ncols = csc.ncols();
85 let nnz = csc.nnz();
86
87 if nnz == 0 {
88 return CsrMatrix::zeros(nrows, ncols);
89 }
90
91 let mut row_counts = vec![0usize; nrows];
93 for &row in csc.row_indices() {
94 row_counts[row] += 1;
95 }
96
97 let mut row_ptrs = vec![0usize; nrows + 1];
99 for i in 0..nrows {
100 row_ptrs[i + 1] = row_ptrs[i] + row_counts[i];
101 }
102
103 let mut col_indices = vec![0usize; nnz];
105 let mut values = vec![T::zero(); nnz];
106 let mut write_pos = row_ptrs.clone();
107
108 for col in 0..ncols {
109 let start = csc.col_ptrs()[col];
110 let end = csc.col_ptrs()[col + 1];
111
112 for i in start..end {
113 let row = csc.row_indices()[i];
114 let pos = write_pos[row];
115
116 col_indices[pos] = col;
117 values[pos] = csc.values()[i].clone();
118
119 write_pos[row] += 1;
120 }
121 }
122
123 unsafe { CsrMatrix::new_unchecked(nrows, ncols, row_ptrs, col_indices, values) }
125}
126
127fn flush_coo_group<T: Scalar<Real = T> + Clone + Field + Real>(
143 buf_indices: &mut Vec<usize>,
144 buf_values: &mut Vec<T>,
145 out_indices: &mut Vec<usize>,
146 out_values: &mut Vec<T>,
147) {
148 let eps = <T as Scalar>::epsilon();
149 for (idx, val) in buf_indices.drain(..).zip(buf_values.drain(..)) {
150 if Scalar::abs(val.clone()) > eps {
151 out_indices.push(idx);
152 out_values.push(val);
153 }
154 }
155}
156
157pub fn coo_to_csr<T: Scalar<Real = T> + Clone + Field + Real>(coo: &CooMatrix<T>) -> CsrMatrix<T> {
162 let nrows = coo.nrows();
163 let ncols = coo.ncols();
164
165 if coo.is_empty() {
166 return CsrMatrix::zeros(nrows, ncols);
167 }
168
169 let mut indices: Vec<usize> = (0..coo.len()).collect();
171 indices.sort_by_key(|&i| (coo.row_indices()[i], coo.col_indices()[i]));
172
173 let mut row_ptrs = Vec::with_capacity(nrows + 1);
179 let mut col_indices = Vec::with_capacity(coo.len());
180 let mut values: Vec<T> = Vec::with_capacity(coo.len());
181
182 row_ptrs.push(0);
183 let mut current_row = 0usize;
184
185 let mut row_col_buf: Vec<usize> = Vec::new();
186 let mut row_val_buf: Vec<T> = Vec::new();
187
188 for &idx in &indices {
189 let row = coo.row_indices()[idx];
190 let col = coo.col_indices()[idx];
191 let val = coo.values()[idx].clone();
192
193 if row != current_row {
194 flush_coo_group(
197 &mut row_col_buf,
198 &mut row_val_buf,
199 &mut col_indices,
200 &mut values,
201 );
202 row_ptrs.push(values.len());
203 current_row += 1;
204
205 while current_row < row {
206 row_ptrs.push(values.len());
207 current_row += 1;
208 }
209 }
210
211 if row_col_buf.last() == Some(&col) {
214 let last = row_val_buf.len() - 1;
215 row_val_buf[last] = row_val_buf[last].clone() + val;
216 } else {
217 row_col_buf.push(col);
218 row_val_buf.push(val);
219 }
220 }
221
222 flush_coo_group(
224 &mut row_col_buf,
225 &mut row_val_buf,
226 &mut col_indices,
227 &mut values,
228 );
229 row_ptrs.push(values.len());
230 current_row += 1;
231
232 while current_row < nrows {
234 row_ptrs.push(values.len());
235 current_row += 1;
236 }
237
238 unsafe { CsrMatrix::new_unchecked(nrows, ncols, row_ptrs, col_indices, values) }
240}
241
242pub fn coo_to_csc<T: Scalar<Real = T> + Clone + Field + Real>(coo: &CooMatrix<T>) -> CscMatrix<T> {
247 let nrows = coo.nrows();
248 let ncols = coo.ncols();
249
250 if coo.is_empty() {
251 return CscMatrix::zeros(nrows, ncols);
252 }
253
254 let mut indices: Vec<usize> = (0..coo.len()).collect();
256 indices.sort_by_key(|&i| (coo.col_indices()[i], coo.row_indices()[i]));
257
258 let mut col_ptrs = Vec::with_capacity(ncols + 1);
264 let mut row_indices = Vec::with_capacity(coo.len());
265 let mut values: Vec<T> = Vec::with_capacity(coo.len());
266
267 col_ptrs.push(0);
268 let mut current_col = 0usize;
269
270 let mut col_row_buf: Vec<usize> = Vec::new();
271 let mut col_val_buf: Vec<T> = Vec::new();
272
273 for &idx in &indices {
274 let row = coo.row_indices()[idx];
275 let col = coo.col_indices()[idx];
276 let val = coo.values()[idx].clone();
277
278 if col != current_col {
279 flush_coo_group(
282 &mut col_row_buf,
283 &mut col_val_buf,
284 &mut row_indices,
285 &mut values,
286 );
287 col_ptrs.push(values.len());
288 current_col += 1;
289
290 while current_col < col {
291 col_ptrs.push(values.len());
292 current_col += 1;
293 }
294 }
295
296 if col_row_buf.last() == Some(&row) {
299 let last = col_val_buf.len() - 1;
300 col_val_buf[last] = col_val_buf[last].clone() + val;
301 } else {
302 col_row_buf.push(row);
303 col_val_buf.push(val);
304 }
305 }
306
307 flush_coo_group(
309 &mut col_row_buf,
310 &mut col_val_buf,
311 &mut row_indices,
312 &mut values,
313 );
314 col_ptrs.push(values.len());
315 current_col += 1;
316
317 while current_col < ncols {
319 col_ptrs.push(values.len());
320 current_col += 1;
321 }
322
323 unsafe { CscMatrix::new_unchecked(nrows, ncols, col_ptrs, row_indices, values) }
325}
326
327pub fn csr_to_coo<T: Scalar + Clone>(csr: &CsrMatrix<T>) -> CooMatrix<T> {
329 let nrows = csr.nrows();
330 let ncols = csr.ncols();
331 let nnz = csr.nnz();
332
333 let mut row_indices = Vec::with_capacity(nnz);
334 let mut col_indices = Vec::with_capacity(nnz);
335 let mut values = Vec::with_capacity(nnz);
336
337 for row in 0..nrows {
338 let start = csr.row_ptrs()[row];
339 let end = csr.row_ptrs()[row + 1];
340
341 for i in start..end {
342 row_indices.push(row);
343 col_indices.push(csr.col_indices()[i]);
344 values.push(csr.values()[i].clone());
345 }
346 }
347
348 unsafe { CooMatrix::new_unchecked(nrows, ncols, row_indices, col_indices, values) }
350}
351
352pub fn csc_to_coo<T: Scalar + Clone>(csc: &CscMatrix<T>) -> CooMatrix<T> {
354 let nrows = csc.nrows();
355 let ncols = csc.ncols();
356 let nnz = csc.nnz();
357
358 let mut row_indices = Vec::with_capacity(nnz);
359 let mut col_indices = Vec::with_capacity(nnz);
360 let mut values = Vec::with_capacity(nnz);
361
362 for col in 0..ncols {
363 let start = csc.col_ptrs()[col];
364 let end = csc.col_ptrs()[col + 1];
365
366 for i in start..end {
367 row_indices.push(csc.row_indices()[i]);
368 col_indices.push(col);
369 values.push(csc.values()[i].clone());
370 }
371 }
372
373 unsafe { CooMatrix::new_unchecked(nrows, ncols, row_indices, col_indices, values) }
375}
376
377pub fn csr_to_dia<T: Scalar + Clone + Field>(
390 csr: &CsrMatrix<T>,
391 offsets: Option<Vec<isize>>,
392) -> DiaMatrix<T> {
393 let (nrows, ncols) = csr.shape();
394 let eps = <T as Scalar>::epsilon();
395
396 let offsets = offsets.unwrap_or_else(|| {
398 let mut found = std::collections::HashSet::new();
399 for (row, col, val) in csr.iter() {
400 if Scalar::abs(val.clone()) > eps {
401 found.insert(col as isize - row as isize);
402 }
403 }
404 let mut offsets: Vec<_> = found.into_iter().collect();
405 offsets.sort();
406 offsets
407 });
408
409 if offsets.is_empty() {
410 return DiaMatrix::zeros(nrows, ncols);
411 }
412
413 let diag_len = nrows.min(ncols);
414 let mut data = Vec::with_capacity(offsets.len());
415
416 for &offset in &offsets {
417 let mut diag = vec![T::zero(); diag_len];
418
419 for (row, col, val) in csr.iter() {
423 let expected_col = (row as isize + offset) as usize;
424 if col == expected_col && row < nrows && col < ncols {
425 let idx = (row as isize + offset) as usize;
427 if idx < diag_len {
428 diag[idx] = val.clone();
429 }
430 }
431 }
432
433 data.push(diag);
434 }
435
436 unsafe { DiaMatrix::new_unchecked(nrows, ncols, offsets, data) }
438}
439
440pub fn dia_to_csr<T: Scalar + Clone + Field>(dia: &DiaMatrix<T>) -> CsrMatrix<T> {
444 dia.to_csr()
445}
446
447pub fn csr_to_ell<T: Scalar + Clone + Field>(
460 csr: &CsrMatrix<T>,
461 max_width: Option<usize>,
462) -> Result<EllMatrix<T>, crate::ell::EllError> {
463 EllMatrix::from_csr(csr, max_width)
464}
465
466pub fn ell_to_csr<T: Scalar + Clone + Field>(ell: &EllMatrix<T>) -> CsrMatrix<T> {
470 ell.to_csr()
471}
472
473pub fn csr_to_bsr<T: Scalar + Clone + Field>(
487 csr: &CsrMatrix<T>,
488 block_rows: usize,
489 block_cols: usize,
490) -> BsrMatrix<T> {
491 BsrMatrix::from_csr(csr, block_rows, block_cols)
492}
493
494pub fn bsr_to_csr<T: Scalar + Clone + Field>(bsr: &BsrMatrix<T>) -> CsrMatrix<T> {
498 bsr.to_csr()
499}
500
501pub fn dia_to_ell<T: Scalar + Clone + Field>(
507 dia: &DiaMatrix<T>,
508 max_width: Option<usize>,
509) -> Result<EllMatrix<T>, crate::ell::EllError> {
510 let csr = dia.to_csr();
511 EllMatrix::from_csr(&csr, max_width)
512}
513
514pub fn ell_to_dia<T: Scalar + Clone + Field>(
516 ell: &EllMatrix<T>,
517 offsets: Option<Vec<isize>>,
518) -> DiaMatrix<T> {
519 let csr = ell.to_csr();
520 csr_to_dia(&csr, offsets)
521}
522
523pub fn dia_to_bsr<T: Scalar + Clone + Field>(
525 dia: &DiaMatrix<T>,
526 block_rows: usize,
527 block_cols: usize,
528) -> BsrMatrix<T> {
529 let csr = dia.to_csr();
530 BsrMatrix::from_csr(&csr, block_rows, block_cols)
531}
532
533pub fn bsr_to_dia<T: Scalar + Clone + Field>(
535 bsr: &BsrMatrix<T>,
536 offsets: Option<Vec<isize>>,
537) -> DiaMatrix<T> {
538 let csr = bsr.to_csr();
539 csr_to_dia(&csr, offsets)
540}
541
542pub fn ell_to_bsr<T: Scalar + Clone + Field>(
544 ell: &EllMatrix<T>,
545 block_rows: usize,
546 block_cols: usize,
547) -> BsrMatrix<T> {
548 let csr = ell.to_csr();
549 BsrMatrix::from_csr(&csr, block_rows, block_cols)
550}
551
552pub fn bsr_to_ell<T: Scalar + Clone + Field>(
554 bsr: &BsrMatrix<T>,
555 max_width: Option<usize>,
556) -> Result<EllMatrix<T>, crate::ell::EllError> {
557 let csr = bsr.to_csr();
558 EllMatrix::from_csr(&csr, max_width)
559}
560
561pub fn csr_to_bsc<T: Scalar + Clone + Field>(
573 csr: &CsrMatrix<T>,
574 block_rows: usize,
575 block_cols: usize,
576) -> BscMatrix<T> {
577 let bsr = BsrMatrix::from_csr(csr, block_rows, block_cols);
578 BscMatrix::from_bsr(&bsr)
579}
580
581pub fn bsc_to_csr<T: Scalar + Clone + Field>(bsc: &BscMatrix<T>) -> CsrMatrix<T> {
583 let bsr = bsc.to_bsr();
584 bsr.to_csr()
585}
586
587pub fn bsc_to_bsr<T: Scalar + Clone + Field>(bsc: &BscMatrix<T>) -> BsrMatrix<T> {
589 bsc.to_bsr()
590}
591
592pub fn bsr_to_bsc<T: Scalar + Clone + Field>(bsr: &BsrMatrix<T>) -> BscMatrix<T> {
594 BscMatrix::from_bsr(bsr)
595}
596
597pub fn csr_to_hyb<T: Scalar + Clone + Field>(
608 csr: &CsrMatrix<T>,
609 strategy: HybWidthStrategy,
610) -> HybMatrix<T> {
611 HybMatrix::from_csr(csr, strategy)
612}
613
614pub fn hyb_to_csr<T: Scalar + Clone + Field>(hyb: &HybMatrix<T>) -> CsrMatrix<T> {
616 hyb.to_csr()
617}
618
619pub fn ell_to_hyb<T: Scalar + Clone + Field>(ell: &EllMatrix<T>) -> HybMatrix<T> {
621 HybMatrix::from_ell(ell)
622}
623
624pub fn hyb_to_ell<T: Scalar + Clone + Field>(hyb: &HybMatrix<T>) -> EllMatrix<T> {
626 hyb.to_ell()
627}
628
629pub fn csr_to_sell<T: Scalar + Clone + Field>(
640 csr: &CsrMatrix<T>,
641 slice_size: SliceSize,
642) -> SellMatrix<T> {
643 SellMatrix::from_csr(csr, slice_size)
644}
645
646pub fn sell_to_csr<T: Scalar + Clone + Field>(sell: &SellMatrix<T>) -> CsrMatrix<T> {
648 sell.to_csr()
649}
650
651#[derive(Debug, Clone, Copy, PartialEq, Eq)]
657pub enum RecommendedFormat {
658 Csr,
660 Csc,
662 Dia,
664 Ell,
666 Hyb,
668 Sell,
670 Bsr,
672 Bsc,
674}
675
676#[derive(Debug, Clone)]
678pub struct SparsityAnalysis {
679 pub nrows: usize,
681 pub ncols: usize,
683 pub nnz: usize,
685 pub density: f64,
687 pub max_row_length: usize,
689 pub min_row_length: usize,
691 pub avg_row_length: f64,
693 pub row_length_stddev: f64,
695 pub num_diagonals: usize,
697 pub has_block_structure: bool,
699 pub detected_block_size: Option<(usize, usize)>,
701 pub recommended_format: RecommendedFormat,
703}
704
705pub fn analyze_sparsity_pattern<T: Scalar + Clone + Field>(csr: &CsrMatrix<T>) -> SparsityAnalysis {
711 let (nrows, ncols) = csr.shape();
712 let nnz = csr.nnz();
713 let eps = <T as Scalar>::epsilon();
714
715 if nrows == 0 || ncols == 0 {
716 return SparsityAnalysis {
717 nrows,
718 ncols,
719 nnz,
720 density: 0.0,
721 max_row_length: 0,
722 min_row_length: 0,
723 avg_row_length: 0.0,
724 row_length_stddev: 0.0,
725 num_diagonals: 0,
726 has_block_structure: false,
727 detected_block_size: None,
728 recommended_format: RecommendedFormat::Csr,
729 };
730 }
731
732 let mut row_lengths = Vec::with_capacity(nrows);
734 for row in 0..nrows {
735 let mut count = 0;
736 for (_, val) in csr.row_iter(row) {
737 if Scalar::abs(val.clone()) > eps {
738 count += 1;
739 }
740 }
741 row_lengths.push(count);
742 }
743
744 let max_row_length = row_lengths.iter().max().copied().unwrap_or(0);
745 let min_row_length = row_lengths.iter().min().copied().unwrap_or(0);
746 let avg_row_length = if nrows > 0 {
747 row_lengths.iter().sum::<usize>() as f64 / nrows as f64
748 } else {
749 0.0
750 };
751
752 let variance: f64 = row_lengths
754 .iter()
755 .map(|&x| {
756 let diff = x as f64 - avg_row_length;
757 diff * diff
758 })
759 .sum::<f64>()
760 / nrows.max(1) as f64;
761 let row_length_stddev = variance.sqrt();
762
763 let mut diagonals = std::collections::HashSet::new();
765 for (row, col, val) in csr.iter() {
766 if Scalar::abs(val.clone()) > eps {
767 diagonals.insert(col as isize - row as isize);
768 }
769 }
770 let num_diagonals = diagonals.len();
771
772 let (has_block_structure, detected_block_size) = detect_block_structure(csr);
774
775 let density = if nrows * ncols > 0 {
776 nnz as f64 / (nrows * ncols) as f64
777 } else {
778 0.0
779 };
780
781 let recommended_format = determine_recommended_format(
783 nrows,
784 ncols,
785 nnz,
786 max_row_length,
787 min_row_length,
788 row_length_stddev,
789 num_diagonals,
790 has_block_structure,
791 );
792
793 SparsityAnalysis {
794 nrows,
795 ncols,
796 nnz,
797 density,
798 max_row_length,
799 min_row_length,
800 avg_row_length,
801 row_length_stddev,
802 num_diagonals,
803 has_block_structure,
804 detected_block_size,
805 recommended_format,
806 }
807}
808
809fn detect_block_structure<T: Scalar + Clone + Field>(
811 csr: &CsrMatrix<T>,
812) -> (bool, Option<(usize, usize)>) {
813 let (nrows, ncols) = csr.shape();
814 let eps = <T as Scalar>::epsilon();
815
816 if nrows < 4 || ncols < 4 {
817 return (false, None);
818 }
819
820 for block_size in [2, 3, 4, 6, 8] {
822 if nrows % block_size != 0 || ncols % block_size != 0 {
823 continue;
824 }
825
826 let _num_block_rows = nrows / block_size;
827 let _num_block_cols = ncols / block_size;
828
829 let block_aligned = true;
831 let mut blocks_found = std::collections::HashSet::new();
832
833 for (row, col, val) in csr.iter() {
834 if Scalar::abs(val.clone()) > eps {
835 let block_row = row / block_size;
836 let block_col = col / block_size;
837 blocks_found.insert((block_row, block_col));
838 }
839 }
840
841 let mut dense_blocks = 0;
843 for &(br, bc) in &blocks_found {
844 let mut count = 0;
845 for i in 0..block_size {
846 for j in 0..block_size {
847 let row = br * block_size + i;
848 let col = bc * block_size + j;
849 if let Some(val) = csr.get(row, col) {
850 if Scalar::abs(val.clone()) > eps {
851 count += 1;
852 }
853 }
854 }
855 }
856 if count * 2 >= block_size * block_size {
858 dense_blocks += 1;
859 }
860 }
861
862 if !blocks_found.is_empty() && dense_blocks * 10 >= blocks_found.len() * 7 {
864 return (true, Some((block_size, block_size)));
865 }
866 if !block_aligned {
867 continue;
869 }
870 }
871
872 (false, None)
873}
874
875fn determine_recommended_format(
877 nrows: usize,
878 ncols: usize,
879 nnz: usize,
880 max_row_length: usize,
881 min_row_length: usize,
882 row_length_stddev: f64,
883 num_diagonals: usize,
884 has_block_structure: bool,
885) -> RecommendedFormat {
886 if nnz == 0 || nrows <= 10 || ncols <= 10 {
888 return RecommendedFormat::Csr;
889 }
890
891 let avg_row_length = nnz as f64 / nrows.max(1) as f64;
892
893 if has_block_structure {
895 return RecommendedFormat::Bsr;
896 }
897
898 if num_diagonals <= 10 && num_diagonals * 2 <= nrows.max(1) {
901 return RecommendedFormat::Dia;
902 }
903
904 let coefficient_of_variation = row_length_stddev / avg_row_length.max(1.0);
906
907 if coefficient_of_variation < 0.3 {
908 return RecommendedFormat::Ell;
910 }
911
912 if coefficient_of_variation < 0.8 {
913 return RecommendedFormat::Hyb;
915 }
916
917 if max_row_length > min_row_length * 10 {
919 return RecommendedFormat::Sell;
921 }
922
923 RecommendedFormat::Csr
925}
926
927#[cfg(test)]
928mod tests {
929 use super::*;
930 use crate::bsr::DenseBlock;
931
932 #[test]
933 fn test_csr_to_csc() {
934 let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
938 let col_indices = vec![0, 2, 1, 0, 2];
939 let row_ptrs = vec![0, 2, 3, 5];
940
941 let csr = CsrMatrix::new(3, 3, row_ptrs, col_indices, values).unwrap();
942 let csc = csr_to_csc(&csr);
943
944 assert_eq!(csc.nnz(), 5);
945 assert_eq!(csc.get(0, 0), Some(&1.0));
946 assert_eq!(csc.get(0, 2), Some(&2.0));
947 assert_eq!(csc.get(1, 1), Some(&3.0));
948 assert_eq!(csc.get(2, 0), Some(&4.0));
949 assert_eq!(csc.get(2, 2), Some(&5.0));
950 }
951
952 #[test]
953 fn test_csc_to_csr() {
954 let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
958 let row_indices = vec![0, 2, 1, 0, 2];
959 let col_ptrs = vec![0, 2, 3, 5];
960
961 let csc = CscMatrix::new(3, 3, col_ptrs, row_indices, values).unwrap();
962 let csr = csc_to_csr(&csc);
963
964 assert_eq!(csr.nnz(), 5);
965 assert_eq!(csr.get(0, 0), Some(&1.0));
966 assert_eq!(csr.get(0, 2), Some(&4.0));
967 assert_eq!(csr.get(1, 1), Some(&3.0));
968 assert_eq!(csr.get(2, 0), Some(&2.0));
969 assert_eq!(csr.get(2, 2), Some(&5.0));
970 }
971
972 #[test]
973 fn test_coo_to_csr() {
974 let row_indices = vec![0, 1, 2, 0, 2];
975 let col_indices = vec![0, 1, 0, 2, 2];
976 let values = vec![1.0f64, 3.0, 4.0, 2.0, 5.0];
977
978 let coo = CooMatrix::new(3, 3, row_indices, col_indices, values).unwrap();
979 let csr = coo_to_csr(&coo);
980
981 assert_eq!(csr.nnz(), 5);
982 assert_eq!(csr.get(0, 0), Some(&1.0));
983 assert_eq!(csr.get(0, 2), Some(&2.0));
984 assert_eq!(csr.get(1, 1), Some(&3.0));
985 assert_eq!(csr.get(2, 0), Some(&4.0));
986 assert_eq!(csr.get(2, 2), Some(&5.0));
987 }
988
989 #[test]
990 fn test_coo_to_csr_duplicates() {
991 let row_indices = vec![0, 0, 1];
993 let col_indices = vec![0, 0, 1];
994 let values = vec![1.0f64, 2.0, 3.0];
995
996 let coo = CooMatrix::new(2, 2, row_indices, col_indices, values).unwrap();
997 let csr = coo_to_csr(&coo);
998
999 assert_eq!(csr.nnz(), 2);
1000 assert_eq!(csr.get(0, 0), Some(&3.0)); assert_eq!(csr.get(1, 1), Some(&3.0));
1002 }
1003
1004 #[test]
1005 fn test_coo_to_csr_row_ptrs_length() {
1006 let nrows = 4;
1009 let row_indices = vec![0, 1, 3];
1010 let col_indices = vec![0, 1, 0];
1011 let values = vec![1.0f64, 2.0, 3.0];
1012
1013 let coo = CooMatrix::new(nrows, 2, row_indices, col_indices, values).unwrap();
1014 let csr = coo_to_csr(&coo);
1015
1016 assert_eq!(csr.row_ptrs().len(), nrows + 1);
1017 }
1018
1019 #[test]
1020 fn test_coo_to_csr_zero_cancellation_does_not_misattribute_row() {
1021 let row_indices = vec![0, 0, 2];
1027 let col_indices = vec![0, 0, 1];
1028 let values = vec![5.0f64, -5.0, 9.0];
1029
1030 let coo = CooMatrix::new(3, 2, row_indices, col_indices, values).unwrap();
1031 let csr = coo_to_csr(&coo);
1032
1033 assert_eq!(csr.row_ptrs().len(), 3 + 1);
1034 assert_eq!(csr.nnz(), 1);
1035 assert_eq!(csr.get(0, 0), None);
1036 assert_eq!(csr.get(1, 1), None);
1037 assert_eq!(csr.get(2, 1), Some(&9.0));
1038 }
1039
1040 #[test]
1041 fn test_coo_to_csc_col_ptrs_length() {
1042 let ncols = 4;
1045 let row_indices = vec![0, 1, 0];
1046 let col_indices = vec![0, 1, 3];
1047 let values = vec![1.0f64, 2.0, 3.0];
1048
1049 let coo = CooMatrix::new(2, ncols, row_indices, col_indices, values).unwrap();
1050 let csc = coo_to_csc(&coo);
1051
1052 assert_eq!(csc.col_ptrs().len(), ncols + 1);
1053 }
1054
1055 #[test]
1056 fn test_coo_to_csc_zero_cancellation_does_not_misattribute_column() {
1057 let row_indices = vec![0, 0, 1];
1061 let col_indices = vec![0, 0, 2];
1062 let values = vec![5.0f64, -5.0, 9.0];
1063
1064 let coo = CooMatrix::new(2, 3, row_indices, col_indices, values).unwrap();
1065 let csc = coo_to_csc(&coo);
1066
1067 assert_eq!(csc.col_ptrs().len(), 3 + 1);
1068 assert_eq!(csc.nnz(), 1);
1069 assert_eq!(csc.get(0, 0), None);
1070 assert_eq!(csc.get(0, 1), None);
1071 assert_eq!(csc.get(1, 2), Some(&9.0));
1072 }
1073
1074 #[test]
1075 fn test_coo_to_csc() {
1076 let row_indices = vec![0, 1, 2, 0, 2];
1077 let col_indices = vec![0, 1, 0, 2, 2];
1078 let values = vec![1.0f64, 3.0, 4.0, 2.0, 5.0];
1079
1080 let coo = CooMatrix::new(3, 3, row_indices, col_indices, values).unwrap();
1081 let csc = coo_to_csc(&coo);
1082
1083 assert_eq!(csc.nnz(), 5);
1084 assert_eq!(csc.get(0, 0), Some(&1.0));
1085 assert_eq!(csc.get(0, 2), Some(&2.0));
1086 assert_eq!(csc.get(1, 1), Some(&3.0));
1087 assert_eq!(csc.get(2, 0), Some(&4.0));
1088 assert_eq!(csc.get(2, 2), Some(&5.0));
1089 }
1090
1091 #[test]
1092 fn test_roundtrip_csr_csc_csr() {
1093 let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
1094 let col_indices = vec![0, 2, 1, 0, 2];
1095 let row_ptrs = vec![0, 2, 3, 5];
1096
1097 let csr1 = CsrMatrix::new(3, 3, row_ptrs, col_indices, values).unwrap();
1098 let csc = csr_to_csc(&csr1);
1099 let csr2 = csc_to_csr(&csc);
1100
1101 assert_eq!(csr1.nnz(), csr2.nnz());
1102 for row in 0..3 {
1103 for col in 0..3 {
1104 assert_eq!(csr1.get(row, col), csr2.get(row, col));
1105 }
1106 }
1107 }
1108
1109 #[test]
1110 fn test_csr_to_coo() {
1111 let values = vec![1.0f64, 2.0, 3.0];
1112 let col_indices = vec![0, 1, 2];
1113 let row_ptrs = vec![0, 1, 2, 3];
1114
1115 let csr = CsrMatrix::new(3, 3, row_ptrs, col_indices, values).unwrap();
1116 let coo = csr_to_coo(&csr);
1117
1118 assert_eq!(coo.len(), 3);
1119 let entries: Vec<_> = coo.iter().map(|(r, c, v)| (r, c, *v)).collect();
1120 assert_eq!(entries, vec![(0, 0, 1.0), (1, 1, 2.0), (2, 2, 3.0)]);
1121 }
1122
1123 #[test]
1124 fn test_empty_matrix_conversion() {
1125 let csr: CsrMatrix<f64> = CsrMatrix::zeros(5, 3);
1126 let csc = csr_to_csc(&csr);
1127
1128 assert_eq!(csc.nrows(), 5);
1129 assert_eq!(csc.ncols(), 3);
1130 assert_eq!(csc.nnz(), 0);
1131 }
1132
1133 #[test]
1138 fn test_csr_to_dia_tridiagonal() {
1139 let values = vec![4.0f64, 1.0, 2.0, 5.0, 1.0, 3.0, 6.0];
1144 let col_indices = vec![0, 1, 0, 1, 2, 1, 2];
1145 let row_ptrs = vec![0, 2, 5, 7];
1146
1147 let csr = CsrMatrix::new(3, 3, row_ptrs, col_indices, values).unwrap();
1148 let dia = csr_to_dia(&csr, None);
1149
1150 assert_eq!(dia.ndiag(), 3);
1151 assert_eq!(dia.get(0, 0), Some(&4.0));
1152 assert_eq!(dia.get(0, 1), Some(&1.0));
1153 assert_eq!(dia.get(1, 0), Some(&2.0));
1154 assert_eq!(dia.get(1, 1), Some(&5.0));
1155 assert_eq!(dia.get(2, 2), Some(&6.0));
1156 }
1157
1158 #[test]
1159 fn test_dia_to_csr() {
1160 let offsets = vec![-1, 0, 1];
1161 let data = vec![
1162 vec![2.0, 3.0, 0.0],
1163 vec![4.0, 5.0, 6.0],
1164 vec![0.0, 1.0, 1.0],
1165 ];
1166
1167 let dia = DiaMatrix::new(3, 3, offsets, data).unwrap();
1168 let csr = dia_to_csr(&dia);
1169
1170 assert_eq!(csr.nrows(), 3);
1171 assert_eq!(csr.get(0, 0), Some(&4.0));
1172 assert_eq!(csr.get(1, 0), Some(&2.0));
1173 }
1174
1175 #[test]
1176 fn test_csr_dia_roundtrip() {
1177 let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
1178 let col_indices = vec![0, 1, 1, 0, 2];
1179 let row_ptrs = vec![0, 2, 3, 5];
1180
1181 let csr1 = CsrMatrix::new(3, 3, row_ptrs, col_indices, values).unwrap();
1182 let dia = csr_to_dia(&csr1, None);
1183 let csr2 = dia_to_csr(&dia);
1184
1185 for row in 0..3 {
1186 for col in 0..3 {
1187 let v1 = csr1.get(row, col).cloned().unwrap_or(0.0);
1188 let v2 = csr2.get(row, col).cloned().unwrap_or(0.0);
1189 assert!((v1 - v2).abs() < 1e-10);
1190 }
1191 }
1192 }
1193
1194 #[test]
1199 fn test_csr_to_ell() {
1200 let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0, 6.0];
1201 let col_indices = vec![0, 1, 1, 2, 0, 3];
1202 let row_ptrs = vec![0, 2, 4, 6];
1203
1204 let csr = CsrMatrix::new(3, 4, row_ptrs, col_indices, values).unwrap();
1205 let ell = csr_to_ell(&csr, None).unwrap();
1206
1207 assert_eq!(ell.width(), 2);
1208 assert_eq!(ell.get(0, 0), Some(&1.0));
1209 assert_eq!(ell.get(1, 2), Some(&4.0));
1210 }
1211
1212 #[test]
1213 fn test_ell_to_csr() {
1214 let data = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
1215 let indices = vec![vec![0, 1], vec![1, 2]];
1216
1217 let ell = EllMatrix::new(2, 3, 2, data, indices).unwrap();
1218 let csr = ell_to_csr(&ell);
1219
1220 assert_eq!(csr.nrows(), 2);
1221 assert_eq!(csr.get(0, 0), Some(&1.0));
1222 assert_eq!(csr.get(1, 2), Some(&4.0));
1223 }
1224
1225 #[test]
1226 fn test_csr_ell_roundtrip() {
1227 let values = vec![1.0f64, 2.0, 3.0, 4.0];
1228 let col_indices = vec![0, 1, 1, 2];
1229 let row_ptrs = vec![0, 2, 4];
1230
1231 let csr1 = CsrMatrix::new(2, 3, row_ptrs, col_indices, values).unwrap();
1232 let ell = csr_to_ell(&csr1, None).unwrap();
1233 let csr2 = ell_to_csr(&ell);
1234
1235 for row in 0..2 {
1236 for col in 0..3 {
1237 let v1 = csr1.get(row, col).cloned().unwrap_or(0.0);
1238 let v2 = csr2.get(row, col).cloned().unwrap_or(0.0);
1239 assert!((v1 - v2).abs() < 1e-10);
1240 }
1241 }
1242 }
1243
1244 #[test]
1249 fn test_csr_to_bsr() {
1250 let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
1252 let col_indices = vec![0, 1, 0, 1, 2, 3, 2, 3];
1253 let row_ptrs = vec![0, 2, 4, 6, 8];
1254
1255 let csr = CsrMatrix::new(4, 4, row_ptrs, col_indices, values).unwrap();
1256 let bsr = csr_to_bsr(&csr, 2, 2);
1257
1258 assert_eq!(bsr.nblocks(), 2);
1259 assert_eq!(bsr.get(0, 0), Some(1.0));
1260 assert_eq!(bsr.get(3, 3), Some(8.0));
1261 }
1262
1263 #[test]
1264 fn test_bsr_to_csr() {
1265 let block1 = DenseBlock::new(2, 2, vec![1.0, 2.0, 3.0, 4.0]);
1266 let block2 = DenseBlock::new(2, 2, vec![5.0, 6.0, 7.0, 8.0]);
1267
1268 let bsr =
1269 BsrMatrix::new(4, 4, 2, 2, vec![0, 1, 2], vec![0, 1], vec![block1, block2]).unwrap();
1270
1271 let csr = bsr_to_csr(&bsr);
1272
1273 assert_eq!(csr.nrows(), 4);
1274 assert_eq!(csr.get(0, 0), Some(&1.0));
1275 assert_eq!(csr.get(2, 2), Some(&5.0));
1276 }
1277
1278 #[test]
1279 fn test_csr_bsr_roundtrip() {
1280 let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
1281 let col_indices = vec![0, 1, 0, 1, 2, 3, 2, 3];
1282 let row_ptrs = vec![0, 2, 4, 6, 8];
1283
1284 let csr1 = CsrMatrix::new(4, 4, row_ptrs, col_indices, values).unwrap();
1285 let bsr = csr_to_bsr(&csr1, 2, 2);
1286 let csr2 = bsr_to_csr(&bsr);
1287
1288 for row in 0..4 {
1289 for col in 0..4 {
1290 let v1 = csr1.get(row, col).cloned().unwrap_or(0.0);
1291 let v2 = csr2.get(row, col).cloned().unwrap_or(0.0);
1292 assert!((v1 - v2).abs() < 1e-10);
1293 }
1294 }
1295 }
1296
1297 #[test]
1302 fn test_dia_to_ell() {
1303 let offsets = vec![0];
1304 let data = vec![vec![1.0, 2.0, 3.0]];
1305
1306 let dia = DiaMatrix::new(3, 3, offsets, data).unwrap();
1307 let ell = dia_to_ell(&dia, None).unwrap();
1308
1309 assert_eq!(ell.width(), 1);
1310 assert_eq!(ell.get(0, 0), Some(&1.0));
1311 assert_eq!(ell.get(1, 1), Some(&2.0));
1312 }
1313
1314 #[test]
1315 fn test_dia_to_bsr() {
1316 let offsets = vec![0];
1317 let data = vec![vec![1.0, 2.0, 3.0, 4.0]];
1318
1319 let dia = DiaMatrix::new(4, 4, offsets, data).unwrap();
1320 let bsr = dia_to_bsr(&dia, 2, 2);
1321
1322 assert_eq!(bsr.get(0, 0), Some(1.0));
1323 assert_eq!(bsr.get(1, 1), Some(2.0));
1324 }
1325
1326 #[test]
1327 fn test_ell_to_bsr() {
1328 let data = vec![
1329 vec![1.0, 2.0],
1330 vec![3.0, 4.0],
1331 vec![5.0, 6.0],
1332 vec![7.0, 8.0],
1333 ];
1334 let indices = vec![vec![0, 1], vec![0, 1], vec![2, 3], vec![2, 3]];
1335
1336 let ell = EllMatrix::new(4, 4, 2, data, indices).unwrap();
1337 let bsr = ell_to_bsr(&ell, 2, 2);
1338
1339 assert_eq!(bsr.nrows(), 4);
1340 assert_eq!(bsr.get(0, 0), Some(1.0));
1341 assert_eq!(bsr.get(3, 3), Some(8.0));
1342 }
1343}