Skip to main content

numopt/matrix/
csr.rs

1//! Sparse matrix in compressed sparse row format.
2
3use crate::matrix::item::MatItem;
4
5/// Sparse matrix in compressed sparse row format.
6#[derive(Debug, Clone)]
7pub struct CsrMat<T> {
8    shape: (usize, usize),
9    indptr: Vec<usize>,
10    indices: Vec<usize>,
11    data: Vec<T>,
12}
13
14impl<T: MatItem> CsrMat<T> {
15
16    /// Creates [CsrMat](struct.CsrMat.html) from raw data.
17    pub fn new(shape: (usize, usize), 
18               indptr: Vec<usize>,
19               indices: Vec<usize>,
20               data: Vec<T>) -> Self {
21        assert_eq!(indptr.len(), shape.0+1);
22        assert_eq!(indices.len(), data.len());
23        assert_eq!(*indptr.last().unwrap(), data.len());
24        Self {
25            shape: shape,
26            indptr: indptr,
27            indices: indices,
28            data: data,
29        }
30    }
31
32    /// Number of rows.
33    pub fn rows(&self) -> usize { self.shape.0 }
34
35    /// Number of columns.
36    pub fn cols(&self) -> usize { self.shape.1 }
37
38    /// Number of nonzero elements.
39    pub fn nnz(&self) -> usize { self.indices.len() }
40
41    /// Vector of index pointers.
42    pub fn indptr(&self) -> &[usize] { &self.indptr }
43
44    /// Vector of column indices.
45    pub fn indices(&self) -> &[usize] { &self.indices }
46
47    /// Vector of data values.
48    pub fn data(&self) -> &[T] { &self.data }
49
50    /// Sums duplicate entries in-place.
51    pub fn sum_duplicates(&mut self) -> () {
52
53        let mut colseen: Vec<bool> = vec![false; self.cols()];
54        let mut colrow: Vec<usize> = vec![0; self.cols()];
55        let mut colnewk: Vec<usize> = vec![0; self.cols()];
56
57        let mut d: T;
58        let mut col: usize;
59        let mut new_k: usize = 0;
60        let mut new_counter: Vec<usize> = vec![0; self.rows()];
61        let mut new_indices: Vec<usize> = Vec::new();
62        let mut new_data: Vec<T> = Vec::new();
63        for row in 0..self.rows() {
64            for k in self.indptr[row]..self.indptr[row+1] {
65                
66                col = self.indices[k];
67                d = self.data[k];
68
69                // New column in row
70                if !colseen[col] || colrow[col] != row {        
71                    colnewk[col] = new_k;
72                    new_counter[row] += 1;
73                    new_indices.push(col);
74                    new_data.push(d);
75                    new_k += 1;
76                }
77                
78                // Duplicate column in row
79                else { 
80                    new_data[colnewk[col]] += d;
81                }
82
83                // Update
84                colseen[col] = true;
85                colrow[col] = row;
86            }
87
88        }
89
90        let mut offset: usize = 0;
91        let mut new_indptr: Vec<usize> = vec![0; self.rows()+1];
92        for (row, c) in new_counter.iter().enumerate() {
93            new_indptr[row+1] = offset + c;
94            offset += c;
95        }
96
97        self.indptr = new_indptr;
98        self.indices = new_indices;
99        self.data = new_data;
100
101        assert_eq!(self.indptr.len(), self.rows()+1);
102        assert_eq!(self.indices.len(), self.indptr[self.rows()]);
103        assert_eq!(self.indices.len(), self.data.len());
104    }
105}
106
107#[cfg(test)]
108mod tests {
109
110    use crate::matrix::coo::CooMat;
111    use crate::assert_vec_approx_eq;
112
113    #[test]
114    fn csr_sum_dublicates() {
115
116        // 6 2 1 0 0
117        // 3 1 0 7 0
118        // 4 6 0 0 1
119
120        let a = CooMat::new(
121            (3, 5),
122            vec![0 ,2 ,0 ,0 ,1 ,2  ,1 ,1 ,2 ,0 ,2],
123            vec![0 ,1 ,2 ,0 ,0 ,4  ,1 ,3 ,0 ,1 ,4],
124            vec![5.,6.,1.,1.,3.,-2.,1.,7.,4.,2.,3.],
125        );
126
127        let mut b = a.to_csr();
128        b.sum_duplicates();
129
130        assert_eq!(b.rows(), 3);
131        assert_eq!(b.cols(), 5);
132        assert_eq!(b.nnz(), 9);
133        assert_vec_approx_eq!(b.indptr(),
134                              vec![0, 3, 6, 9],
135                              epsilon=0);
136        assert_vec_approx_eq!(b.indices(),
137                              vec![0, 2, 1, 0, 1, 3, 1, 4, 0],
138                              epsilon=0);
139        assert_vec_approx_eq!(b.data(),
140                              vec![6., 1., 2., 3., 1., 7., 6., 1., 4.],
141                              epsilon=1e-8);
142    }
143}
144