Skip to main content

rusparse/
formats.rs

1use crate::{CsrMatrixOwned, IndexBase, SparseError};
2
3#[derive(Debug, Clone, Copy)]
4pub struct CooMatrix<'a> {
5    rows: usize,
6    columns: usize,
7    row_indices: &'a [u32],
8    column_indices: &'a [u32],
9    values: &'a [f32],
10    index_base: IndexBase,
11}
12
13#[derive(Debug, Clone, PartialEq)]
14pub struct CooMatrixOwned {
15    rows: usize,
16    columns: usize,
17    row_indices: Vec<u32>,
18    column_indices: Vec<u32>,
19    values: Vec<f32>,
20    index_base: IndexBase,
21}
22
23impl CooMatrixOwned {
24    pub fn new(
25        rows: usize,
26        columns: usize,
27        row_indices: Vec<u32>,
28        column_indices: Vec<u32>,
29        values: Vec<f32>,
30        index_base: IndexBase,
31    ) -> Result<Self, SparseError> {
32        CooMatrix::new(
33            rows,
34            columns,
35            &row_indices,
36            &column_indices,
37            &values,
38            index_base,
39        )?;
40        Ok(Self {
41            rows,
42            columns,
43            row_indices,
44            column_indices,
45            values,
46            index_base,
47        })
48    }
49
50    pub fn as_ref(&self) -> CooMatrix<'_> {
51        CooMatrix {
52            rows: self.rows,
53            columns: self.columns,
54            row_indices: &self.row_indices,
55            column_indices: &self.column_indices,
56            values: &self.values,
57            index_base: self.index_base,
58        }
59    }
60}
61
62impl<'a> CooMatrix<'a> {
63    pub fn new(
64        rows: usize,
65        columns: usize,
66        row_indices: &'a [u32],
67        column_indices: &'a [u32],
68        values: &'a [f32],
69        index_base: IndexBase,
70    ) -> Result<Self, SparseError> {
71        validate_dimension(rows, "COO rows")?;
72        validate_dimension(columns, "COO columns")?;
73        check_length("COO row indices", values.len(), row_indices.len())?;
74        check_length("COO column indices", values.len(), column_indices.len())?;
75        let base = index_base.value();
76        let row_limit = index_limit(base, rows, "COO row range")?;
77        let column_limit = index_limit(base, columns, "COO column range")?;
78        if row_indices
79            .iter()
80            .any(|&row| row < base || row >= row_limit)
81        {
82            return Err(SparseError::InvalidSparseIndex("COO row"));
83        }
84        if column_indices
85            .iter()
86            .any(|&column| column < base || column >= column_limit)
87        {
88            return Err(SparseError::InvalidSparseIndex("COO column"));
89        }
90        Ok(Self {
91            rows,
92            columns,
93            row_indices,
94            column_indices,
95            values,
96            index_base,
97        })
98    }
99
100    pub const fn rows(&self) -> usize {
101        self.rows
102    }
103
104    pub const fn columns(&self) -> usize {
105        self.columns
106    }
107
108    pub const fn nnz(&self) -> usize {
109        self.values.len()
110    }
111
112    pub const fn row_indices(&self) -> &'a [u32] {
113        self.row_indices
114    }
115
116    pub const fn column_indices(&self) -> &'a [u32] {
117        self.column_indices
118    }
119
120    pub const fn values(&self) -> &'a [f32] {
121        self.values
122    }
123
124    pub const fn index_base(&self) -> IndexBase {
125        self.index_base
126    }
127
128    /// Convert COO to row-sorted CSR. Entries in the same row retain their
129    /// original relative order, including duplicate coordinates.
130    pub fn to_csr(self) -> Result<CsrMatrixOwned, SparseError> {
131        let base = self.index_base.value();
132        let mut row_offsets = vec![0_u32; self.rows + 1];
133        for &row in self.row_indices {
134            let row = (row - base) as usize;
135            row_offsets[row] = row_offsets[row]
136                .checked_add(1)
137                .ok_or(SparseError::SizeOverflow("COO row counts"))?;
138        }
139        prefix_offsets(&mut row_offsets, base, "COO row offsets")?;
140        let mut column_indices = vec![0_u32; self.nnz()];
141        let mut values = vec![0.0_f32; self.nnz()];
142        for entry in (0..self.nnz()).rev() {
143            let row = (self.row_indices[entry] - base) as usize;
144            row_offsets[row] -= 1;
145            let destination = (row_offsets[row] - base) as usize;
146            column_indices[destination] = self.column_indices[entry];
147            values[destination] = self.values[entry];
148        }
149        CsrMatrixOwned::new(
150            self.rows,
151            self.columns,
152            row_offsets,
153            column_indices,
154            values,
155            self.index_base,
156        )
157    }
158}
159
160#[derive(Debug, Clone, Copy)]
161pub struct CscMatrix<'a> {
162    rows: usize,
163    columns: usize,
164    column_offsets: &'a [u32],
165    row_indices: &'a [u32],
166    values: &'a [f32],
167    index_base: IndexBase,
168}
169
170#[derive(Debug, Clone, PartialEq)]
171pub struct CscMatrixOwned {
172    rows: usize,
173    columns: usize,
174    column_offsets: Vec<u32>,
175    row_indices: Vec<u32>,
176    values: Vec<f32>,
177    index_base: IndexBase,
178}
179
180impl CscMatrixOwned {
181    pub fn new(
182        rows: usize,
183        columns: usize,
184        column_offsets: Vec<u32>,
185        row_indices: Vec<u32>,
186        values: Vec<f32>,
187        index_base: IndexBase,
188    ) -> Result<Self, SparseError> {
189        CscMatrix::new(
190            rows,
191            columns,
192            &column_offsets,
193            &row_indices,
194            &values,
195            index_base,
196        )?;
197        Ok(Self {
198            rows,
199            columns,
200            column_offsets,
201            row_indices,
202            values,
203            index_base,
204        })
205    }
206
207    pub fn as_ref(&self) -> CscMatrix<'_> {
208        CscMatrix {
209            rows: self.rows,
210            columns: self.columns,
211            column_offsets: &self.column_offsets,
212            row_indices: &self.row_indices,
213            values: &self.values,
214            index_base: self.index_base,
215        }
216    }
217}
218
219impl<'a> CscMatrix<'a> {
220    pub fn new(
221        rows: usize,
222        columns: usize,
223        column_offsets: &'a [u32],
224        row_indices: &'a [u32],
225        values: &'a [f32],
226        index_base: IndexBase,
227    ) -> Result<Self, SparseError> {
228        validate_dimension(rows, "CSC rows")?;
229        validate_dimension(columns, "CSC columns")?;
230        let expected_offsets = columns
231            .checked_add(1)
232            .ok_or(SparseError::SizeOverflow("CSC column offsets"))?;
233        check_length("CSC column offsets", expected_offsets, column_offsets.len())?;
234        check_length("CSC row indices", values.len(), row_indices.len())?;
235        validate_offsets("CSC column", column_offsets, values.len(), index_base)?;
236        let base = index_base.value();
237        let row_limit = index_limit(base, rows, "CSC row range")?;
238        if row_indices
239            .iter()
240            .any(|&row| row < base || row >= row_limit)
241        {
242            return Err(SparseError::InvalidSparseIndex("CSC row"));
243        }
244        Ok(Self {
245            rows,
246            columns,
247            column_offsets,
248            row_indices,
249            values,
250            index_base,
251        })
252    }
253
254    pub const fn rows(&self) -> usize {
255        self.rows
256    }
257
258    pub const fn columns(&self) -> usize {
259        self.columns
260    }
261
262    pub const fn nnz(&self) -> usize {
263        self.values.len()
264    }
265
266    pub const fn column_offsets(&self) -> &'a [u32] {
267        self.column_offsets
268    }
269
270    pub const fn row_indices(&self) -> &'a [u32] {
271        self.row_indices
272    }
273
274    pub const fn values(&self) -> &'a [f32] {
275        self.values
276    }
277
278    pub const fn index_base(&self) -> IndexBase {
279        self.index_base
280    }
281
282    pub fn to_csr(self) -> Result<CsrMatrixOwned, SparseError> {
283        let base = self.index_base.value();
284        let mut row_offsets = vec![0_u32; self.rows + 1];
285        for &row in self.row_indices {
286            let row = (row - base) as usize;
287            row_offsets[row + 1] = row_offsets[row + 1]
288                .checked_add(1)
289                .ok_or(SparseError::SizeOverflow("CSC to CSR row counts"))?;
290        }
291        prefix_offsets(&mut row_offsets, base, "CSC to CSR row offsets")?;
292        let mut positions = row_offsets[..self.rows]
293            .iter()
294            .map(|offset| offset - base)
295            .collect::<Vec<_>>();
296        let mut column_indices = vec![0_u32; self.nnz()];
297        let mut values = vec![0.0_f32; self.nnz()];
298        for column in 0..self.columns {
299            let start = (self.column_offsets[column] - base) as usize;
300            let end = (self.column_offsets[column + 1] - base) as usize;
301            for entry in start..end {
302                let row = (self.row_indices[entry] - base) as usize;
303                let destination = positions[row] as usize;
304                column_indices[destination] = validate_dimension(column, "CSC column")? + base;
305                values[destination] = self.values[entry];
306                positions[row] += 1;
307            }
308        }
309        CsrMatrixOwned::new(
310            self.rows,
311            self.columns,
312            row_offsets,
313            column_indices,
314            values,
315            self.index_base,
316        )
317    }
318}
319
320#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
321pub enum BlockDirection {
322    RowMajor,
323    ColumnMajor,
324}
325
326#[derive(Debug, Clone, Copy)]
327pub struct BsrMatrix<'a> {
328    block_rows: usize,
329    block_columns: usize,
330    block_dimension: usize,
331    row_offsets: &'a [u32],
332    column_indices: &'a [u32],
333    values: &'a [f32],
334    direction: BlockDirection,
335    index_base: IndexBase,
336}
337
338impl<'a> BsrMatrix<'a> {
339    #[allow(clippy::too_many_arguments)]
340    pub fn new(
341        block_rows: usize,
342        block_columns: usize,
343        block_dimension: usize,
344        row_offsets: &'a [u32],
345        column_indices: &'a [u32],
346        values: &'a [f32],
347        direction: BlockDirection,
348        index_base: IndexBase,
349    ) -> Result<Self, SparseError> {
350        validate_dimension(block_rows, "BSR block rows")?;
351        validate_dimension(block_columns, "BSR block columns")?;
352        if block_dimension == 0 {
353            return Err(SparseError::InvalidBlockDimension);
354        }
355        validate_dimension(block_dimension, "BSR block dimension")?;
356        let expected_offsets = block_rows
357            .checked_add(1)
358            .ok_or(SparseError::SizeOverflow("BSR row offsets"))?;
359        check_length("BSR row offsets", expected_offsets, row_offsets.len())?;
360        validate_offsets("BSR row", row_offsets, column_indices.len(), index_base)?;
361        let block_elements = block_dimension
362            .checked_mul(block_dimension)
363            .ok_or(SparseError::SizeOverflow("BSR block"))?;
364        let expected_values = column_indices
365            .len()
366            .checked_mul(block_elements)
367            .ok_or(SparseError::SizeOverflow("BSR values"))?;
368        check_length("BSR values", expected_values, values.len())?;
369        let base = index_base.value();
370        let column_limit = index_limit(base, block_columns, "BSR block-column range")?;
371        if column_indices
372            .iter()
373            .any(|&column| column < base || column >= column_limit)
374        {
375            return Err(SparseError::InvalidSparseIndex("BSR block column"));
376        }
377        block_rows
378            .checked_mul(block_dimension)
379            .ok_or(SparseError::SizeOverflow("BSR rows"))?;
380        block_columns
381            .checked_mul(block_dimension)
382            .ok_or(SparseError::SizeOverflow("BSR columns"))?;
383        Ok(Self {
384            block_rows,
385            block_columns,
386            block_dimension,
387            row_offsets,
388            column_indices,
389            values,
390            direction,
391            index_base,
392        })
393    }
394
395    pub const fn block_rows(&self) -> usize {
396        self.block_rows
397    }
398
399    pub const fn block_columns(&self) -> usize {
400        self.block_columns
401    }
402
403    pub const fn block_dimension(&self) -> usize {
404        self.block_dimension
405    }
406
407    pub const fn nnzb(&self) -> usize {
408        self.column_indices.len()
409    }
410
411    pub const fn direction(&self) -> BlockDirection {
412        self.direction
413    }
414
415    pub const fn index_base(&self) -> IndexBase {
416        self.index_base
417    }
418
419    pub fn rows(&self) -> usize {
420        self.block_rows * self.block_dimension
421    }
422
423    pub fn columns(&self) -> usize {
424        self.block_columns * self.block_dimension
425    }
426
427    pub fn to_csr(self) -> Result<CsrMatrixOwned, SparseError> {
428        let base = self.index_base.value();
429        let rows = self.rows();
430        let columns = self.columns();
431        let entries_per_block_row = self.block_dimension;
432        let nnz = self
433            .nnzb()
434            .checked_mul(self.block_dimension)
435            .and_then(|value| value.checked_mul(self.block_dimension))
436            .ok_or(SparseError::SizeOverflow("BSR to CSR nnz"))?;
437        let mut row_offsets = vec![base; rows + 1];
438        let mut column_indices = Vec::with_capacity(nnz);
439        let mut values = Vec::with_capacity(nnz);
440        for block_row in 0..self.block_rows {
441            let block_start = (self.row_offsets[block_row] - base) as usize;
442            let block_end = (self.row_offsets[block_row + 1] - base) as usize;
443            for row_in_block in 0..entries_per_block_row {
444                for block_entry in block_start..block_end {
445                    let block_column = (self.column_indices[block_entry] - base) as usize;
446                    for column_in_block in 0..self.block_dimension {
447                        let value_offset = match self.direction {
448                            BlockDirection::RowMajor => {
449                                row_in_block * self.block_dimension + column_in_block
450                            }
451                            BlockDirection::ColumnMajor => {
452                                column_in_block * self.block_dimension + row_in_block
453                            }
454                        };
455                        column_indices.push(
456                            validate_dimension(
457                                block_column * self.block_dimension + column_in_block,
458                                "BSR to CSR column",
459                            )? + base,
460                        );
461                        values.push(
462                            self.values[block_entry * self.block_dimension * self.block_dimension
463                                + value_offset],
464                        );
465                    }
466                }
467                let row = block_row * self.block_dimension + row_in_block;
468                row_offsets[row + 1] = validate_dimension(column_indices.len(), "BSR to CSR nnz")?
469                    .checked_add(base)
470                    .ok_or(SparseError::SizeOverflow("BSR to CSR row offsets"))?;
471            }
472        }
473        CsrMatrixOwned::new(
474            rows,
475            columns,
476            row_offsets,
477            column_indices,
478            values,
479            self.index_base,
480        )
481    }
482}
483
484#[derive(Debug, Clone, Copy)]
485pub struct EllMatrix<'a> {
486    rows: usize,
487    columns: usize,
488    width: usize,
489    column_indices: &'a [u32],
490    values: &'a [f32],
491    index_base: IndexBase,
492}
493
494impl<'a> EllMatrix<'a> {
495    pub fn new(
496        rows: usize,
497        columns: usize,
498        width: usize,
499        column_indices: &'a [u32],
500        values: &'a [f32],
501        index_base: IndexBase,
502    ) -> Result<Self, SparseError> {
503        validate_dimension(rows, "ELL rows")?;
504        validate_dimension(columns, "ELL columns")?;
505        validate_dimension(width, "ELL width")?;
506        let expected = rows
507            .checked_mul(width)
508            .ok_or(SparseError::SizeOverflow("ELL storage"))?;
509        check_length("ELL column indices", expected, column_indices.len())?;
510        check_length("ELL values", expected, values.len())?;
511        Ok(Self {
512            rows,
513            columns,
514            width,
515            column_indices,
516            values,
517            index_base,
518        })
519    }
520
521    pub const fn rows(&self) -> usize {
522        self.rows
523    }
524
525    pub const fn columns(&self) -> usize {
526        self.columns
527    }
528
529    pub const fn width(&self) -> usize {
530        self.width
531    }
532
533    pub const fn column_indices(&self) -> &'a [u32] {
534        self.column_indices
535    }
536
537    pub const fn values(&self) -> &'a [f32] {
538        self.values
539    }
540
541    pub const fn index_base(&self) -> IndexBase {
542        self.index_base
543    }
544
545    pub const fn padding_index() -> u32 {
546        u32::MAX
547    }
548
549    /// Convert rocSPARSE's column-major ELL storage (`slot * rows + row`)
550    /// to CSR, ignoring out-of-range padding columns.
551    pub fn to_csr(self) -> Result<CsrMatrixOwned, SparseError> {
552        let base = self.index_base.value();
553        let column_limit = index_limit(base, self.columns, "ELL column range")?;
554        let mut row_offsets = Vec::with_capacity(self.rows + 1);
555        let mut column_indices = Vec::new();
556        let mut values = Vec::new();
557        row_offsets.push(base);
558        for row in 0..self.rows {
559            for slot in 0..self.width {
560                let entry = slot * self.rows + row;
561                let column = self.column_indices[entry];
562                if column >= base && column < column_limit {
563                    column_indices.push(column);
564                    values.push(self.values[entry]);
565                }
566            }
567            row_offsets.push(
568                validate_dimension(column_indices.len(), "ELL to CSR nnz")?
569                    .checked_add(base)
570                    .ok_or(SparseError::SizeOverflow("ELL to CSR row offsets"))?,
571            );
572        }
573        CsrMatrixOwned::new(
574            self.rows,
575            self.columns,
576            row_offsets,
577            column_indices,
578            values,
579            self.index_base,
580        )
581    }
582}
583
584fn validate_dimension(value: usize, name: &'static str) -> Result<u32, SparseError> {
585    u32::try_from(value).map_err(|_| SparseError::DimensionTooLarge(name))
586}
587
588fn check_length(name: &'static str, expected: usize, actual: usize) -> Result<(), SparseError> {
589    if expected != actual {
590        return Err(SparseError::BufferLength {
591            name,
592            expected,
593            actual,
594        });
595    }
596    Ok(())
597}
598
599fn index_limit(base: u32, dimension: usize, name: &'static str) -> Result<u32, SparseError> {
600    base.checked_add(validate_dimension(dimension, name)?)
601        .ok_or(SparseError::SizeOverflow(name))
602}
603
604fn validate_offsets(
605    format: &'static str,
606    offsets: &[u32],
607    entries: usize,
608    index_base: IndexBase,
609) -> Result<(), SparseError> {
610    let base = index_base.value();
611    if offsets.first().copied() != Some(base) {
612        return Err(SparseError::InvalidSparseOffsets {
613            format,
614            message: "first offset does not equal the index base",
615        });
616    }
617    if offsets.windows(2).any(|pair| pair[0] > pair[1]) {
618        return Err(SparseError::InvalidSparseOffsets {
619            format,
620            message: "offsets are not nondecreasing",
621        });
622    }
623    let terminal = offsets.last().copied().unwrap_or(base);
624    if terminal.checked_sub(base).map(|value| value as usize) != Some(entries) {
625        return Err(SparseError::InvalidSparseOffsets {
626            format,
627            message: "terminal offset does not match the entry count",
628        });
629    }
630    Ok(())
631}
632
633fn prefix_offsets(offsets: &mut [u32], base: u32, name: &'static str) -> Result<(), SparseError> {
634    for index in 0..offsets.len().saturating_sub(1) {
635        offsets[index + 1] = offsets[index + 1]
636            .checked_add(offsets[index])
637            .ok_or(SparseError::SizeOverflow(name))?;
638    }
639    for offset in offsets {
640        *offset = offset
641            .checked_add(base)
642            .ok_or(SparseError::SizeOverflow(name))?;
643    }
644    Ok(())
645}