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}