use crate::EigenProjsT2;
use super::Tensor2;
use russell_lab::{StrError, small_mat_inv, small_mat_mat_mul, small_mat_svd, small_mat_t_mat_mul};
pub(crate) fn polar_decomp_eigen(rr: &mut Tensor2<9>, uu: &mut Tensor2<6>, ff: &Tensor2<9>) -> Result<(), StrError> {
let mut f = [[0.0; 3]; 3];
ff.to_std_matrix_slice(&mut f);
let mut cc = [[0.0; 3]; 3];
small_mat_t_mat_mul(&mut cc, 1.0, &f, &f, 0.0, 3);
let c = Tensor2::<6>::from_std_matrix(&cc)?;
let mut eig = EigenProjsT2::new();
let mut ll = [0.0; 3];
let mut projs = [Tensor2::<6>::new(), Tensor2::<6>::new(), Tensor2::<6>::new()];
eig.calculate(&mut ll, &mut projs, &c)?;
let sqrt_l0 = ll[0].sqrt();
let sqrt_l1 = ll[1].sqrt();
let sqrt_l2 = ll[2].sqrt();
for m in 0..6 {
uu.vec[m] = sqrt_l0 * projs[0].vec[m] + sqrt_l1 * projs[1].vec[m] + sqrt_l2 * projs[2].vec[m];
}
let mut u3 = [[0.0; 3]; 3];
uu.to_std_matrix_slice(&mut u3);
let mut ui = [[0.0; 3]; 3];
small_mat_inv(&mut ui, &u3, 3)?;
let mut r = [[0.0; 3]; 3];
small_mat_mat_mul(&mut r, 1.0, &f, &ui, 0.0, 3);
rr.set_std_matrix(&r)?;
Ok(())
}
pub(crate) fn polar_decomp_svd(rr: &mut Tensor2<9>, uu: &mut Tensor2<6>, ff: &Tensor2<9>) -> Result<(), StrError> {
let mut f = [[0.0; 3]; 3];
ff.to_std_matrix_slice(&mut f);
let mut s = [0.0; 3];
let mut p = [[0.0; 3]; 3];
let mut qt = [[0.0; 3]; 3]; small_mat_svd(&mut s, &mut p, &mut qt, &f)?;
let mut u = [[0.0; 3]; 3];
for i in 0..3 {
for j in i..3 {
let mut sum = 0.0;
for k in 0..3 {
sum += qt[k][i] * s[k] * qt[k][j];
}
u[i][j] = sum;
u[j][i] = sum;
}
}
let mut r = [[0.0; 3]; 3];
small_mat_mat_mul(&mut r, 1.0, &p, &qt, 0.0, 3);
rr.set_std_matrix(&r)?;
uu.set_std_matrix(&u)?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::{polar_decomp_eigen, polar_decomp_svd};
use crate::Tensor2;
use russell_lab::{Matrix, mat_approx_eq, mat_mat_mul, mat_t_mat_mul};
const TOL: f64 = 1e-12;
fn check_polar(ff: &Tensor2<9>, rr: &Tensor2<9>, uu: &Tensor2<6>, tol: f64) {
let f = ff.as_std_matrix();
let r = rr.as_std_matrix();
let u = uu.as_std_matrix();
let mut ru = Matrix::new(3, 3);
mat_mat_mul(&mut ru, 1.0, &r, &u, 0.0).unwrap();
mat_approx_eq(&ru, &f, tol);
let mut rtr = Matrix::new(3, 3);
mat_t_mat_mul(&mut rtr, 1.0, &r, &r, 0.0).unwrap();
mat_approx_eq(&rtr, &Matrix::identity(3), tol);
for i in 0..3 {
for j in 0..3 {
assert!((u.get(i, j) - u.get(j, i)).abs() < tol);
}
}
}
fn check_ref(ff: &Tensor2<9>, r_ref: &[[f64; 3]; 3], u_ref: &[[f64; 3]; 3], tol: f64) {
let mut rr = Tensor2::<9>::new();
let mut uu = Tensor2::<6>::new();
polar_decomp_eigen(&mut rr, &mut uu, ff).unwrap();
check_polar(ff, &rr, &uu, tol);
mat_approx_eq(&rr.as_std_matrix(), r_ref, tol);
mat_approx_eq(&uu.as_std_matrix(), u_ref, tol);
let mut rr = Tensor2::<9>::new();
let mut uu = Tensor2::<6>::new();
polar_decomp_svd(&mut rr, &mut uu, ff).unwrap();
check_polar(ff, &rr, &uu, tol);
mat_approx_eq(&rr.as_std_matrix(), r_ref, tol);
mat_approx_eq(&uu.as_std_matrix(), u_ref, tol);
}
#[test]
fn polar_decomp_classic_works() {
#[rustfmt::skip]
let ff = Tensor2::<9>::from_std_matrix(&[
[ 1.0, 0.495, 0.5 ],
[-0.333, 1.0, -0.247],
[ 0.959, 0.0, 1.5 ],
]).unwrap();
#[rustfmt::skip]
let r_ref = [
[ 0.9143288659766733, 0.3769304925214390, -0.1480747400786342],
[-0.3738918867839085, 0.9261806148757240, 0.0489318467421089],
[ 0.1555878589060806, 0.0106241440111795, 0.9877649243241274],
];
#[rustfmt::skip]
let u_ref = [
[1.1880436209666460, 0.0787009018745443, 0.7828975173830820],
[0.0787009018745443, 1.1127612086738370, -0.0243651495968139],
[0.7828975173830820, -0.0243651495968139, 1.3955238503015750],
];
check_ref(&ff, &r_ref, &u_ref, TOL);
#[rustfmt::skip]
let ff = Tensor2::<9>::from_std_matrix(&[
[2.0, 1.0, 0.5],
[0.3, 3.0, 1.0],
[0.7, -0.2, 2.5],
]).unwrap();
#[rustfmt::skip]
let r_ref = [
[ 0.9812167785042613, 0.1797641567869235, -0.0699891528481767],
[-0.1577488954874382, 0.9565376734603720, 0.2452569371567530],
[ 0.1110356679369841, -0.2296095102248696, 0.9669284116521154],
];
#[rustfmt::skip]
let u_ref = [
[1.9928338559181810, 0.4857629584545492, 0.6104486636071538],
[0.4857629584545494, 3.0952990792130150, 0.4723959762916591],
[0.6104486636071539, 0.4723959762916588, 2.6275833898629540],
];
check_ref(&ff, &r_ref, &u_ref, TOL);
#[rustfmt::skip]
let ff = Tensor2::<9>::from_std_matrix(&[
[1.0, 0.5, 0.0],
[0.2, 1.0, 0.3],
[0.0, -0.4, -1.0],
]).unwrap();
#[rustfmt::skip]
let r_ref = [
[ 0.9883141202500686, 0.1508649447384873, -0.0217937643234417],
[-0.1517146219954520, 0.9874085229738411, -0.0448004713300558],
[-0.0147605280091824, -0.0475833711255416, -0.9987582037736755],
];
#[rustfmt::skip]
let u_ref = [
[ 0.9579711958509785, 0.3483466493332560, -0.0307538585894529],
[ 0.3483466493332560, 1.0818743437933010, 0.3438059280176936],
[-0.0307538585894528, 0.3438059280176936, 0.9853180623746591],
];
check_ref(&ff, &r_ref, &u_ref, TOL);
}
}