Skip to main content

rusparse/
host.rs

1use crate::{CsrMatrix, CsrMatrixOwned, DenseMatrix, DenseMatrixOwned, DenseOrder, SparseError};
2
3mod binary;
4mod sampled_sparse;
5mod indexing;
6pub use indexing::{csr_gather, csr_scatter_add};
7pub use binary::{csrgeam, csrgemm, csr_sum_pattern, csr_product_pattern};
8pub use sampled_sparse::sampled_csrgemm;
9
10pub fn csr_matmul(
11    matrix: CsrMatrix<'_>,
12    rhs: DenseMatrix<'_>,
13    transpose: bool,
14) -> Result<DenseMatrixOwned, SparseError> {
15    let transposed;
16    let matrix = if transpose {
17        transposed = matrix.transpose()?;
18        transposed.as_ref()
19    } else {
20        matrix
21    };
22    if matrix.columns() != rhs.rows() {
23        return Err(SparseError::DimensionMismatch("CSR matmul inner dimensions differ"));
24    }
25    let length = matrix.rows().checked_mul(rhs.columns())
26        .ok_or(SparseError::SizeOverflow("CSR matmul output"))?;
27    let mut output = vec![0.0_f32; length];
28    let base = matrix.index_base().value();
29    for row in 0..matrix.rows() {
30        let start = (matrix.row_offsets()[row] - base) as usize;
31        let end = (matrix.row_offsets()[row + 1] - base) as usize;
32        for entry in start..end {
33            let inner = (matrix.column_indices()[entry] - base) as usize;
34            let value = matrix.values()[entry];
35            for column in 0..rhs.columns() {
36                let destination = row * rhs.columns() + column;
37                output[destination] = value.mul_add(dense_at(rhs, inner, column), output[destination]);
38            }
39        }
40    }
41    DenseMatrixOwned::new(output, matrix.rows(), rhs.columns(), DenseOrder::RowMajor)
42}
43
44pub fn csr_sampled_matmul(
45    pattern: CsrMatrix<'_>,
46    lhs: DenseMatrix<'_>,
47    rhs: DenseMatrix<'_>,
48) -> Result<CsrMatrixOwned, SparseError> {
49    if lhs.rows() != pattern.rows() || rhs.columns() != pattern.columns() || lhs.columns() != rhs.rows() {
50        return Err(SparseError::DimensionMismatch("sampled matmul dimensions differ from CSR pattern"));
51    }
52    let mut values = vec![0.0_f32; pattern.nnz()];
53    let base = pattern.index_base().value();
54    for row in 0..pattern.rows() {
55        let start = (pattern.row_offsets()[row] - base) as usize;
56        let end = (pattern.row_offsets()[row + 1] - base) as usize;
57        for entry in start..end {
58            let column = (pattern.column_indices()[entry] - base) as usize;
59            let mut sum = 0.0_f32;
60            for inner in 0..lhs.columns() {
61                sum = dense_at(lhs, row, inner).mul_add(dense_at(rhs, inner, column), sum);
62            }
63            values[entry] = sum;
64        }
65    }
66    CsrMatrixOwned::new(
67        pattern.rows(), pattern.columns(), pattern.row_offsets().to_vec(),
68        pattern.column_indices().to_vec(), values, pattern.index_base(),
69    )
70}
71
72fn dense_at(matrix: DenseMatrix<'_>, row: usize, column: usize) -> f32 {
73    let offset = match matrix.order() {
74        DenseOrder::RowMajor => row * matrix.columns() + column,
75        DenseOrder::ColumnMajor => column * matrix.rows() + row,
76    };
77    matrix.values()[offset]
78}
79
80pub fn csr_to_dense_backward(pattern: CsrMatrix<'_>, grad: DenseMatrix<'_>) -> Result<Vec<f32>, SparseError> {
81    if grad.rows() != pattern.rows() || grad.columns() != pattern.columns() {
82        return Err(SparseError::DimensionMismatch("CSR to dense gradient shape mismatch"));
83    }
84    let mut values = vec![0.0_f32; pattern.nnz()];
85    let mut columns_seen = std::collections::BTreeSet::new();
86    let base = pattern.index_base().value();
87    for row in 0..pattern.rows() {
88        columns_seen.clear();
89        let start = (pattern.row_offsets()[row] - base) as usize;
90        let end = (pattern.row_offsets()[row + 1] - base) as usize;
91        for entry in (start..end).rev() {
92            let column = (pattern.column_indices()[entry] - base) as usize;
93            if columns_seen.insert(column) {
94                values[entry] = dense_at(grad, row, column);
95            }
96        }
97    }
98    Ok(values)
99}