Skip to main content

conspire/math/sparse/matrix/
mod.rs

1#[cfg(test)]
2mod test;
3
4mod amd;
5
6use crate::math::{Scalar, Vector};
7use std::ops::Mul;
8
9/// A sparse matrix in compressed sparse column format.
10#[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    /// Builds the sparsity structure from a list of nonzero (row, column) positions,
23    /// with all values initialized to zero.
24    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    /// Fills the values from a source, summing duplicate positions in the pattern.
55    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    /// Iterates over the nonzero entries of a column as (row, value).
63    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    /// Iterates over the nonzero entries as (row, column, value), in column-major order.
70    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    /// The nonzero (row, column) positions this structure was built from.
77    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    /// A column-to-row matching pairing every column with a structurally
109    /// nonzero row (a maximum transversal), or None if structurally singular.
110    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}