use crate::polar_brannon::polar_rotation_brannon;
use crate::polar_classic::{polar_decomp_eigen, polar_decomp_svd};
use crate::polar_higham::polar_quaternion_higham;
use crate::{Tensor2, t2_gen_dot_gen_tra_chop, t2_gen_tra_dot_gen_chop};
use russell_lab::StrError;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum PolarAlgo {
Eigen,
SVD,
Iterative,
Quaternion,
}
#[inline]
pub fn polar_decomp(
rr: &mut Tensor2<9>,
uu: &mut Tensor2<6>,
vv: Option<&mut Tensor2<6>>,
ff: &Tensor2<9>,
) -> Result<usize, StrError> {
polar_decomp_mx(rr, uu, vv, PolarAlgo::Quaternion, ff)
}
pub fn polar_decomp_mx(
rr: &mut Tensor2<9>,
uu: &mut Tensor2<6>,
vv: Option<&mut Tensor2<6>>,
algo: PolarAlgo,
ff: &Tensor2<9>,
) -> Result<usize, StrError> {
let nit = match algo {
PolarAlgo::Eigen => {
polar_decomp_eigen(rr, uu, ff)?; 0
}
PolarAlgo::SVD => {
polar_decomp_svd(rr, uu, ff)?; 0
}
PolarAlgo::Iterative => {
let nit = polar_rotation_brannon(rr, ff)?;
t2_gen_tra_dot_gen_chop(uu.as_mut_vec(), 1.0, rr.as_vec(), ff.as_vec()); nit
}
PolarAlgo::Quaternion => {
polar_quaternion_higham(rr, uu, ff)?; 0
}
};
if let Some(v) = vv {
t2_gen_dot_gen_tra_chop(v.as_mut_vec(), 1.0, ff.as_vec(), rr.as_vec());
}
Ok(nit)
}
#[cfg(test)]
mod tests {
use super::{PolarAlgo, polar_decomp, polar_decomp_mx};
use crate::Tensor2;
use crate::testing::{ReferencePolarDecomp, check_agree, check_polar};
use russell_lab::{Matrix, mat_approx_eq, mat_mat_mul};
#[test]
fn polar_decomp_default_works() {
let ff = ReferencePolarDecomp::example03();
let mut rr = Tensor2::<9>::new();
let mut uu = Tensor2::<6>::new();
let mut vv = Tensor2::<6>::new();
let _ = polar_decomp(&mut rr, &mut uu, Some(&mut vv), &ff).unwrap();
check_polar(&ff, &rr, &uu, 1e-13);
}
#[test]
fn polar_decomp_brannon_works() {
let ff = ReferencePolarDecomp::example03();
let mut rr = Tensor2::<9>::new();
let mut uu = Tensor2::<6>::new();
let mut vv = Tensor2::<6>::new();
let nit = polar_decomp_mx(&mut rr, &mut uu, Some(&mut vv), PolarAlgo::Iterative, &ff).unwrap();
assert!(nit > 0);
check_polar(&ff, &rr, &uu, 1e-13);
let f = ff.as_std_matrix();
let r = rr.as_std_matrix();
let v = vv.as_std_matrix();
let mut vr = Matrix::new(3, 3);
mat_mat_mul(&mut vr, 1.0, &v, &r, 0.0).unwrap();
mat_approx_eq(&vr, &f, 1e-13);
mat_approx_eq(&r, &ReferencePolarDecomp::example03_rotation(), 1e-3);
mat_approx_eq(&uu.as_std_matrix(), &ReferencePolarDecomp::example03_stretch(), 1e-3);
}
#[test]
fn polar_decomp_brannon_on_higham_cases() {
check_agree(&ReferencePolarDecomp::case51());
for y in [1.0f64, 1e-2, 1e-4, 1e-6, 1e-8] {
let a = ReferencePolarDecomp::case52(y);
let mut rr = Tensor2::<9>::new();
let mut uu = Tensor2::<6>::new();
let mut vv = Tensor2::<6>::new();
polar_decomp_mx(&mut rr, &mut uu, Some(&mut vv), PolarAlgo::Iterative, &a).unwrap();
let tol = if y == 1.0 { 1e-13 } else { 1e-8 };
check_polar(&a, &rr, &uu, tol);
if y == 1.0 {
check_agree(&a);
}
}
}
#[test]
fn polar_decomp_higham_algo_works() {
let a = ReferencePolarDecomp::case51();
let mut rr = Tensor2::<9>::new();
let mut uu = Tensor2::<6>::new();
let nit = polar_decomp_mx(&mut rr, &mut uu, None, PolarAlgo::Quaternion, &a).unwrap();
assert_eq!(nit, 0); check_polar(&a, &rr, &uu, 1e-13);
}
#[test]
fn polar_decomp_eigen_works() {
let ff = ReferencePolarDecomp::example03();
let mut rr = Tensor2::<9>::new();
let mut uu = Tensor2::<6>::new();
let mut vv = Tensor2::<6>::new();
let nit = polar_decomp_mx(&mut rr, &mut uu, Some(&mut vv), PolarAlgo::Eigen, &ff).unwrap();
assert_eq!(nit, 0);
check_polar(&ff, &rr, &uu, 1e-13);
let f = ff.as_std_matrix();
let r = rr.as_std_matrix();
let v = vv.as_std_matrix();
let mut vr = Matrix::new(3, 3);
mat_mat_mul(&mut vr, 1.0, &v, &r, 0.0).unwrap();
mat_approx_eq(&vr, &f, 1e-13);
mat_approx_eq(&r, &ReferencePolarDecomp::example03_rotation(), 1e-3);
mat_approx_eq(&uu.as_std_matrix(), &ReferencePolarDecomp::example03_stretch(), 1e-3);
}
#[test]
fn polar_decomp_svd_works() {
let ff = ReferencePolarDecomp::example03();
let mut rr = Tensor2::<9>::new();
let mut uu = Tensor2::<6>::new();
let mut vv = Tensor2::<6>::new();
let nit = polar_decomp_mx(&mut rr, &mut uu, Some(&mut vv), PolarAlgo::SVD, &ff).unwrap();
assert_eq!(nit, 0);
check_polar(&ff, &rr, &uu, 1e-13);
let f = ff.as_std_matrix();
let r = rr.as_std_matrix();
let v = vv.as_std_matrix();
let mut vr = Matrix::new(3, 3);
mat_mat_mul(&mut vr, 1.0, &v, &r, 0.0).unwrap();
mat_approx_eq(&vr, &f, 1e-13);
mat_approx_eq(&r, &ReferencePolarDecomp::example03_rotation(), 1e-3);
mat_approx_eq(&uu.as_std_matrix(), &ReferencePolarDecomp::example03_stretch(), 1e-3);
}
#[test]
fn polar_decomp_eigen_on_higham_cases() {
let a = ReferencePolarDecomp::case51();
let mut r_e = Tensor2::<9>::new();
let mut u_e = Tensor2::<6>::new();
polar_decomp_mx(&mut r_e, &mut u_e, None, PolarAlgo::Eigen, &a).unwrap();
check_polar(&a, &r_e, &u_e, 1e-13);
let mut r_h = Tensor2::<9>::new();
let mut u_h = Tensor2::<6>::new();
polar_decomp_mx(&mut r_h, &mut u_h, None, PolarAlgo::Quaternion, &a).unwrap();
mat_approx_eq(&r_e.as_std_matrix(), &r_h.as_std_matrix(), 1e-13);
mat_approx_eq(&u_e.as_std_matrix(), &u_h.as_std_matrix(), 1e-13);
let a = ReferencePolarDecomp::case52(1.0);
let mut r_e = Tensor2::<9>::new();
let mut u_e = Tensor2::<6>::new();
polar_decomp_mx(&mut r_e, &mut u_e, None, PolarAlgo::Eigen, &a).unwrap();
check_polar(&a, &r_e, &u_e, 1e-13);
}
#[test]
fn polar_decomp_svd_on_higham_cases() {
let a = ReferencePolarDecomp::case51();
let mut r_s = Tensor2::<9>::new();
let mut u_s = Tensor2::<6>::new();
polar_decomp_mx(&mut r_s, &mut u_s, None, PolarAlgo::SVD, &a).unwrap();
check_polar(&a, &r_s, &u_s, 1e-13);
let mut r_h = Tensor2::<9>::new();
let mut u_h = Tensor2::<6>::new();
polar_decomp_mx(&mut r_h, &mut u_h, None, PolarAlgo::Quaternion, &a).unwrap();
mat_approx_eq(&r_s.as_std_matrix(), &r_h.as_std_matrix(), 1e-13);
mat_approx_eq(&u_s.as_std_matrix(), &u_h.as_std_matrix(), 1e-13);
for y in [1.0f64, 1e-2, 1e-4, 1e-6, 1e-8] {
let a = ReferencePolarDecomp::case52(y);
let tol = if y == 1.0 { 1e-13 } else { 1e-8 };
let mut r_s = Tensor2::<9>::new();
let mut u_s = Tensor2::<6>::new();
polar_decomp_mx(&mut r_s, &mut u_s, None, PolarAlgo::SVD, &a).unwrap();
check_polar(&a, &r_s, &u_s, tol);
}
}
}