use deep_causality_num::RealField;
pub(crate) fn solve_linear<T: RealField>(a: &mut [T], b: &mut [T], n: usize) {
for col in 0..n {
let mut pivot_row = col;
let mut best = a[col * n + col].abs();
for row in (col + 1)..n {
let mag = a[row * n + col].abs();
if mag > best {
best = mag;
pivot_row = row;
}
}
if pivot_row != col {
for j in 0..n {
a.swap(col * n + j, pivot_row * n + j);
}
b.swap(col, pivot_row);
}
let pivot = a[col * n + col];
for row in (col + 1)..n {
let factor = a[row * n + col] / pivot;
for j in (col + 1)..n {
let above = a[col * n + j];
a[row * n + j] -= factor * above;
}
let b_col = b[col];
b[row] -= factor * b_col;
}
}
for i in (0..n).rev() {
let mut s = b[i];
for j in (i + 1)..n {
s -= a[i * n + j] * b[j];
}
b[i] = s / a[i * n + i];
}
}
#[cfg(test)]
mod tests {
use super::solve_linear;
#[test]
fn solves_a_known_2x2_system() {
let mut a = vec![4.0_f64, 1.0, 1.0, 3.0];
let mut b = vec![1.0_f64, 2.0];
solve_linear(&mut a, &mut b, 2);
assert!((b[0] - 1.0 / 11.0).abs() < 1e-12);
assert!((b[1] - 7.0 / 11.0).abs() < 1e-12);
}
#[test]
fn solves_identity_to_itself() {
let mut a = vec![1.0_f64, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0];
let mut b = vec![3.0_f64, -2.0, 5.0];
solve_linear(&mut a, &mut b, 3);
assert_eq!(b, vec![3.0, -2.0, 5.0]);
}
#[test]
fn partial_pivoting_handles_a_zero_leading_pivot() {
let mut a = vec![0.0_f64, 1.0, 1.0, 0.0];
let mut b = vec![2.0_f64, 3.0];
solve_linear(&mut a, &mut b, 2);
assert!((b[0] - 3.0).abs() < 1e-12);
assert!((b[1] - 2.0).abs() < 1e-12);
}
#[test]
fn solves_an_extreme_scale_system() {
let mut a = vec![1e8_f64, 0.0, 0.0, 1e-4];
let mut b = vec![3.0_f64, 5.0];
solve_linear(&mut a, &mut b, 2);
assert!((b[0] - 3e-8).abs() < 1e-20);
assert!((b[1] - 5e4).abs() < 1e-6);
}
}