conspire/math/sparse/matrix/
mod.rs1#[cfg(test)]
2mod test;
3
4mod amd;
5
6use crate::math::{Scalar, TensorRank1Vec, TensorRank2, Vector};
7use std::ops::Mul;
8
9#[derive(Clone, Debug, PartialEq)]
11pub struct CscMatrix {
12 height: usize,
13 width: usize,
14 col_ptr: Vec<usize>,
15 row_idx: Vec<usize>,
16 values: Vec<Scalar>,
17 pattern: Vec<(usize, usize)>,
18 scatter: Vec<usize>,
19}
20
21impl CscMatrix {
22 pub fn from_pattern(height: usize, width: usize, pattern: Vec<(usize, usize)>) -> Self {
25 assert!(!pattern.is_empty(), "Matrix must have at least one entry.");
26 let mut order: Vec<usize> = (0..pattern.len()).collect();
27 order.sort_unstable_by_key(|&k| (pattern[k].1, pattern[k].0));
28 let mut col_ptr = vec![0; width + 1];
29 let mut row_idx = Vec::with_capacity(pattern.len());
30 let mut scatter = vec![0; pattern.len()];
31 let mut last = (usize::MAX, usize::MAX);
32 order.into_iter().for_each(|k| {
33 let (i, j) = pattern[k];
34 assert!(i < height && j < width, "Position out of bounds.");
35 if (j, i) != last {
36 last = (j, i);
37 row_idx.push(i);
38 col_ptr[j + 1] += 1;
39 }
40 scatter[k] = row_idx.len() - 1;
41 });
42 (0..width).for_each(|j| col_ptr[j + 1] += col_ptr[j]);
43 let values = vec![0.0; row_idx.len()];
44 Self {
45 height,
46 width,
47 col_ptr,
48 row_idx,
49 values,
50 pattern,
51 scatter,
52 }
53 }
54 pub fn fill(&mut self, mut source: impl FnMut(usize, usize) -> Scalar) {
56 self.values.fill(0.0);
57 self.pattern
58 .iter()
59 .zip(self.scatter.iter())
60 .for_each(|(&(i, j), &k)| self.values[k] += source(i, j));
61 }
62 pub fn column(&self, j: usize) -> impl Iterator<Item = (usize, &Scalar)> {
64 (self.col_ptr[j]..self.col_ptr[j + 1]).map(move |k| (self.row_idx[k], &self.values[k]))
65 }
66 pub fn height(&self) -> usize {
67 self.height
68 }
69 pub fn iter(&self) -> impl Iterator<Item = (usize, usize, &Scalar)> {
71 (0..self.width).flat_map(move |j| self.column(j).map(move |(i, value)| (i, j, value)))
72 }
73 pub fn nonzeros(&self) -> usize {
74 self.row_idx.len()
75 }
76 pub fn pattern(&self) -> &[(usize, usize)] {
78 &self.pattern
79 }
80 pub fn transpose(&self) -> Self {
81 let nnz = self.row_idx.len();
82 let mut col_ptr = vec![0; self.height + 1];
83 self.row_idx.iter().for_each(|&i| col_ptr[i + 1] += 1);
84 (0..self.height).for_each(|i| col_ptr[i + 1] += col_ptr[i]);
85 let mut next = col_ptr.clone();
86 let mut row_idx = vec![0; nnz];
87 let mut values = vec![0.0; nnz];
88 let mut perm = vec![0; nnz];
89 (0..self.width).for_each(|j| {
90 (self.col_ptr[j]..self.col_ptr[j + 1]).for_each(|k| {
91 let p = next[self.row_idx[k]];
92 next[self.row_idx[k]] += 1;
93 row_idx[p] = j;
94 values[p] = self.values[k];
95 perm[k] = p;
96 })
97 });
98 Self {
99 height: self.width,
100 width: self.height,
101 col_ptr,
102 row_idx,
103 values,
104 pattern: self.pattern.iter().map(|&(i, j)| (j, i)).collect(),
105 scatter: self.scatter.iter().map(|&k| perm[k]).collect(),
106 }
107 }
108 pub(crate) fn maxtrans(&self) -> Option<Vec<usize>> {
111 const NONE: usize = usize::MAX;
112 let n = self.width;
113 assert_eq!(n, self.height);
114 if (0..n).all(|c| self.row_idx[self.col_ptr[c]..self.col_ptr[c + 1]].contains(&c)) {
115 return Some((0..n).collect());
116 }
117 let mut cmatch = vec![NONE; n];
118 let mut rmatch = vec![NONE; n];
119 (0..n).for_each(|c| {
120 if self.row_idx[self.col_ptr[c]..self.col_ptr[c + 1]].contains(&c) {
121 rmatch[c] = c;
122 cmatch[c] = c;
123 }
124 });
125 (0..n).for_each(|c| {
126 if cmatch[c] == NONE {
127 for &r in &self.row_idx[self.col_ptr[c]..self.col_ptr[c + 1]] {
128 if rmatch[r] == NONE {
129 rmatch[r] = c;
130 cmatch[c] = r;
131 break;
132 }
133 }
134 }
135 });
136 let mut visited = vec![NONE; n];
137 let mut cstack = vec![0; n];
138 let mut rstack = vec![0; n];
139 let mut estack = vec![0; n];
140 for root in 0..n {
141 if cmatch[root] != NONE {
142 continue;
143 }
144 let mut head = 0;
145 cstack[0] = root;
146 estack[0] = self.col_ptr[root];
147 let mut found = false;
148 'dfs: loop {
149 let c = cstack[head];
150 let end = self.col_ptr[c + 1];
151 if estack[head] == self.col_ptr[c] {
152 for &r in &self.row_idx[self.col_ptr[c]..end] {
153 if rmatch[r] == NONE {
154 rstack[head] = r;
155 (0..=head).for_each(|level| {
156 cmatch[cstack[level]] = rstack[level];
157 rmatch[rstack[level]] = cstack[level];
158 });
159 found = true;
160 break 'dfs;
161 }
162 }
163 }
164 let mut descended = false;
165 while estack[head] < end {
166 let r = self.row_idx[estack[head]];
167 estack[head] += 1;
168 if visited[r] == root || rmatch[r] == NONE {
169 continue;
170 }
171 visited[r] = root;
172 rstack[head] = r;
173 head += 1;
174 cstack[head] = rmatch[r];
175 estack[head] = self.col_ptr[rmatch[r]];
176 descended = true;
177 break;
178 }
179 if !descended {
180 if head == 0 {
181 break 'dfs;
182 }
183 head -= 1;
184 }
185 }
186 if !found {
187 return None;
188 }
189 }
190 Some(cmatch)
191 }
192 pub fn width(&self) -> usize {
193 self.width
194 }
195}
196
197impl CscMatrix {
198 pub fn entry(&self, row: usize, column: usize) -> Scalar {
200 (self.col_ptr[column]..self.col_ptr[column + 1])
201 .find(|&k| self.row_idx[k] == row)
202 .map_or(0.0, |k| self.values[k])
203 }
204 fn multiply(&self, entry: impl Fn(usize) -> Scalar) -> Vector {
207 let mut output = Vector::zero(self.height);
208 (0..self.width).for_each(|j| {
209 let entry_j = entry(j);
210 (self.col_ptr[j]..self.col_ptr[j + 1])
211 .for_each(|k| output[self.row_idx[k]] += self.values[k] * entry_j)
212 });
213 output
214 }
215}
216
217impl Mul<&Vector> for &CscMatrix {
218 type Output = Vector;
219 fn mul(self, vector: &Vector) -> Self::Output {
220 self.multiply(|j| vector[j])
221 }
222}
223
224impl<const D: usize, I, J> Mul<&TensorRank2<D, I, J>> for &CscMatrix {
225 type Output = Vector;
226 fn mul(self, tensor_rank_2: &TensorRank2<D, I, J>) -> Self::Output {
227 self.multiply(|j| tensor_rank_2[j / D][j % D].value())
228 }
229}
230
231impl<const D: usize, I> Mul<&TensorRank1Vec<D, I>> for &CscMatrix {
232 type Output = Vector;
233 fn mul(self, tensor_rank_1_vec: &TensorRank1Vec<D, I>) -> Self::Output {
234 self.multiply(|j| tensor_rank_1_vec[j / D][j % D].value())
235 }
236}
237
238impl Mul<&CscMatrix> for &Vector {
239 type Output = Vector;
240 fn mul(self, csc_matrix: &CscMatrix) -> Self::Output {
241 let mut output = Vector::zero(csc_matrix.width);
242 (0..csc_matrix.width).for_each(|j| {
243 output[j] = (csc_matrix.col_ptr[j]..csc_matrix.col_ptr[j + 1])
244 .map(|k| csc_matrix.values[k] * self[csc_matrix.row_idx[k]])
245 .sum()
246 });
247 output
248 }
249}