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