1#![forbid(unsafe_code)]
2
3use core::hint::cold_path;
11
12use crate::matrix::Matrix;
13use crate::scaled_product::{ScaledProduct, range_checked_product};
14use crate::vector::Vector;
15use crate::{ArithmeticOperation, FactorizationKind, LaError, Tolerance};
16
17#[must_use]
24#[derive(Clone, Copy, Debug, PartialEq)]
25pub struct Lu<const D: usize> {
26 factors: LuFactors<D>,
27 permutation: RowPermutation<D>,
28}
29
30#[derive(Clone, Copy, Debug, PartialEq)]
35struct LuFactors<const D: usize> {
36 storage: [[f64; D]; D],
37}
38
39impl<const D: usize> LuFactors<D> {
40 #[inline]
42 const fn try_from_computation(storage: [[f64; D]; D]) -> Result<Self, LaError> {
43 let mut row = 0;
44 while row < D {
45 let mut col = 0;
46 while col < D {
47 if !storage[row][col].is_finite() {
48 return Err(LaError::non_finite_computation_matrix(
49 ArithmeticOperation::LuFactorization,
50 row,
51 col,
52 ));
53 }
54 col += 1;
55 }
56 row += 1;
57 }
58
59 Ok(Self { storage })
60 }
61
62 #[inline]
64 #[must_use]
65 const fn row(&self, index: usize) -> &[f64; D] {
66 &self.storage[index]
67 }
68
69 #[inline]
71 #[must_use]
72 const fn diag(&self, index: usize) -> f64 {
73 self.storage[index][index]
74 }
75}
76
77#[derive(Clone, Copy, Debug, PartialEq, Eq)]
83struct RowPermutation<const D: usize> {
84 source_rows: [usize; D],
85 odd: bool,
86}
87
88impl<const D: usize> RowPermutation<D> {
89 const fn identity() -> Self {
91 let mut source_rows = [0; D];
92 let mut row = 0;
93 while row < D {
94 source_rows[row] = row;
95 row += 1;
96 }
97 Self {
98 source_rows,
99 odd: false,
100 }
101 }
102
103 const fn swap(&mut self, left: usize, right: usize) {
105 if left != right {
106 let source_row = self.source_rows[left];
107 self.source_rows[left] = self.source_rows[right];
108 self.source_rows[right] = source_row;
109 self.odd = !self.odd;
110 }
111 }
112
113 const fn source_row(&self, row: usize) -> usize {
115 self.source_rows[row]
116 }
117
118 const fn is_odd(&self) -> bool {
120 self.odd
121 }
122}
123
124impl<const D: usize> Lu<D> {
125 #[inline]
134 pub(crate) fn factor_finite(a: Matrix<D>, tol: Tolerance) -> Result<Self, LaError> {
135 let mut rows = a.into_rows();
136 let tolerance = tol.get();
137 let mut permutation = RowPermutation::identity();
138
139 {
140 let rows = &mut rows;
141
142 for k in 0..D {
143 let mut pivot_row = k;
145 let mut pivot_abs = rows[k][k].abs();
146
147 #[expect(
148 clippy::needless_range_loop,
149 reason = "the row index identifies the pivot later used for synchronized matrix and permutation swaps"
150 )]
151 for r in (k + 1)..D {
152 let v = rows[r][k].abs();
153 if v > pivot_abs {
154 pivot_abs = v;
155 pivot_row = r;
156 }
157 }
158
159 if pivot_abs <= tolerance {
160 cold_path();
161
162 for (row, values) in rows.iter().enumerate() {
166 for (col, value) in values.iter().enumerate() {
167 if !value.is_finite() {
168 return Err(LaError::non_finite_computation_matrix(
169 ArithmeticOperation::LuFactorization,
170 row,
171 col,
172 ));
173 }
174 }
175 }
176
177 return Err(LaError::singular_numerical(
178 k,
179 FactorizationKind::Lu,
180 pivot_abs,
181 tolerance,
182 ));
183 }
184
185 if pivot_row != k {
186 rows.swap(k, pivot_row);
187 permutation.swap(k, pivot_row);
188 }
189
190 let pivot = rows[k][k];
191
192 for r in (k + 1)..D {
194 let mult = rows[r][k] / pivot;
195 rows[r][k] = mult;
196
197 #[expect(
198 clippy::needless_range_loop,
199 reason = "the column index pairs pivot-row reads with eliminated-row writes in the in-place update"
200 )]
201 for c in (k + 1)..D {
202 let updated = (-mult).mul_add(rows[k][c], rows[r][c]);
203 rows[r][c] = updated;
204 }
205 }
206 }
207 }
208
209 let factors = LuFactors::try_from_computation(rows)?;
210
211 Ok(Self {
212 factors,
213 permutation,
214 })
215 }
216
217 #[inline]
245 pub const fn solve(&self, b: Vector<D>) -> Result<Vector<D>, LaError> {
246 let mut x = [0.0; D];
247 let b = b.as_array();
248 let mut i = 0;
249
250 if D <= 4 {
251 while i < D {
252 x[i] = b[self.permutation.source_row(i)];
253 i += 1;
254 }
255
256 i = 0;
259 while i < D {
260 let mut sum = x[i];
261 let row = self.factors.row(i);
262 let mut j = 0;
263 while j < i {
264 sum = (-row[j]).mul_add(x[j], sum);
265 j += 1;
266 }
267 if !sum.is_finite() {
268 cold_path();
269 return Err(LaError::non_finite_computation_step(
270 ArithmeticOperation::LuSolve,
271 i,
272 ));
273 }
274 x[i] = sum;
275 i += 1;
276 }
277 } else {
278 while i < D {
281 let mut sum = b[self.permutation.source_row(i)];
282 let row = self.factors.row(i);
283 let mut j = 0;
284 while j < i {
285 sum = (-row[j]).mul_add(x[j], sum);
286 j += 1;
287 }
288 if !sum.is_finite() {
289 cold_path();
290 return Err(LaError::non_finite_computation_step(
291 ArithmeticOperation::LuSolve,
292 i,
293 ));
294 }
295 x[i] = sum;
296 i += 1;
297 }
298 }
299
300 let mut ii = 0;
302 while ii < D {
303 let i = D - 1 - ii;
304 let mut sum = x[i];
305 let row = self.factors.row(i);
306 let mut j = i + 1;
307 while j < D {
308 sum = (-row[j]).mul_add(x[j], sum);
309 j += 1;
310 }
311
312 let diag = row[i];
313 if !sum.is_finite() {
314 cold_path();
315 return Err(LaError::non_finite_computation_step(
316 ArithmeticOperation::LuSolve,
317 i,
318 ));
319 }
320
321 let quotient = sum / diag;
322 if !quotient.is_finite() {
323 cold_path();
324 return Err(LaError::non_finite_computation_step(
325 ArithmeticOperation::LuSolve,
326 i,
327 ));
328 }
329 x[i] = quotient;
330 ii += 1;
331 }
332
333 Vector::from_computation(x, ArithmeticOperation::LuSolve)
334 }
335
336 #[inline]
364 pub const fn det(&self) -> Result<f64, LaError> {
365 let mut det = if self.permutation.is_odd() { -1.0 } else { 1.0 };
366 let mut i = 0;
367
368 if D <= 4 {
369 while i < D {
372 let step = range_checked_product(det, self.factors.diag(i));
373 if !step.range_preserved() {
374 cold_path();
375 return self.scaled_det();
376 }
377 det = step.product();
378 i += 1;
379 }
380 return Ok(det);
381 }
382
383 let mut range_preserved = true;
384 while i < D {
385 let factor = self.factors.diag(i);
386 let step = range_checked_product(det, factor);
387 det = step.product();
388 range_preserved &= step.range_preserved();
389 i += 1;
390 }
391 if range_preserved {
392 Ok(det)
393 } else {
394 cold_path();
395 self.scaled_det()
396 }
397 }
398
399 #[cold]
401 const fn scaled_det(&self) -> Result<f64, LaError> {
402 let mut product = ScaledProduct::new(self.permutation.is_odd());
403 let mut i = 0;
404 while i < D {
405 product.multiply(self.factors.diag(i));
406 i += 1;
407 }
408
409 if let Some(det) = product.finish() {
410 Ok(det)
411 } else {
412 Err(LaError::non_finite_computation_step(
413 ArithmeticOperation::Determinant,
414 D.saturating_sub(1),
415 ))
416 }
417 }
418}
419
420#[cfg(test)]
421mod tests {
422 use core::hint::black_box;
423
424 use approx::assert_abs_diff_eq;
425 use pastey::paste;
426
427 use super::*;
428 use crate::DEFAULT_SINGULAR_TOL;
429
430 const TWO_NEG_800: f64 = f64::from_bits(223_u64 << 52);
431 const TWO_POS_800: f64 = f64::from_bits(1823_u64 << 52);
432
433 #[test]
434 fn row_permutation_keeps_mapping_and_parity_synchronized() {
435 let mut permutation = RowPermutation::<4>::identity();
436 assert_eq!(
437 core::array::from_fn(|row| permutation.source_row(row)),
438 [0, 1, 2, 3]
439 );
440 assert!(!permutation.is_odd());
441
442 permutation.swap(0, 3);
443 assert_eq!(
444 core::array::from_fn(|row| permutation.source_row(row)),
445 [3, 1, 2, 0]
446 );
447 assert!(permutation.is_odd());
448
449 permutation.swap(1, 2);
450 assert_eq!(
451 core::array::from_fn(|row| permutation.source_row(row)),
452 [3, 2, 1, 0]
453 );
454 assert!(!permutation.is_odd());
455 }
456
457 macro_rules! gen_pivoting_solve_and_det_tests {
458 ($d:literal) => {
459 paste! {
460 #[test]
461 fn [<lu_solve_pivoting_ $d d>]() {
462 let mut rows = [[0.0f64; $d]; $d];
468 for i in 0..$d {
469 rows[i][i] = 1.0;
470 }
471 rows.swap(0, 1);
472
473 let a = Matrix::<$d>::try_from_rows(black_box(rows)).unwrap();
474 let lu_fn: fn(Matrix<$d>, Tolerance) -> Result<Lu<$d>, LaError> =
475 black_box(Matrix::<$d>::lu);
476 let lu = lu_fn(a, DEFAULT_SINGULAR_TOL).unwrap();
477
478 let b_arr = {
480 let mut arr = [0.0f64; $d];
481 let mut val = 1.0f64;
482 for dst in arr.iter_mut() {
483 *dst = val;
484 val += 1.0;
485 }
486 arr
487 };
488 let mut expected = b_arr;
489 expected.swap(0, 1);
490 let b = Vector::<$d>::new(black_box(b_arr));
491
492 let solve_fn: fn(&Lu<$d>, Vector<$d>) -> Result<Vector<$d>, LaError> =
493 black_box(Lu::<$d>::solve);
494 let x = solve_fn(&lu, b).unwrap().into_array();
495
496 for i in 0..$d {
497 assert_abs_diff_eq!(x[i], expected[i], epsilon = 1e-12);
498 }
499 }
500
501 #[test]
502 fn [<lu_det_pivoting_ $d d>]() {
503 let mut rows = [[0.0f64; $d]; $d];
508 for i in 0..$d {
509 rows[i][i] = 1.0;
510 }
511 rows.swap(0, 1);
512
513 let a = Matrix::<$d>::try_from_rows(black_box(rows)).unwrap();
514 let lu_fn: fn(Matrix<$d>, Tolerance) -> Result<Lu<$d>, LaError> =
515 black_box(Matrix::<$d>::lu);
516 let lu = lu_fn(a, DEFAULT_SINGULAR_TOL).unwrap();
517
518 let det_fn: fn(&Lu<$d>) -> Result<f64, LaError> =
520 black_box(Lu::<$d>::det);
521 assert_abs_diff_eq!(det_fn(&lu).unwrap(), -1.0, epsilon = 1e-12);
522 }
523 }
524 };
525 }
526
527 gen_pivoting_solve_and_det_tests!(2);
528 gen_pivoting_solve_and_det_tests!(3);
529 gen_pivoting_solve_and_det_tests!(4);
530 gen_pivoting_solve_and_det_tests!(5);
531
532 macro_rules! gen_tridiagonal_smoke_solve_and_det_tests {
533 ($d:literal $(, #[$stack_array_expectation:meta])?) => {
534 paste! {
535 #[test]
536 fn [<lu_solve_tridiagonal_smoke_ $d d>]() {
537 $(#[$stack_array_expectation])?
542 let mut rows = [[0.0f64; $d]; $d];
543 for i in 0..$d {
544 rows[i][i] = 2.0;
545 if i > 0 {
546 rows[i][i - 1] = -1.0;
547 }
548 if i + 1 < $d {
549 rows[i][i + 1] = -1.0;
550 }
551 }
552
553 let a = Matrix::<$d>::try_from_rows(black_box(rows)).unwrap();
554 let lu_fn: fn(Matrix<$d>, Tolerance) -> Result<Lu<$d>, LaError> =
555 black_box(Matrix::<$d>::lu);
556 let lu = lu_fn(a, DEFAULT_SINGULAR_TOL).unwrap();
557
558 let mut b_arr = [0.0f64; $d];
560 b_arr[0] = 1.0;
561 b_arr[$d - 1] = 1.0;
562 let b = Vector::<$d>::new(black_box(b_arr));
563
564 let solve_fn: fn(&Lu<$d>, Vector<$d>) -> Result<Vector<$d>, LaError> =
565 black_box(Lu::<$d>::solve);
566 let x = solve_fn(&lu, b).unwrap().into_array();
567
568 for &x_i in &x {
569 assert_abs_diff_eq!(x_i, 1.0, epsilon = 1e-9);
570 }
571 }
572
573 #[test]
574 fn [<lu_det_tridiagonal_smoke_ $d d>]() {
575 $(#[$stack_array_expectation])?
581 let mut rows = [[0.0f64; $d]; $d];
582 for i in 0..$d {
583 rows[i][i] = 2.0;
584 if i > 0 {
585 rows[i][i - 1] = -1.0;
586 }
587 if i + 1 < $d {
588 rows[i][i + 1] = -1.0;
589 }
590 }
591
592 let a = Matrix::<$d>::try_from_rows(black_box(rows)).unwrap();
593 let lu_fn: fn(Matrix<$d>, Tolerance) -> Result<Lu<$d>, LaError> =
594 black_box(Matrix::<$d>::lu);
595 let lu = lu_fn(a, DEFAULT_SINGULAR_TOL).unwrap();
596
597 let det_fn: fn(&Lu<$d>) -> Result<f64, LaError> =
598 black_box(Lu::<$d>::det);
599 assert_abs_diff_eq!(det_fn(&lu).unwrap(), f64::from($d) + 1.0, epsilon = 1e-8);
600 }
601 }
602 };
603 }
604
605 gen_tridiagonal_smoke_solve_and_det_tests!(16);
606 gen_tridiagonal_smoke_solve_and_det_tests!(32);
607 gen_tridiagonal_smoke_solve_and_det_tests!(
608 64,
609 #[expect(
610 clippy::large_stack_arrays,
611 reason = "the test deliberately exercises the crate's stack-allocated matrix storage"
612 )]
613 );
614
615 #[test]
616 fn solve_0x0_returns_empty_vector_and_unit_det() {
617 let a = Matrix::<0>::zero();
618 let lu = a.lu(DEFAULT_SINGULAR_TOL).unwrap();
619
620 assert_eq!(lu.det(), Ok(1.0));
621 assert!(
622 lu.solve(Vector::<0>::zero())
623 .unwrap()
624 .into_array()
625 .is_empty()
626 );
627 }
628
629 #[test]
630 fn solve_1x1() {
631 let a = Matrix::<1>::try_from_rows(black_box([[2.0]])).unwrap();
632 let lu = a.lu(DEFAULT_SINGULAR_TOL).unwrap();
633
634 let b = Vector::<1>::new(black_box([6.0]));
635 let solve_fn: fn(&Lu<1>, Vector<1>) -> Result<Vector<1>, LaError> =
636 black_box(Lu::<1>::solve);
637 let x = solve_fn(&lu, b).unwrap().into_array();
638 assert_abs_diff_eq!(x[0], 3.0, epsilon = 1e-12);
639
640 let det_fn: fn(&Lu<1>) -> Result<f64, LaError> = black_box(Lu::<1>::det);
641 assert_abs_diff_eq!(det_fn(&lu).unwrap(), 2.0, epsilon = 0.0);
642 }
643
644 #[test]
645 fn solve_2x2_basic() {
646 let a = Matrix::<2>::try_from_rows(black_box([[1.0, 2.0], [3.0, 4.0]])).unwrap();
647 let lu = a.lu(DEFAULT_SINGULAR_TOL).unwrap();
648 let b = Vector::<2>::new(black_box([5.0, 11.0]));
649
650 let solve_fn: fn(&Lu<2>, Vector<2>) -> Result<Vector<2>, LaError> =
651 black_box(Lu::<2>::solve);
652 let x = solve_fn(&lu, b).unwrap().into_array();
653
654 assert_abs_diff_eq!(x[0], 1.0, epsilon = 1e-12);
655 assert_abs_diff_eq!(x[1], 2.0, epsilon = 1e-12);
656 }
657
658 #[test]
659 fn det_2x2_basic() {
660 let a = Matrix::<2>::try_from_rows(black_box([[1.0, 2.0], [3.0, 4.0]])).unwrap();
661 let lu = a.lu(DEFAULT_SINGULAR_TOL).unwrap();
662
663 let det_fn: fn(&Lu<2>) -> Result<f64, LaError> = black_box(Lu::<2>::det);
664 assert_abs_diff_eq!(det_fn(&lu).unwrap(), -2.0, epsilon = 1e-12);
665 }
666
667 #[test]
668 fn det_ordinary_factors_matches_direct_product_bits() {
669 let diagonal = [1.5, -2.0, 0.25, 8.0];
670 let mut rows = [[0.0; 4]; 4];
671 let mut expected = 1.0;
672 for (i, factor) in diagonal.into_iter().enumerate() {
673 rows[i][i] = factor;
674 expected *= factor;
675 }
676
677 let lu = Matrix::<4>::try_from_rows(rows)
678 .unwrap()
679 .lu(DEFAULT_SINGULAR_TOL)
680 .unwrap();
681 assert_eq!(lu.det().unwrap().to_bits(), expected.to_bits());
682 }
683
684 #[test]
685 fn singular_detected() {
686 let a = Matrix::<2>::try_from_rows(black_box([[1.0, 2.0], [2.0, 4.0]])).unwrap();
687 let err = a.lu(DEFAULT_SINGULAR_TOL).unwrap_err();
688 assert_eq!(
689 err,
690 LaError::singular_numerical(1, FactorizationKind::Lu, 0.0, DEFAULT_SINGULAR_TOL.get())
691 );
692 }
693
694 #[test]
695 fn singular_due_to_tolerance_at_first_pivot() {
696 let a = Matrix::<2>::try_from_rows(black_box([[1e-13, 0.0], [0.0, 1.0]])).unwrap();
698 let err = a.lu(DEFAULT_SINGULAR_TOL).unwrap_err();
699 assert_eq!(
700 err,
701 LaError::singular_numerical(
702 0,
703 FactorizationKind::Lu,
704 1e-13,
705 DEFAULT_SINGULAR_TOL.get()
706 )
707 );
708 }
709
710 #[test]
711 fn non_finite_detected_in_trailing_update() {
712 let a = Matrix::<3>::try_from_rows([
713 [1.0, f64::MAX, 0.0],
714 [-1.0, f64::MAX, 0.0],
715 [0.0, 0.0, 1.0],
716 ])
717 .unwrap();
718
719 let err = a.lu(DEFAULT_SINGULAR_TOL).unwrap_err();
720 assert_eq!(
721 err,
722 LaError::non_finite_computation_matrix(ArithmeticOperation::LuFactorization, 1, 1)
723 );
724 }
725
726 #[test]
727 fn generated_non_finite_takes_precedence_over_later_singular_pivot() {
728 let a = Matrix::<4>::try_from_rows([
731 [1.0, f64::MAX, 0.0, 0.0],
732 [1.0, f64::MAX, 0.0, 0.0],
733 [-1.0, f64::MAX, 0.0, 0.0],
734 [-1.0, f64::MAX, 0.0, 0.0],
735 ])
736 .unwrap();
737
738 let err = a.lu(DEFAULT_SINGULAR_TOL).unwrap_err();
739 assert_eq!(
740 err,
741 LaError::non_finite_computation_matrix(ArithmeticOperation::LuFactorization, 1, 1,)
742 );
743 }
744
745 #[test]
746 fn solve_non_finite_forward_substitution_overflow() {
747 let a = Matrix::<3>::try_from_rows([[1.0, 0.0, 0.0], [-1.0, 1.0, 0.0], [0.0, 0.0, 1.0]])
749 .unwrap();
750 let lu = a.lu(DEFAULT_SINGULAR_TOL).unwrap();
751
752 let b = Vector::<3>::new([1.0e308, 1.0e308, 0.0]);
753 let err = lu.solve(b).unwrap_err();
754 assert_eq!(
755 err,
756 LaError::non_finite_computation_step(ArithmeticOperation::LuSolve, 1)
757 );
758 }
759
760 #[test]
761 fn solve_non_finite_forward_substitution_overflow_fused_branch_5d() {
762 let a = Matrix::<5>::try_from_rows([
765 [1.0, 0.0, 0.0, 0.0, 0.0],
766 [-1.0, 1.0, 0.0, 0.0, 0.0],
767 [0.0, 0.0, 1.0, 0.0, 0.0],
768 [0.0, 0.0, 0.0, 1.0, 0.0],
769 [0.0, 0.0, 0.0, 0.0, 1.0],
770 ])
771 .unwrap();
772 let lu = a.lu(DEFAULT_SINGULAR_TOL).unwrap();
773
774 let b = Vector::<5>::new([1.0e308, 1.0e308, 0.0, 0.0, 0.0]);
775 let err = lu.solve(b).unwrap_err();
776 assert_eq!(
777 err,
778 LaError::non_finite_computation_step(ArithmeticOperation::LuSolve, 1)
779 );
780 }
781
782 #[test]
783 fn solve_non_finite_back_substitution_overflow() {
784 let a = Matrix::<2>::try_from_rows([[1.0, 1.0], [0.0, 2.0e-12]]).unwrap();
786 let lu = a.lu(DEFAULT_SINGULAR_TOL).unwrap();
787
788 let b = Vector::<2>::new([0.0, 1.0e300]);
789 let err = lu.solve(b).unwrap_err();
790 assert_eq!(
791 err,
792 LaError::non_finite_computation_step(ArithmeticOperation::LuSolve, 1)
793 );
794 }
795
796 #[test]
797 fn solve_non_finite_back_substitution_sum_overflow() {
798 let a = Matrix::<3>::try_from_rows([[1.0, 0.0, 0.0], [0.0, 1.0, 1.0e200], [0.0, 0.0, 1.0]])
805 .unwrap();
806 let lu = a.lu(DEFAULT_SINGULAR_TOL).unwrap();
807
808 let b = Vector::<3>::new([0.0, 0.0, 1.0e200]);
809 let err = lu.solve(b).unwrap_err();
810 assert_eq!(
811 err,
812 LaError::non_finite_computation_step(ArithmeticOperation::LuSolve, 1)
813 );
814 }
815
816 #[test]
817 fn det_rejects_product_overflow() {
818 let a = Matrix::<5>::try_from_rows([
819 [1.0e100, 0.0, 0.0, 0.0, 0.0],
820 [0.0, 1.0e100, 0.0, 0.0, 0.0],
821 [0.0, 0.0, 1.0e100, 0.0, 0.0],
822 [0.0, 0.0, 0.0, 1.0e100, 0.0],
823 [0.0, 0.0, 0.0, 0.0, 1.0e100],
824 ])
825 .unwrap();
826 let lu = a.lu(DEFAULT_SINGULAR_TOL).unwrap();
827 assert_eq!(
828 lu.det(),
829 Err(LaError::non_finite_computation_step(
830 ArithmeticOperation::Determinant,
831 4
832 ))
833 );
834 }
835
836 #[test]
837 fn det_balances_extreme_diagonals_independently_of_storage_order() {
838 let zero_tolerance = Tolerance::try_new(0.0).unwrap();
839 for diagonal in [
840 [TWO_NEG_800, TWO_NEG_800, TWO_POS_800, TWO_POS_800],
841 [TWO_POS_800, TWO_POS_800, TWO_NEG_800, TWO_NEG_800],
842 ] {
843 let mut rows = [[0.0; 4]; 4];
844 for (i, value) in diagonal.into_iter().enumerate() {
845 rows[i][i] = value;
846 }
847
848 let lu = Matrix::<4>::try_from_rows(rows)
849 .unwrap()
850 .lu(zero_tolerance)
851 .unwrap();
852 assert_eq!(lu.det(), Ok(1.0));
853 }
854 }
855
856 #[test]
857 fn matrix_det_fallback_inherits_balanced_extreme_accumulation() {
858 let zero_tolerance = Tolerance::try_new(0.0).unwrap();
859 for diagonal in [
860 [TWO_NEG_800, TWO_NEG_800, TWO_POS_800, TWO_POS_800, 1.0, 1.0],
861 [TWO_POS_800, TWO_POS_800, TWO_NEG_800, TWO_NEG_800, 1.0, 1.0],
862 ] {
863 let mut rows = [[0.0; 6]; 6];
864 for (i, value) in diagonal.into_iter().enumerate() {
865 rows[i][i] = value;
866 }
867
868 let matrix = Matrix::<6>::try_from_rows(rows).unwrap();
869 assert_eq!(matrix.det(), Ok(1.0));
870 assert_eq!(matrix.lu(zero_tolerance).unwrap().det(), Ok(1.0));
871 }
872 }
873
874 #[test]
875 fn det_rounds_final_tiny_magnitude_to_zero() {
876 let zero_tolerance = Tolerance::try_new(0.0).unwrap();
877 let positive =
878 Matrix::<2>::try_from_rows([[TWO_NEG_800, 0.0], [0.0, TWO_NEG_800]]).unwrap();
879 let positive_det = positive.lu(zero_tolerance).unwrap().det().unwrap();
880 assert_eq!(positive_det.to_bits(), 0.0f64.to_bits());
881
882 let negative =
883 Matrix::<2>::try_from_rows([[-TWO_NEG_800, 0.0], [0.0, TWO_NEG_800]]).unwrap();
884 let negative_det = negative.lu(zero_tolerance).unwrap().det().unwrap();
885 assert_eq!(negative_det.to_bits(), (-0.0f64).to_bits());
886 }
887
888 #[test]
898 fn lu_det_const_eval_d2() {
899 const DET: Result<f64, LaError> = {
900 let Ok(factors) = LuFactors::try_from_computation([[2.0, 0.0], [0.0, 3.0]]) else {
902 panic!("LU test factors must be finite");
903 };
904 let lu = Lu::<2> {
905 factors,
906 permutation: RowPermutation::identity(),
907 };
908 lu.det()
909 };
910 assert_eq!(DET, Ok(6.0));
911 }
912
913 #[test]
914 fn lu_det_const_eval_d3_row_swap() {
915 const DET: Result<f64, LaError> = {
916 let Ok(factors) = LuFactors::try_from_computation(Matrix::<3>::identity().into_rows())
919 else {
920 panic!("LU test factors must be usable");
921 };
922 let mut permutation = RowPermutation::identity();
923 permutation.swap(0, 1);
924 let lu = Lu::<3> {
925 factors,
926 permutation,
927 };
928 lu.det()
929 };
930 assert_eq!(DET, Ok(-1.0));
931 }
932
933 #[test]
934 fn lu_solve_const_eval_d2() {
935 const X: Result<Vector<2>, LaError> = {
937 let Ok(factors) = LuFactors::try_from_computation(Matrix::<2>::identity().into_rows())
938 else {
939 panic!("LU test factors must be usable");
940 };
941 let lu = Lu::<2> {
942 factors,
943 permutation: RowPermutation::identity(),
944 };
945 let b = Vector::<2>::new([1.0, 2.0]);
946 lu.solve(b)
947 };
948 let x = X.unwrap().into_array();
949 assert!((x[0] - 1.0).abs() <= 1e-12);
950 assert!((x[1] - 2.0).abs() <= 1e-12);
951 }
952}