use approx::assert_abs_diff_eq;
use ferrolearn_core::traits::{Fit, Transform};
use ferrolearn_decomp::PCA;
use ndarray::{Array2, array};
fn fixture_a() -> Array2<f64> {
array![
[1.0, 0.1, 8.0],
[2.0, 0.4, 1.0],
[3.0, 0.0, 5.0],
[4.0, 0.9, 2.0],
[5.0, 0.2, 9.0],
[6.0, 0.7, 0.5],
]
}
#[allow(
clippy::excessive_precision,
reason = "live sklearn 1.5.2 oracle (R-CHAR-3)"
)]
const SK_COMPONENTS: [[f64; 3]; 3] = [
[-0.161336219074942, -0.070935617130925, 0.984346871096189],
[0.984685837730652, 0.055197220315576, 0.165369488848810],
[-0.066063797856960, 0.995952511464120, 0.060943986750344],
];
#[allow(
clippy::excessive_precision,
reason = "live sklearn 1.5.2 oracle (R-CHAR-3)"
)]
const SK_TRANSFORM: [[f64; 3]; 6] = [
[4.114739739151824, -1.857218223566339, 0.111512900041358],
[-2.958305262735716, -2.013559641682686, -0.082373051628776],
[0.846120249426467, -0.389474736683023, -0.303041907070008],
[-2.332098638354874, 0.148780132785217, 0.344419595139708],
[4.446648172235152, 2.252414338236639, 0.007796946510272],
[-4.117104259722856, 1.859058130910190, -0.078314482992554],
];
#[allow(
clippy::excessive_precision,
reason = "live sklearn 1.5.2 oracle (R-CHAR-3)"
)]
const SK_EXPLAINED_VARIANCE: [f64; 3] = [13.712096827412307, 3.241395108836682, 0.047174730417675];
#[allow(
clippy::excessive_precision,
reason = "live sklearn 1.5.2 oracle (R-CHAR-3)"
)]
const SK_EXPLAINED_VARIANCE_RATIO: [f64; 3] =
[0.806562301130092, 0.190662823546332, 0.002774875323576];
#[allow(
clippy::excessive_precision,
reason = "live sklearn 1.5.2 oracle (R-CHAR-3)"
)]
const SK_SINGULAR_VALUES: [f64; 3] = [8.280125852730835, 4.025788810181603, 0.485668253119733];
#[test]
fn divergence_components_sign_value_parity() {
let x = fixture_a();
let fitted = PCA::<f64>::new(3)
.fit(&x, &())
.expect("fit must succeed on non-degenerate fixture");
let c = fitted.components();
for (i, sk_row) in SK_COMPONENTS.iter().enumerate() {
for (j, &sk_val) in sk_row.iter().enumerate() {
assert_abs_diff_eq!(c[[i, j]], sk_val, epsilon = 1e-6);
}
}
}
#[test]
fn divergence_transform_sign_value_parity() {
let x = fixture_a();
let fitted = PCA::<f64>::new(3)
.fit(&x, &())
.expect("fit must succeed on non-degenerate fixture");
let projected = fitted
.transform(&x)
.expect("transform must succeed on training data");
for (i, sk_row) in SK_TRANSFORM.iter().enumerate() {
for (j, &sk_val) in sk_row.iter().enumerate() {
assert_abs_diff_eq!(projected[[i, j]], sk_val, epsilon = 1e-6);
}
}
}
#[test]
fn divergence_components_sign_convention_max_abs_positive() {
let x = fixture_a();
let fitted = PCA::<f64>::new(3)
.fit(&x, &())
.expect("fit must succeed on non-degenerate fixture");
let c = fitted.components();
for i in 0..c.nrows() {
let row = c.row(i);
let mut max_j = 0usize;
let mut max_abs = 0.0f64;
for (j, &v) in row.iter().enumerate() {
if v.abs() > max_abs {
max_abs = v.abs();
max_j = j;
}
}
assert!(
row[max_j] > 0.0,
"svd_flip invariant violated: component row {i} max-abs entry \
(idx {max_j}) = {} is not positive",
row[max_j]
);
}
}
#[test]
fn green_explained_variance_matches_sklearn() {
let x = fixture_a();
let fitted = PCA::<f64>::new(3).fit(&x, &()).expect("fit");
let ev = fitted.explained_variance();
for (k, &sk) in SK_EXPLAINED_VARIANCE.iter().enumerate() {
assert_abs_diff_eq!(ev[k], sk, epsilon = 1e-6);
}
}
#[test]
fn green_explained_variance_ratio_matches_sklearn() {
let x = fixture_a();
let fitted = PCA::<f64>::new(3).fit(&x, &()).expect("fit");
let ratio = fitted.explained_variance_ratio();
for (k, &sk) in SK_EXPLAINED_VARIANCE_RATIO.iter().enumerate() {
assert_abs_diff_eq!(ratio[k], sk, epsilon = 1e-6);
}
let sum: f64 = ratio.iter().sum();
assert!(sum <= 1.0 + 1e-9, "ratio sum {sum} exceeds 1");
assert_abs_diff_eq!(sum, 1.0, epsilon = 1e-9);
}
#[test]
fn green_explained_variance_ratio_partial_le_1() {
let x = fixture_a();
let fitted = PCA::<f64>::new(2).fit(&x, &()).expect("fit");
let sum: f64 = fitted.explained_variance_ratio().iter().sum();
assert!(sum <= 1.0 + 1e-9 && sum > 0.0, "partial ratio sum = {sum}");
assert_abs_diff_eq!(
sum,
SK_EXPLAINED_VARIANCE_RATIO[0] + SK_EXPLAINED_VARIANCE_RATIO[1],
epsilon = 1e-6
);
}
#[test]
fn green_singular_values_matches_sklearn() {
let x = fixture_a();
let fitted = PCA::<f64>::new(3).fit(&x, &()).expect("fit");
let sv = fitted.singular_values();
for (k, &sk) in SK_SINGULAR_VALUES.iter().enumerate() {
assert_abs_diff_eq!(sv[k], sk, epsilon = 1e-6);
}
}
#[test]
fn green_components_orthonormal() {
let x = fixture_a();
let fitted = PCA::<f64>::new(3).fit(&x, &()).expect("fit");
let c = fitted.components();
for i in 0..c.nrows() {
let norm: f64 = c.row(i).iter().map(|v| v * v).sum::<f64>().sqrt();
assert_abs_diff_eq!(norm, 1.0, epsilon = 1e-8);
}
for i in 0..c.nrows() {
for j in (i + 1)..c.nrows() {
let dot: f64 = c
.row(i)
.iter()
.zip(c.row(j).iter())
.map(|(a, b)| a * b)
.sum();
assert_abs_diff_eq!(dot, 0.0, epsilon = 1e-8);
}
}
}
#[test]
fn green_inverse_transform_roundtrip_exact() {
let x = fixture_a();
let fitted = PCA::<f64>::new(3).fit(&x, &()).expect("fit");
let projected = fitted.transform(&x).expect("transform");
let recovered = fitted
.inverse_transform(&projected)
.expect("inverse_transform");
for (a, b) in x.iter().zip(recovered.iter()) {
assert_abs_diff_eq!(a, b, epsilon = 1e-8);
}
}
#[test]
fn green_fit_is_deterministic() {
let x = fixture_a();
let a = PCA::<f64>::new(3).fit(&x, &()).expect("fit a");
let b = PCA::<f64>::new(3).fit(&x, &()).expect("fit b");
let ca = a.components();
let cb = b.components();
assert_eq!(ca.dim(), cb.dim());
for (va, vb) in ca.iter().zip(cb.iter()) {
assert_abs_diff_eq!(va, vb, epsilon = 0.0);
}
}
#[test]
fn green_err_n_components_zero() {
let x = fixture_a();
assert!(PCA::<f64>::new(0).fit(&x, &()).is_err());
}
#[test]
fn green_err_n_components_too_large() {
let x = fixture_a();
assert!(PCA::<f64>::new(4).fit(&x, &()).is_err());
}
#[test]
fn green_err_insufficient_samples() {
let x = array![[1.0, 2.0, 3.0]];
assert!(PCA::<f64>::new(1).fit(&x, &()).is_err());
}
#[test]
fn green_err_transform_col_mismatch() {
let x = fixture_a();
let fitted = PCA::<f64>::new(3).fit(&x, &()).expect("fit");
let bad = array![[1.0, 2.0]];
assert!(fitted.transform(&bad).is_err());
}
#[test]
fn green_err_inverse_transform_col_mismatch() {
let x = fixture_a();
let fitted = PCA::<f64>::new(3).fit(&x, &()).expect("fit");
let bad = array![[1.0, 2.0]];
assert!(fitted.inverse_transform(&bad).is_err());
}