use log::debug;
use nalgebra::Scalar;
use yui_core::abst::{Ring, RingOps, Field, FieldOps};
use crate::{MatTrait, Perm};
use crate::dense::Mat;
pub struct Pluq<R> {
pub p: Perm,
pub q: Perm,
pub l: Mat<R>,
pub u: Mat<R>,
pub s: Mat<R>,
}
impl<R> Pluq<R> {
pub fn rank(&self) -> usize { self.l.n_cols() }
}
impl<R: Scalar> Pluq<R> {
pub fn transpose(self) -> Self {
let Pluq { p, q, l, u, s } = self;
Pluq { p: q, q: p, l: u.transpose(), u: l.transpose(), s: s.transpose() }
}
}
pub fn pluq<R>(a: &Mat<R>) -> Pluq<R>
where R: Ring, for<'x> &'x R: RingOps<R> {
debug!("compute dense pluq: {:?}", a.shape());
let (m, n) = a.shape();
let mut work = a.clone();
let mut col_of: Vec<usize> = (0..n).collect();
let (pivot_rows, u) = reduce(&mut work, &mut col_of);
let rank = pivot_rows.len();
let p = Perm::forward_indices(m, pivot_rows.iter().copied());
let q = Perm::from_indices(col_of).inv();
let p_inv = p.inv();
let l = build_l(&work, &p_inv, rank);
let s = build_s(&work, &p_inv, rank);
Pluq { p, q, l, u, s }
}
pub fn solve_pluq<R>(a: &Mat<R>, y: &[R]) -> Option<Vec<R>>
where R: Field, for<'x> &'x R: FieldOps<R> {
debug!("dense solve: {:?}", a.shape());
assert_eq!(y.len(), a.n_rows());
if y.iter().all(|yi| yi.is_zero()) {
return Some(vec![R::zero(); a.n_cols()]); }
let Pluq { p, q, l, u, .. } = pluq(a);
let yp = p.apply_to(y.to_vec());
debug!("forward sub: {:?}", l.shape());
let z = forward_sub(&l, &yp);
if !check_consistent(&l, &yp, &z) { return None; }
debug!("back sub: {:?}", u.shape());
let xp = back_sub(&u, &z);
Some((0..xp.len()).map(|j| xp[q.at(j)].clone()).collect())
}
fn forward_sub<R>(l: &Mat<R>, yp: &[R]) -> Vec<R>
where R: Field, for<'x> &'x R: FieldOps<R> {
(0..l.n_cols()).fold(vec![], |mut z, k| {
let pivot_inv = l[(k, k)].inv().unwrap();
let val = (0..k).fold(yp[k].clone(), |v, j| v - &l[(k, j)] * &z[j]) * pivot_inv;
z.push(val);
z
})
}
fn check_consistent<R>(l: &Mat<R>, yp: &[R], z: &[R]) -> bool
where R: Ring + PartialEq, for<'x> &'x R: RingOps<R> {
let rank = z.len();
(rank..yp.len()).all(|i| {
(0..rank).fold(R::zero(), |acc, j| acc + &l[(i, j)] * &z[j]) == yp[i]
})
}
fn back_sub<R>(u: &Mat<R>, z: &[R]) -> Vec<R>
where R: Field, for<'x> &'x R: FieldOps<R> {
let (rank, n) = (z.len(), u.n_cols());
(0..rank).rev().fold(vec![R::zero(); n], |mut xp, k| {
xp[k] = (k + 1..rank).fold(z[k].clone(), |v, j| v - &u[(k, j)] * &xp[j]);
xp
})
}
fn reduce<R>(work: &mut Mat<R>, col_of: &mut Vec<usize>) -> (Vec<usize>, Mat<R>)
where R: Ring, for<'x> &'x R: RingOps<R> {
let (m, n) = work.shape();
let mut pivot_rows = Vec::new();
let mut u_rows: Vec<Vec<R>> = Vec::new();
let mut c = 0;
for i in 0..m {
if c >= n { break; }
let Some(pivot_pos) = (c..n).find(|&j| work[(i, j)].is_unit()) else { continue; };
if pivot_pos != c {
work.swap_cols(c, pivot_pos);
col_of.swap(c, pivot_pos);
u_rows.iter_mut().for_each(|row| row.swap(c, pivot_pos));
}
let u_row = build_u_row(work, i, c);
eliminate_right(work, &u_row, c);
u_rows.push(u_row);
pivot_rows.push(i);
c += 1;
}
let rank = pivot_rows.len();
let u = Mat::generate((rank, n), |k, j| u_rows[k][j].clone());
(pivot_rows, u)
}
fn build_u_row<R>(work: &Mat<R>, i: usize, c: usize) -> Vec<R>
where R: Ring, for<'x> &'x R: RingOps<R> {
use std::cmp::Ordering::*;
let pivot_inv = work[(i, c)].inv().unwrap();
(0..work.n_cols()).map(|j| match j.cmp(&c) {
Less => R::zero(),
Equal => R::one(),
Greater => work[(i, j)].clone() * pivot_inv.clone(),
}).collect()
}
fn eliminate_right<R>(work: &mut Mat<R>, u_row: &[R], c: usize)
where R: Ring, for<'x> &'x R: RingOps<R> {
let n = work.n_cols();
(c + 1..n)
.filter(|&j| !u_row[j].is_zero())
.for_each(|j| work.add_col_to(c, j, &-u_row[j].clone()));
}
fn build_l<R>(work: &Mat<R>, p_inv: &Perm, rank: usize) -> Mat<R>
where R: Ring, for<'x> &'x R: RingOps<R> {
Mat::generate((work.n_rows(), rank), |i, k| work[(p_inv.at(i), k)].clone())
}
fn build_s<R>(work: &Mat<R>, p_inv: &Perm, rank: usize) -> Mat<R>
where R: Ring, for<'x> &'x R: RingOps<R> {
let (m, n) = work.shape();
Mat::generate((m - rank, n - rank), |i, j| work[(p_inv.at(i + rank), rank + j)].clone())
}
#[cfg(test)]
mod tests {
use num_traits::Zero;
use yui_core::num::Ratio;
use super::*;
#[test]
fn test_pluq_transpose() {
type R = Ratio<i64>;
let r = |n: i64| R::from(n);
let a = Mat::from_row_major((2, 3), [r(1),r(2),r(3),r(4),r(5),r(6)]);
let at = a.transpose();
let dp = pluq(&a).transpose();
let rank = dp.rank();
let (m, n) = at.shape(); assert_eq!(dp.l.shape(), (m, rank));
assert_eq!(dp.u.shape(), (rank, n));
let paq = apply_perms(&at, &dp.p, &dp.q);
let rem_full = Mat::generate((m, n), |i, j| {
if i >= rank && j >= rank { dp.s[(i - rank, j - rank)] } else { R::zero() }
});
assert_eq!(paq, &dp.l * &dp.u + &rem_full);
}
type R = Ratio<i64>;
fn r(n: i64) -> R { R::from(n) }
fn rf(n: i64, d: i64) -> R { R::new(n, d) }
fn sample() -> Mat<R> {
Mat::from_row_major((3, 4), [
r(1), r(2), r(3), r(4),
r(2), r(4), r(5), r(6),
r(3), r(6), r(7), r(8),
])
}
fn apply_perms(a: &Mat<R>, p: &Perm, q: &Perm) -> Mat<R> {
let (m, n) = a.shape();
let mut out = Mat::zero((m, n));
for i in 0..m {
for j in 0..n {
out[(p.at(i), q.at(j))] = a[(i, j)];
}
}
out
}
fn check(a: &Mat<R>) -> Pluq<R> {
let (m, n) = a.shape();
let pp = pluq(a);
let rank = pp.rank();
assert_eq!(pp.l.shape(), (m, rank), "L shape");
assert_eq!(pp.u.shape(), (rank, n), "U shape");
assert_eq!(pp.s.shape(), (m - rank, n - rank), "s shape");
for k in 0..rank {
for i in 0..k {
assert_eq!(pp.l[(i, k)], r(0), "L[{i},{k}] should be 0 (above diagonal)");
}
}
for k in 0..rank {
assert_eq!(pp.u[(k, k)], r(1), "U[{k},{k}] should be 1");
for i in (k + 1)..rank {
assert_eq!(pp.u[(i, k)], r(0), "U[{i},{k}] should be 0 (below diagonal)");
}
}
let paq = apply_perms(a, &pp.p, &pp.q);
let rem_full = Mat::generate((m, n), |i, j| {
if i >= rank && j >= rank { pp.s[(i - rank, j - rank)] } else { R::zero() }
});
assert_eq!(paq, &pp.l * &pp.u + &rem_full, "p*A*q should equal L*U + s");
pp
}
#[test]
fn test_sample() {
let pp = check(&sample());
assert_eq!(pp.rank(), 2);
assert!(pp.s.is_zero());
}
#[test]
fn test_zero() {
let pp = check(&Mat::<R>::zero((3, 4)));
assert_eq!(pp.rank(), 0);
assert!(pp.s.is_zero());
}
#[test]
fn test_identity() {
let pp = check(&Mat::id(3));
assert_eq!(pp.rank(), 3);
assert!(pp.s.is_zero());
}
#[test]
fn test_full_row_rank() {
let a = Mat::from_row_major((2, 3), [
r(1), r(0), r(2),
r(0), r(1), r(3),
]);
let pp = check(&a);
assert_eq!(pp.rank(), 2);
assert!(pp.s.is_zero());
}
#[test]
fn test_full_col_rank() {
let a = Mat::from_row_major((3, 2), [
r(1), r(2),
r(3), r(4),
r(5), r(6),
]);
let pp = check(&a);
assert_eq!(pp.rank(), 2);
assert!(pp.s.is_zero());
}
#[test]
fn test_rank_deficient_cols() {
let a = Mat::from_row_major((3, 3), [
r(1), r(0), r(2),
r(2), r(1), r(4),
r(3), r(2), r(6),
]);
let pp = check(&a);
assert_eq!(pp.rank(), 2);
assert!(pp.s.is_zero());
}
#[test]
fn test_pivot_not_in_first_col() {
let a = Mat::from_row_major((2, 3), [
r(0), r(1), r(2),
r(0), r(3), r(4),
]);
let pp = check(&a);
assert_eq!(pp.rank(), 2);
assert!(pp.q.at(0) >= pp.rank(), "col 0 is non-pivot");
assert!(pp.s.is_zero());
}
#[test]
fn test_col_swap() {
let a = Mat::from_row_major((3, 3), [
r(0), r(1), r(2),
r(1), r(0), r(3),
r(2), r(1), r(4),
]);
let pp = check(&a);
assert_eq!(pp.rank(), 3);
assert_eq!(pp.q.at(1), 0, "original col 1 should move to position 0");
assert!(pp.s.is_zero());
}
#[test]
fn test_fractions() {
let a = Mat::from_row_major((2, 2), [
rf(1, 2), rf(1, 3),
rf(1, 4), rf(1, 5),
]);
let pp = check(&a);
assert_eq!(pp.rank(), 2);
assert!(pp.s.is_zero());
}
#[test]
fn test_single_row() {
let a = Mat::from_row_major((1, 4), [r(0), r(2), r(0), r(3)]);
let pp = check(&a);
assert_eq!(pp.rank(), 1);
assert!(pp.s.is_zero());
}
#[test]
fn test_single_col() {
let a = Mat::from_row_major((3, 1), [r(2), r(0), r(4)]);
let pp = check(&a);
assert_eq!(pp.rank(), 1);
assert!(pp.s.is_zero());
}
#[test]
fn test_ring_nonzero_rem() {
let a = Mat::<i32>::from_row_major((2, 2), [2, 3, 1, 4]);
let pp = pluq(&a);
assert_eq!(pp.rank(), 1);
assert_eq!(pp.l.shape(), (2, 1));
assert_eq!(pp.u.shape(), (1, 2));
assert_eq!(pp.s.shape(), (1, 1));
let paq: Mat<i32> = {
let (m, n) = a.shape();
let mut out = Mat::zero((m, n));
for i in 0..m { for j in 0..n { out[(pp.p.at(i), pp.q.at(j))] = a[(i, j)]; } }
out
};
let rem_full = Mat::generate((2, 2), |i, j| {
if i >= 1 && j >= 1 { pp.s[(i - 1, j - 1)] } else { 0 }
});
assert_eq!(paq, &pp.l * &pp.u + &rem_full);
assert!(!pp.s.is_zero());
}
fn solve_check(a: &Mat<R>, y: &[R]) -> Vec<R> {
let x = solve_pluq(a, y).expect("expected a solution");
let (m, n) = a.shape();
assert_eq!(x.len(), n);
for i in 0..m {
let ax_i: R = (0..n).fold(R::zero(), |acc, j| acc + a[(i, j)] * x[j]);
assert_eq!(ax_i, y[i], "row {i}: (A*x)[{i}] != y[{i}]");
}
x
}
#[test]
fn test_solve_square_full_rank() {
let a = Mat::from_row_major((2, 2), [r(1), r(2), r(3), r(4)]);
let y = vec![r(5), r(6)];
solve_check(&a, &y);
}
#[test]
fn test_solve_overdetermined_consistent() {
let a = Mat::from_row_major((3, 2), [r(1), r(0), r(0), r(1), r(1), r(1)]);
let y = vec![r(2), r(3), r(5)]; solve_check(&a, &y);
}
#[test]
fn test_solve_overdetermined_inconsistent() {
let a = Mat::from_row_major((3, 2), [r(1), r(0), r(0), r(1), r(1), r(1)]);
let y = vec![r(1), r(1), r(0)]; assert!(solve_pluq(&a, &y).is_none());
}
#[test]
fn test_solve_underdetermined() {
let a = Mat::from_row_major((2, 3), [r(1), r(0), r(2), r(0), r(1), r(3)]);
let y = vec![r(4), r(5)];
solve_check(&a, &y);
}
#[test]
fn test_solve_zero_rhs() {
let a = Mat::from_row_major((2, 2), [r(1), r(2), r(3), r(4)]);
let y = vec![r(0), r(0)];
let x = solve_check(&a, &y);
assert_eq!(x, vec![r(0), r(0)]);
}
#[test]
fn test_solve_no_solution_rank_deficient() {
let a = Mat::from_row_major((2, 2), [r(1), r(2), r(2), r(4)]);
let y = vec![r(1), r(0)]; assert!(solve_pluq(&a, &y).is_none());
}
}