use ndarray::{Array1, Array2};
use gam_terms::construction::ReparamResult;
use super::penalty_logdet::PenaltyPseudologdet;
use super::reml_outer_engine::{HessianDerivativeProvider, PenaltyLogdetDerivs};
pub struct RawInnerReparamContext<'a> {
pub hessian: &'a Array2<f64>,
pub beta: &'a Array1<f64>,
pub penalties_embedded: &'a [Array2<f64>],
pub lambdas: &'a [f64],
}
pub struct ReparameterizedInner<'dp> {
pub hessian_transformed: Array2<f64>,
pub beta_transformed: Array1<f64>,
pub penalty_logdet: PenaltyLogdetDerivs,
pub deriv_provider: Option<Box<dyn HessianDerivativeProvider + 'dp>>,
pub qs: Array2<f64>,
}
impl std::fmt::Debug for ReparameterizedInner<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ReparameterizedInner")
.field(
"hessian_transformed",
&format_args!(
"{}x{}",
self.hessian_transformed.nrows(),
self.hessian_transformed.ncols()
),
)
.field("beta_transformed_len", &self.beta_transformed.len())
.field("penalty_logdet_value", &self.penalty_logdet.value)
.field("has_deriv_provider", &self.deriv_provider.is_some())
.field("qs", &format_args!("{}x{}", self.qs.nrows(), self.qs.ncols()))
.finish()
}
}
pub fn assemble_reparameterized_inner<'dp>(
ctx: &RawInnerReparamContext<'_>,
deriv_provider: Option<Box<dyn HessianDerivativeProvider + 'dp>>,
reparam: &ReparamResult,
) -> Result<ReparameterizedInner<'dp>, String> {
let qs = &reparam.qs;
let p = qs.nrows();
if qs.ncols() != p {
return Err(format!(
"reparameterized inner: Qs must be square, got {}x{}",
qs.nrows(),
qs.ncols()
));
}
if ctx.hessian.nrows() != p || ctx.hessian.ncols() != p {
return Err(format!(
"reparameterized inner: Hessian must be {p}x{p} to match Qs, got {}x{}",
ctx.hessian.nrows(),
ctx.hessian.ncols()
));
}
if ctx.beta.len() != p {
return Err(format!(
"reparameterized inner: beta length {} must match Qs dimension {p}",
ctx.beta.len()
));
}
if ctx.lambdas.len() != ctx.penalties_embedded.len() {
return Err(format!(
"reparameterized inner: {} lambdas but {} penalty blocks",
ctx.lambdas.len(),
ctx.penalties_embedded.len()
));
}
for (k, s_k) in ctx.penalties_embedded.iter().enumerate() {
if s_k.nrows() != p || s_k.ncols() != p {
return Err(format!(
"reparameterized inner: penalty block {k} must be {p}x{p} (embedded), got {}x{}",
s_k.nrows(),
s_k.ncols()
));
}
}
let orth_tol = 128.0 * (p.max(1) as f64) * f64::EPSILON;
let orth_residual = max_abs_orthogonality_defect(qs);
if orth_residual.is_nan() || orth_residual > orth_tol {
return Err(format!(
"reparameterized inner: Qs is not orthogonal — ‖QsᵀQs − I‖_∞ = {orth_residual:.3e} \
exceeds tolerance {orth_tol:.3e} (p = {p})"
));
}
if reparam.penalty_shrinkage_ridge != 0.0 {
return Err(format!(
"reparameterized inner: ReparamResult carries a nonzero shrinkage ridge \
({:.3e}); a ρ-independent shrinkage ridge changes the prior normalizer and \
must be resolved upstream, never silently conjugated",
reparam.penalty_shrinkage_ridge
));
}
let hessian_transformed = qs.t().dot(ctx.hessian).dot(qs);
let beta_transformed = qs.t().dot(ctx.beta);
let pld = PenaltyPseudologdet::from_components(ctx.penalties_embedded, ctx.lambdas, 0.0)
.map_err(|e| format!("reparameterized inner: joint penalty logdet failed: {e}"))?;
let (det1, det2) = pld.rho_derivatives(ctx.penalties_embedded, ctx.lambdas);
let penalty_logdet = PenaltyLogdetDerivs {
value: pld.value(),
first: det1,
second: Some(det2),
};
let deriv_provider: Option<Box<dyn HessianDerivativeProvider + 'dp>> =
deriv_provider.map(|inner| {
Box::new(ConjugatedDerivProvider {
inner,
qs: qs.clone(),
}) as Box<dyn HessianDerivativeProvider + 'dp>
});
Ok(ReparameterizedInner {
hessian_transformed,
beta_transformed,
penalty_logdet,
deriv_provider,
qs: qs.clone(),
})
}
fn max_abs_orthogonality_defect(qs: &Array2<f64>) -> f64 {
let n = qs.ncols();
let mut worst = 0.0_f64;
for i in 0..n {
let col_i = qs.column(i);
for j in i..n {
let col_j = qs.column(j);
let dot: f64 = col_i.iter().zip(col_j.iter()).map(|(a, b)| a * b).sum();
let target = if i == j { 1.0 } else { 0.0 };
worst = worst.max((dot - target).abs());
}
}
worst
}
struct ConjugatedDerivProvider<'dp> {
inner: Box<dyn HessianDerivativeProvider + 'dp>,
qs: Array2<f64>,
}
impl<'dp> ConjugatedDerivProvider<'dp> {
#[inline]
fn to_raw(&self, v_prime: &Array1<f64>) -> Array1<f64> {
self.qs.dot(v_prime)
}
#[inline]
fn conjugate(&self, c: &Array2<f64>) -> Array2<f64> {
self.qs.t().dot(c).dot(&self.qs)
}
}
impl<'dp> HessianDerivativeProvider for ConjugatedDerivProvider<'dp> {
fn hessian_derivative_correction(
&self,
v_k: &Array1<f64>,
) -> Result<Option<Array2<f64>>, String> {
let v_raw = self.to_raw(v_k);
Ok(self
.inner
.hessian_derivative_correction(&v_raw)?
.map(|c| self.conjugate(&c)))
}
fn hessian_second_derivative_correction(
&self,
v_k: &Array1<f64>,
v_l: &Array1<f64>,
u_kl: &Array1<f64>,
) -> Result<Option<Array2<f64>>, String> {
let v_k_raw = self.to_raw(v_k);
let v_l_raw = self.to_raw(v_l);
let u_kl_raw = self.to_raw(u_kl);
Ok(self
.inner
.hessian_second_derivative_correction(&v_k_raw, &v_l_raw, &u_kl_raw)?
.map(|c| self.conjugate(&c)))
}
fn has_corrections(&self) -> bool {
self.inner.has_corrections()
}
}
#[cfg(test)]
mod tests {
use super::*;
use faer::Side;
use gam_linalg::faer_ndarray::FaerEigh;
fn deterministic_orthogonal(p: usize) -> Array2<f64> {
let sym = Array2::<f64>::from_shape_fn((p, p), |(i, j)| {
(((i * j) as f64) * 0.1).sin() + (((i + j) as f64) * 0.3).cos()
});
let (_, evecs) = sym.eigh(Side::Lower).expect("eigh for orthogonal factor");
evecs
}
fn deterministic_spd(p: usize) -> Array2<f64> {
let b = Array2::<f64>::from_shape_fn((p, p), |(i, j)| {
(((i * 31 + j * 17) % 13) as f64 - 6.0) * 0.25
});
let mut h = b.dot(&b.t());
for i in 0..p {
h[[i, i]] += 1.0;
}
h
}
fn embed(block: &Array2<f64>, start: usize, p: usize) -> Array2<f64> {
let m = block.nrows();
let mut s = Array2::<f64>::zeros((p, p));
for i in 0..m {
for j in 0..m {
s[[start + i, start + j]] = block[[i, j]];
}
}
s
}
fn diff_block(m: usize) -> Array2<f64> {
let mut s = Array2::<f64>::zeros((m, m));
for i in 0..m {
s[[i, i]] = 2.0;
if i + 1 < m {
s[[i, i + 1]] = -1.0;
s[[i + 1, i]] = -1.0;
}
}
for i in 0..m {
s[[i, i]] += 0.1;
}
s
}
fn sorted_eigenvalues(m: &Array2<f64>) -> Vec<f64> {
let (evals, _) = m.eigh(Side::Lower).expect("eigh");
let mut e: Vec<f64> = evals.to_vec();
e.sort_by(|a, b| a.partial_cmp(b).unwrap());
e
}
fn reparam_with(qs: Array2<f64>, shrinkage_ridge: f64) -> ReparamResult {
let p = qs.nrows();
ReparamResult {
s_transformed: Array2::zeros((p, p)),
log_det: 0.0,
det1: Array1::zeros(0),
qs,
canonical_transformed: Vec::new(),
e_transformed: Array2::zeros((0, p)),
u_truncated: Array2::zeros((p, 0)),
penalty_shrinkage_ridge: shrinkage_ridge,
}
}
#[test]
fn disjoint_blocks_joint_matches_per_block_and_eigenvalues_invariant() {
let p = 8;
let s0 = embed(&diff_block(4), 0, p);
let s1 = embed(&diff_block(4), 4, p);
let penalties = vec![s0.clone(), s1.clone()];
let lambdas = vec![2.5, 0.4];
let h = deterministic_spd(p);
let beta = Array1::from_shape_fn(p, |i| (i as f64 - 3.5) * 0.3);
let qs = deterministic_orthogonal(p);
let reparam = reparam_with(qs.clone(), 0.0);
let ctx = RawInnerReparamContext {
hessian: &h,
beta: &beta,
penalties_embedded: &penalties,
lambdas: &lambdas,
};
let out = assemble_reparameterized_inner(&ctx, None, &reparam)
.expect("disjoint reparam should succeed");
let mut per_block_value = 0.0;
let mut per_block_det1 = Array1::<f64>::zeros(2);
for (k, s_k) in penalties.iter().enumerate() {
let pld_k =
PenaltyPseudologdet::from_components(&[s_k.clone()], &[lambdas[k]], 0.0).unwrap();
let (d1_k, _) = pld_k.rho_derivatives(&[s_k.clone()], &[lambdas[k]]);
per_block_value += pld_k.value();
per_block_det1[k] = d1_k[0];
}
assert!(
(out.penalty_logdet.value - per_block_value).abs() <= 1e-10,
"joint value {} vs per-block sum {}",
out.penalty_logdet.value,
per_block_value
);
for k in 0..2 {
assert!(
(out.penalty_logdet.first[k] - per_block_det1[k]).abs() <= 1e-10,
"det1[{k}] joint {} vs per-block {}",
out.penalty_logdet.first[k],
per_block_det1[k]
);
}
let det2 = out.penalty_logdet.second.as_ref().unwrap();
assert!(
det2[[0, 1]].abs() <= 1e-10 && det2[[1, 0]].abs() <= 1e-10,
"cross-block det2 nonzero: {} / {}",
det2[[0, 1]],
det2[[1, 0]]
);
let ev_h = sorted_eigenvalues(&h);
let ev_hp = sorted_eigenvalues(&out.hessian_transformed);
for (a, b) in ev_h.iter().zip(ev_hp.iter()) {
assert!(
(a - b).abs() <= 1e-11 * (1.0 + a.abs()),
"eigenvalue drift under conjugation: {a} vs {b}"
);
}
let recovered = qs.dot(&out.beta_transformed);
for i in 0..p {
assert!((recovered[i] - beta[i]).abs() <= 1e-11);
}
}
#[test]
fn full_span_ridge_overlap_joint_matches_assembled_and_differs_from_per_block_sum() {
let p = 6;
let s_block = embed(&diff_block(4), 0, p);
let s_ridge = Array2::<f64>::eye(p);
let penalties = vec![s_block.clone(), s_ridge.clone()];
let lambdas = vec![3.0, 0.05];
let h = deterministic_spd(p);
let beta = Array1::from_shape_fn(p, |i| (i as f64) * 0.11 - 0.3);
let qs = deterministic_orthogonal(p);
let reparam = reparam_with(qs, 0.0);
let ctx = RawInnerReparamContext {
hessian: &h,
beta: &beta,
penalties_embedded: &penalties,
lambdas: &lambdas,
};
let out = assemble_reparameterized_inner(&ctx, None, &reparam)
.expect("ridge-overlap reparam should succeed");
let joint = out.penalty_logdet.value;
let mut assembled = Array2::<f64>::zeros((p, p));
assembled.scaled_add(lambdas[0], &s_block);
assembled.scaled_add(lambdas[1], &s_ridge);
let ref_assembled = PenaltyPseudologdet::from_assembled(assembled, None)
.unwrap()
.value();
assert!(
(joint - ref_assembled).abs() <= 1e-9,
"joint {joint} vs directly-assembled reference {ref_assembled}"
);
let v_block = PenaltyPseudologdet::from_components(&[s_block], &[lambdas[0]], 0.0)
.unwrap()
.value();
let v_ridge = PenaltyPseudologdet::from_components(&[s_ridge], &[lambdas[1]], 0.0)
.unwrap()
.value();
let per_block_sum = v_block + v_ridge;
assert!(
(joint - per_block_sum).abs() > 1e-6,
"joint {joint} indistinguishable from per-block sum {per_block_sum}"
);
}
#[test]
fn non_orthogonal_qs_is_rejected() {
let p = 4;
let h = deterministic_spd(p);
let beta = Array1::zeros(p);
let penalties = vec![embed(&diff_block(4), 0, p)];
let lambdas = vec![1.0];
let bad_qs = 2.0 * Array2::<f64>::eye(p);
let reparam = reparam_with(bad_qs, 0.0);
let ctx = RawInnerReparamContext {
hessian: &h,
beta: &beta,
penalties_embedded: &penalties,
lambdas: &lambdas,
};
let err = assemble_reparameterized_inner(&ctx, None, &reparam)
.expect_err("non-orthogonal Qs must be rejected");
assert!(err.contains("not orthogonal"), "unexpected error: {err}");
}
#[test]
fn nonzero_shrinkage_ridge_is_rejected() {
let p = 4;
let h = deterministic_spd(p);
let beta = Array1::zeros(p);
let penalties = vec![embed(&diff_block(4), 0, p)];
let lambdas = vec![1.0];
let qs = deterministic_orthogonal(p);
let reparam = reparam_with(qs, 1e-3);
let ctx = RawInnerReparamContext {
hessian: &h,
beta: &beta,
penalties_embedded: &penalties,
lambdas: &lambdas,
};
let err = assemble_reparameterized_inner(&ctx, None, &reparam)
.expect_err("nonzero shrinkage ridge must be rejected");
assert!(err.contains("shrinkage ridge"), "unexpected error: {err}");
}
struct FixedCorrectionProvider {
m0: Array2<f64>,
}
impl HessianDerivativeProvider for FixedCorrectionProvider {
fn hessian_derivative_correction(
&self,
v_k: &Array1<f64>,
) -> Result<Option<Array2<f64>>, String> {
let n = v_k.len();
let mut c = self.m0.clone();
for i in 0..n {
for j in 0..n {
c[[i, j]] += v_k[i] * v_k[j];
}
}
Ok(Some(c))
}
fn has_corrections(&self) -> bool {
true
}
}
fn trace(m: &Array2<f64>) -> f64 {
(0..m.nrows().min(m.ncols())).map(|i| m[[i, i]]).sum()
}
#[test]
fn conjugated_provider_preserves_trace_and_quadratic_form() {
let p = 5;
let qs = deterministic_orthogonal(p);
let m0 = deterministic_spd(p);
let inner = FixedCorrectionProvider { m0: m0.clone() };
let conj = ConjugatedDerivProvider {
inner: Box::new(FixedCorrectionProvider { m0 }),
qs: qs.clone(),
};
let v_prime = Array1::from_shape_fn(p, |i| (i as f64 - 2.0) * 0.37 + 0.1);
let v_raw = qs.dot(&v_prime);
let c_prime = conj
.hessian_derivative_correction(&v_prime)
.unwrap()
.unwrap();
let c_raw = inner
.hessian_derivative_correction(&v_raw)
.unwrap()
.unwrap();
assert!(
(trace(&c_prime) - trace(&c_raw)).abs() <= 1e-10 * (1.0 + trace(&c_raw).abs()),
"trace not preserved: {} vs {}",
trace(&c_prime),
trace(&c_raw)
);
let q_prime = v_prime.dot(&c_prime.dot(&v_prime));
let q_raw = v_raw.dot(&c_raw.dot(&v_raw));
assert!(
(q_prime - q_raw).abs() <= 1e-10 * (1.0 + q_raw.abs()),
"quadratic form not preserved: {q_prime} vs {q_raw}"
);
let expected = qs.t().dot(&c_raw).dot(&qs);
for i in 0..p {
for j in 0..p {
assert!(
(c_prime[[i, j]] - expected[[i, j]]).abs() <= 1e-10,
"conjugation mismatch at ({i},{j})"
);
}
}
}
}