1use crate::matrix::item::MatItem;
4
5#[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 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 pub fn rows(&self) -> usize { self.shape.0 }
34
35 pub fn cols(&self) -> usize { self.shape.1 }
37
38 pub fn nnz(&self) -> usize { self.indices.len() }
40
41 pub fn indptr(&self) -> &[usize] { &self.indptr }
43
44 pub fn indices(&self) -> &[usize] { &self.indices }
46
47 pub fn data(&self) -> &[T] { &self.data }
49
50 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 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 else {
80 new_data[colnewk[col]] += d;
81 }
82
83 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 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