use crate::error::{CsError, CsResult};
use crate::linalg::cholesky::{cholesky_factor, cholesky_solve};
use crate::linalg::{mat_t_vec, mat_vec, norm2, submat_columns};
use crate::sbl::SblResult;
pub fn fast_marginal_likelihood(
phi: &[f64],
m: usize,
n: usize,
y: &[f64],
max_iter: usize,
tol: f64,
) -> CsResult<SblResult> {
if phi.len() != m * n {
return Err(CsError::ShapeMismatch {
expected: vec![m, n],
got: vec![phi.len()],
});
}
if y.len() != m {
return Err(CsError::DimensionMismatch { a: y.len(), b: m });
}
let mut sigma2 = (norm2(y) / (m as f64).sqrt()).max(1.0e-6);
sigma2 *= sigma2;
let mut corr = vec![0.0_f64; n];
for j in 0..n {
let mut s = 0.0_f64;
for i in 0..m {
s += phi[i * n + j] * y[i];
}
corr[j] = s;
}
let mut active: Vec<usize> = Vec::new();
let mut alphas: Vec<f64> = Vec::new();
let mut best_j = 0usize;
let mut best_score = f64::NEG_INFINITY;
for j in 0..n {
let mut s_sq = 0.0_f64;
for i in 0..m {
let v = phi[i * n + j];
s_sq += v * v;
}
let q_sq = corr[j] * corr[j];
let score = q_sq / s_sq.max(1.0e-300);
if score > best_score {
best_score = score;
best_j = j;
}
}
let mut col_norm_sq_first = 0.0_f64;
for i in 0..m {
let v = phi[i * n + best_j];
col_norm_sq_first += v * v;
}
let init_alpha = col_norm_sq_first / sigma2;
active.push(best_j);
alphas.push(init_alpha);
let mut mu = vec![0.0_f64; n];
let mut iter = 0usize;
for _ in 0..max_iter {
let k_act = active.len();
let phi_a = submat_columns(phi, m, n, &active)?;
let mut g = vec![0.0_f64; k_act * k_act];
for kk in 0..m {
for a in 0..k_act {
let pai = phi_a[kk * k_act + a];
for b in 0..k_act {
g[a * k_act + b] += pai * phi_a[kk * k_act + b];
}
}
}
for j in 0..k_act {
g[j * k_act + j] = g[j * k_act + j] / sigma2 + alphas[j];
}
for j in 0..k_act {
for k_idx in 0..k_act {
if k_idx != j {
g[j * k_act + k_idx] /= sigma2;
}
}
}
let l = cholesky_factor(&g, k_act)?;
let phi_a_t_y = mat_t_vec(&phi_a, m, k_act, y)?;
let mut rhs = vec![0.0_f64; k_act];
for j in 0..k_act {
rhs[j] = phi_a_t_y[j] / sigma2;
}
let mu_a = cholesky_solve(&l, k_act, &rhs)?;
mu.fill(0.0);
for (i, &j) in active.iter().enumerate() {
mu[j] = mu_a[i];
}
let phi_a_mu = mat_vec(&phi_a, m, k_act, &mu_a)?;
let mut resid_sq = 0.0_f64;
for i in 0..m {
let d = y[i] - phi_a_mu[i];
resid_sq += d * d;
}
let mut sigma_diag = vec![0.0_f64; k_act];
for j in 0..k_act {
let mut ej = vec![0.0_f64; k_act];
ej[j] = 1.0;
let sj_vec = cholesky_solve(&l, k_act, &ej)?;
sigma_diag[j] = sj_vec[j];
}
let mut denom = m as f64;
for j in 0..k_act {
denom -= 1.0 - alphas[j] * sigma_diag[j];
}
if denom > 1.0e-6 {
sigma2 = (resid_sq / denom).max(1.0e-12);
}
for j in 0..k_act {
let mu2 = mu_a[j] * mu_a[j];
let denom_alpha = (mu2 + sigma_diag[j]).max(1.0e-300);
alphas[j] = 1.0 / denom_alpha;
}
iter += 1;
let mut delta = 0.0_f64;
for (i, &j) in active.iter().enumerate() {
let d = mu_a[i] - mu[j];
delta += d * d;
}
if delta.sqrt() / norm2(&mu).max(1.0e-300) < tol {
break;
}
}
let gamma: Vec<f64> = (0..n)
.map(|j| {
if let Some(pos) = active.iter().position(|&a| a == j) {
1.0 / alphas[pos].max(1.0e-300)
} else {
0.0
}
})
.collect();
Ok(SblResult {
x: mu,
gamma,
sigma2,
iterations: iter,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn fml_runs() {
let phi = vec![1.0, 0.0, 0.0, 1.0];
let y = vec![1.0, 0.0];
let r = fast_marginal_likelihood(&phi, 2, 2, &y, 20, 1.0e-7).expect("ok");
assert!(r.iterations > 0);
}
}