conspire/math/sparse/matrix/
mod.rs1#[cfg(test)]
2mod test;
3
4mod amd;
5
6use crate::math::{Scalar, 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 Mul<&Vector> for &CscMatrix {
198 type Output = Vector;
199 fn mul(self, vector: &Vector) -> Self::Output {
200 let mut output = Vector::zero(self.height);
201 (0..self.width).for_each(|j| {
202 (self.col_ptr[j]..self.col_ptr[j + 1])
203 .for_each(|k| output[self.row_idx[k]] += self.values[k] * vector[j])
204 });
205 output
206 }
207}