Skip to main content

oxiblas_sparse/
convert.rs

1//! Format conversion utilities for sparse matrices.
2//!
3//! Provides efficient conversions between:
4//! - CSR ↔ CSC
5//! - COO → CSR/CSC
6//! - CSR ↔ DIA
7//! - CSR ↔ ELL
8//! - CSR ↔ BSR
9//! - CSR ↔ BSC
10//! - CSR ↔ HYB
11//! - CSR ↔ SELL
12//! - Sparse ↔ Dense
13//!
14//! # Automatic Format Selection
15//!
16//! Use [`analyze_sparsity_pattern`] to determine the optimal format for a matrix.
17
18use 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
29/// Converts a CSR matrix to CSC format.
30///
31/// Time complexity: O(nnz)
32/// Space complexity: O(nnz) for the output
33pub 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    // Count entries per column
43    let mut col_counts = vec![0usize; ncols];
44    for &col in csr.col_indices() {
45        col_counts[col] += 1;
46    }
47
48    // Build column pointers
49    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    // Fill in values and row indices
55    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    // SAFETY: We've constructed valid CSC data
75    unsafe { CscMatrix::new_unchecked(nrows, ncols, col_ptrs, row_indices, values) }
76}
77
78/// Converts a CSC matrix to CSR format.
79///
80/// Time complexity: O(nnz)
81/// Space complexity: O(nnz) for the output
82pub 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    // Count entries per row
92    let mut row_counts = vec![0usize; nrows];
93    for &row in csc.row_indices() {
94        row_counts[row] += 1;
95    }
96
97    // Build row pointers
98    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    // Fill in values and column indices
104    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    // SAFETY: We've constructed valid CSR data
124    unsafe { CsrMatrix::new_unchecked(nrows, ncols, row_ptrs, col_indices, values) }
125}
126
127/// Flushes a per-row (or per-column) COO accumulation buffer into the
128/// output CSR/CSC arrays.
129///
130/// Duplicate entries at the same (row, col) have already been summed into
131/// `buf_indices`/`buf_values` while the row (or column) was being scanned.
132/// This drops any entries that summed to exactly zero (within epsilon) and
133/// appends the survivors to the output arrays.
134///
135/// Crucially, this is called exactly once per row/column, *after* every
136/// entry belonging to that row/column has been accumulated. That keeps
137/// zero-pruning entirely local to one row/column: it can never run after
138/// `values.len()` has already been captured into a `row_ptrs`/`col_ptrs`
139/// boundary for a *later* row/column, which is what previously let a
140/// pruned zero corrupt the pointer array and misattribute later entries
141/// to the wrong row/column.
142fn 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
157/// Converts a COO matrix to CSR format, summing duplicate entries.
158///
159/// Time complexity: O(nnz log nnz) due to sorting
160/// Space complexity: O(nnz) for the output
161pub 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    // Sort entries by (row, col)
170    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    // Build CSR data, summing duplicates. Each row is accumulated into a
174    // scratch buffer first; only once the *entire* row has been summed do
175    // we prune exact-zero results and append the survivors to the output
176    // arrays, then push the row_ptrs boundary. Zero-pruning therefore
177    // never straddles a row boundary (see `flush_coo_group`).
178    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            // Row boundary: flush the completed row, then fill any fully
195            // empty rows between it and the new row.
196            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        // Accumulate into the current row's buffer, summing duplicates at
212        // the same column (adjacent, since entries are sorted).
213        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 the final row's buffer.
223    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    // Fill any trailing empty rows.
233    while current_row < nrows {
234        row_ptrs.push(values.len());
235        current_row += 1;
236    }
237
238    // SAFETY: We've constructed valid CSR data
239    unsafe { CsrMatrix::new_unchecked(nrows, ncols, row_ptrs, col_indices, values) }
240}
241
242/// Converts a COO matrix to CSC format, summing duplicate entries.
243///
244/// Time complexity: O(nnz log nnz) due to sorting
245/// Space complexity: O(nnz) for the output
246pub 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    // Sort entries by (col, row)
255    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    // Build CSC data, summing duplicates. Mirrors `coo_to_csr`: each column
259    // is accumulated into a scratch buffer first; only once the *entire*
260    // column has been summed do we prune exact-zero results and append the
261    // survivors, then push the col_ptrs boundary. Zero-pruning therefore
262    // never straddles a column boundary (see `flush_coo_group`).
263    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            // Column boundary: flush the completed column, then fill any
280            // fully empty columns between it and the new column.
281            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        // Accumulate into the current column's buffer, summing duplicates
297        // at the same row (adjacent, since entries are sorted).
298        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 the final column's buffer.
308    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    // Fill any trailing empty columns.
318    while current_col < ncols {
319        col_ptrs.push(values.len());
320        current_col += 1;
321    }
322
323    // SAFETY: We've constructed valid CSC data
324    unsafe { CscMatrix::new_unchecked(nrows, ncols, col_ptrs, row_indices, values) }
325}
326
327/// Converts a CSR matrix to COO format.
328pub 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    // SAFETY: Valid COO data derived from valid CSR
349    unsafe { CooMatrix::new_unchecked(nrows, ncols, row_indices, col_indices, values) }
350}
351
352/// Converts a CSC matrix to COO format.
353pub 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    // SAFETY: Valid COO data derived from valid CSC
374    unsafe { CooMatrix::new_unchecked(nrows, ncols, row_indices, col_indices, values) }
375}
376
377// ============================================================================
378// DIA Conversions
379// ============================================================================
380
381/// Converts a CSR matrix to DIA format.
382///
383/// # Arguments
384///
385/// * `csr` - Source CSR matrix
386/// * `offsets` - Optional list of diagonal offsets to extract. If None, all non-empty diagonals are extracted.
387///
388/// Time complexity: O(nnz)
389pub 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    // Find all non-empty diagonals if not specified
397    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        // Fill diagonal from CSR
420        // Element A[row, col] where col = row + offset goes to data index (row + offset)
421        // This matches DiaMatrix::data_index which uses (row as isize + offset) as usize
422        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                // data_index = row + offset (accounting for padding)
426                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    // Safety: we constructed valid DIA data
437    unsafe { DiaMatrix::new_unchecked(nrows, ncols, offsets, data) }
438}
439
440/// Converts a DIA matrix to CSR format.
441///
442/// Time complexity: O(nrows * ndiag)
443pub fn dia_to_csr<T: Scalar + Clone + Field>(dia: &DiaMatrix<T>) -> CsrMatrix<T> {
444    dia.to_csr()
445}
446
447// ============================================================================
448// ELL Conversions
449// ============================================================================
450
451/// Converts a CSR matrix to ELL format.
452///
453/// # Arguments
454///
455/// * `csr` - Source CSR matrix
456/// * `max_width` - Optional maximum width (if None, uses actual max non-zeros per row)
457///
458/// Time complexity: O(nnz)
459pub 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
466/// Converts an ELL matrix to CSR format.
467///
468/// Time complexity: O(nrows * width)
469pub fn ell_to_csr<T: Scalar + Clone + Field>(ell: &EllMatrix<T>) -> CsrMatrix<T> {
470    ell.to_csr()
471}
472
473// ============================================================================
474// BSR Conversions
475// ============================================================================
476
477/// Converts a CSR matrix to BSR format.
478///
479/// # Arguments
480///
481/// * `csr` - Source CSR matrix
482/// * `block_rows` - Block row size
483/// * `block_cols` - Block column size
484///
485/// Time complexity: O(nnz)
486pub 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
494/// Converts a BSR matrix to CSR format.
495///
496/// Time complexity: O(nblocks * block_size)
497pub fn bsr_to_csr<T: Scalar + Clone + Field>(bsr: &BsrMatrix<T>) -> CsrMatrix<T> {
498    bsr.to_csr()
499}
500
501// ============================================================================
502// Cross-format conversions
503// ============================================================================
504
505/// Converts a DIA matrix to ELL format.
506pub 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
514/// Converts an ELL matrix to DIA format.
515pub 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
523/// Converts a DIA matrix to BSR format.
524pub 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
533/// Converts a BSR matrix to DIA format.
534pub 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
542/// Converts an ELL matrix to BSR format.
543pub 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
552/// Converts a BSR matrix to ELL format.
553pub 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
561// ============================================================================
562// BSC Conversions
563// ============================================================================
564
565/// Converts a CSR matrix to BSC format.
566///
567/// # Arguments
568///
569/// * `csr` - Source CSR matrix
570/// * `block_rows` - Block row size
571/// * `block_cols` - Block column size
572pub 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
581/// Converts a BSC matrix to CSR format.
582pub 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
587/// Converts a BSC matrix to BSR format.
588pub fn bsc_to_bsr<T: Scalar + Clone + Field>(bsc: &BscMatrix<T>) -> BsrMatrix<T> {
589    bsc.to_bsr()
590}
591
592/// Converts a BSR matrix to BSC format.
593pub fn bsr_to_bsc<T: Scalar + Clone + Field>(bsr: &BsrMatrix<T>) -> BscMatrix<T> {
594    BscMatrix::from_bsr(bsr)
595}
596
597// ============================================================================
598// HYB Conversions
599// ============================================================================
600
601/// Converts a CSR matrix to HYB format.
602///
603/// # Arguments
604///
605/// * `csr` - Source CSR matrix
606/// * `strategy` - Strategy for determining ELL width
607pub 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
614/// Converts a HYB matrix to CSR format.
615pub fn hyb_to_csr<T: Scalar + Clone + Field>(hyb: &HybMatrix<T>) -> CsrMatrix<T> {
616    hyb.to_csr()
617}
618
619/// Converts an ELL matrix to HYB format (no COO overflow).
620pub fn ell_to_hyb<T: Scalar + Clone + Field>(ell: &EllMatrix<T>) -> HybMatrix<T> {
621    HybMatrix::from_ell(ell)
622}
623
624/// Converts a HYB matrix to ELL format.
625pub fn hyb_to_ell<T: Scalar + Clone + Field>(hyb: &HybMatrix<T>) -> EllMatrix<T> {
626    hyb.to_ell()
627}
628
629// ============================================================================
630// SELL Conversions
631// ============================================================================
632
633/// Converts a CSR matrix to SELL (Sliced ELLPACK) format.
634///
635/// # Arguments
636///
637/// * `csr` - Source CSR matrix
638/// * `slice_size` - Size of each slice (typically 32 or 64 for GPU)
639pub 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
646/// Converts a SELL matrix to CSR format.
647pub fn sell_to_csr<T: Scalar + Clone + Field>(sell: &SellMatrix<T>) -> CsrMatrix<T> {
648    sell.to_csr()
649}
650
651// ============================================================================
652// Format Detection and Analysis
653// ============================================================================
654
655/// Recommended sparse matrix format based on sparsity analysis.
656#[derive(Debug, Clone, Copy, PartialEq, Eq)]
657pub enum RecommendedFormat {
658    /// CSR: General purpose, good for row-wise operations.
659    Csr,
660    /// CSC: Good for column-wise operations and direct solvers.
661    Csc,
662    /// DIA: Optimal for banded/diagonal matrices.
663    Dia,
664    /// ELL: Good for matrices with uniform row lengths.
665    Ell,
666    /// HYB: Good for matrices with mostly uniform rows but some outliers.
667    Hyb,
668    /// SELL: Good for GPU computation with variable row lengths.
669    Sell,
670    /// BSR: Good for block-structured matrices.
671    Bsr,
672    /// BSC: Good for column-oriented block-structured matrices.
673    Bsc,
674}
675
676/// Analysis of a sparse matrix's sparsity pattern.
677#[derive(Debug, Clone)]
678pub struct SparsityAnalysis {
679    /// Number of rows.
680    pub nrows: usize,
681    /// Number of columns.
682    pub ncols: usize,
683    /// Number of non-zeros.
684    pub nnz: usize,
685    /// Density (nnz / (nrows * ncols)).
686    pub density: f64,
687    /// Maximum row length.
688    pub max_row_length: usize,
689    /// Minimum row length.
690    pub min_row_length: usize,
691    /// Average row length.
692    pub avg_row_length: f64,
693    /// Standard deviation of row lengths.
694    pub row_length_stddev: f64,
695    /// Number of distinct diagonals with entries.
696    pub num_diagonals: usize,
697    /// True if matrix appears to have block structure.
698    pub has_block_structure: bool,
699    /// Detected block size (if any).
700    pub detected_block_size: Option<(usize, usize)>,
701    /// Recommended format for this matrix.
702    pub recommended_format: RecommendedFormat,
703}
704
705/// Analyzes the sparsity pattern of a CSR matrix and recommends a format.
706///
707/// # Returns
708///
709/// A `SparsityAnalysis` containing statistics and a recommended format.
710pub 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    // Compute row lengths
733    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    // Compute standard deviation
753    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    // Count distinct diagonals
764    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    // Check for block structure (simple heuristic)
773    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    // Determine recommended format
782    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
809/// Detects if a matrix has block structure.
810fn 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    // Try common block sizes
821    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        // Check if entries align with blocks
830        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        // Verify that within each block, we have dense or near-dense entries
842        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            // Consider block dense if > 50% full
857            if count * 2 >= block_size * block_size {
858                dense_blocks += 1;
859            }
860        }
861
862        // Consider it block-structured if > 70% of found blocks are dense
863        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            // Just to avoid warnings, this is always true
868            continue;
869        }
870    }
871
872    (false, None)
873}
874
875/// Determines the recommended format based on matrix characteristics.
876fn 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    // Empty or very small matrix
887    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    // Block structure
894    if has_block_structure {
895        return RecommendedFormat::Bsr;
896    }
897
898    // Diagonal/banded structure
899    // If number of diagonals is small relative to matrix size
900    if num_diagonals <= 10 && num_diagonals * 2 <= nrows.max(1) {
901        return RecommendedFormat::Dia;
902    }
903
904    // Uniform row lengths (low variance)
905    let coefficient_of_variation = row_length_stddev / avg_row_length.max(1.0);
906
907    if coefficient_of_variation < 0.3 {
908        // Very uniform - ELL is efficient
909        return RecommendedFormat::Ell;
910    }
911
912    if coefficient_of_variation < 0.8 {
913        // Moderately uniform but with some variation - HYB is good
914        return RecommendedFormat::Hyb;
915    }
916
917    // High variance in row lengths
918    if max_row_length > min_row_length * 10 {
919        // Very irregular - SELL handles this well for GPU
920        return RecommendedFormat::Sell;
921    }
922
923    // Default to CSR
924    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        // [1 0 2]
935        // [0 3 0]
936        // [4 0 5]
937        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        // [1 0 4]
955        // [0 3 0]
956        // [2 0 5]
957        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        // Duplicate entries at (0,0)
992        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)); // 1 + 2
1001        assert_eq!(csr.get(1, 1), Some(&3.0));
1002    }
1003
1004    #[test]
1005    fn test_coo_to_csr_row_ptrs_length() {
1006        // Regression test: row_ptrs must have exactly nrows + 1 entries.
1007        // A stray trailing push used to make it nrows + 2.
1008        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        // Row 0 holds two entries at the same column that cancel to
1022        // exactly zero; row 1 is empty; row 2 holds a single surviving
1023        // entry. Before the fix, deferring zero-pruning past the row
1024        // boundary corrupted row_ptrs so that row 2's entry was
1025        // misattributed to row 0 (and row 2 appeared empty).
1026        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        // Regression test: col_ptrs must have exactly ncols + 1 entries.
1043        // A stray trailing push used to make it ncols + 2.
1044        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        // Mirror of the CSR regression test: column 0 cancels to exactly
1058        // zero, column 1 is empty, column 2 holds a single surviving
1059        // entry.
1060        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    // ========================================================================
1134    // DIA conversion tests
1135    // ========================================================================
1136
1137    #[test]
1138    fn test_csr_to_dia_tridiagonal() {
1139        // Tridiagonal matrix:
1140        // [4 1 0]
1141        // [2 5 1]
1142        // [0 3 6]
1143        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    // ========================================================================
1195    // ELL conversion tests
1196    // ========================================================================
1197
1198    #[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    // ========================================================================
1245    // BSR conversion tests
1246    // ========================================================================
1247
1248    #[test]
1249    fn test_csr_to_bsr() {
1250        // 4x4 matrix with 2x2 block structure
1251        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    // ========================================================================
1298    // Cross-format conversion tests
1299    // ========================================================================
1300
1301    #[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}