use crate::construction::{kronecker_marginal_eigensystems, kronecker_multi_index_advance};
use gam_problem::EstimationError;
use ndarray::{Array1, Array2};
use std::sync::Arc;
pub use gam_problem::penalty_matrix::kronecker_product;
#[derive(Clone, Debug)]
pub struct KroneckerInvariantStructure {
pub marginal_eigenvalues: Arc<Vec<Array1<f64>>>,
pub marginal_qs: Arc<Vec<Array2<f64>>>,
pub reparameterized_marginals: Arc<Vec<Array2<f64>>>,
pub max_balanced_eigenvalue: f64,
}
impl KroneckerInvariantStructure {
pub fn compute(
marginal_designs: &[Array2<f64>],
marginal_penalties: &[Array2<f64>],
marginal_dims: &[usize],
) -> Result<Self, EstimationError> {
let d = marginal_dims.len();
let mut marginal_eigenvalues = Vec::with_capacity(d);
let mut marginal_qs = Vec::with_capacity(d);
for (evals, evecs) in kronecker_marginal_eigensystems(
marginal_penalties,
"kronecker_reparameterization_engine",
)? {
marginal_eigenvalues.push(evals);
marginal_qs.push(evecs);
}
let reparameterized_marginals: Vec<Array2<f64>> = marginal_designs
.iter()
.zip(marginal_qs.iter())
.map(|(b_k, u_k)| gam_linalg::faer_ndarray::fast_ab(b_k, u_k))
.collect();
let mut max_balanced_eigenvalue = 0.0_f64;
let mut multi_idx = vec![0usize; d];
let frob_norms: Vec<f64> = marginal_penalties
.iter()
.map(|s| s.iter().map(|v| v * v).sum::<f64>().sqrt())
.collect();
loop {
let mut sigma = 0.0;
for k in 0..d {
if frob_norms[k] > 0.0 {
sigma += marginal_eigenvalues[k][multi_idx[k]] / frob_norms[k];
}
}
max_balanced_eigenvalue = max_balanced_eigenvalue.max(sigma);
if kronecker_multi_index_advance(&mut multi_idx, marginal_dims) {
break;
}
}
Ok(Self {
marginal_eigenvalues: Arc::new(marginal_eigenvalues),
marginal_qs: Arc::new(marginal_qs),
reparameterized_marginals: Arc::new(reparameterized_marginals),
max_balanced_eigenvalue,
})
}
}