Skip to main content

rusparse/
lib.rs

1//! Rust sparse formats, conversions and device execution.
2
3use std::error::Error;
4use std::fmt::{Display, Formatter};
5
6
7pub mod conversion;
8pub mod host;
9mod formats;
10#[cfg(feature = "serde")]
11mod serde;
12mod symbolic;
13pub use formats::{
14    BlockDirection, BsrMatrix, CooMatrix, CooMatrixOwned, CscMatrix, CscMatrixOwned, EllMatrix,
15};
16
17
18
19
20#[cfg(feature = "tensor")]
21pub mod tensor;
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
24pub enum IndexBase {
25    Zero,
26    One,
27}
28
29impl IndexBase {
30    const fn value(self) -> u32 {
31        match self {
32            Self::Zero => 0,
33            Self::One => 1,
34        }
35    }
36}
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
39pub enum Operation {
40    None,
41    Transpose,
42    ConjugateTranspose,
43}
44
45#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
46pub enum DenseOrder {
47    RowMajor,
48    ColumnMajor,
49}
50
51#[derive(Debug, Clone, Copy)]
52pub struct DenseMatrix<'a> {
53    values: &'a [f32],
54    rows: usize,
55    columns: usize,
56    order: DenseOrder,
57}
58
59#[derive(Debug, Clone, PartialEq)]
60pub struct DenseMatrixOwned {
61    values: Vec<f32>,
62    rows: usize,
63    columns: usize,
64    order: DenseOrder,
65}
66
67impl DenseMatrixOwned {
68    pub fn new(
69        values: Vec<f32>,
70        rows: usize,
71        columns: usize,
72        order: DenseOrder,
73    ) -> Result<Self, SparseError> {
74        DenseMatrix::new(&values, rows, columns, order)?;
75        Ok(Self {
76            values,
77            rows,
78            columns,
79            order,
80        })
81    }
82
83    pub fn as_ref(&self) -> DenseMatrix<'_> {
84        DenseMatrix {
85            values: &self.values,
86            rows: self.rows,
87            columns: self.columns,
88            order: self.order,
89        }
90    }
91}
92
93#[derive(Debug, Clone, Copy)]
94pub struct SparseVector<'a> {
95    indices: &'a [u32],
96    values: &'a [f32],
97    size: usize,
98    index_base: IndexBase,
99}
100
101#[derive(Debug, Clone, PartialEq)]
102pub struct SparseVectorOwned {
103    indices: Vec<u32>,
104    values: Vec<f32>,
105    size: usize,
106    index_base: IndexBase,
107}
108
109impl SparseVectorOwned {
110    pub fn new(
111        size: usize,
112        indices: Vec<u32>,
113        values: Vec<f32>,
114        index_base: IndexBase,
115    ) -> Result<Self, SparseError> {
116        SparseVector::new(size, &indices, &values, index_base)?;
117        Ok(Self { indices, values, size, index_base })
118    }
119
120    pub fn as_ref(&self) -> SparseVector<'_> {
121        SparseVector {
122            indices: &self.indices,
123            values: &self.values,
124            size: self.size,
125            index_base: self.index_base,
126        }
127    }
128}
129
130impl<'a> SparseVector<'a> {
131    pub fn new(
132        size: usize,
133        indices: &'a [u32],
134        values: &'a [f32],
135        index_base: IndexBase,
136    ) -> Result<Self, SparseError> {
137        dimension(size, "sparse vector size")?;
138        if indices.len() != values.len() {
139            return Err(SparseError::BufferLength {
140                name: "sparse vector indices",
141                expected: values.len(),
142                actual: indices.len(),
143            });
144        }
145        let base = index_base.value();
146        let limit = base
147            .checked_add(dimension(size, "sparse vector index range")?)
148            .ok_or(SparseError::SizeOverflow("sparse vector index range"))?;
149        if indices.iter().any(|&index| index < base || index >= limit) {
150            return Err(SparseError::InvalidSparseIndex("sparse vector"));
151        }
152        Ok(Self {
153            indices,
154            values,
155            size,
156            index_base,
157        })
158    }
159
160    pub const fn size(&self) -> usize {
161        self.size
162    }
163
164    pub const fn nnz(&self) -> usize {
165        self.values.len()
166    }
167
168    pub const fn indices(&self) -> &'a [u32] {
169        self.indices
170    }
171
172    pub const fn values(&self) -> &'a [f32] {
173        self.values
174    }
175
176    pub const fn index_base(&self) -> IndexBase {
177        self.index_base
178    }
179}
180
181impl<'a> DenseMatrix<'a> {
182    pub fn new(
183        values: &'a [f32],
184        rows: usize,
185        columns: usize,
186        order: DenseOrder,
187    ) -> Result<Self, SparseError> {
188        let expected = rows
189            .checked_mul(columns)
190            .ok_or(SparseError::SizeOverflow("dense matrix"))?;
191        if values.len() != expected {
192            return Err(SparseError::BufferLength {
193                name: "dense matrix",
194                expected,
195                actual: values.len(),
196            });
197        }
198        Ok(Self {
199            values,
200            rows,
201            columns,
202            order,
203        })
204    }
205
206    pub const fn values(&self) -> &'a [f32] {
207        self.values
208    }
209
210    pub const fn rows(&self) -> usize {
211        self.rows
212    }
213
214    pub const fn columns(&self) -> usize {
215        self.columns
216    }
217
218    pub const fn order(&self) -> DenseOrder {
219        self.order
220    }
221
222    fn physical_strides(self) -> Result<(u32, u32), SparseError> {
223        let rows = dimension(self.rows, "dense rows")?;
224        let columns = dimension(self.columns, "dense columns")?;
225        Ok(match self.order {
226            DenseOrder::RowMajor => (columns, 1),
227            DenseOrder::ColumnMajor => (1, rows),
228        })
229    }
230
231    fn operation_shape(self, operation: Operation) -> (usize, usize) {
232        match operation {
233            Operation::None => (self.rows, self.columns),
234            Operation::Transpose | Operation::ConjugateTranspose => (self.columns, self.rows),
235        }
236    }
237
238    fn operation_strides(self, operation: Operation) -> Result<(u32, u32), SparseError> {
239        let (row, column) = self.physical_strides()?;
240        Ok(match operation {
241            Operation::None => (row, column),
242            Operation::Transpose | Operation::ConjugateTranspose => (column, row),
243        })
244    }
245}
246
247#[derive(Debug, Clone, Copy)]
248pub struct CsrMatrix<'a> {
249    values: &'a [f32],
250    row_offsets: &'a [u32],
251    column_indices: &'a [u32],
252    rows: usize,
253    columns: usize,
254    index_base: IndexBase,
255}
256
257impl<'a> CsrMatrix<'a> {
258    pub fn new(
259        rows: usize,
260        columns: usize,
261        row_offsets: &'a [u32],
262        column_indices: &'a [u32],
263        values: &'a [f32],
264        index_base: IndexBase,
265    ) -> Result<Self, SparseError> {
266        let matrix = Self {
267            values,
268            row_offsets,
269            column_indices,
270            rows,
271            columns,
272            index_base,
273        };
274        matrix.validate()?;
275        Ok(matrix)
276    }
277
278    pub const fn rows(&self) -> usize {
279        self.rows
280    }
281
282    pub const fn columns(&self) -> usize {
283        self.columns
284    }
285
286    pub const fn nnz(&self) -> usize {
287        self.values.len()
288    }
289
290    pub const fn index_base(&self) -> IndexBase {
291        self.index_base
292    }
293
294    pub const fn row_offsets(&self) -> &'a [u32] {
295        self.row_offsets
296    }
297
298    pub const fn column_indices(&self) -> &'a [u32] {
299        self.column_indices
300    }
301
302    pub const fn values(&self) -> &'a [f32] {
303        self.values
304    }
305
306    fn validate(self) -> Result<(), SparseError> {
307        dimension(self.rows, "CSR rows")?;
308        dimension(self.columns, "CSR columns")?;
309        let expected_offsets = self
310            .rows
311            .checked_add(1)
312            .ok_or(SparseError::SizeOverflow("CSR row offsets"))?;
313        if self.row_offsets.len() != expected_offsets {
314            return Err(SparseError::BufferLength {
315                name: "CSR row offsets",
316                expected: expected_offsets,
317                actual: self.row_offsets.len(),
318            });
319        }
320        if self.column_indices.len() != self.values.len() {
321            return Err(SparseError::BufferLength {
322                name: "CSR column indices",
323                expected: self.values.len(),
324                actual: self.column_indices.len(),
325            });
326        }
327        let base = self.index_base.value();
328        if self.row_offsets.first().copied() != Some(base) {
329            return Err(SparseError::InvalidRowOffsets(
330                "first row offset does not equal the index base",
331            ));
332        }
333        for offsets in self.row_offsets.windows(2) {
334            if offsets[0] > offsets[1] {
335                return Err(SparseError::InvalidRowOffsets(
336                    "row offsets are not nondecreasing",
337                ));
338            }
339        }
340        let terminal = self.row_offsets.last().copied().unwrap_or(base);
341        let encoded_nnz = terminal
342            .checked_sub(base)
343            .ok_or(SparseError::InvalidRowOffsets(
344                "terminal row offset is below the index base",
345            ))?;
346        if encoded_nnz as usize != self.values.len() {
347            return Err(SparseError::InvalidRowOffsets(
348                "terminal row offset does not match nnz",
349            ));
350        }
351        let column_limit = base
352            .checked_add(dimension(self.columns, "CSR columns")?)
353            .ok_or(SparseError::SizeOverflow("CSR column index range"))?;
354        if self
355            .column_indices
356            .iter()
357            .any(|&column| column < base || column >= column_limit)
358        {
359            return Err(SparseError::InvalidColumnIndex);
360        }
361        Ok(())
362    }
363
364    pub fn transpose(self) -> Result<CsrMatrixOwned, SparseError> {
365        self.transpose_impl(false).map(|(matrix, _)| matrix)
366    }
367
368    pub fn transpose_with_permutation(self) -> Result<(CsrMatrixOwned, Vec<u32>), SparseError> {
369        self.transpose_impl(true)
370    }
371
372    fn transpose_impl(self, capture_permutation: bool) -> Result<(CsrMatrixOwned, Vec<u32>), SparseError> {
373        let base = self.index_base.value();
374        let mut row_offsets = vec![0_u32; self.columns + 1];
375        for &encoded_column in self.column_indices {
376            let column = (encoded_column - base) as usize;
377            row_offsets[column + 1] = row_offsets[column + 1]
378                .checked_add(1)
379                .ok_or(SparseError::SizeOverflow("transposed CSR row counts"))?;
380        }
381        for row in 0..self.columns {
382            row_offsets[row + 1] = row_offsets[row + 1]
383                .checked_add(row_offsets[row])
384                .ok_or(SparseError::SizeOverflow("transposed CSR row offsets"))?;
385        }
386        let mut positions = row_offsets[..self.columns].to_vec();
387        let mut column_indices = vec![0_u32; self.nnz()];
388        let mut values = vec![0.0_f32; self.nnz()];
389        let mut permutation = if capture_permutation { vec![0u32; self.nnz()] } else { Vec::new() };
390        for row in 0..self.rows {
391            let start = (self.row_offsets[row] - base) as usize;
392            let end = (self.row_offsets[row + 1] - base) as usize;
393            for entry in start..end {
394                let column = (self.column_indices[entry] - base) as usize;
395                let destination = positions[column] as usize;
396                column_indices[destination] = dimension(row, "transposed CSR column")? + base;
397                values[destination] = self.values[entry];
398                if capture_permutation {
399                    permutation[destination] = dimension(entry, "transposed CSR permutation")?;
400                }
401                positions[column] += 1;
402            }
403        }
404        for offset in &mut row_offsets {
405            *offset = offset
406                .checked_add(base)
407                .ok_or(SparseError::SizeOverflow("transposed CSR index base"))?;
408        }
409        let matrix = CsrMatrixOwned::new(
410            self.columns,
411            self.rows,
412            row_offsets,
413            column_indices,
414            values,
415            self.index_base,
416        )?;
417        Ok((matrix, permutation))
418    }
419}
420
421#[derive(Debug, Clone, PartialEq)]
422pub struct CsrMatrixOwned {
423    values: Vec<f32>,
424    row_offsets: Vec<u32>,
425    column_indices: Vec<u32>,
426    rows: usize,
427    columns: usize,
428    index_base: IndexBase,
429}
430
431impl CsrMatrixOwned {
432    pub fn new(
433        rows: usize,
434        columns: usize,
435        row_offsets: Vec<u32>,
436        column_indices: Vec<u32>,
437        values: Vec<f32>,
438        index_base: IndexBase,
439    ) -> Result<Self, SparseError> {
440        CsrMatrix::new(
441            rows,
442            columns,
443            &row_offsets,
444            &column_indices,
445            &values,
446            index_base,
447        )?;
448        Ok(Self {
449            values,
450            row_offsets,
451            column_indices,
452            rows,
453            columns,
454            index_base,
455        })
456    }
457
458    pub fn as_ref(&self) -> CsrMatrix<'_> {
459        CsrMatrix {
460            values: &self.values,
461            row_offsets: &self.row_offsets,
462            column_indices: &self.column_indices,
463            rows: self.rows,
464            columns: self.columns,
465            index_base: self.index_base,
466        }
467    }
468}
469
470#[derive(Debug, Clone, Copy, PartialEq, Eq)]
471pub enum SparseAlgorithm {
472    RowSplitWavefront32,
473}
474
475#[derive(Debug, Clone, Copy, PartialEq, Eq)]
476pub struct SparsePlan {
477    pub algorithm: SparseAlgorithm,
478    pub output_elements: u32,
479    pub block_threads: u32,
480    pub grid_blocks: u32,
481}
482
483#[derive(Debug)]
484pub enum SparseError {
485    BufferLength {
486        name: &'static str,
487        expected: usize,
488        actual: usize,
489    },
490    DimensionMismatch(&'static str),
491    DimensionTooLarge(&'static str),
492    InvalidColumnIndex,
493    InvalidSparseIndex(&'static str),
494    InvalidSparseOffsets {
495        format: &'static str,
496        message: &'static str,
497    },
498    InvalidBlockDimension,
499    InvalidRowOffsets(&'static str),
500    SizeOverflow(&'static str),
501    
502    #[cfg(feature = "tensor")]
503    TensorExecution(ruda_core::tensor::execution::ExecutionError),
504    #[cfg(feature = "tensor")]
505    TensorData(ruda_core::tensor::data::DataError),
506    Device(&'static str),
507}
508
509impl Display for SparseError {
510    fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
511        match self {
512            Self::BufferLength {
513                name,
514                expected,
515                actual,
516            } => write!(formatter, "{name} needs {expected} elements, got {actual}"),
517            Self::DimensionMismatch(message) => formatter.write_str(message),
518            Self::DimensionTooLarge(name) => write!(formatter, "{name} does not fit u32"),
519            Self::InvalidColumnIndex => formatter.write_str("CSR column index is out of bounds"),
520            Self::InvalidSparseIndex(name) => write!(formatter, "{name} index is out of bounds"),
521            Self::InvalidSparseOffsets { format, message } => {
522                write!(formatter, "invalid {format} offsets: {message}")
523            }
524            Self::InvalidBlockDimension => {
525                formatter.write_str("BSR block dimension must be greater than zero")
526            }
527            Self::InvalidRowOffsets(message) => {
528                write!(formatter, "invalid CSR row offsets: {message}")
529            }
530            Self::SizeOverflow(name) => write!(formatter, "{name} size overflows"),
531            
532            #[cfg(feature = "tensor")]
533            Self::TensorExecution(error) => Display::fmt(error, formatter),
534            #[cfg(feature = "tensor")]
535            Self::TensorData(error) => Display::fmt(error, formatter),
536            Self::Device(message) => formatter.write_str(message),
537        }
538    }
539}
540
541impl Error for SparseError {
542    fn source(&self) -> Option<&(dyn Error + 'static)> {
543        match self {
544            
545            #[cfg(feature = "tensor")]
546            Self::TensorExecution(error) => Some(error),
547            #[cfg(feature = "tensor")]
548            Self::TensorData(error) => Some(error),
549            _ => None,
550        }
551    }
552}
553
554
555
556fn dimension(value: usize, name: &'static str) -> Result<u32, SparseError> {
557    u32::try_from(value).map_err(|_| SparseError::DimensionTooLarge(name))
558}
559
560fn dense_strides(
561    rows: usize,
562    columns: usize,
563    order: DenseOrder,
564    name: &'static str,
565) -> Result<(u32, u32), SparseError> {
566    let rows = dimension(rows, name)?;
567    let columns = dimension(columns, name)?;
568    Ok(match order {
569        DenseOrder::RowMajor => (columns, 1),
570        DenseOrder::ColumnMajor => (1, rows),
571    })
572}