1#[cfg(target_arch = "x86_64")]
2mod avx;
3#[cfg(test)]
4mod test;
5
6use super::super::{SparseError, matrix::CscMatrix};
7use super::gemm::{CHUNK, NONE, axpy, etree, gemm_wide, max_below, reach_sorted, supernodes};
8use crate::{
9 ABS_TOL,
10 math::{Scalar, Vector},
11};
12
13const PIVOT_TOL: Scalar = 0.001;
16
17pub struct CscLu {
19 pub(super) fill: usize,
20 pub(super) sn_of: Vec<usize>,
21 pub(super) sn_start: Vec<usize>,
22 pub(super) sn_rows_ptr: Vec<usize>,
23 pub(super) sn_rows: Vec<usize>,
24 pub(super) sn_panel_ptr: Vec<usize>,
25 pub(super) sn_values: Vec<Scalar>,
26 pub(super) u_col_ptr: Vec<usize>,
27 pub(super) u_row_idx: Vec<usize>,
28 pub(super) u_values: Vec<Scalar>,
29 pub(super) pinv: Vec<usize>,
30 pub(super) q: Vec<usize>,
31}
32
33impl CscMatrix {
34 pub fn lu(&self) -> Result<CscLu, SparseError> {
36 self.factor((0..self.height()).collect())
37 }
38 pub fn lu_amd(&self) -> Result<CscLu, SparseError> {
41 self.factor(self.amd())
42 }
43 pub fn lu_symbolic(&self) -> Result<CscLu, SparseError> {
49 let n = self.height();
50 assert_eq!(n, self.width());
51 let matching = self.maxtrans().ok_or(SparseError::Singular)?;
52 let q = if matching.iter().enumerate().all(|(c, &r)| r == c) {
53 self.amd()
54 } else {
55 let mut rinv = vec![0; n];
56 matching.iter().enumerate().for_each(|(c, &r)| rinv[r] = c);
57 let mut permuted = Vec::with_capacity(self.nonzeros());
58 (0..n).for_each(|c| {
59 self.column(c)
60 .for_each(|(r, _)| permuted.push((rinv[r], c)))
61 });
62 CscMatrix::from_pattern(n, n, permuted).amd()
63 };
64 let mut pinv = vec![NONE; n];
65 q.iter()
66 .enumerate()
67 .for_each(|(j, &q_j)| pinv[matching[q_j]] = j);
68 Ok(self.symbolic(q, pinv, &vec![false; n]))
69 }
70 pub(super) fn symbolic(&self, q: Vec<usize>, pinv: Vec<usize>, locked: &[bool]) -> CscLu {
71 let n = self.height();
72 let mut qinv = vec![NONE; n];
73 q.iter().enumerate().for_each(|(j, &q_j)| qinv[q_j] = j);
74 let (parent, adj_ptr, adj) = etree(self, n, n, |c| qinv[c], |r| pinv[r]);
75 let mut mark = vec![NONE; n];
76 let mut row = Vec::new();
77 let mut l_count = vec![1_usize; n];
78 let mut u_col_ptr = Vec::with_capacity(n + 1);
79 let mut u_row_idx = Vec::new();
80 u_col_ptr.push(0);
81 (0..n).for_each(|k| {
82 reach_sorted(k, &adj_ptr, &adj, &parent, &mut mark, &mut row);
83 row.iter().for_each(|&j| l_count[j] += 1);
84 u_row_idx.extend_from_slice(&row);
85 u_row_idx.push(k);
86 u_col_ptr.push(u_row_idx.len());
87 });
88 let mut l_col_ptr = Vec::with_capacity(n + 1);
89 l_col_ptr.push(0);
90 l_count.iter().for_each(|count| {
91 l_col_ptr.push(l_col_ptr.last().unwrap() + count);
92 });
93 let mut l_row_idx = vec![0; l_col_ptr[n]];
94 let mut next = vec![0; n];
95 (0..n).for_each(|j| {
96 l_row_idx[l_col_ptr[j]] = j;
97 next[j] = l_col_ptr[j] + 1;
98 });
99 (0..n).for_each(|k| {
100 u_row_idx[u_col_ptr[k]..u_col_ptr[k + 1] - 1]
101 .iter()
102 .for_each(|&j| {
103 l_row_idx[next[j]] = k;
104 next[j] += 1;
105 });
106 });
107 let fill = l_row_idx.len() + u_row_idx.len();
108 let l_values = vec![0.0; l_row_idx.len()];
109 let (sn_of, sn_start, sn_rows_ptr, sn_rows, sn_panel_ptr, sn_values) =
110 supernodes(&l_col_ptr, &l_row_idx, &l_values, n, locked);
111 CscLu {
112 fill,
113 sn_of,
114 sn_start,
115 sn_rows_ptr,
116 sn_rows,
117 sn_panel_ptr,
118 sn_values,
119 u_col_ptr,
120 u_row_idx,
121 u_values: vec![0.0; fill - l_values.len()],
122 pinv,
123 q,
124 }
125 }
126 fn factor(&self, q: Vec<usize>) -> Result<CscLu, SparseError> {
127 let n = self.height();
128 assert_eq!(n, self.width());
129 let mut pinv = vec![NONE; n];
130 let mut l_cols = Vec::<Vec<(usize, Scalar)>>::with_capacity(n);
131 let mut u_cols = Vec::<Vec<(usize, Scalar)>>::with_capacity(n);
132 let mut x = vec![0.0; n];
133 let mut mark = vec![0; n];
134 let mut order = vec![0; n];
135 let mut stack = vec![0; n];
136 let mut pstack = vec![0; n];
137 let mut top;
138 for (j, &q_j) in q.iter().enumerate() {
139 top = reach(
140 self.column(q_j).map(|(i, _)| i),
141 &l_cols,
142 &pinv,
143 j + 1,
144 &mut mark,
145 &mut order,
146 &mut stack,
147 &mut pstack,
148 );
149 self.column(q_j).for_each(|(i, value)| x[i] = *value);
150 order[top..n].iter().for_each(|&i| {
151 if pinv[i] != NONE {
152 let x_i = x[i];
153 l_cols[pinv[i]]
154 .iter()
155 .skip(1)
156 .for_each(|&(row, value)| x[row] -= value * x_i);
157 }
158 });
159 let mut pivot_row = NONE;
160 let mut pivot_abs = 0.0;
161 order[top..n].iter().for_each(|&i| {
162 if pinv[i] == NONE && x[i].abs() > pivot_abs {
163 pivot_abs = x[i].abs();
164 pivot_row = i;
165 }
166 });
167 if pivot_row == NONE || pivot_abs < ABS_TOL {
168 return Err(SparseError::Singular);
169 }
170 if pinv[q_j] == NONE && x[q_j].abs() >= PIVOT_TOL * pivot_abs {
171 pivot_row = q_j;
172 }
173 let pivot = x[pivot_row];
174 pinv[pivot_row] = j;
175 let mut l_col = vec![(pivot_row, 1.0)];
176 let mut u_col = Vec::new();
177 order[top..n].iter().for_each(|&i| {
178 if pinv[i] == NONE {
179 l_col.push((i, x[i] / pivot));
180 } else if i != pivot_row {
181 u_col.push((pinv[i], x[i]));
182 }
183 x[i] = 0.0;
184 });
185 u_col.push((j, pivot));
186 u_col.sort_unstable_by_key(|&(row, _)| row);
187 l_cols.push(l_col);
188 u_cols.push(u_col);
189 }
190 let (l_col_ptr, l_row_idx, l_values) = compress(l_cols, &pinv);
191 let (u_col_ptr, u_row_idx, u_values) = compress(u_cols, &(0..n).collect::<Vec<usize>>());
192 let fill = l_values.len() + u_values.len();
193 let (sn_of, sn_start, sn_rows_ptr, sn_rows, sn_panel_ptr, sn_values) =
194 supernodes(&l_col_ptr, &l_row_idx, &l_values, n, &vec![false; n]);
195 Ok(CscLu {
196 fill,
197 sn_of,
198 sn_start,
199 sn_rows_ptr,
200 sn_rows,
201 sn_panel_ptr,
202 sn_values,
203 u_col_ptr,
204 u_row_idx,
205 u_values,
206 pinv,
207 q,
208 })
209 }
210}
211
212impl CscLu {
213 pub fn solve(&self, b: &Vector) -> Vector {
215 let n = self.pinv.len();
216 let mut x = vec![0.0; n];
217 let mut below = vec![0.0; self.max_below()];
218 self.pinv
219 .iter()
220 .enumerate()
221 .for_each(|(i, &p_i)| x[p_i] = b[i]);
222 (0..self.sn_start.len() - 1).for_each(|s| {
223 let t1 = self.sn_start[s];
224 let t2 = self.sn_start[s + 1];
225 let width = t2 - t1;
226 let rows = &self.sn_rows[self.sn_rows_ptr[s]..self.sn_rows_ptr[s + 1]];
227 let m = rows.len();
228 let panel = &self.sn_values[self.sn_panel_ptr[s]..self.sn_panel_ptr[s + 1]];
229 (0..width).for_each(|c| {
230 let x_c = x[t1 + c];
231 if x_c != 0.0 {
232 let column = &panel[c * m..(c + 1) * m];
233 x[t1 + c + 1..t2]
234 .iter_mut()
235 .zip(column[c + 1..width].iter())
236 .for_each(|(x_r, value)| *x_r -= value * x_c);
237 below[..m - width]
238 .iter_mut()
239 .zip(column[width..].iter())
240 .for_each(|(below_r, value)| *below_r += value * x_c);
241 }
242 });
243 rows[width..]
244 .iter()
245 .zip(below[..m - width].iter_mut())
246 .for_each(|(&row, below_r)| {
247 x[row] -= *below_r;
248 *below_r = 0.0;
249 });
250 });
251 (0..n).rev().for_each(|j| {
252 let end = self.u_col_ptr[j + 1];
253 x[j] /= self.u_values[end - 1];
254 let x_j = x[j];
255 if x_j != 0.0 {
256 self.u_row_idx[self.u_col_ptr[j]..end - 1]
257 .iter()
258 .zip(self.u_values[self.u_col_ptr[j]..end - 1].iter())
259 .for_each(|(&row, value)| x[row] -= value * x_j);
260 }
261 });
262 let mut solution = Vector::zero(n);
263 self.q
264 .iter()
265 .enumerate()
266 .for_each(|(j, &q_j)| solution[q_j] = x[j]);
267 solution
268 }
269 pub fn nonzeros(&self) -> usize {
271 self.fill
272 }
273 pub fn refactor(&mut self, matrix: &CscMatrix) -> Result<(), SparseError> {
277 let n = self.pinv.len();
278 assert_eq!(n, matrix.height());
279 let mut work = vec![0.0; n * CHUNK];
280 let mut temp = vec![0.0; CHUNK * self.max_below()];
281 let mut tile = vec![
282 0.0;
283 CHUNK
284 * (0..self.sn_start.len() - 1)
285 .map(|s| self.sn_start[s + 1] - self.sn_start[s])
286 .max()
287 .unwrap_or(0)
288 ];
289 let mut pointers = [0; CHUNK];
290 for s in 0..self.sn_start.len() - 1 {
291 let s1 = self.sn_start[s];
292 let s2 = self.sn_start[s + 1];
293 let s_width = s2 - s1;
294 let s_rows_start = self.sn_rows_ptr[s];
295 let s_m = self.sn_rows_ptr[s + 1] - s_rows_start;
296 let mut c1 = s1;
297 while c1 < s2 {
298 let c2 = s2.min(c1 + CHUNK);
299 let chunk = c2 - c1;
300 (c1..c2).zip(pointers.iter_mut()).for_each(|(j, pointer)| {
301 let pinv = &self.pinv;
302 let column = &mut work[(j - c1) * n..(j - c1 + 1) * n];
303 matrix
304 .column(self.q[j])
305 .for_each(|(i, value)| column[pinv[i]] = *value);
306 *pointer = self.u_col_ptr[j];
307 });
308 loop {
309 let mut t = NONE;
310 (c1..c2).zip(pointers.iter()).for_each(|(j, &pointer)| {
311 if pointer < self.u_col_ptr[j + 1] - 1 {
312 let row = self.u_row_idx[pointer];
313 if row < c1 && (t == NONE || self.sn_of[row] < t) {
314 t = self.sn_of[row];
315 }
316 }
317 });
318 if t == NONE {
319 break;
320 }
321 let t1 = self.sn_start[t];
322 let t2 = self.sn_start[t + 1];
323 let width = t2 - t1;
324 let consumed = width.min(c1 - t1);
325 let rows = &self.sn_rows[self.sn_rows_ptr[t]..self.sn_rows_ptr[t + 1]];
326 let m = rows.len();
327 let below = m - width;
328 let panel = &self.sn_values[self.sn_panel_ptr[t]..self.sn_panel_ptr[t + 1]];
329 if consumed >= 4 {
330 if chunk < CHUNK {
331 tile[..width * CHUNK].fill(0.0);
332 }
333 (0..chunk).for_each(|b| {
334 work[b * n + t1..b * n + t2]
335 .iter()
336 .zip(tile.chunks_exact_mut(CHUNK))
337 .for_each(|(&value, row)| row[b] = value);
338 });
339 trisolve(&mut tile[..width * CHUNK], panel, m, consumed, width);
340 (0..chunk).for_each(|b| {
341 work[b * n + t1..b * n + t2]
342 .iter_mut()
343 .zip(tile.chunks_exact(CHUNK))
344 .for_each(|(value, row)| *value = row[b]);
345 });
346 (c1..c2).zip(pointers.iter_mut()).for_each(|(j, pointer)| {
347 let p_end = self.u_col_ptr[j + 1] - 1;
348 let column = &work[(j - c1) * n..(j - c1 + 1) * n];
349 while *pointer < p_end {
350 let k = self.u_row_idx[*pointer];
351 if k >= t1 + consumed {
352 break;
353 }
354 self.u_values[*pointer] = column[k];
355 *pointer += 1;
356 }
357 });
358 } else {
359 (c1..c2).zip(pointers.iter_mut()).for_each(|(j, pointer)| {
360 let p_end = self.u_col_ptr[j + 1] - 1;
361 let column = &mut work[(j - c1) * n..(j - c1 + 1) * n];
362 if consumed == width
363 && width <= 3
364 && *pointer + width <= p_end
365 && self.u_row_idx[*pointer] == t1
366 && self.u_row_idx[*pointer + width - 1] == t1 + width - 1
367 {
368 match width {
369 1 => self.u_values[*pointer] = column[t1],
370 2 => {
371 let u_0 = column[t1];
372 let u_1 = column[t1 + 1] - panel[1] * u_0;
373 column[t1 + 1] = u_1;
374 self.u_values[*pointer] = u_0;
375 self.u_values[*pointer + 1] = u_1;
376 }
377 _ => {
378 let u_0 = column[t1];
379 let u_1 = column[t1 + 1] - panel[1] * u_0;
380 let u_2 =
381 column[t1 + 2] - panel[2] * u_0 - panel[m + 2] * u_1;
382 column[t1 + 1] = u_1;
383 column[t1 + 2] = u_2;
384 self.u_values[*pointer] = u_0;
385 self.u_values[*pointer + 1] = u_1;
386 self.u_values[*pointer + 2] = u_2;
387 }
388 }
389 *pointer += width;
390 } else {
391 while *pointer < p_end {
392 let k = self.u_row_idx[*pointer];
393 if k >= c1 || k >= t2 {
394 break;
395 }
396 let c = k - t1;
397 let u = column[k];
398 self.u_values[*pointer] = u;
399 if u != 0.0 {
400 column[k + 1..t2]
401 .iter_mut()
402 .zip(panel[c * m + c + 1..c * m + width].iter())
403 .for_each(|(work_r, value)| *work_r -= value * u);
404 }
405 *pointer += 1;
406 }
407 }
408 });
409 }
410 if below > 0 {
411 temp[..CHUNK * below].fill(0.0);
412 gemm_wide(
413 &mut temp[..CHUNK * below],
414 &work,
415 n,
416 panel,
417 m,
418 width,
419 t1,
420 consumed,
421 below,
422 chunk,
423 );
424 (0..chunk).for_each(|b| {
425 let column = &mut work[b * n..(b + 1) * n];
426 rows[width..]
427 .iter()
428 .zip(temp[b * below..(b + 1) * below].iter())
429 .for_each(|(&row, value)| column[row] -= value);
430 });
431 }
432 (0..chunk).for_each(|c| work[c * n + t1..c * n + t1 + consumed].fill(0.0));
433 }
434 for j in c1..c2 {
435 let p_end = self.u_col_ptr[j + 1] - 1;
436 let offset = (j - c1) * n;
437 let mut pointer = pointers[j - c1];
438 while pointer < p_end {
439 let k = self.u_row_idx[pointer];
440 let c = k - s1;
441 let u = work[offset + k];
442 self.u_values[pointer] = u;
443 work[offset + k] = 0.0;
444 if u != 0.0 {
445 let start = self.sn_panel_ptr[s] + c * s_m;
446 axpy(
447 &mut work[offset + k + 1..offset + s2],
448 &self.sn_values[start + c + 1..start + s_width],
449 u,
450 );
451 self.sn_rows[s_rows_start + s_width..s_rows_start + s_m]
452 .iter()
453 .zip(self.sn_values[start + s_width..start + s_m].iter())
454 .for_each(|(&row, value)| work[offset + row] -= value * u);
455 }
456 pointer += 1;
457 }
458 let pivot = work[offset + j];
459 work[offset + j] = 0.0;
460 if pivot.abs() < ABS_TOL {
461 return Err(SparseError::Singular);
462 }
463 self.u_values[p_end] = pivot;
464 let c = j - s1;
465 let start = self.sn_panel_ptr[s] + c * s_m;
466 self.sn_values[start + c] = 1.0;
467 (c + 1..s_width).for_each(|local| {
468 self.sn_values[start + local] = work[offset + s1 + local] / pivot;
469 work[offset + s1 + local] = 0.0;
470 });
471 (s_width..s_m).for_each(|local| {
472 let row = self.sn_rows[s_rows_start + local];
473 self.sn_values[start + local] = work[offset + row] / pivot;
474 work[offset + row] = 0.0;
475 });
476 }
477 c1 = c2;
478 }
479 }
480 Ok(())
481 }
482 pub(super) fn max_below(&self) -> usize {
483 max_below(&self.sn_start, &self.sn_rows_ptr)
484 }
485}
486
487fn trisolve(tile: &mut [Scalar], panel: &[Scalar], m: usize, consumed: usize, width: usize) {
491 #[cfg(target_arch = "x86_64")]
492 if super::simd() {
493 return unsafe { avx::trisolve(tile, panel, m, consumed, width) };
494 }
495 (0..consumed).for_each(|c| {
496 let (row, rest) = tile[c * CHUNK..].split_at_mut(CHUNK);
497 if row.iter().any(|&u| u != 0.0) {
498 rest[..(width - c - 1) * CHUNK]
499 .chunks_exact_mut(CHUNK)
500 .zip(panel[c * m + c + 1..c * m + width].iter())
501 .for_each(|(target, &value)| {
502 target
503 .iter_mut()
504 .zip(row.iter())
505 .for_each(|(entry, &u)| *entry -= value * u)
506 });
507 }
508 });
509}
510
511#[allow(clippy::too_many_arguments)]
514fn reach(
515 starts: impl Iterator<Item = usize>,
516 l_cols: &[Vec<(usize, Scalar)>],
517 pinv: &[usize],
518 tag: usize,
519 mark: &mut [usize],
520 order: &mut [usize],
521 stack: &mut [usize],
522 pstack: &mut [usize],
523) -> usize {
524 let mut top = order.len();
525 starts.for_each(|start| {
526 if mark[start] != tag {
527 mark[start] = tag;
528 stack[0] = start;
529 pstack[0] = 0;
530 let mut head = 0;
531 loop {
532 let i = stack[head];
533 let mut descended = false;
534 if pinv[i] != NONE {
535 let column = &l_cols[pinv[i]];
536 while pstack[head] < column.len() {
537 let child = column[pstack[head]].0;
538 pstack[head] += 1;
539 if mark[child] != tag {
540 mark[child] = tag;
541 head += 1;
542 stack[head] = child;
543 pstack[head] = 0;
544 descended = true;
545 break;
546 }
547 }
548 }
549 if !descended {
550 top -= 1;
551 order[top] = i;
552 if head == 0 {
553 break;
554 }
555 head -= 1;
556 }
557 }
558 }
559 });
560 top
561}
562
563fn compress(
564 cols: Vec<Vec<(usize, Scalar)>>,
565 row_map: &[usize],
566) -> (Vec<usize>, Vec<usize>, Vec<Scalar>) {
567 let nnz = cols.iter().map(|col| col.len()).sum();
568 let mut col_ptr = Vec::with_capacity(cols.len() + 1);
569 let mut row_idx = Vec::with_capacity(nnz);
570 let mut values = Vec::with_capacity(nnz);
571 col_ptr.push(0);
572 cols.into_iter().for_each(|col| {
573 let mut mapped: Vec<(usize, Scalar)> = col
574 .into_iter()
575 .map(|(i, value)| (row_map[i], value))
576 .collect();
577 mapped.sort_unstable_by_key(|&(row, _)| row);
578 mapped.into_iter().for_each(|(row, value)| {
579 row_idx.push(row);
580 values.push(value);
581 });
582 col_ptr.push(row_idx.len());
583 });
584 (col_ptr, row_idx, values)
585}