use faer::Side;
use gam_runtime::warm_start::{Fingerprint, Fingerprinter};
use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use serde::{Deserialize, Serialize};
use crate::arrow_schur::{ArrowFactorCache, ArrowSchurSystem};
use crate::priority_selection::{PriorityCandidate, rank_priority_candidates};
use gam_linalg::faer_ndarray::FaerEigh;
use gam_linalg::lanczos::{
SymmetricLanczosOptions, symmetric_lanczos_eigenpairs, symmetric_lanczos_log_quadrature,
};
use gam_linalg::pairwise_reduce::{BASE_CHUNK, pairwise_sum};
use gam_linalg::triangular::cholesky_solve_vector;
use gam_math::special::bessel_i0_log_minus_abs_and_ratio;
pub const ANALYTIC_LOGDET_DENSE_DIM_THRESHOLD: usize = 1024;
const EVIDENCE_LOGDET_SLQ_PROBES: usize = 16;
const EVIDENCE_LOGDET_LANCZOS_STEPS: usize = 32;
const EVIDENCE_HVP_SYMMETRY_REL_TOL: f64 = 1e-8;
const EVIDENCE_HVP_SYMMETRY_PROBES: usize = 4;
#[derive(Clone, Copy)]
pub struct EvidenceHvpLogDet<'a> {
pub dim: usize,
pub apply: &'a dyn Fn(&[f64]) -> Vec<f64>,
}
#[derive(Clone, Copy)]
pub enum EvidenceLogDetSource<'a> {
FactoredArrow {
cache: &'a ArrowFactorCache,
fallback_hvp: Option<EvidenceHvpLogDet<'a>>,
},
Hvp(EvidenceHvpLogDet<'a>),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum TopologyKind {
Periodic,
Flat,
Sphere,
Torus,
}
impl TopologyKind {
pub fn complexity_rank(self) -> u8 {
match self {
TopologyKind::Flat => 0,
TopologyKind::Periodic => 1,
TopologyKind::Sphere => 2,
TopologyKind::Torus => 3,
}
}
}
#[derive(Debug, Clone)]
pub struct TopologyCandidate {
pub kind: TopologyKind,
pub negative_log_evidence: f64,
pub effective_dim: f64,
pub n_obs: usize,
pub converged: bool,
pub exclusion_reason: Option<String>,
}
#[derive(Debug, Clone)]
pub struct SelectedTopology {
pub winner: TopologyKind,
pub ranking: Vec<TopologyCandidate>,
pub tie: bool,
}
#[derive(Debug, Clone, Copy)]
pub struct TopologySelectOptions {
pub tie_tolerance: f64,
pub score_scale: TopologyScoreScale,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TopologyScoreScale {
PerObservation,
PerEffectiveDim,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub struct StackingConfig {
pub max_iter: usize,
pub kkt_tol: f64,
}
impl Default for StackingConfig {
fn default() -> Self {
Self {
max_iter: 256,
kkt_tol: f64::EPSILON.sqrt(),
}
}
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub struct StackingCertificate {
pub mean_log_score: f64,
pub duality_gap: f64,
pub simplex_residual: f64,
pub multiplier_residual: f64,
pub complementarity_residual: f64,
}
impl StackingCertificate {
pub fn residual(&self) -> f64 {
self.duality_gap
.max(self.simplex_residual)
.max(self.multiplier_residual)
.max(self.complementarity_residual)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StackingCheckpoint {
pub weights: Array1<f64>,
pub completed_iterations: usize,
density_fingerprint: Fingerprint,
}
#[derive(Debug, Clone)]
pub enum StackingError {
InvalidInput {
message: String,
},
NumericalFailure {
message: String,
certificate: Option<StackingCertificate>,
checkpoint: Option<StackingCheckpoint>,
},
DidNotConverge {
max_iterations: usize,
tolerance: f64,
certificate: StackingCertificate,
checkpoint: StackingCheckpoint,
},
}
impl std::fmt::Display for StackingError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::InvalidInput { message } => write!(f, "invalid stacking problem: {message}"),
Self::NumericalFailure {
message,
certificate,
checkpoint,
} => write!(
f,
"stacking numerical failure: {message} (certificate residual {}, checkpoint iterations {})",
certificate.map_or(f64::NAN, |value| value.residual()),
checkpoint
.as_ref()
.map_or(0, |value| value.completed_iterations)
),
Self::DidNotConverge {
max_iterations,
tolerance,
certificate,
checkpoint,
} => write!(
f,
"stacking did not certify after {max_iterations} additional iterations (total {}): KKT residual {:.6e} exceeds tolerance {:.3e}; resume from the carried weights checkpoint",
checkpoint.completed_iterations,
certificate.residual(),
tolerance
),
}
}
}
impl std::error::Error for StackingError {}
#[derive(Debug, Clone)]
pub struct StackingWeights {
pub weights: Array1<f64>,
pub iterations: usize,
pub certificate: StackingCertificate,
}
impl StackingWeights {
pub fn mean_log_score(&self) -> f64 {
self.certificate.mean_log_score
}
}
struct StackingProblem {
scaled_density: Array2<f64>,
row_log_scale: Array1<f64>,
}
impl StackingProblem {
fn from_log_density(log_density: ArrayView2<'_, f64>) -> Result<Self, StackingError> {
let n_obs = log_density.nrows();
let n_cand = log_density.ncols();
if n_cand == 0 || n_obs == 0 {
return Err(StackingError::InvalidInput {
message: "at least one candidate and one held-out row are required".to_string(),
});
}
if let Some(((row, col), value)) = log_density
.indexed_iter()
.find(|(_, value)| value.is_nan() || **value == f64::INFINITY)
{
return Err(StackingError::InvalidInput {
message: format!(
"log density at row {row}, candidate {col} is {value}; NaN and +infinity are not predictive densities"
),
});
}
let mut scaled_density = Array2::<f64>::zeros((n_obs, n_cand));
let mut row_log_scale = Array1::<f64>::zeros(n_obs);
for row in 0..n_obs {
let row_max = (0..n_cand)
.map(|col| log_density[[row, col]])
.fold(f64::NEG_INFINITY, f64::max);
if !row_max.is_finite() {
return Err(StackingError::InvalidInput {
message: format!(
"held-out row {row} has zero density under every candidate; deleting it would change the stacking target"
),
});
}
row_log_scale[row] = row_max;
for col in 0..n_cand {
let value = log_density[[row, col]];
if value.is_finite() {
scaled_density[[row, col]] = (value - row_max).exp();
}
}
}
Ok(Self {
scaled_density,
row_log_scale,
})
}
fn evaluate(
&self,
weights: ArrayView1<'_, f64>,
) -> Result<(Array1<f64>, StackingCertificate, f64), String> {
let n = self.scaled_density.nrows();
let k = self.scaled_density.ncols();
let mass = weights.sum();
if weights.len() != k
|| weights
.iter()
.any(|value| !value.is_finite() || *value < 0.0)
|| !(mass.is_finite() && mass > 0.0)
{
return Err(
"checkpoint weights are not a finite nonnegative simplex vector".to_string(),
);
}
let mut gradient = Array1::<f64>::zeros(k);
let mut centered_objective = 0.0_f64;
let mut mean_log_score = 0.0_f64;
for row in 0..n {
let mut mixture = 0.0_f64;
for col in 0..k {
mixture += weights[col] * self.scaled_density[[row, col]];
}
if !(mixture.is_finite() && mixture > 0.0) {
return Err(format!(
"candidate mixture lost held-out row {row} (scaled density {mixture})"
));
}
let log_mixture = mixture.ln();
centered_objective += log_mixture / n as f64;
let log_score = self.row_log_scale[row] + log_mixture;
let count = (row + 1) as f64;
mean_log_score = mean_log_score * ((count - 1.0) / count) + log_score / count;
for col in 0..k {
gradient[col] += self.scaled_density[[row, col]] / mixture / n as f64;
}
}
if !centered_objective.is_finite()
|| !mean_log_score.is_finite()
|| gradient.iter().any(|value| !value.is_finite())
{
return Err("objective or analytic gradient became non-finite".to_string());
}
let multiplier = weights.dot(&gradient);
let max_gradient = gradient.iter().copied().fold(f64::NEG_INFINITY, f64::max);
let certificate = StackingCertificate {
mean_log_score,
duality_gap: (max_gradient - multiplier).max(0.0),
simplex_residual: (mass - 1.0).abs(),
multiplier_residual: (multiplier - 1.0).abs(),
complementarity_residual: weights
.iter()
.zip(gradient.iter())
.map(|(&weight, &gain)| weight * (gain - multiplier).abs())
.fold(0.0_f64, f64::max),
};
Ok((gradient, certificate, centered_objective))
}
fn centered_objective(&self, weights: ArrayView1<'_, f64>) -> Option<f64> {
let n = self.scaled_density.nrows();
let mut objective = 0.0_f64;
for row in 0..n {
let mixture = self.scaled_density.row(row).dot(&weights);
if !(mixture.is_finite() && mixture > 0.0) {
return None;
}
objective += mixture.ln() / n as f64;
}
objective.is_finite().then_some(objective)
}
}
pub fn solve_stacking_weights(
log_density: ArrayView2<'_, f64>,
config: StackingConfig,
) -> Result<StackingWeights, StackingError> {
solve_stacking_weights_impl(log_density, config, None)
}
pub fn resume_stacking_weights(
log_density: ArrayView2<'_, f64>,
config: StackingConfig,
checkpoint: &StackingCheckpoint,
) -> Result<StackingWeights, StackingError> {
solve_stacking_weights_impl(log_density, config, Some(checkpoint))
}
fn solve_stacking_weights_impl(
log_density: ArrayView2<'_, f64>,
config: StackingConfig,
checkpoint: Option<&StackingCheckpoint>,
) -> Result<StackingWeights, StackingError> {
if config.max_iter == 0 {
return Err(StackingError::InvalidInput {
message: "max_iter must be positive".to_string(),
});
}
let numerical_floor = f64::EPSILON.sqrt();
if !config.kkt_tol.is_finite() || config.kkt_tol < numerical_floor {
return Err(StackingError::InvalidInput {
message: format!(
"kkt_tol must be finite and at least the floating-point resolution floor {numerical_floor:.3e}"
),
});
}
let density_fingerprint = evidence_matrix_fingerprint("stacking-log-density-v1", log_density);
let problem = StackingProblem::from_log_density(log_density)?;
let k = problem.scaled_density.ncols();
let (mut weights, completed_before) = if let Some(checkpoint) = checkpoint {
if checkpoint.density_fingerprint != density_fingerprint {
return Err(StackingError::InvalidInput {
message: "checkpoint belongs to a different held-out density table".to_string(),
});
}
if checkpoint.weights.len() != k {
return Err(StackingError::InvalidInput {
message: format!(
"checkpoint has {} weights but the density table has {k} candidates",
checkpoint.weights.len()
),
});
}
let mut weights = checkpoint.weights.clone();
let mass = weights.sum();
if weights
.iter()
.any(|value| !value.is_finite() || *value < 0.0)
|| !mass.is_finite()
|| (mass - 1.0).abs() > config.kkt_tol
{
return Err(StackingError::InvalidInput {
message: "checkpoint weights must be a finite nonnegative simplex vector"
.to_string(),
});
}
weights.mapv_inplace(|value| value / mass);
(weights, checkpoint.completed_iterations)
} else {
(Array1::<f64>::from_elem(k, 1.0 / k as f64), 0)
};
for additional_iterations in 0..=config.max_iter {
let completed_iterations = completed_before + additional_iterations;
let checkpoint = StackingCheckpoint {
weights: weights.clone(),
completed_iterations,
density_fingerprint,
};
let (gradient, certificate, objective) =
problem.evaluate(weights.view()).map_err(|message| {
StackingError::NumericalFailure {
message,
certificate: None,
checkpoint: Some(checkpoint.clone()),
}
})?;
if certificate.residual() <= config.kkt_tol {
return Ok(StackingWeights {
weights,
iterations: completed_iterations,
certificate,
});
}
if additional_iterations == config.max_iter {
return Err(StackingError::DidNotConverge {
max_iterations: config.max_iter,
tolerance: config.kkt_tol,
certificate,
checkpoint,
});
}
let max_gradient_col = gradient
.iter()
.enumerate()
.max_by(|left, right| left.1.total_cmp(right.1))
.map(|(index, _)| index)
.expect("stacking has at least one candidate");
let candidate = stacking_newton_step(&problem, weights.view(), gradient.view(), objective)
.or_else(|| {
stacking_vertex_step(&problem, weights.view(), max_gradient_col, objective)
})
.ok_or_else(|| StackingError::NumericalFailure {
message: "positive KKT gap remained but neither the analytic Newton direction nor the exact vertex line solve produced a representable ascent step".to_string(),
certificate: Some(certificate),
checkpoint: Some(checkpoint),
})?;
weights = candidate;
}
Err(StackingError::NumericalFailure {
message: format!(
"stacking solver exhausted its inclusive iteration budget ({}) without producing a \
terminal verdict",
config.max_iter
),
certificate: None,
checkpoint: None,
})
}
fn stacking_newton_step(
problem: &StackingProblem,
weights: ArrayView1<'_, f64>,
gradient: ArrayView1<'_, f64>,
objective: f64,
) -> Option<Array1<f64>> {
let active: Vec<usize> = weights
.iter()
.enumerate()
.filter_map(|(index, &weight)| (weight > 0.0).then_some(index))
.collect();
if active.len() < 2 {
return None;
}
let reference_position = active
.iter()
.enumerate()
.max_by(|left, right| weights[*left.1].total_cmp(&weights[*right.1]))
.map(|(position, _)| position)?;
let reference = active[reference_position];
let free: Vec<usize> = active
.iter()
.copied()
.filter(|&index| index != reference)
.collect();
let dimension = free.len();
let n = problem.scaled_density.nrows();
let mut information = Array2::<f64>::zeros((dimension, dimension));
for row in 0..n {
let mixture = problem.scaled_density.row(row).dot(&weights);
if !(mixture.is_finite() && mixture > 0.0) {
return None;
}
let reference_density = problem.scaled_density[[row, reference]];
let contrasts: Vec<f64> = free
.iter()
.map(|&col| (problem.scaled_density[[row, col]] - reference_density) / mixture)
.collect();
for left in 0..dimension {
for right in 0..=left {
information[[left, right]] += contrasts[left] * contrasts[right] / n as f64;
information[[right, left]] = information[[left, right]];
}
}
}
let reduced_gradient =
Array1::from_iter(free.iter().map(|&col| gradient[col] - gradient[reference]));
let (eigenvalues, eigenvectors) = information.eigh(Side::Lower).ok()?;
let spectral_scale = eigenvalues.iter().copied().fold(0.0_f64, f64::max);
if !(spectral_scale.is_finite() && spectral_scale > 0.0) {
return None;
}
let rank_tolerance = f64::EPSILON * (dimension as f64) * spectral_scale.max(f64::MIN_POSITIVE);
let projected = eigenvectors.t().dot(&reduced_gradient);
let mut spectral_step = Array1::<f64>::zeros(dimension);
for index in 0..dimension {
if eigenvalues[index] > rank_tolerance {
spectral_step[index] = projected[index] / eigenvalues[index];
}
}
let reduced_step = eigenvectors.dot(&spectral_step);
let ascent = reduced_gradient.dot(&reduced_step);
if !(ascent.is_finite() && ascent > 0.0) {
return None;
}
let mut direction = Array1::<f64>::zeros(weights.len());
for (position, &col) in free.iter().enumerate() {
direction[col] = reduced_step[position];
}
direction[reference] = -reduced_step.sum();
let mut step = 1.0_f64;
let mut boundary = None;
for col in 0..weights.len() {
if direction[col] < 0.0 {
let candidate = -weights[col] / direction[col];
if candidate < step {
step = candidate;
boundary = Some(col);
}
}
}
loop {
let mut candidate = &weights + &(direction.mapv(|value| step * value));
if let Some(col) = boundary {
if step == -weights[col] / direction[col] {
candidate[col] = 0.0;
}
}
for value in candidate.iter_mut() {
if *value < 0.0 && *value >= -f64::EPSILON {
*value = 0.0;
}
}
let mass = candidate.sum();
if mass.is_finite() && mass > 0.0 {
candidate.mapv_inplace(|value| value / mass);
if problem
.centered_objective(candidate.view())
.is_some_and(|value| value > objective)
{
return Some(candidate);
}
}
let next_step = 0.5 * step;
if next_step == step || next_step == 0.0 {
return None;
}
step = next_step;
boundary = None;
}
}
fn stacking_vertex_step(
problem: &StackingProblem,
weights: ArrayView1<'_, f64>,
vertex: usize,
objective: f64,
) -> Option<Array1<f64>> {
let derivative = |step: f64| -> f64 {
let mut value = 0.0_f64;
let n = problem.scaled_density.nrows();
for row in 0..n {
let current = problem.scaled_density.row(row).dot(&weights);
let target = problem.scaled_density[[row, vertex]];
let mixture = (1.0 - step) * current + step * target;
if mixture <= 0.0 {
return f64::NEG_INFINITY;
}
value += (target - current) / mixture / n as f64;
}
value
};
if derivative(0.0) <= 0.0 {
return None;
}
let mut step = if derivative(1.0) >= 0.0 {
1.0
} else {
let mut lower = 0.0_f64;
let mut upper = 1.0_f64;
while upper - lower > f64::EPSILON.sqrt() {
let middle = 0.5 * (lower + upper);
if derivative(middle) > 0.0 {
lower = middle;
} else {
upper = middle;
}
}
0.5 * (lower + upper)
};
loop {
let mut candidate = weights.mapv(|weight| (1.0 - step) * weight);
candidate[vertex] += step;
if problem
.centered_objective(candidate.view())
.is_some_and(|value| value > objective)
{
return Some(candidate);
}
let next_step = 0.5 * step;
if next_step == step || next_step == 0.0 {
return None;
}
step = next_step;
}
}
pub fn stacked_predictive_mean(
weights: &Array1<f64>,
candidate_means: &[Array1<f64>],
) -> Result<Array1<f64>, String> {
if candidate_means.len() != weights.len() {
return Err(format!(
"stacked_predictive_mean: {} weights but {} candidate mean vectors",
weights.len(),
candidate_means.len()
));
}
let Some(first) = candidate_means.first() else {
return Err("stacked_predictive_mean requires at least one candidate".to_string());
};
let n_rows = first.len();
if candidate_means.iter().any(|means| means.len() != n_rows) {
return Err(
"stacked_predictive_mean: candidate mean vectors disagree on row count".to_string(),
);
}
let mut out = Array1::<f64>::zeros(n_rows);
for (weight, means) in weights.iter().zip(candidate_means) {
if *weight != 0.0 {
out.scaled_add(*weight, means);
}
}
Ok(out)
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub struct GaussianMixtureConfig {
pub max_iter: usize,
pub loglik_tol: f64,
pub parameter_tol: f64,
pub covariance_floor: f64,
pub kmeans_max_iter: usize,
}
impl Default for GaussianMixtureConfig {
fn default() -> Self {
Self {
max_iter: 1000,
loglik_tol: f64::EPSILON.sqrt(),
parameter_tol: f64::EPSILON.sqrt(),
covariance_floor: 1e-6,
kmeans_max_iter: 25,
}
}
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub struct GaussianMixtureCertificate {
pub mean_log_likelihood: f64,
pub mean_log_likelihood_gain: f64,
pub monotonicity_uncertainty: f64,
pub objective_residual: f64,
pub objective_tolerance: f64,
pub parameter_residual: f64,
pub parameter_tolerance: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GaussianMixtureCheckpoint {
pub weights: Array1<f64>,
pub means: Array2<f64>,
pub covariances: Vec<Array2<f64>>,
pub mean_log_likelihood: f64,
pub completed_iterations: usize,
data_fingerprint: Fingerprint,
covariance_floor: f64,
}
#[derive(Debug, Clone)]
pub enum GaussianMixtureError {
InvalidInput {
message: String,
},
NumericalFailure {
message: String,
checkpoint: Option<GaussianMixtureCheckpoint>,
},
MonotonicityViolation {
previous_mean_log_likelihood: f64,
next_mean_log_likelihood: f64,
numerical_uncertainty: f64,
checkpoint: GaussianMixtureCheckpoint,
},
DidNotConverge {
max_iterations: usize,
certificate: GaussianMixtureCertificate,
checkpoint: GaussianMixtureCheckpoint,
},
}
impl std::fmt::Display for GaussianMixtureError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::InvalidInput { message } => write!(f, "invalid Gaussian mixture: {message}"),
Self::NumericalFailure {
message,
checkpoint,
} => write!(
f,
"Gaussian-mixture numerical failure: {message} (checkpoint iterations {})",
checkpoint
.as_ref()
.map_or(0, |value| value.completed_iterations)
),
Self::MonotonicityViolation {
previous_mean_log_likelihood,
next_mean_log_likelihood,
numerical_uncertainty,
checkpoint,
} => write!(
f,
"Gaussian-mixture EM violated monotone ascent at iteration {}: mean log likelihood {previous_mean_log_likelihood:.12e} -> {next_mean_log_likelihood:.12e} (comparison uncertainty {numerical_uncertainty:.3e}); resume from the carried checkpoint only after diagnosing the numerical failure",
checkpoint.completed_iterations
),
Self::DidNotConverge {
max_iterations,
certificate,
checkpoint,
} => write!(
f,
"Gaussian-mixture EM did not certify after {max_iterations} additional iterations (total {}): signed mean-log-likelihood gain {:.6e} (numerical uncertainty {:.3e}), objective residual {:.6e}/{:.3e}, parameter-map residual {:.6e}/{:.3e}; resume from the carried checkpoint, which is not comparable evidence",
checkpoint.completed_iterations,
certificate.mean_log_likelihood_gain,
certificate.monotonicity_uncertainty,
certificate.objective_residual,
certificate.objective_tolerance,
certificate.parameter_residual,
certificate.parameter_tolerance
),
}
}
}
impl std::error::Error for GaussianMixtureError {}
#[derive(Debug, Clone)]
pub struct GaussianMixtureFit {
weights: Array1<f64>,
means: Array2<f64>,
covariances: Vec<Array2<f64>>,
k: usize,
d: usize,
n_obs: usize,
loglik: f64,
iterations: usize,
certificate: GaussianMixtureCertificate,
}
impl GaussianMixtureFit {
pub fn weights(&self) -> ArrayView1<'_, f64> {
self.weights.view()
}
pub fn means(&self) -> ArrayView2<'_, f64> {
self.means.view()
}
pub fn covariances(&self) -> &[Array2<f64>] {
&self.covariances
}
pub fn iterations(&self) -> usize {
self.iterations
}
pub fn certificate(&self) -> GaussianMixtureCertificate {
self.certificate
}
pub fn num_free_parameters(&self) -> usize {
let cov_per = self.d * (self.d + 1) / 2;
(self.k - 1) + self.k * self.d + self.k * cov_per
}
pub fn per_point_log_density(&self, data: ArrayView2<'_, f64>) -> Result<Array1<f64>, String> {
if data.ncols() != self.d {
return Err(format!(
"mixture log-density expects {} columns, got {}",
self.d,
data.ncols()
));
}
let n = data.nrows();
let mut comp = Vec::with_capacity(self.k);
for j in 0..self.k {
comp.push(GaussianComponentEval::factor(
self.means.row(j),
&self.covariances[j],
)?);
}
let mut out = Array1::<f64>::zeros(n);
let log_w: Vec<f64> = self.weights.iter().map(|w| w.ln()).collect();
for i in 0..n {
let row = data.row(i);
let mut log_terms = vec![f64::NEG_INFINITY; self.k];
let mut max_term = f64::NEG_INFINITY;
for j in 0..self.k {
let lt = log_w[j] + comp[j].log_density(row);
log_terms[j] = lt;
if lt > max_term {
max_term = lt;
}
}
out[i] = log_sum_exp(&log_terms, max_term);
}
Ok(out)
}
pub fn bic(&self) -> f64 {
-self.loglik + 0.5 * self.num_free_parameters() as f64 * (self.n_obs as f64).ln()
}
}
#[derive(Debug, Clone)]
struct GaussianComponentEval {
residual_origin: Array1<f64>,
residual_scale: Array1<f64>,
residual_normalized_offset: Array1<f64>,
precision: Array2<f64>,
log_norm: f64,
d: usize,
}
impl GaussianComponentEval {
fn factor(mean: ArrayView1<'_, f64>, cov: &Array2<f64>) -> Result<Self, String> {
let d = mean.len();
if mean.iter().any(|value| !value.is_finite()) {
return Err("mixture component mean must be finite".to_string());
}
if cov.nrows() != d || cov.ncols() != d {
return Err(format!(
"mixture component covariance must be {d}x{d}, got {}x{}",
cov.nrows(),
cov.ncols()
));
}
let (evals, evecs) = cov
.eigh(Side::Lower)
.map_err(|e| format!("mixture component covariance eigendecomposition failed: {e}"))?;
let mut log_det = 0.0_f64;
let mut inv_evals = Array1::<f64>::zeros(d);
for (idx, &ev) in evals.iter().enumerate() {
if !ev.is_finite() || ev <= 0.0 {
return Err(format!(
"mixture component covariance is not SPD: eigenvalue {idx} is {ev:.3e}"
));
}
log_det += ev.ln();
let inverse = ev.recip();
if !inverse.is_finite() {
return Err(format!(
"mixture component precision is not representable: eigenvalue {idx} is {ev:.3e}"
));
}
inv_evals[idx] = inverse;
}
let mut precision = Array2::<f64>::zeros((d, d));
for a in 0..d {
for b in 0..d {
let mut acc = 0.0_f64;
for m in 0..d {
acc += evecs[[a, m]] * inv_evals[m] * evecs[[b, m]];
}
precision[[a, b]] = acc;
}
}
let log_norm = -0.5 * (d as f64 * (2.0 * std::f64::consts::PI).ln() + log_det);
if precision.iter().any(|value| !value.is_finite()) || !log_norm.is_finite() {
return Err(
"mixture component factorization produced non-finite precision or log normalizer"
.to_string(),
);
}
Ok(Self {
residual_origin: mean.to_owned(),
residual_scale: Array1::zeros(d),
residual_normalized_offset: Array1::zeros(d),
precision,
log_norm,
d,
})
}
fn isotropic(charts: &[StableScalarMeanChart], variance: f64) -> Result<Self, String> {
let d = charts.len();
if d == 0 {
return Err("isotropic Gaussian density requires positive dimension".to_string());
}
if !(variance.is_finite() && variance > 0.0) {
return Err(format!(
"isotropic Gaussian variance must be finite and positive, got {variance}"
));
}
let inverse_variance = variance.recip();
if !inverse_variance.is_finite() {
return Err(format!(
"isotropic Gaussian precision is non-finite for variance {variance}"
));
}
let mut precision = Array2::<f64>::zeros((d, d));
for axis in 0..d {
precision[[axis, axis]] = inverse_variance;
}
let log_norm = -0.5 * d as f64 * ((2.0 * std::f64::consts::PI).ln() + variance.ln());
if !log_norm.is_finite() {
return Err("isotropic Gaussian log normalizer is non-finite".to_string());
}
Ok(Self {
residual_origin: Array1::from_iter(charts.iter().map(|chart| chart.origin)),
residual_scale: Array1::from_iter(charts.iter().map(|chart| chart.scale)),
residual_normalized_offset: Array1::from_iter(
charts.iter().map(|chart| chart.normalized_offset),
),
precision,
log_norm,
d,
})
}
#[inline]
fn log_density(&self, y: ArrayView1<'_, f64>) -> f64 {
let residual = self.residual(y);
let pv = self.precision_times_residual(&residual);
let mut quad = 0.0_f64;
for c in 0..self.d {
quad += residual[c] * pv[c];
}
self.log_norm - 0.5 * quad
}
#[inline]
fn residual(&self, y: ArrayView1<'_, f64>) -> Vec<f64> {
let mut residual = vec![0.0_f64; self.d];
for axis in 0..self.d {
residual[axis] = (-self.residual_normalized_offset[axis]).mul_add(
self.residual_scale[axis],
y[axis] - self.residual_origin[axis],
);
}
residual
}
#[inline]
fn precision_times_residual(&self, residual: &[f64]) -> Vec<f64> {
let mut out = vec![0.0_f64; self.d];
for a in 0..self.d {
let mut acc = 0.0_f64;
for b in 0..self.d {
acc += self.precision[[a, b]] * residual[b];
}
out[a] = acc;
}
out
}
}
#[inline]
fn log_sum_exp(terms: &[f64], max_term: f64) -> f64 {
if !max_term.is_finite() {
return f64::NEG_INFINITY;
}
let mut acc = 0.0_f64;
for &t in terms {
acc += (t - max_term).exp();
}
max_term + acc.ln()
}
fn evidence_matrix_fingerprint(namespace: &str, values: ArrayView2<'_, f64>) -> Fingerprint {
let mut hasher = Fingerprinter::new();
hasher.write_str(namespace);
hasher.write_usize(values.nrows());
hasher.write_usize(values.ncols());
for &value in values {
hasher.write_f64(value);
}
hasher.finalize()
}
fn mixture_data_fingerprint(data: ArrayView2<'_, f64>) -> Fingerprint {
evidence_matrix_fingerprint("gaussian-mixture-em-v1", data)
}
pub fn fit_gaussian_mixture(
data: ArrayView2<'_, f64>,
k: usize,
config: GaussianMixtureConfig,
) -> Result<GaussianMixtureFit, GaussianMixtureError> {
validate_gaussian_mixture_problem(data, k, config)?;
let means = gam_terms::basis::select_centers_by_strategy(
data,
&gam_terms::basis::CenterStrategy::KMeans {
num_centers: k,
max_iter: config.kmeans_max_iter,
},
)
.map_err(|error| GaussianMixtureError::NumericalFailure {
message: format!("deterministic k-means seeding failed: {error}"),
checkpoint: None,
})?;
if means.nrows() != k || means.ncols() != data.ncols() {
return Err(GaussianMixtureError::NumericalFailure {
message: format!(
"seeding returned {}x{} centers, expected {k}x{}",
means.nrows(),
means.ncols(),
data.ncols()
),
checkpoint: None,
});
}
let global_covariance =
constrained_data_covariance(data, config.covariance_floor).map_err(|message| {
GaussianMixtureError::NumericalFailure {
message,
checkpoint: None,
}
})?;
let weights = Array1::<f64>::from_elem(k, 1.0 / k as f64);
let covariances = vec![global_covariance; k];
let initial_e_step =
mixture_e_step(data, &weights, &means, &covariances).map_err(|message| {
GaussianMixtureError::NumericalFailure {
message,
checkpoint: None,
}
})?;
let data_fingerprint = mixture_data_fingerprint(data);
let checkpoint = GaussianMixtureCheckpoint {
weights,
means,
covariances,
mean_log_likelihood: initial_e_step.mean_log_likelihood,
completed_iterations: 0,
data_fingerprint,
covariance_floor: config.covariance_floor,
};
run_gaussian_mixture_em(data, config, checkpoint)
}
pub fn resume_gaussian_mixture(
data: ArrayView2<'_, f64>,
config: GaussianMixtureConfig,
checkpoint: GaussianMixtureCheckpoint,
) -> Result<GaussianMixtureFit, GaussianMixtureError> {
let k = checkpoint.weights.len();
validate_gaussian_mixture_problem(data, k, config)?;
validate_gaussian_mixture_checkpoint(data, config.covariance_floor, &checkpoint)?;
run_gaussian_mixture_em(data, config, checkpoint)
}
fn validate_gaussian_mixture_problem(
data: ArrayView2<'_, f64>,
k: usize,
config: GaussianMixtureConfig,
) -> Result<(), GaussianMixtureError> {
let n = data.nrows();
let d = data.ncols();
if k == 0 {
return Err(GaussianMixtureError::InvalidInput {
message: "k must be positive".to_string(),
});
}
if d == 0 {
return Err(GaussianMixtureError::InvalidInput {
message: "at least one data column is required".to_string(),
});
}
if k > n {
return Err(GaussianMixtureError::InvalidInput {
message: format!("requested {k} components but data has {n} rows"),
});
}
if data.iter().any(|value| !value.is_finite()) {
return Err(GaussianMixtureError::InvalidInput {
message: "data must be finite".to_string(),
});
}
if config.max_iter == 0 || config.kmeans_max_iter == 0 {
return Err(GaussianMixtureError::InvalidInput {
message: "max_iter and kmeans_max_iter must be positive".to_string(),
});
}
let numerical_floor = f64::EPSILON.sqrt();
if !config.loglik_tol.is_finite()
|| config.loglik_tol < numerical_floor
|| !config.parameter_tol.is_finite()
|| config.parameter_tol < numerical_floor
|| !config.covariance_floor.is_finite()
|| config.covariance_floor <= 0.0
{
return Err(GaussianMixtureError::InvalidInput {
message: format!(
"loglik_tol and parameter_tol must be finite and >= {numerical_floor:.3e}, and covariance_floor must be finite and positive"
),
});
}
Ok(())
}
fn validate_gaussian_mixture_checkpoint(
data: ArrayView2<'_, f64>,
covariance_floor: f64,
checkpoint: &GaussianMixtureCheckpoint,
) -> Result<(), GaussianMixtureError> {
let d = data.ncols();
let k = checkpoint.weights.len();
let mass = checkpoint.weights.sum();
if k == 0
|| checkpoint.data_fingerprint != mixture_data_fingerprint(data)
|| checkpoint.covariance_floor.to_bits() != covariance_floor.to_bits()
|| checkpoint.means.dim() != (k, d)
|| checkpoint.covariances.len() != k
|| checkpoint
.covariances
.iter()
.any(|covariance| covariance.dim() != (d, d))
|| checkpoint
.weights
.iter()
.chain(checkpoint.means.iter())
.chain(checkpoint.covariances.iter().flat_map(|value| value.iter()))
.any(|value| !value.is_finite())
|| checkpoint.weights.iter().any(|value| *value <= 0.0)
|| !mass.is_finite()
|| (mass - 1.0).abs() > f64::EPSILON.sqrt()
|| !checkpoint.mean_log_likelihood.is_finite()
{
return Err(GaussianMixtureError::InvalidInput {
message: "checkpoint problem identity, dimensions, interior parameters, likelihood, or simplex mass are invalid".to_string(),
});
}
Ok(())
}
fn run_gaussian_mixture_em(
data: ArrayView2<'_, f64>,
config: GaussianMixtureConfig,
mut checkpoint: GaussianMixtureCheckpoint,
) -> Result<GaussianMixtureFit, GaussianMixtureError> {
validate_gaussian_mixture_checkpoint(data, config.covariance_floor, &checkpoint)?;
let k = checkpoint.weights.len();
let d = data.ncols();
let data_fingerprint = mixture_data_fingerprint(data);
for additional_updates in 0..=config.max_iter {
let current = mixture_e_step(
data,
&checkpoint.weights,
&checkpoint.means,
&checkpoint.covariances,
)
.map_err(|message| GaussianMixtureError::NumericalFailure {
message,
checkpoint: Some(checkpoint.clone()),
})?;
if (checkpoint.mean_log_likelihood - current.mean_log_likelihood).abs()
> current.mean_log_likelihood_roundoff
{
return Err(GaussianMixtureError::InvalidInput {
message: format!(
"checkpoint mean log likelihood {:.12e} disagrees with its parameters ({:.12e} +/- {:.3e})",
checkpoint.mean_log_likelihood,
current.mean_log_likelihood,
current.mean_log_likelihood_roundoff
),
});
}
checkpoint.mean_log_likelihood = current.mean_log_likelihood;
let (next_weights, next_means, next_covariances) = mixture_m_step(
data,
current.responsibilities.view(),
config.covariance_floor,
)
.map_err(|message| GaussianMixtureError::NumericalFailure {
message,
checkpoint: Some(checkpoint.clone()),
})?;
let next = mixture_e_step(data, &next_weights, &next_means, &next_covariances).map_err(
|message| GaussianMixtureError::NumericalFailure {
message,
checkpoint: Some(checkpoint.clone()),
},
)?;
let objective_scale = current
.mean_log_likelihood
.abs()
.max(next.mean_log_likelihood.abs())
.max(1.0);
let objective_step = next.mean_log_likelihood - current.mean_log_likelihood;
let objective_residual = objective_step.abs() / objective_scale;
let parameter_residual = mixture_parameter_residual(
&checkpoint.weights,
&checkpoint.means,
&checkpoint.covariances,
&next_weights,
&next_means,
&next_covariances,
);
let monotonicity_uncertainty = gaussian_mixture_monotonicity_uncertainty(
objective_scale,
current.mean_log_likelihood_roundoff,
next.mean_log_likelihood_roundoff,
);
let certificate = GaussianMixtureCertificate {
mean_log_likelihood: current.mean_log_likelihood,
mean_log_likelihood_gain: objective_step,
monotonicity_uncertainty,
objective_residual,
objective_tolerance: config.loglik_tol,
parameter_residual,
parameter_tolerance: config.parameter_tol,
};
if objective_step < -monotonicity_uncertainty {
return Err(GaussianMixtureError::MonotonicityViolation {
previous_mean_log_likelihood: current.mean_log_likelihood,
next_mean_log_likelihood: next.mean_log_likelihood,
numerical_uncertainty: monotonicity_uncertainty,
checkpoint,
});
}
if objective_residual <= config.loglik_tol && parameter_residual <= config.parameter_tol {
let loglik = current.mean_log_likelihood * data.nrows() as f64;
if !loglik.is_finite() {
return Err(GaussianMixtureError::NumericalFailure {
message: "certified mean log likelihood overflows as a total likelihood"
.to_string(),
checkpoint: Some(checkpoint),
});
}
return Ok(GaussianMixtureFit {
weights: checkpoint.weights,
means: checkpoint.means,
covariances: checkpoint.covariances,
k,
d,
n_obs: data.nrows(),
loglik,
iterations: checkpoint.completed_iterations,
certificate,
});
}
if additional_updates == config.max_iter {
return Err(GaussianMixtureError::DidNotConverge {
max_iterations: config.max_iter,
certificate,
checkpoint,
});
}
checkpoint = GaussianMixtureCheckpoint {
weights: next_weights,
means: next_means,
covariances: next_covariances,
mean_log_likelihood: next.mean_log_likelihood,
completed_iterations: checkpoint.completed_iterations + 1,
data_fingerprint,
covariance_floor: config.covariance_floor,
};
}
Err(GaussianMixtureError::NumericalFailure {
message: format!(
"EM refinement exhausted its inclusive update budget ({}) without producing a \
terminal verdict",
config.max_iter
),
checkpoint: Some(checkpoint),
})
}
struct GaussianMixtureEStep {
responsibilities: Array2<f64>,
mean_log_likelihood: f64,
mean_log_likelihood_roundoff: f64,
}
fn gaussian_mixture_monotonicity_uncertainty(
objective_scale: f64,
current_reduction_roundoff: f64,
next_reduction_roundoff: f64,
) -> f64 {
let reduction_roundoff = current_reduction_roundoff + next_reduction_roundoff;
let composite_map_resolution = f64::EPSILON.sqrt() * objective_scale;
reduction_roundoff.max(composite_map_resolution)
}
fn pairwise_sum_max_depth(term_count: usize) -> usize {
if term_count <= 1 {
return 0;
}
let within_block = term_count.min(BASE_CHUNK) - 1;
let blocks = term_count.div_ceil(BASE_CHUNK);
let tree_levels = if blocks <= 1 {
0
} else {
(usize::BITS - (blocks - 1).leading_zeros()) as usize
};
within_block.saturating_add(tree_levels)
}
fn pairwise_mean_with_roundoff(mut values: Vec<f64>) -> Result<(f64, f64), String> {
if values.is_empty() || values.iter().any(|value| !value.is_finite()) {
return Err("mean log-likelihood terms must be nonempty and finite".to_string());
}
let sum = pairwise_sum(&values);
for value in &mut values {
*value = value.abs();
}
let magnitude_sum = pairwise_sum(&values);
let unit_roundoff = 0.5 * f64::EPSILON;
let accumulated = pairwise_sum_max_depth(values.len()) as f64 * unit_roundoff;
let addition_bound = if accumulated < 1.0 {
accumulated / (1.0 - accumulated) * magnitude_sum
} else {
f64::INFINITY
};
let count = values.len() as f64;
let mean = sum / count;
let roundoff = addition_bound / count + unit_roundoff * mean.abs();
if !(mean.is_finite() && roundoff.is_finite()) {
return Err("mean mixture log likelihood or its rounding bound is non-finite".to_string());
}
Ok((mean, roundoff))
}
fn mixture_e_step(
data: ArrayView2<'_, f64>,
weights: &Array1<f64>,
means: &Array2<f64>,
covariances: &[Array2<f64>],
) -> Result<GaussianMixtureEStep, String> {
let n = data.nrows();
let k = weights.len();
if weights
.iter()
.any(|weight| !weight.is_finite() || *weight <= 0.0)
{
return Err("mixture E-step requires strictly positive finite weights".to_string());
}
let mut components = Vec::with_capacity(k);
for component in 0..k {
components.push(GaussianComponentEval::factor(
means.row(component),
&covariances[component],
)?);
}
let log_weights: Vec<f64> = weights.iter().map(|weight| weight.ln()).collect();
let mut responsibilities = Array2::<f64>::zeros((n, k));
let mut row_log_likelihoods = Vec::with_capacity(n);
for row in 0..n {
let observation = data.row(row);
let mut log_terms = vec![f64::NEG_INFINITY; k];
let mut max_term = f64::NEG_INFINITY;
for component in 0..k {
let term = log_weights[component] + components[component].log_density(observation);
log_terms[component] = term;
max_term = max_term.max(term);
}
let log_mixture = log_sum_exp(&log_terms, max_term);
if !log_mixture.is_finite() {
return Err(format!(
"mixture density is non-finite at training row {row}"
));
}
row_log_likelihoods.push(log_mixture);
for component in 0..k {
responsibilities[[row, component]] = (log_terms[component] - log_mixture).exp();
}
}
let (mean_log_likelihood, mean_log_likelihood_roundoff) =
pairwise_mean_with_roundoff(row_log_likelihoods)?;
Ok(GaussianMixtureEStep {
responsibilities,
mean_log_likelihood,
mean_log_likelihood_roundoff,
})
}
fn mixture_m_step(
data: ArrayView2<'_, f64>,
responsibilities: ArrayView2<'_, f64>,
covariance_floor: f64,
) -> Result<(Array1<f64>, Array2<f64>, Vec<Array2<f64>>), String> {
let n = data.nrows();
let d = data.ncols();
let k = responsibilities.ncols();
let mut component_mass = Array1::<f64>::zeros(k);
for component in 0..k {
component_mass[component] = responsibilities.column(component).sum();
}
if component_mass
.iter()
.any(|mass| !mass.is_finite() || *mass <= 0.0)
{
return Err(
"M-step reached a zero-mass component; the requested mixture order has no interior fitted density"
.to_string(),
);
}
let mut weights = component_mass.mapv(|mass| mass / n as f64);
let total_weight = weights.sum();
if !(total_weight.is_finite() && total_weight > 0.0) {
return Err("M-step produced invalid mixture-weight mass".to_string());
}
weights.mapv_inplace(|weight| weight / total_weight);
let mut means = Array2::<f64>::zeros((k, d));
let mut covariances = Vec::with_capacity(k);
for component in 0..k {
let mass = component_mass[component];
let mut mean = Array1::<f64>::zeros(d);
for row in 0..n {
let responsibility = responsibilities[[row, component]];
for col in 0..d {
mean[col] += responsibility * data[[row, col]];
}
}
mean.mapv_inplace(|value| value / mass);
means.row_mut(component).assign(&mean);
let mut covariance = Array2::<f64>::zeros((d, d));
for row in 0..n {
let responsibility = responsibilities[[row, component]];
for left in 0..d {
let left_residual = data[[row, left]] - mean[left];
for right in 0..d {
covariance[[left, right]] +=
responsibility * left_residual * (data[[row, right]] - mean[right]);
}
}
}
covariance.mapv_inplace(|value| value / mass);
covariances.push(constrain_covariance(covariance, covariance_floor)?);
}
Ok((weights, means, covariances))
}
fn relative_parameter_step(previous: f64, next: f64) -> f64 {
(next - previous).abs() / previous.abs().max(next.abs()).max(1.0)
}
fn labeled_gaussian_component_measure_residual(
previous_weights: &Array1<f64>,
previous_means: &Array2<f64>,
previous_covariance: impl Fn(usize, usize, usize) -> f64,
next_weights: &Array1<f64>,
next_means: &Array2<f64>,
next_covariance: impl Fn(usize, usize, usize) -> f64,
) -> f64 {
let k = previous_weights.len();
let d = previous_means.ncols();
let mut residual = 0.0_f64;
for component in 0..k {
let previous_weight = previous_weights[component];
let next_weight = next_weights[component];
residual = residual.max(relative_parameter_step(previous_weight, next_weight));
for left in 0..d {
let previous_first = previous_weight * previous_means[[component, left]];
let next_first = next_weight * next_means[[component, left]];
residual = residual.max(relative_parameter_step(previous_first, next_first));
for right in 0..d {
let previous_second = previous_weight
* (previous_covariance(component, left, right)
+ previous_means[[component, left]] * previous_means[[component, right]]);
let next_second = next_weight
* (next_covariance(component, left, right)
+ next_means[[component, left]] * next_means[[component, right]]);
residual = residual.max(relative_parameter_step(previous_second, next_second));
}
}
}
residual
}
fn mixture_parameter_residual(
previous_weights: &Array1<f64>,
previous_means: &Array2<f64>,
previous_covariances: &[Array2<f64>],
next_weights: &Array1<f64>,
next_means: &Array2<f64>,
next_covariances: &[Array2<f64>],
) -> f64 {
labeled_gaussian_component_measure_residual(
previous_weights,
previous_means,
|component, left, right| previous_covariances[component][[left, right]],
next_weights,
next_means,
|component, left, right| next_covariances[component][[left, right]],
)
}
fn constrain_covariance(covariance: Array2<f64>, floor: f64) -> Result<Array2<f64>, String> {
let (eigenvalues, eigenvectors) = covariance
.eigh(Side::Lower)
.map_err(|error| format!("covariance eigendecomposition failed: {error}"))?;
let d = covariance.nrows();
let mut constrained = Array2::<f64>::zeros((d, d));
for row in 0..d {
for col in 0..d {
let mut value = 0.0_f64;
for index in 0..d {
value += eigenvectors[[row, index]]
* eigenvalues[index].max(floor)
* eigenvectors[[col, index]];
}
constrained[[row, col]] = value;
}
}
if constrained.iter().any(|value| !value.is_finite()) {
return Err("constrained covariance became non-finite".to_string());
}
Ok(constrained)
}
fn constrained_data_covariance(
data: ArrayView2<'_, f64>,
floor: f64,
) -> Result<Array2<f64>, String> {
let n = data.nrows();
let d = data.ncols();
let mut mean = Array1::<f64>::zeros(d);
for i in 0..n {
for c in 0..d {
mean[c] += data[[i, c]];
}
}
mean.mapv_inplace(|v| v / n.max(1) as f64);
let mut cov = Array2::<f64>::zeros((d, d));
for i in 0..n {
for a in 0..d {
let da = data[[i, a]] - mean[a];
for b in 0..d {
cov[[a, b]] += da * (data[[i, b]] - mean[b]);
}
}
}
let inv = 1.0 / n as f64;
cov.mapv_inplace(|v| v * inv);
constrain_covariance(cov, floor)
}
#[derive(Debug, Clone)]
pub struct RingGaussianMixtureFit {
weights: Array1<f64>,
center: Array1<f64>,
radius: f64,
directions: Array2<f64>,
variance: f64,
k: usize,
n_obs: usize,
loglik: f64,
iterations: usize,
certificate: GaussianMixtureCertificate,
}
impl RingGaussianMixtureFit {
pub fn weights(&self) -> ArrayView1<'_, f64> {
self.weights.view()
}
pub fn center(&self) -> ArrayView1<'_, f64> {
self.center.view()
}
pub fn radius(&self) -> f64 {
self.radius
}
pub fn directions(&self) -> ArrayView2<'_, f64> {
self.directions.view()
}
pub fn variance(&self) -> f64 {
self.variance
}
pub fn iterations(&self) -> usize {
self.iterations
}
pub fn certificate(&self) -> GaussianMixtureCertificate {
self.certificate
}
pub fn num_free_parameters(&self) -> usize {
2 * self.k + 3
}
pub fn per_point_log_density(&self, data: ArrayView2<'_, f64>) -> Result<Array1<f64>, String> {
if data.ncols() != 2 {
return Err(format!(
"ring-of-clusters density expects two columns, got {}",
data.ncols()
));
}
ring_mixture_log_density(
data,
&self.weights,
&self.center,
self.radius,
&self.directions,
self.variance,
)
}
pub fn bic(&self) -> f64 {
-self.loglik + 0.5 * self.num_free_parameters() as f64 * (self.n_obs as f64).ln()
}
}
#[derive(Debug, Clone)]
struct RingMixtureState {
weights: Array1<f64>,
center: Array1<f64>,
radius: f64,
directions: Array2<f64>,
variance: f64,
mean_log_likelihood: f64,
completed_iterations: usize,
}
fn ring_component_means(
center: &Array1<f64>,
radius: f64,
directions: &Array2<f64>,
) -> Array2<f64> {
let mut means = Array2::<f64>::zeros((directions.nrows(), 2));
for component in 0..directions.nrows() {
means[[component, 0]] = center[0] + radius * directions[[component, 0]];
means[[component, 1]] = center[1] + radius * directions[[component, 1]];
}
means
}
fn ring_mixture_log_terms(
data: ArrayView2<'_, f64>,
weights: &Array1<f64>,
center: &Array1<f64>,
radius: f64,
directions: &Array2<f64>,
variance: f64,
) -> Result<(Array2<f64>, Vec<f64>), String> {
if data.ncols() != 2
|| center.len() != 2
|| directions.ncols() != 2
|| directions.nrows() != weights.len()
|| weights
.iter()
.any(|weight| !weight.is_finite() || *weight <= 0.0)
|| !(radius.is_finite() && radius > 0.0)
|| !(variance.is_finite() && variance > 0.0)
{
return Err("invalid ring-of-clusters parameter state".to_string());
}
let means = ring_component_means(center, radius, directions);
let log_normalizer = -(std::f64::consts::TAU).ln() - variance.ln();
let mut terms = Array2::<f64>::zeros((data.nrows(), weights.len()));
let mut row_log_likelihoods = Vec::with_capacity(data.nrows());
for row in 0..data.nrows() {
let mut max_term = f64::NEG_INFINITY;
for component in 0..weights.len() {
let dx = data[[row, 0]] - means[[component, 0]];
let dy = data[[row, 1]] - means[[component, 1]];
let term =
weights[component].ln() + log_normalizer - 0.5 * (dx * dx + dy * dy) / variance;
terms[[row, component]] = term;
max_term = max_term.max(term);
}
let values = terms.row(row).to_vec();
let log_likelihood = log_sum_exp(&values, max_term);
if !log_likelihood.is_finite() {
return Err(format!(
"ring-of-clusters density is non-finite at training row {row}"
));
}
row_log_likelihoods.push(log_likelihood);
}
Ok((terms, row_log_likelihoods))
}
fn ring_mixture_e_step(
data: ArrayView2<'_, f64>,
state: &RingMixtureState,
) -> Result<(Array2<f64>, f64, f64), String> {
let (terms, row_log_likelihoods) = ring_mixture_log_terms(
data,
&state.weights,
&state.center,
state.radius,
&state.directions,
state.variance,
)?;
let mut responsibilities = Array2::<f64>::zeros(terms.raw_dim());
for row in 0..terms.nrows() {
for component in 0..terms.ncols() {
responsibilities[[row, component]] =
(terms[[row, component]] - row_log_likelihoods[row]).exp();
}
}
let (mean, roundoff) = pairwise_mean_with_roundoff(row_log_likelihoods)?;
Ok((responsibilities, mean, roundoff))
}
fn ring_mixture_log_density(
data: ArrayView2<'_, f64>,
weights: &Array1<f64>,
center: &Array1<f64>,
radius: f64,
directions: &Array2<f64>,
variance: f64,
) -> Result<Array1<f64>, String> {
let (_, row_log_likelihoods) =
ring_mixture_log_terms(data, weights, center, radius, directions, variance)?;
Ok(Array1::from_vec(row_log_likelihoods))
}
fn ring_identifiable_parameter_residual(
previous: &RingMixtureState,
next: &RingMixtureState,
) -> f64 {
let previous_means =
ring_component_means(&previous.center, previous.radius, &previous.directions);
let next_means = ring_component_means(&next.center, next.radius, &next.directions);
let noise_scale = previous
.variance
.sqrt()
.max(next.variance.sqrt())
.max(f64::MIN_POSITIVE);
let weight_residual = previous
.weights
.iter()
.zip(next.weights.iter())
.map(|(&left, &right)| (right - left).abs())
.fold(0.0, f64::max);
let mean_residual = previous_means
.rows()
.into_iter()
.zip(next_means.rows())
.zip(previous.weights.iter().zip(next.weights.iter()))
.map(|((left, right), (&previous_weight, &next_weight))| {
previous_weight.max(next_weight) * (right[0] - left[0]).hypot(right[1] - left[1])
/ noise_scale
})
.fold(0.0, f64::max);
let variance_residual = (next.variance / previous.variance).ln().abs();
weight_residual.max(mean_residual).max(variance_residual)
}
fn fit_weighted_component_circle(
component_means: &Array2<f64>,
component_mass: &Array1<f64>,
initial_center: &Array1<f64>,
initial_radius: f64,
parameter_tol: f64,
max_iter: usize,
) -> Result<(Array1<f64>, f64, Array2<f64>), String> {
let k = component_means.nrows();
let total_mass = component_mass.sum();
if component_means.ncols() != 2
|| component_mass.len() != k
|| component_mass
.iter()
.any(|mass| !mass.is_finite() || *mass <= 0.0)
|| !(total_mass.is_finite() && total_mass > 0.0)
{
return Err("ring M-step requires positive component masses and 2-D means".to_string());
}
let mut center = initial_center.clone();
let mut radius = initial_radius;
let mut directions = Array2::<f64>::zeros((k, 2));
for _ in 0..max_iter {
for component in 0..k {
let dx = component_means[[component, 0]] - center[0];
let dy = component_means[[component, 1]] - center[1];
let norm = dx.hypot(dy);
if !(norm.is_finite() && norm > 0.0) {
return Err(
"ring M-step reached a component centroid at the circle center; its angle is unidentified"
.to_string(),
);
}
directions[[component, 0]] = dx / norm;
directions[[component, 1]] = dy / norm;
}
let mut mean_point = Array1::<f64>::zeros(2);
let mut mean_direction = Array1::<f64>::zeros(2);
for component in 0..k {
let weight = component_mass[component] / total_mass;
for axis in 0..2 {
mean_point[axis] += weight * component_means[[component, axis]];
mean_direction[axis] += weight * directions[[component, axis]];
}
}
let mut numerator = 0.0;
let mut denominator = 0.0;
for component in 0..k {
let mass = component_mass[component];
let dux = directions[[component, 0]] - mean_direction[0];
let duy = directions[[component, 1]] - mean_direction[1];
numerator += mass
* (dux * (component_means[[component, 0]] - mean_point[0])
+ duy * (component_means[[component, 1]] - mean_point[1]));
denominator += mass * (dux * dux + duy * duy);
}
if !(denominator.is_finite() && denominator > 0.0) {
return Err(
"ring M-step component directions are identical; radius and center are unidentified"
.to_string(),
);
}
let mut next_radius = numerator / denominator;
if !next_radius.is_finite() || next_radius == 0.0 {
return Err("ring M-step produced an unidentified zero radius".to_string());
}
if next_radius < 0.0 {
next_radius = -next_radius;
directions.mapv_inplace(|value| -value);
}
let next_center = Array1::from_vec(vec![
mean_point[0] - next_radius * mean_direction[0],
mean_point[1] - next_radius * mean_direction[1],
]);
let residual = center
.iter()
.zip(next_center.iter())
.map(|(&left, &right)| relative_parameter_step(left, right))
.chain(std::iter::once(relative_parameter_step(
radius,
next_radius,
)))
.fold(0.0, f64::max);
center = next_center;
radius = next_radius;
if residual <= parameter_tol {
for component in 0..k {
let dx = component_means[[component, 0]] - center[0];
let dy = component_means[[component, 1]] - center[1];
let norm = dx.hypot(dy);
if !(norm.is_finite() && norm > 0.0) {
return Err("ring M-step terminal component angle is unidentified".to_string());
}
directions[[component, 0]] = dx / norm;
directions[[component, 1]] = dy / norm;
}
return Ok((center, radius, directions));
}
}
Err(format!(
"ring M-step did not certify its constrained center/radius fixed point after {max_iter} iterations"
))
}
fn ring_mixture_m_step(
data: ArrayView2<'_, f64>,
responsibilities: ArrayView2<'_, f64>,
previous: &RingMixtureState,
config: GaussianMixtureConfig,
) -> Result<RingMixtureState, String> {
let n = data.nrows();
let k = responsibilities.ncols();
let mut component_mass = Array1::<f64>::zeros(k);
let mut component_means = Array2::<f64>::zeros((k, 2));
for component in 0..k {
let mass = responsibilities.column(component).sum();
if !(mass.is_finite() && mass > 0.0) {
return Err(
"ring M-step reached a zero-mass component; the requested order is singular"
.to_string(),
);
}
component_mass[component] = mass;
for row in 0..n {
for axis in 0..2 {
component_means[[component, axis]] +=
responsibilities[[row, component]] * data[[row, axis]];
}
}
for axis in 0..2 {
component_means[[component, axis]] /= mass;
}
}
let mut weights = component_mass.mapv(|mass| mass / n as f64);
let weight_sum = weights.sum();
weights.mapv_inplace(|weight| weight / weight_sum);
let (center, radius, directions) = fit_weighted_component_circle(
&component_means,
&component_mass,
&previous.center,
previous.radius,
config.parameter_tol,
config.max_iter,
)?;
let means = ring_component_means(¢er, radius, &directions);
let mut expected_squared_error = 0.0;
for row in 0..n {
for component in 0..k {
let dx = data[[row, 0]] - means[[component, 0]];
let dy = data[[row, 1]] - means[[component, 1]];
expected_squared_error += responsibilities[[row, component]] * (dx * dx + dy * dy);
}
}
let variance = (expected_squared_error / (2 * n) as f64).max(config.covariance_floor);
if !variance.is_finite() {
return Err("ring M-step produced non-finite shared variance".to_string());
}
Ok(RingMixtureState {
weights,
center,
radius,
directions,
variance,
mean_log_likelihood: f64::NAN,
completed_iterations: previous.completed_iterations + 1,
})
}
pub fn fit_ring_gaussian_mixture(
data: ArrayView2<'_, f64>,
k: usize,
config: GaussianMixtureConfig,
) -> Result<RingGaussianMixtureFit, String> {
validate_gaussian_mixture_problem(data, k, config).map_err(|error| error.to_string())?;
if data.ncols() != 2 {
return Err(format!(
"ring-of-clusters fitting requires exactly two columns, got {}",
data.ncols()
));
}
if k < 3 {
return Err(format!(
"ring-of-clusters fitting requires at least three component centers, got {k}"
));
}
let seeded_means = gam_terms::basis::select_centers_by_strategy(
data,
&gam_terms::basis::CenterStrategy::KMeans {
num_centers: k,
max_iter: config.kmeans_max_iter,
},
)
.map_err(|error| format!("ring-of-clusters deterministic seeding failed: {error}"))?;
let component_mass = Array1::<f64>::ones(k);
let mut initial_center = Array1::<f64>::zeros(2);
for component in 0..k {
initial_center[0] += seeded_means[[component, 0]] / k as f64;
initial_center[1] += seeded_means[[component, 1]] / k as f64;
}
let mut initial_radius = 0.0;
for component in 0..k {
initial_radius += (seeded_means[[component, 0]] - initial_center[0])
.hypot(seeded_means[[component, 1]] - initial_center[1])
/ k as f64;
}
if !(initial_radius.is_finite() && initial_radius > 0.0) {
return Err("ring-of-clusters seed has an unidentified zero radius".to_string());
}
let (center, radius, directions) = fit_weighted_component_circle(
&seeded_means,
&component_mass,
&initial_center,
initial_radius,
config.parameter_tol,
config.max_iter,
)?;
let means = ring_component_means(¢er, radius, &directions);
let mut squared_error = 0.0;
for row in 0..data.nrows() {
let mut nearest = f64::INFINITY;
for component in 0..k {
let dx = data[[row, 0]] - means[[component, 0]];
let dy = data[[row, 1]] - means[[component, 1]];
nearest = nearest.min(dx * dx + dy * dy);
}
squared_error += nearest;
}
let variance = (squared_error / (2 * data.nrows()) as f64).max(config.covariance_floor);
let mut state = RingMixtureState {
weights: Array1::from_elem(k, 1.0 / k as f64),
center,
radius,
directions,
variance,
mean_log_likelihood: f64::NAN,
completed_iterations: 0,
};
for additional_updates in 0..=config.max_iter {
let (responsibilities, current_mean, current_roundoff) = ring_mixture_e_step(data, &state)?;
state.mean_log_likelihood = current_mean;
let mut next = ring_mixture_m_step(data, responsibilities.view(), &state, config)?;
let (_, next_mean, next_roundoff) = ring_mixture_e_step(data, &next)?;
next.mean_log_likelihood = next_mean;
let objective_scale = current_mean.abs().max(next_mean.abs()).max(1.0);
let objective_step = next_mean - current_mean;
let objective_residual = objective_step.abs() / objective_scale;
let parameter_residual = ring_identifiable_parameter_residual(&state, &next);
let monotonicity_uncertainty = gaussian_mixture_monotonicity_uncertainty(
objective_scale,
current_roundoff,
next_roundoff,
);
let certificate = GaussianMixtureCertificate {
mean_log_likelihood: current_mean,
mean_log_likelihood_gain: objective_step,
monotonicity_uncertainty,
objective_residual,
objective_tolerance: config.loglik_tol,
parameter_residual,
parameter_tolerance: config.parameter_tol,
};
if objective_step < -monotonicity_uncertainty {
return Err(format!(
"ring-of-clusters generalized EM violated monotone ascent at iteration {}: {current_mean:.12e} -> {next_mean:.12e} (comparison uncertainty {monotonicity_uncertainty:.3e})",
state.completed_iterations
));
}
if objective_residual <= config.loglik_tol && parameter_residual <= config.parameter_tol {
let loglik = current_mean * data.nrows() as f64;
if !loglik.is_finite() {
return Err("ring-of-clusters total log likelihood overflowed".to_string());
}
return Ok(RingGaussianMixtureFit {
weights: state.weights,
center: state.center,
radius: state.radius,
directions: state.directions,
variance: state.variance,
k,
n_obs: data.nrows(),
loglik,
iterations: state.completed_iterations,
certificate,
});
}
if additional_updates == config.max_iter {
return Err(format!(
"ring-of-clusters generalized EM did not certify after {} iterations: objective residual {:.6e}/{:.3e}, parameter-map residual {:.6e}/{:.3e}",
config.max_iter,
objective_residual,
config.loglik_tol,
parameter_residual,
config.parameter_tol,
));
}
state = next;
}
Err("ring-of-clusters generalized EM exhausted without a terminal certificate".to_string())
}
#[derive(Debug, Clone, Copy)]
pub struct CircularGaussianFit2d {
center: [f64; 2],
radius: f64,
noise_variance: f64,
}
impl CircularGaussianFit2d {
pub const NUM_FREE_PARAMETERS: usize = 4;
pub fn from_parameters(
center: [f64; 2],
radius: f64,
noise_variance: f64,
) -> Result<Self, String> {
if !center.iter().all(|value| value.is_finite()) {
return Err("circular Gaussian center must be finite".to_string());
}
if !(radius.is_finite() && radius >= 0.0) {
return Err("circular Gaussian radius must be finite and nonnegative".to_string());
}
if !(noise_variance.is_finite() && noise_variance > 0.0) {
return Err("circular Gaussian noise variance must be finite and positive".to_string());
}
Ok(Self {
center,
radius,
noise_variance,
})
}
pub fn fit(coords: ArrayView2<'_, f64>, rows: &[usize]) -> Result<Self, String> {
if coords.ncols() != 2 {
return Err(format!(
"circular Gaussian requires 2-D data, got {} columns",
coords.ncols()
));
}
if rows.is_empty() {
return Err("circular Gaussian requires a nonempty training set".to_string());
}
if rows.iter().any(|&row| row >= coords.nrows()) {
return Err("circular Gaussian row index is out of bounds".to_string());
}
if rows
.iter()
.any(|&row| !coords[[row, 0]].is_finite() || !coords[[row, 1]].is_finite())
{
return Err("circular Gaussian requires finite training coordinates".to_string());
}
let anchor_row = rows[0];
let anchor = [coords[[anchor_row, 0]], coords[[anchor_row, 1]]];
let mut scale = 0.0_f64;
for &row in rows {
let dx = coords[[row, 0]] - anchor[0];
let dy = coords[[row, 1]] - anchor[1];
if !(dx.is_finite() && dy.is_finite()) {
return Err("circular Gaussian coordinate range exceeds f64".to_string());
}
scale = scale.max(dx.hypot(dy));
}
if !(scale.is_finite() && scale > 0.0) {
return Err("circular Gaussian requires nonzero spatial extent".to_string());
}
let mut points = Vec::with_capacity(rows.len());
let mut mean = [0.0_f64; 2];
for &row in rows {
let point = [
(coords[[row, 0]] - anchor[0]) / scale,
(coords[[row, 1]] - anchor[1]) / scale,
];
points.push(point);
mean[0] += point[0];
mean[1] += point[1];
}
let count = rows.len() as f64;
mean[0] /= count;
mean[1] /= count;
let mut squared_radii = Vec::with_capacity(rows.len());
let mut mean_squared_radius = 0.0_f64;
for point in &points {
let dx = point[0] - mean[0];
let dy = point[1] - mean[1];
let squared_radius = dx * dx + dy * dy;
squared_radii.push(squared_radius);
mean_squared_radius += squared_radius;
}
mean_squared_radius /= count;
let mut squared_radius_variance = 0.0_f64;
for squared_radius in squared_radii {
squared_radius_variance += (squared_radius - mean_squared_radius).powi(2);
}
squared_radius_variance /= count;
let variance_floor = (64.0 * f64::EPSILON * mean_squared_radius).max(f64::MIN_POSITIVE);
let radius_squared = (mean_squared_radius * mean_squared_radius - squared_radius_variance)
.max(0.0)
.sqrt();
let mut radius = radius_squared.sqrt();
let mut noise_variance = (0.5 * (mean_squared_radius - radius_squared)).max(variance_floor);
let mut center = mean;
const MAX_EM_ITERATIONS: usize = 4096;
const EM_TOLERANCE: f64 = 2.0e-12;
let mut posterior_means = vec![[0.0_f64; 2]; points.len()];
let mut converged = false;
for _ in 0..MAX_EM_ITERATIONS {
let mut posterior_mean = [0.0_f64; 2];
for (point, latent_mean) in points.iter().zip(&mut posterior_means) {
let dx = point[0] - center[0];
let dy = point[1] - center[1];
let observed_radius = dx.hypot(dy);
if observed_radius == 0.0 || radius == 0.0 {
*latent_mean = [0.0, 0.0];
} else {
let (_, bessel_ratio) =
circular_gaussian_bessel_terms(radius, observed_radius, noise_variance);
if !(bessel_ratio.is_finite() && (0.0..=1.0).contains(&bessel_ratio)) {
return Err("circular Gaussian Bessel ratio left [0, 1]".to_string());
}
let multiplier = bessel_ratio / observed_radius;
*latent_mean = [multiplier * dx, multiplier * dy];
}
posterior_mean[0] += latent_mean[0];
posterior_mean[1] += latent_mean[1];
}
posterior_mean[0] /= count;
posterior_mean[1] /= count;
let denominator =
1.0 - posterior_mean[0] * posterior_mean[0] - posterior_mean[1] * posterior_mean[1];
if !(denominator.is_finite() && denominator > 0.0) {
return Err("circular Gaussian EM radius update is singular".to_string());
}
let mut radius_numerator = 0.0_f64;
for (point, latent_mean) in points.iter().zip(&posterior_means) {
radius_numerator +=
latent_mean[0] * (point[0] - mean[0]) + latent_mean[1] * (point[1] - mean[1]);
}
let next_radius = (radius_numerator / (count * denominator)).max(0.0);
let next_center = [
mean[0] - next_radius * posterior_mean[0],
mean[1] - next_radius * posterior_mean[1],
];
let mut residual_sum = 0.0_f64;
for (point, latent_mean) in points.iter().zip(&posterior_means) {
let dx = point[0] - next_center[0];
let dy = point[1] - next_center[1];
let ex = dx - next_radius * latent_mean[0];
let ey = dy - next_radius * latent_mean[1];
let latent_norm_squared =
latent_mean[0] * latent_mean[0] + latent_mean[1] * latent_mean[1];
residual_sum += ex * ex
+ ey * ey
+ next_radius * next_radius * (1.0 - latent_norm_squared).max(0.0);
}
let next_noise_variance = (residual_sum / (2.0 * count)).max(variance_floor);
let parameter_change = (next_center[0] - center[0])
.hypot(next_center[1] - center[1])
.max((next_radius - radius).abs())
.max(
(next_noise_variance - noise_variance).abs()
/ (next_noise_variance + noise_variance),
);
center = next_center;
radius = next_radius;
noise_variance = next_noise_variance;
if parameter_change <= EM_TOLERANCE {
converged = true;
break;
}
}
if !converged {
return Err("circular Gaussian maximum-likelihood fit did not converge".to_string());
}
let fitted_noise_sd = scale * noise_variance.sqrt();
Self::from_parameters(
[anchor[0] + scale * center[0], anchor[1] + scale * center[1]],
scale * radius,
fitted_noise_sd * fitted_noise_sd,
)
.map_err(|error| format!("circular Gaussian fit produced invalid parameters: {error}"))
}
pub const fn center(self) -> [f64; 2] {
self.center
}
pub const fn radius(self) -> f64 {
self.radius
}
pub const fn noise_variance(self) -> f64 {
self.noise_variance
}
pub fn log_density(self, x: f64, y: f64) -> f64 {
let observed_radius = (x - self.center[0]).hypot(y - self.center[1]);
let (log_i0_minus_kappa, _) =
circular_gaussian_bessel_terms(self.radius, observed_radius, self.noise_variance);
let standardized_radial_residual =
(observed_radius - self.radius) / self.noise_variance.sqrt();
-std::f64::consts::TAU.ln()
- self.noise_variance.ln()
- 0.5 * standardized_radial_residual.powi(2)
+ log_i0_minus_kappa
}
pub fn log_likelihood(
self,
coords: ArrayView2<'_, f64>,
rows: &[usize],
) -> Result<f64, String> {
if coords.ncols() != 2 || rows.iter().any(|&row| row >= coords.nrows()) {
return Err(
"circular Gaussian likelihood received invalid coordinates or rows".to_string(),
);
}
let mut log_densities = Vec::with_capacity(rows.len());
for &row in rows {
let value = self.log_density(coords[[row, 0]], coords[[row, 1]]);
if !value.is_finite() {
return Err("circular Gaussian likelihood is not finite".to_string());
}
log_densities.push(value);
}
let log_likelihood = pairwise_sum(&log_densities);
if !log_likelihood.is_finite() {
return Err("circular Gaussian likelihood sum is not finite".to_string());
}
Ok(log_likelihood)
}
pub fn fit_with_bic(
coords: ArrayView2<'_, f64>,
rows: &[usize],
) -> Result<(Self, f64), String> {
let fit = Self::fit(coords, rows)?;
let log_likelihood = fit.log_likelihood(coords, rows)?;
let bic =
-log_likelihood + 0.5 * Self::NUM_FREE_PARAMETERS as f64 * (rows.len() as f64).ln();
if !bic.is_finite() {
return Err("circular Gaussian BIC is not finite".to_string());
}
Ok((fit, bic))
}
}
fn circular_gaussian_bessel_terms(
radius: f64,
observed_radius: f64,
noise_variance: f64,
) -> (f64, f64) {
if radius == 0.0 || observed_radius == 0.0 {
return (0.0, 0.0);
}
let kappa = radius * observed_radius / noise_variance;
if kappa.is_finite() {
return bessel_i0_log_minus_abs_and_ratio(kappa);
}
let log_kappa = radius.ln() + observed_radius.ln() - noise_variance.ln();
if log_kappa <= f64::MAX.ln() {
return bessel_i0_log_minus_abs_and_ratio(log_kappa.exp());
}
(-0.5 * (std::f64::consts::TAU.ln() + log_kappa), 1.0)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum UnionStructure {
CircleCircle,
CirclePointCluster,
LineCluster,
}
pub const UNION_STRUCTURE_LADDER: &[UnionStructure] = &[
UnionStructure::CircleCircle,
UnionStructure::CirclePointCluster,
UnionStructure::LineCluster,
];
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum UnionComponentKind {
Circle,
Line,
PointCluster,
}
impl UnionStructure {
pub const fn as_str(self) -> &'static str {
match self {
UnionStructure::CircleCircle => "union_circle+circle",
UnionStructure::CirclePointCluster => "union_circle+cluster",
UnionStructure::LineCluster => "union_line+cluster",
}
}
pub const fn components(self) -> &'static [UnionComponentKind] {
match self {
UnionStructure::CircleCircle => {
&[UnionComponentKind::Circle, UnionComponentKind::Circle]
}
UnionStructure::CirclePointCluster => {
&[UnionComponentKind::Circle, UnionComponentKind::PointCluster]
}
UnionStructure::LineCluster => {
&[UnionComponentKind::Line, UnionComponentKind::PointCluster]
}
}
}
pub const fn num_components(self) -> usize {
self.components().len()
}
}
#[derive(Debug, Clone)]
pub struct UnionComponentFit {
pub kind: UnionComponentKind,
pub row_count: usize,
pub num_parameters: usize,
pub mixing_weight: f64,
}
#[derive(Debug, Clone)]
pub struct UnionStructureFit {
pub structure: UnionStructure,
pub components: Vec<UnionComponentFit>,
pub log_likelihood: f64,
pub bic: f64,
pub total_parameters: usize,
}
pub fn union_responsibility_split(
data: ArrayView2<'_, f64>,
m: usize,
config: GaussianMixtureConfig,
) -> Result<Vec<Vec<usize>>, String> {
let n = data.nrows();
if m == 0 {
return Err("union split requires at least one component".to_string());
}
if m > n {
return Err(format!(
"union split requested {m} groups but data has {n} rows"
));
}
if m == 1 {
return Ok(vec![(0..n).collect()]);
}
let fit = fit_gaussian_mixture(data, m, config).map_err(|error| error.to_string())?;
let mut groups: Vec<Vec<usize>> = vec![Vec::new(); m];
let mut comp = Vec::with_capacity(m);
for j in 0..m {
comp.push(GaussianComponentEval::factor(
fit.means.row(j),
&fit.covariances[j],
)?);
}
let log_w = fit
.weights
.iter()
.enumerate()
.map(|(component, &weight)| {
if weight.is_finite() && weight > 0.0 {
Ok(weight.ln())
} else {
Err(format!(
"union split received invalid fitted weight {weight} for component {component}"
))
}
})
.collect::<Result<Vec<_>, _>>()?;
for i in 0..n {
let row = data.row(i);
let mut best_j = 0usize;
let mut best_lt = f64::NEG_INFINITY;
for j in 0..m {
let lt = log_w[j] + comp[j].log_density(row);
if lt > best_lt {
best_lt = lt;
best_j = j;
}
}
if !best_lt.is_finite() {
return Err(format!(
"union split produced no finite component score at row {i}"
));
}
groups[best_j].push(i);
}
Ok(groups)
}
pub fn fit_union_structure(
data: ArrayView2<'_, f64>,
structure: UnionStructure,
config: GaussianMixtureConfig,
) -> Result<UnionStructureFit, String> {
let fitted = fit_union_density(data, structure, config)?;
Ok(UnionStructureFit {
structure,
components: fitted
.components
.iter()
.map(UnionComponentDensity::summary)
.collect(),
log_likelihood: fitted.log_likelihood,
bic: fitted.bic,
total_parameters: fitted.total_parameters,
})
}
pub fn fit_union_ladder(
data: ArrayView2<'_, f64>,
config: GaussianMixtureConfig,
) -> Result<Vec<UnionStructureFit>, String> {
let mut fits = Vec::new();
let mut errors = Vec::new();
for &structure in UNION_STRUCTURE_LADDER {
match fit_union_structure(data, structure, config) {
Ok(fit) => fits.push(fit),
Err(e) => errors.push(format!("{}: {e}", structure.as_str())),
}
}
if !errors.is_empty() {
return Err(format!(
"union ladder comparison failed; every declared structure must fit ({})",
errors.join("; ")
));
}
if fits.is_empty() {
return Err("union ladder is empty".to_string());
}
let ranked = rank_priority_candidates(
fits.into_iter()
.enumerate()
.map(|(idx, row)| {
let score = row.bic;
let tie = row.total_parameters; PriorityCandidate::new(row, idx, score, tie)
})
.collect(),
)
.into_iter()
.map(|row| row.item)
.collect::<Vec<_>>();
Ok(ranked)
}
fn gather_union_rows(data: ArrayView2<'_, f64>, idx: &[usize]) -> Array2<f64> {
let d = data.ncols();
let mut out = Array2::<f64>::zeros((idx.len(), d));
for (r, &i) in idx.iter().enumerate() {
for c in 0..d {
out[[r, c]] = data[[i, c]];
}
}
out
}
fn union_circle_rows(group: ArrayView2<'_, f64>) -> Result<Vec<usize>, String> {
let minimum_rows = CircularGaussianFit2d::NUM_FREE_PARAMETERS + 1;
if group.nrows() < minimum_rows {
return Err(format!(
"union circle component needs at least {minimum_rows} rows, got {}",
group.nrows()
));
}
Ok((0..group.nrows()).collect())
}
#[derive(Debug, Clone)]
enum UnionDensityModel {
Gaussian(GaussianComponentEval),
Circle(CircularGaussianFit2d),
}
#[derive(Debug, Clone)]
struct UnionComponentDensity {
kind: UnionComponentKind,
row_count: usize,
num_parameters: usize,
mixing_weight: f64,
log_weight: f64,
model: UnionDensityModel,
}
impl UnionComponentDensity {
fn summary(&self) -> UnionComponentFit {
UnionComponentFit {
kind: self.kind,
row_count: self.row_count,
num_parameters: self.num_parameters,
mixing_weight: self.mixing_weight,
}
}
fn dimension(&self) -> usize {
match &self.model {
UnionDensityModel::Gaussian(eval) => eval.d,
UnionDensityModel::Circle(_) => 2,
}
}
fn weighted_log_density(&self, y: ArrayView1<'_, f64>) -> f64 {
let component_log_density = match &self.model {
UnionDensityModel::Gaussian(eval) => eval.log_density(y),
UnionDensityModel::Circle(fit) => fit.log_density(y[0], y[1]),
};
self.log_weight + component_log_density
}
}
#[derive(Debug, Clone)]
struct FittedUnionDensity {
components: Vec<UnionComponentDensity>,
log_likelihood: f64,
bic: f64,
total_parameters: usize,
}
fn fit_union_density(
train: ArrayView2<'_, f64>,
structure: UnionStructure,
config: GaussianMixtureConfig,
) -> Result<FittedUnionDensity, String> {
let groups = union_responsibility_split(train, structure.num_components(), config)?;
fit_union_density_from_groups(train, structure, &groups, config)
}
fn fit_union_density_from_groups(
train: ArrayView2<'_, f64>,
structure: UnionStructure,
groups: &[Vec<usize>],
config: GaussianMixtureConfig,
) -> Result<FittedUnionDensity, String> {
validate_union_partition(train.nrows(), structure.num_components(), groups)?;
let assignments = unique_union_role_assignments(structure.components());
let mut best: Option<FittedUnionDensity> = None;
let mut errors = Vec::new();
for roles in assignments {
let candidate = (|| {
let mut components = Vec::with_capacity(groups.len());
let n_train = train.nrows() as f64;
for (&kind, rows) in roles.iter().zip(groups) {
let group = gather_union_rows(train, rows);
let mixing_weight = rows.len() as f64 / n_train;
components.push(fit_union_component_density(
group.view(),
kind,
mixing_weight,
config,
)?);
}
let component_parameters = components.iter().try_fold(0usize, |sum, component| {
sum.checked_add(component.num_parameters)
.ok_or_else(|| "union component parameter count overflowed usize".to_string())
})?;
let mixing_parameters = components.len() - 1;
let total_parameters = component_parameters
.checked_add(mixing_parameters)
.ok_or_else(|| "union total parameter count overflowed usize".to_string())?;
let per_point = score_union_components(&components, train)?;
let log_likelihood = pairwise_sum(
per_point
.as_slice()
.expect("owned union score vector must be contiguous"),
);
if !log_likelihood.is_finite() {
return Err("union training log likelihood is non-finite".to_string());
}
let bic = -log_likelihood + 0.5 * total_parameters as f64 * (train.nrows() as f64).ln();
if !bic.is_finite() {
return Err("union normalized soft-mixture BIC is non-finite".to_string());
}
Ok(FittedUnionDensity {
components,
log_likelihood,
bic,
total_parameters,
})
})();
match candidate {
Ok(candidate) => {
let replace = match &best {
Some(current) => candidate.bic.total_cmp(¤t.bic).is_lt(),
None => true,
};
if replace {
best = Some(candidate);
}
}
Err(error) => errors.push(format!("{roles:?}: {error}")),
}
}
best.ok_or_else(|| {
format!(
"union {} has no finite role assignment ({})",
structure.as_str(),
errors.join("; ")
)
})
}
fn validate_union_partition(
n_rows: usize,
expected_groups: usize,
groups: &[Vec<usize>],
) -> Result<(), String> {
if n_rows == 0 {
return Err("union fitting requires at least one training row".to_string());
}
if groups.len() != expected_groups {
return Err(format!(
"union partition has {} groups, expected {expected_groups}",
groups.len()
));
}
let mut seen = vec![false; n_rows];
for (group_index, rows) in groups.iter().enumerate() {
if rows.is_empty() {
return Err(format!("union partition group {group_index} is empty"));
}
for &row in rows {
if row >= n_rows {
return Err(format!(
"union partition group {group_index} contains out-of-range row {row} for {n_rows} rows"
));
}
if std::mem::replace(&mut seen[row], true) {
return Err(format!("union partition contains duplicate row {row}"));
}
}
}
if let Some(missing) = seen.iter().position(|included| !included) {
return Err(format!("union partition omits row {missing}"));
}
Ok(())
}
fn unique_union_role_assignments(roles: &[UnionComponentKind]) -> Vec<Vec<UnionComponentKind>> {
fn visit(
roles: &[UnionComponentKind],
used: &mut [bool],
assignment: &mut Vec<UnionComponentKind>,
out: &mut Vec<Vec<UnionComponentKind>>,
) {
if assignment.len() == roles.len() {
out.push(assignment.clone());
return;
}
let mut used_at_depth = Vec::new();
for (index, &role) in roles.iter().enumerate() {
if used[index] || used_at_depth.contains(&role) {
continue;
}
used_at_depth.push(role);
used[index] = true;
assignment.push(role);
visit(roles, used, assignment, out);
assignment.pop();
used[index] = false;
}
}
let mut out = Vec::new();
visit(
roles,
&mut vec![false; roles.len()],
&mut Vec::with_capacity(roles.len()),
&mut out,
);
out
}
fn fit_union_component_density(
group: ArrayView2<'_, f64>,
kind: UnionComponentKind,
mixing_weight: f64,
config: GaussianMixtureConfig,
) -> Result<UnionComponentDensity, String> {
if !(mixing_weight.is_finite() && mixing_weight > 0.0 && mixing_weight <= 1.0) {
return Err(format!(
"union component mixing weight must be finite and in (0, 1], got {mixing_weight}"
));
}
let row_count = group.nrows();
let (model, num_parameters) = match kind {
UnionComponentKind::Line => {
if group.nrows() < group.ncols() + 1 {
return Err(format!(
"union line component needs >= {} rows, got {}",
group.ncols() + 1,
group.nrows()
));
}
let fit = fit_gaussian_mixture(group, 1, config).map_err(|error| error.to_string())?;
let num_parameters = fit.num_free_parameters();
let eval = GaussianComponentEval::factor(fit.means.row(0), &fit.covariances[0])?;
(UnionDensityModel::Gaussian(eval), num_parameters)
}
UnionComponentKind::PointCluster => {
if group.nrows() < group.ncols() + 1 {
return Err(format!(
"union isotropic point component needs >= {} rows, got {}",
group.ncols() + 1,
group.nrows()
));
}
let eval = fit_isotropic_gaussian_component(group, config.covariance_floor)?;
(
UnionDensityModel::Gaussian(eval),
group
.ncols()
.checked_add(1)
.ok_or_else(|| "union point parameter count overflowed usize".to_string())?,
)
}
UnionComponentKind::Circle => {
let rows = union_circle_rows(group)?;
let fit = CircularGaussianFit2d::fit(group, &rows)?;
(
UnionDensityModel::Circle(fit),
CircularGaussianFit2d::NUM_FREE_PARAMETERS,
)
}
};
Ok(UnionComponentDensity {
kind,
row_count,
num_parameters,
mixing_weight,
log_weight: mixing_weight.ln(),
model,
})
}
#[derive(Debug, Clone, Copy)]
struct StableScalarMeanChart {
origin: f64,
scale: f64,
normalized_offset: f64,
}
impl StableScalarMeanChart {
#[inline]
fn centered(self, value: f64) -> Result<f64, String> {
let relative = value - self.origin;
let centered = (-self.normalized_offset).mul_add(self.scale, relative);
if centered.is_finite() {
Ok(centered)
} else {
Err("union isotropic point residual is not representable".to_string())
}
}
}
fn stable_scalar_mean_chart(values: ArrayView1<'_, f64>) -> Result<StableScalarMeanChart, String> {
if values.is_empty() || values.iter().any(|value| !value.is_finite()) {
return Err("stable scalar mean requires finite nonempty values".to_string());
}
let anchor = values[0];
let anchor_chart_is_representable = values.iter().all(|&value| (value - anchor).is_finite());
let origin = if anchor_chart_is_representable {
anchor
} else {
0.0
};
let scale = values
.iter()
.map(|&value| (value - origin).abs())
.fold(0.0_f64, f64::max);
if scale == 0.0 {
return Ok(StableScalarMeanChart {
origin,
scale: 0.0,
normalized_offset: 0.0,
});
}
let normalized = values
.iter()
.map(|&value| (value - origin) / scale)
.collect::<Vec<_>>();
let normalized_offset = pairwise_sum(&normalized) / values.len() as f64;
let mean = normalized_offset.mul_add(scale, origin);
if !(normalized_offset.is_finite() && mean.is_finite()) {
return Err("union isotropic point mean is not representable".to_string());
}
Ok(StableScalarMeanChart {
origin,
scale,
normalized_offset,
})
}
fn fit_isotropic_gaussian_component(
group: ArrayView2<'_, f64>,
covariance_floor: f64,
) -> Result<GaussianComponentEval, String> {
let n = group.nrows();
let d = group.ncols();
if n == 0 || d == 0 {
return Err("union isotropic point component requires a non-empty matrix".to_string());
}
if !(covariance_floor.is_finite() && covariance_floor > 0.0) {
return Err(format!(
"union isotropic covariance floor must be finite and positive, got {covariance_floor}"
));
}
for row in group.rows() {
for axis in 0..d {
let value = row[axis];
if !value.is_finite() {
return Err(format!(
"union isotropic point data contains non-finite coordinate {value}"
));
}
}
}
let mut charts = Vec::with_capacity(d);
for axis in 0..d {
let chart = stable_scalar_mean_chart(group.column(axis))?;
charts.push(chart);
}
let scalar_count = n
.checked_mul(d)
.ok_or_else(|| "union isotropic residual count overflowed usize".to_string())?;
let mut residuals = Vec::with_capacity(scalar_count);
let mut residual_scale = 0.0_f64;
for row in group.rows() {
for axis in 0..d {
let residual = charts[axis].centered(row[axis])?;
residual_scale = residual_scale.max(residual.abs());
residuals.push(residual);
}
}
let variance = if residual_scale == 0.0 {
covariance_floor
} else {
for residual in &mut residuals {
*residual = (*residual / residual_scale).powi(2);
}
let normalized_mean_square = pairwise_sum(&residuals) / scalar_count as f64;
let rms = residual_scale * normalized_mean_square.sqrt();
let unconstrained = rms * rms;
if !unconstrained.is_finite() {
return Err("union isotropic point variance is non-finite".to_string());
}
unconstrained.max(covariance_floor)
};
GaussianComponentEval::isotropic(&charts, variance)
}
fn score_union_components(
components: &[UnionComponentDensity],
eval: ArrayView2<'_, f64>,
) -> Result<Array1<f64>, String> {
if components.is_empty() {
return Err("union density requires at least one component".to_string());
}
if eval.iter().any(|coordinate| !coordinate.is_finite()) {
return Err("union eval coordinates must be finite".to_string());
}
for component in components {
if component.dimension() != eval.ncols() {
return Err(format!(
"union component {:?} has dimension {}, eval has {} columns",
component.kind,
component.dimension(),
eval.ncols()
));
}
}
let mut out = Array1::<f64>::zeros(eval.nrows());
let mut terms = vec![f64::NEG_INFINITY; components.len()];
for i in 0..eval.nrows() {
let row = eval.row(i);
let mut max_term = f64::NEG_INFINITY;
for (component_index, component) in components.iter().enumerate() {
let term = component.weighted_log_density(row);
terms[component_index] = term;
if term > max_term {
max_term = term;
}
}
let value = log_sum_exp(&terms, max_term);
if !value.is_finite() {
return Err(format!(
"union density produced non-finite log density at eval row {i}"
));
}
out[i] = value;
}
Ok(out)
}
pub fn union_per_point_log_density(
train: ArrayView2<'_, f64>,
eval: ArrayView2<'_, f64>,
structure: UnionStructure,
config: GaussianMixtureConfig,
) -> Result<Array1<f64>, String> {
if train.ncols() != eval.ncols() {
return Err(format!(
"union held-out density: train has {} columns, eval has {}",
train.ncols(),
eval.ncols()
));
}
let fitted = fit_union_density(train, structure, config)?;
score_union_components(&fitted.components, eval)
}
#[derive(Clone, Debug)]
pub struct RemlCandidate {
pub index: usize,
pub name: String,
pub score: f64,
pub edf: Option<f64>,
pub log_lik: Option<f64>,
pub family: Option<String>,
pub n_obs: Option<usize>,
}
impl RemlCandidate {
pub fn ranking_score(&self) -> f64 {
match (self.log_lik, self.edf) {
(Some(log_lik), Some(edf)) if log_lik.is_finite() && edf.is_finite() => {
-2.0 * log_lik + 2.0 * edf
}
_ => self.score,
}
}
}
#[derive(Clone, Debug)]
pub struct RemlComparison {
pub ranking: Vec<RankedRow>,
pub winner: String,
pub evidence_summary: String,
pub score_table: Vec<ScoreRow>,
}
#[derive(Clone, Debug)]
pub struct RankedRow {
pub name: String,
pub score: f64,
pub delta: f64,
pub bayes_factor: f64,
pub edf: Option<f64>,
}
#[derive(Clone, Debug)]
pub struct ScoreRow {
pub name: String,
pub reml_score: f64,
pub delta_reml: f64,
pub bayes_factor_best_over_model: f64,
pub effective_dof: Option<f64>,
}
#[inline]
pub fn log_bayes_factor(reml_score_a: f64, reml_score_b: f64) -> f64 {
reml_score_b - reml_score_a
}
pub fn compare_reml_fits(mut candidates: Vec<RemlCandidate>) -> Result<RemlComparison, String> {
if candidates.is_empty() {
return Err("compare_models requires at least one fit".to_string());
}
{
let mut seen_family: Option<&str> = None;
for cand in &candidates {
if let Some(fam) = cand.family.as_deref() {
match seen_family {
None => seen_family = Some(fam),
Some(prev) if prev != fam => {
return Err(format!(
"compare_models: cannot compare fits of different response families ('{prev}' vs '{fam}'); their REML/LAML evidence scores are on incomparable base measures. Compare models fit to the same response under the same family."
));
}
Some(_) => {}
}
}
}
}
{
let mut seen_n: Option<usize> = None;
for cand in &candidates {
if let Some(n) = cand.n_obs {
match seen_n {
None => seen_n = Some(n),
Some(prev) if prev != n => {
return Err(format!(
"compare_models: cannot compare fits made on a different number of \
observations (n={prev} vs n={n}); AIC / REML-LAML evidence scales \
with the sample size, so their score difference is not a Bayes \
factor. Compare models fit to the same response on the same data."
));
}
Some(_) => {}
}
}
}
}
candidates = rank_priority_candidates(
candidates
.into_iter()
.enumerate()
.map(|(idx, row)| {
let ranking = row.ranking_score();
PriorityCandidate::new(row, idx, ranking, 0)
})
.collect(),
)
.into_iter()
.map(|row| row.item)
.collect();
let winner = candidates[0].name.clone();
let best_ranking_score = candidates[0].ranking_score();
let best_raw_score = candidates
.iter()
.map(|c| c.score)
.fold(f64::INFINITY, f64::min);
let mut ranking = Vec::with_capacity(candidates.len());
let mut score_table = Vec::with_capacity(candidates.len());
for row in &candidates {
let delta = log_bayes_factor(best_ranking_score, row.ranking_score());
let bayes_factor = (0.5 * delta).exp();
let delta_reml = log_bayes_factor(best_raw_score, row.score);
ranking.push(RankedRow {
name: row.name.clone(),
score: row.score,
delta,
bayes_factor,
edf: row.edf,
});
score_table.push(ScoreRow {
name: row.name.clone(),
reml_score: row.score,
delta_reml,
bayes_factor_best_over_model: delta_reml.exp(),
effective_dof: row.edf,
});
}
let evidence_summary = if let Some(runner_up) = candidates.get(1) {
let margin = runner_up.ranking_score() - candidates[0].ranking_score();
format!(
"{} wins by Bayes factor {} over {}",
winner,
format_bayes_factor(0.5 * margin),
runner_up.name
)
} else {
format!("{winner} (single fit; no comparison)")
};
Ok(RemlComparison {
ranking,
winner,
evidence_summary,
score_table,
})
}
pub fn format_bayes_factor(log_bf: f64) -> String {
if !log_bf.is_finite() {
return "inf".to_string();
}
if log_bf.abs() >= std::f64::consts::LN_10 * 3.0 {
return format!("1e{:+.1}", log_bf / std::f64::consts::LN_10);
}
format_three_significant(log_bf.exp())
}
pub fn format_three_significant(value: f64) -> String {
if value == 0.0 {
return "0".to_string();
}
if !value.is_finite() {
return format!("{value}");
}
let exponent = value.abs().log10().floor() as i32;
if exponent >= 3 {
return format!("{value:.2e}");
}
let decimals = (2 - exponent).max(0) as usize;
let scale = 10f64.powi(decimals as i32);
let rounded = (value * scale).abs().round() / scale * value.signum();
format!("{rounded:.decimals$}")
}
impl Default for TopologySelectOptions {
fn default() -> Self {
Self {
tie_tolerance: 1e-3,
score_scale: TopologyScoreScale::PerObservation,
}
}
}
pub fn laplace_evidence(
logdet_source: EvidenceLogDetSource<'_>,
penalty_log_det: f64,
residual_objective: f64,
effective_dim: f64,
penalty_rank: f64,
) -> f64 {
if !(effective_dim.is_finite() && penalty_rank.is_finite()) {
return f64::NAN;
}
let log_det_h = match evidence_hessian_log_det(logdet_source) {
Ok(v) => v,
Err(_) => return f64::NAN,
};
let null_dim = effective_dim - penalty_rank;
if !null_dim.is_finite() || null_dim < -1e-9 {
return f64::NAN;
}
residual_objective + 0.5 * log_det_h
- 0.5 * penalty_log_det
- 0.5 * null_dim.max(0.0) * (2.0 * std::f64::consts::PI).ln()
}
pub fn evidence_hessian_log_det(source: EvidenceLogDetSource<'_>) -> Result<f64, String> {
match source {
EvidenceLogDetSource::FactoredArrow {
cache,
fallback_hvp,
} => match arrow_log_det_from_cache(cache) {
Some(v) => Ok(v),
None => match fallback_hvp {
Some(hvp) => hessian_log_det_from_hvp(hvp),
None => {
Err("evidence Hessian logdet requires exact factors or HVP fallback".into())
}
},
},
EvidenceLogDetSource::Hvp(hvp) => hessian_log_det_from_hvp(hvp),
}
}
pub fn hessian_log_det_from_hvp(hvp: EvidenceHvpLogDet<'_>) -> Result<f64, String> {
if hvp.dim == 0 {
return Ok(0.0);
}
if hvp.dim <= ANALYTIC_LOGDET_DENSE_DIM_THRESHOLD {
let mut dense = Array2::<f64>::zeros((hvp.dim, hvp.dim));
let mut basis = vec![0.0_f64; hvp.dim];
for j in 0..hvp.dim {
basis[j] = 1.0;
let col = (hvp.apply)(&basis);
basis[j] = 0.0;
if col.len() != hvp.dim || col.iter().any(|v| !v.is_finite()) {
return Err(format!(
"evidence HVP logdet expected finite column of length {}, got {}",
hvp.dim,
col.len()
));
}
for i in 0..hvp.dim {
dense[[i, j]] = col[i];
}
}
validate_dense_hvp_symmetry(&dense)?;
for i in 0..hvp.dim {
for j in (i + 1)..hvp.dim {
let avg = 0.5 * (dense[[i, j]] + dense[[j, i]]);
dense[[i, j]] = avg;
dense[[j, i]] = avg;
}
}
dense_spd_log_det(&dense)
} else {
stochastic_hvp_log_det(hvp)
}
}
fn dense_spd_log_det(matrix: &Array2<f64>) -> Result<f64, String> {
if matrix.nrows() != matrix.ncols() {
return Err(format!(
"evidence dense logdet requires square matrix, got {}x{}",
matrix.nrows(),
matrix.ncols()
));
}
if gam_gpu::cuda_selected().map_err(|error| error.to_string())? {
return crate::gpu::reml_gpu::evidence_derivatives_gpu(
crate::gpu::reml_gpu::RemlGpuInput {
penalized_hessian: matrix.view(),
derivative_hessians: Vec::new(),
},
)
.map(|evidence| evidence.logdet_hessian);
}
let (evals, _) = matrix
.eigh(Side::Lower)
.map_err(|e| format!("evidence dense logdet eigendecomposition failed: {e}"))?;
let mut logdet = 0.0_f64;
for (idx, &ev) in evals.iter().enumerate() {
if !ev.is_finite() || ev <= 0.0 {
return Err(format!(
"evidence dense logdet expected SPD Hessian, eigenvalue {idx} is {ev:.3e}"
));
}
logdet += ev.ln();
}
Ok(logdet)
}
fn validate_dense_hvp_symmetry(matrix: &Array2<f64>) -> Result<(), String> {
let n = matrix.nrows();
let mut norm_sq = 0.0_f64;
for &value in matrix.iter() {
norm_sq += value * value;
}
let mut skew_sq = 0.0_f64;
for i in 0..n {
for j in (i + 1)..n {
let skew = matrix[[i, j]] - matrix[[j, i]];
skew_sq += 2.0 * skew * skew;
}
}
let rel_skew = skew_sq.sqrt() / norm_sq.sqrt().max(1.0);
if !rel_skew.is_finite() || rel_skew > EVIDENCE_HVP_SYMMETRY_REL_TOL {
return Err(format!(
"evidence HVP logdet requires symmetric operator, relative skew norm is {rel_skew:.3e}"
));
}
Ok(())
}
fn validate_hvp_randomized_symmetry(hvp: EvidenceHvpLogDet<'_>) -> Result<(), String> {
let inv_norm = 1.0 / (hvp.dim as f64).sqrt();
for probe in 0..EVIDENCE_HVP_SYMMETRY_PROBES.max(1) {
let mut x = vec![0.0_f64; hvp.dim];
let mut y = vec![0.0_f64; hvp.dim];
rademacher_unit_probe_into_slice(&mut x, (2 * probe) as u64, inv_norm);
rademacher_unit_probe_into_slice(&mut y, (2 * probe + 1) as u64, inv_norm);
let hx = (hvp.apply)(&x);
let hy = (hvp.apply)(&y);
if hx.len() != hvp.dim || hx.iter().any(|v| !v.is_finite()) {
return Err(format!(
"evidence HVP symmetry check expected finite vector of length {}, got {}",
hvp.dim,
hx.len()
));
}
if hy.len() != hvp.dim || hy.iter().any(|v| !v.is_finite()) {
return Err(format!(
"evidence HVP symmetry check expected finite vector of length {}, got {}",
hvp.dim,
hy.len()
));
}
let lhs = dot_slice(&x, &hy);
let rhs = dot_slice(&hx, &y);
let scale = (norm2_slice(&hx) * norm2_slice(&y))
.max(norm2_slice(&hy) * norm2_slice(&x))
.max(lhs.abs())
.max(rhs.abs())
.max(1.0);
let rel = (lhs - rhs).abs() / scale;
if !rel.is_finite() || rel > EVIDENCE_HVP_SYMMETRY_REL_TOL {
return Err(format!(
"evidence HVP logdet requires symmetric operator, randomized symmetry probe {probe} has relative bilinear mismatch {rel:.3e}"
));
}
}
Ok(())
}
fn stochastic_hvp_log_det(hvp: EvidenceHvpLogDet<'_>) -> Result<f64, String> {
validate_hvp_randomized_symmetry(hvp)?;
let probes = EVIDENCE_LOGDET_SLQ_PROBES.max(1);
let steps = EVIDENCE_LOGDET_LANCZOS_STEPS.min(hvp.dim).max(1);
let inv_norm = 1.0 / (hvp.dim as f64).sqrt();
let mut estimate = 0.0_f64;
for probe in 0..probes {
let mut q0 = vec![0.0_f64; hvp.dim];
rademacher_unit_probe_into_slice(&mut q0, probe as u64, inv_norm);
let quad = lanczos_log_quadrature_hvp(hvp, q0, steps)?;
estimate += hvp.dim as f64 * quad;
}
Ok(estimate / probes as f64)
}
fn lanczos_log_quadrature_hvp(
hvp: EvidenceHvpLogDet<'_>,
q: Vec<f64>,
max_steps: usize,
) -> Result<f64, String> {
let n = hvp.dim;
let eigen = symmetric_lanczos_eigenpairs(
n,
&q,
SymmetricLanczosOptions {
max_steps,
residual_tol: 1e-12,
local_reorthogonalize: false,
full_reorthogonalize: false,
},
|q, out| {
let applied = (hvp.apply)(q);
if applied.len() != n || applied.iter().any(|v| !v.is_finite()) {
return Err(format!(
"evidence HVP SLQ expected finite vector of length {n}, got {}",
applied.len()
));
}
out.copy_from_slice(&applied);
Ok(())
},
)
.map_err(|e| format!("evidence HVP SLQ Lanczos failed: {e}"))?;
symmetric_lanczos_log_quadrature(&eigen, "evidence HVP SLQ expected SPD Hessian")
}
#[inline]
fn dot_slice(a: &[f64], b: &[f64]) -> f64 {
assert_eq!(a.len(), b.len());
let mut s = 0.0_f64;
for i in 0..a.len() {
s += a[i] * b[i];
}
s
}
#[inline]
fn norm2_slice(a: &[f64]) -> f64 {
dot_slice(a, a).sqrt()
}
fn rademacher_unit_probe_into_slice(z: &mut [f64], probe: u64, scale: f64) {
let mut state = 0x6A09E667F3BCC909_u64 ^ probe.wrapping_mul(0xD1B54A32D192ED03);
let mut bits = 0_u64;
let mut remaining_bits = 0_u32;
for value in z.iter_mut() {
if remaining_bits == 0 {
bits = splitmix64(&mut state);
remaining_bits = 64;
}
*value = if bits & 1 == 0 { scale } else { -scale };
bits >>= 1;
remaining_bits -= 1;
}
}
#[inline]
const fn splitmix64(state: &mut u64) -> u64 {
gam_linalg::utils::splitmix64(state)
}
pub fn arrow_log_det_from_cache(cache: &ArrowFactorCache) -> Option<f64> {
if let Some(log_det) = cache.joint_hessian_log_det {
return log_det.is_finite().then_some(log_det);
}
if cache.ridge_t != 0.0 || cache.ridge_beta != 0.0 {
return None;
}
if cache.k > 0 && !cache.schur_factor_is_undamped {
return None;
}
cache.compute_undamped_arrow_log_det()
}
pub fn ift_du_dbeta(cache: &ArrowFactorCache) -> Array2<f64> {
let n = cache.undamped_factor_count();
let total_len = cache.delta_t_len();
let k = cache.k;
if !cache.htbeta_available() {
return Array2::<f64>::from_elem((total_len, k), f64::NAN);
}
let mut out = Array2::<f64>::zeros((total_len, k));
let mut beta_basis = Array1::<f64>::zeros(k);
let mut rhs = Array1::<f64>::zeros(cache.d);
for i in 0..n {
let di = cache.row_dims[i];
let row_base = cache.row_offsets[i];
let factor = cache.undamped_factor(i);
for col in 0..k {
beta_basis.fill(0.0);
beta_basis[col] = 1.0;
let mut rhs_i = rhs.slice_mut(ndarray::s![..di]).to_owned();
if !cache.apply_htbeta_row(i, beta_basis.view(), &mut rhs_i) {
return Array2::<f64>::from_elem((total_len, k), f64::NAN);
}
let y = cholesky_solve_vector(factor, &rhs_i);
for c in 0..di {
out[[row_base + c, col]] = -y[c];
}
}
}
out
}
pub fn coupling_components(hessian: ArrayView2<'_, f64>) -> Vec<usize> {
let p = hessian.nrows();
if p == 0 || hessian.ncols() != p {
return Vec::new();
}
let mut parent: Vec<usize> = (0..p).collect();
let mut size: Vec<usize> = vec![1; p];
fn find(parent: &mut [usize], mut x: usize) -> usize {
while parent[x] != x {
parent[x] = parent[parent[x]];
x = parent[x];
}
x
}
for i in 0..p {
for j in (i + 1)..p {
if hessian[[i, j]] != 0.0 || hessian[[j, i]] != 0.0 {
let (ri, rj) = (find(&mut parent, i), find(&mut parent, j));
if ri != rj {
let (small, large) = if size[ri] < size[rj] {
(ri, rj)
} else {
(rj, ri)
};
parent[small] = large;
size[large] += size[small];
}
}
}
}
let mut label_of_root: Vec<Option<usize>> = vec![None; p];
let mut next_label = 0usize;
let mut labels = vec![0usize; p];
for idx in 0..p {
let root = find(&mut parent, idx);
let label = match label_of_root[root] {
Some(l) => l,
None => {
let l = next_label;
label_of_root[root] = Some(l);
next_label += 1;
l
}
};
labels[idx] = label;
}
labels
}
pub fn cone_of_influence(labels: &[usize], support: &[usize]) -> Vec<usize> {
if support.is_empty() {
return Vec::new();
}
let mut in_cone_labels: Vec<usize> = support
.iter()
.filter_map(|&idx| labels.get(idx).copied())
.collect();
in_cone_labels.sort_unstable();
in_cone_labels.dedup();
if in_cone_labels.is_empty() {
return Vec::new();
}
(0..labels.len())
.filter(|idx| in_cone_labels.binary_search(&labels[*idx]).is_ok())
.collect()
}
pub fn ift_dbeta_drho(
cache: &ArrowFactorCache,
dg_red_drho: ArrayView2<'_, f64>,
) -> Option<Array2<f64>> {
if !cache.schur_factor_is_undamped {
return None;
}
let schur = cache.schur_factor.as_ref()?;
if dg_red_drho.nrows() != cache.k || schur.nrows() != cache.k {
return None;
}
crate::sensitivity::FitSensitivity::from_lower_triangular(schur).mode_response(dg_red_drho)
}
#[derive(Clone)]
pub struct EvidenceIftGradientTerms<'a> {
pub dbeta_drho: ArrayView2<'a, f64>,
pub du_drho: ArrayView2<'a, f64>,
pub value_beta: ArrayView1<'a, f64>,
pub value_u: ArrayView1<'a, f64>,
pub logdet_h_beta: ArrayView1<'a, f64>,
pub logdet_h_u: ArrayView1<'a, f64>,
}
pub fn evidence_ift_gradient_correction(terms: EvidenceIftGradientTerms<'_>) -> Array1<f64> {
let k = terms.dbeta_drho.nrows();
let nd = terms.du_drho.nrows();
let r = terms.dbeta_drho.ncols();
if terms.du_drho.ncols() != r
|| terms.value_beta.len() != k
|| terms.logdet_h_beta.len() != k
|| terms.value_u.len() != nd
|| terms.logdet_h_u.len() != nd
{
return Array1::<f64>::from_elem(r, f64::NAN);
}
let mut out = Array1::<f64>::zeros(r);
for a in 0..r {
let mut acc = 0.0_f64;
for j in 0..k {
let mode = terms.dbeta_drho[[j, a]];
acc += terms.value_beta[j] * mode;
acc += 0.5 * terms.logdet_h_beta[j] * mode;
}
for j in 0..nd {
let mode = terms.du_drho[[j, a]];
acc += terms.value_u[j] * mode;
acc += 0.5 * terms.logdet_h_u[j] * mode;
}
out[a] = acc;
}
out
}
pub fn evidence_grad_rho(
cache: &ArrowFactorCache,
value_rho: ArrayView1<'_, f64>,
huu_drho: &[Vec<Array2<f64>>],
htbeta_drho: &[Vec<Array2<f64>>],
hbb_drho: &[Array2<f64>],
pen_logdet_drho: ArrayView1<'_, f64>,
ift_terms: EvidenceIftGradientTerms<'_>,
) -> Array1<f64> {
let r = value_rho.len();
let n = cache.undamped_factor_count();
let k = cache.k;
let mut out = Array1::<f64>::zeros(r);
if !cache.htbeta_available()
|| pen_logdet_drho.len() != r
|| huu_drho.len() != n
|| htbeta_drho.len() != n
|| hbb_drho.len() != r
|| huu_drho.iter().any(|row| row.len() != r)
|| htbeta_drho.iter().any(|row| row.len() != r)
|| hbb_drho.iter().any(|m| m.nrows() != k || m.ncols() != k)
|| huu_drho.iter().enumerate().any(|(i, row)| {
let di = cache.row_dims[i];
row.iter().any(|m| m.nrows() != di || m.ncols() != di)
})
|| htbeta_drho.iter().enumerate().any(|(i, row)| {
let di = cache.row_dims[i];
row.iter().any(|m| m.nrows() != di || m.ncols() != k)
})
{
out.fill(f64::NAN);
return out;
}
let ift_correction = evidence_ift_gradient_correction(ift_terms);
if ift_correction.len() != r || ift_correction.iter().any(|v| v.is_nan()) {
out.fill(f64::NAN);
return out;
}
let schur = match cache.schur_factor.as_ref() {
Some(s) => s,
None => {
for a in 0..r {
out[a] = f64::NAN;
}
return out;
}
};
if !cache.schur_factor_is_undamped {
for a in 0..r {
out[a] = f64::NAN;
}
return out;
}
let mut y_blocks: Vec<Array2<f64>> = Vec::with_capacity(n);
let mut beta_basis = Array1::<f64>::zeros(k);
let mut rhs = Array1::<f64>::zeros(cache.d);
for i in 0..n {
let di = cache.row_dims[i];
let factor = cache.undamped_factor(i);
let mut yi = Array2::<f64>::zeros((di, k));
for col in 0..k {
beta_basis.fill(0.0);
beta_basis[col] = 1.0;
let mut rhs_i = rhs.slice_mut(ndarray::s![..di]).to_owned();
if !cache.apply_htbeta_row(i, beta_basis.view(), &mut rhs_i) {
out.fill(f64::NAN);
return out;
}
let v = cholesky_solve_vector(factor, &rhs_i);
for c in 0..di {
yi[[c, col]] = v[c];
}
}
y_blocks.push(yi);
}
let mut trace_rhs = Array1::<f64>::zeros(cache.d);
let mut da_tmp = Array2::<f64>::zeros((cache.d, k));
let mut col_scratch = Array1::<f64>::zeros(k);
for a in 0..r {
let mut grad = value_rho[a];
let mut row_trace_acc = 0.0_f64;
for i in 0..n {
let di = cache.row_dims[i];
let m_i = &huu_drho[i][a];
assert_eq!(m_i.shape(), &[di, di]);
for col in 0..di {
let mut tr_rhs_i = trace_rhs.slice_mut(ndarray::s![..di]).to_owned();
for r0 in 0..di {
tr_rhs_i[r0] = m_i[[r0, col]];
}
let v = cholesky_solve_vector(cache.undamped_factor(i), &tr_rhs_i);
row_trace_acc += v[col];
}
}
let mut da = hbb_drho[a].clone();
assert_eq!(da.shape(), &[k, k]);
for i in 0..n {
let di = cache.row_dims[i];
let dhtb = &htbeta_drho[i][a]; let yi = &y_blocks[i]; for r0 in 0..k {
for c0 in 0..k {
let mut acc = 0.0;
for cc in 0..di {
acc += dhtb[[cc, r0]] * yi[[cc, c0]];
}
da[[r0, c0]] -= acc;
}
}
for r0 in 0..k {
for c0 in 0..k {
let mut acc = 0.0;
for cc in 0..di {
acc += yi[[cc, r0]] * dhtb[[cc, c0]];
}
da[[r0, c0]] -= acc;
}
}
let dhuu = &huu_drho[i][a];
let mut da_tmp_i = da_tmp.slice_mut(ndarray::s![..di, ..]).to_owned();
for r0 in 0..di {
for c0 in 0..k {
let mut acc = 0.0;
for cc in 0..di {
acc += dhuu[[r0, cc]] * yi[[cc, c0]];
}
da_tmp_i[[r0, c0]] = acc;
}
}
for r0 in 0..k {
for c0 in 0..k {
let mut acc = 0.0;
for cc in 0..di {
acc += yi[[cc, r0]] * da_tmp_i[[cc, c0]];
}
da[[r0, c0]] += acc;
}
}
}
let mut schur_trace_acc = 0.0_f64;
for j in 0..k {
for r0 in 0..k {
col_scratch[r0] = da[[r0, j]];
}
let v = cholesky_solve_vector(schur, &col_scratch);
schur_trace_acc += v[j];
}
grad += 0.5 * (row_trace_acc + schur_trace_acc);
grad += ift_correction[a];
grad -= 0.5 * pen_logdet_drho[a];
out[a] = grad;
}
out
}
pub fn select_topology(
candidates: &[TopologyCandidate],
options: TopologySelectOptions,
) -> SelectedTopology {
let mut valid: Vec<TopologyCandidate> = candidates
.iter()
.filter(|c| {
c.converged
&& c.exclusion_reason.is_none()
&& c.negative_log_evidence.is_finite()
&& topology_selection_score(c, options.score_scale).is_finite()
})
.cloned()
.collect();
let mut excluded: Vec<TopologyCandidate> = candidates
.iter()
.filter(|c| {
!(c.converged && c.exclusion_reason.is_none() && c.negative_log_evidence.is_finite())
|| !topology_selection_score(c, options.score_scale).is_finite()
})
.cloned()
.collect();
assert!(
!valid.is_empty(),
"select_topology: no finite valid candidates; proposal §6.11 forbids silent fallback"
);
valid = rank_priority_candidates(
valid
.into_iter()
.enumerate()
.map(|(idx, row)| {
let score = topology_selection_score(&row, options.score_scale);
let tie_break = usize::from(row.kind.complexity_rank());
PriorityCandidate::new(row, idx, score, tie_break)
})
.collect(),
)
.into_iter()
.map(|row| row.item)
.collect();
let tie = if valid.len() >= 2 {
let top = topology_selection_score(&valid[0], options.score_scale);
let next = topology_selection_score(&valid[1], options.score_scale);
(next - top).abs() <= options.tie_tolerance
} else {
false
};
if tie {
let top_score = topology_selection_score(&valid[0], options.score_scale);
let tied_end = valid
.iter()
.position(|c| {
(topology_selection_score(c, options.score_scale) - top_score).abs()
> options.tie_tolerance
})
.unwrap_or(valid.len());
valid[..tied_end].sort_by_key(|c| c.kind.complexity_rank());
}
let winner = valid[0].kind;
valid.append(&mut excluded);
SelectedTopology {
winner,
ranking: valid,
tie,
}
}
fn topology_selection_score(candidate: &TopologyCandidate, scale: TopologyScoreScale) -> f64 {
match scale {
TopologyScoreScale::PerObservation => {
if candidate.n_obs == 0 {
f64::NAN
} else {
candidate.negative_log_evidence / candidate.n_obs as f64
}
}
TopologyScoreScale::PerEffectiveDim => {
if !(candidate.effective_dim.is_finite() && candidate.effective_dim > 0.0) {
f64::NAN
} else {
candidate.negative_log_evidence / candidate.effective_dim
}
}
}
}
pub fn cache_matches_system(cache: &ArrowFactorCache, sys: &ArrowSchurSystem) -> bool {
cache.d == sys.d
&& cache.k == sys.k
&& cache.n_rows() == sys.rows.len()
&& cache.undamped_factor_count() == sys.rows.len()
&& cache.manifold_mode_fingerprint == sys.manifold_mode_fingerprint
&& cache.row_hessian_fingerprint == sys.current_row_hessian_fingerprint()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum HybridAtomParam {
Curved { latent_dim: usize },
Linear,
}
impl HybridAtomParam {
pub const fn as_str(self) -> &'static str {
match self {
HybridAtomParam::Curved { .. } => "curved",
HybridAtomParam::Linear => "linear",
}
}
pub const fn is_linear(self) -> bool {
matches!(self, HybridAtomParam::Linear)
}
}
#[derive(Debug, Clone, Copy)]
pub struct HybridAtomCandidate {
pub param: HybridAtomParam,
pub negative_log_evidence: f64,
pub num_parameters: usize,
pub fitted_turning: Option<f64>,
}
impl HybridAtomCandidate {
pub fn linear(negative_log_evidence: f64, num_parameters: usize) -> Self {
Self {
param: HybridAtomParam::Linear,
negative_log_evidence,
num_parameters,
fitted_turning: Some(0.0),
}
}
pub fn curved(
latent_dim: usize,
negative_log_evidence: f64,
num_parameters: usize,
fitted_turning: Option<f64>,
) -> Self {
Self {
param: HybridAtomParam::Curved { latent_dim },
negative_log_evidence,
num_parameters,
fitted_turning,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct HybridAtomChoice {
pub param: HybridAtomParam,
pub negative_log_evidence: f64,
pub num_parameters: usize,
pub curved_turning: Option<f64>,
pub curved_evidence_margin: f64,
}
pub const HYBRID_LINEAR_TURNING_FLOOR: f64 = 1e-9;
pub fn select_hybrid_atom(candidates: &[HybridAtomCandidate]) -> Option<HybridAtomChoice> {
if candidates.is_empty() {
return None;
}
let linear = candidates.iter().find(|c| c.param.is_linear());
let curved = candidates.iter().find(|c| !c.param.is_linear());
let curved_turning = curved.and_then(|c| c.fitted_turning);
let curved_evidence_margin = match (linear, curved) {
(Some(l), Some(c)) => l.negative_log_evidence - c.negative_log_evidence,
_ => 0.0,
};
if let (Some(l), Some(turning)) = (linear, curved_turning)
&& turning <= HYBRID_LINEAR_TURNING_FLOOR
{
return Some(HybridAtomChoice {
param: l.param,
negative_log_evidence: l.negative_log_evidence,
num_parameters: l.num_parameters,
curved_turning,
curved_evidence_margin,
});
}
let mut best = candidates[0];
for cand in &candidates[1..] {
let better_evidence = cand.negative_log_evidence < best.negative_log_evidence;
let tied = cand.negative_log_evidence == best.negative_log_evidence;
let cheaper_on_tie = tied && cand.num_parameters < best.num_parameters;
if better_evidence || cheaper_on_tie {
best = *cand;
}
}
Some(HybridAtomChoice {
param: best.param,
negative_log_evidence: best.negative_log_evidence,
num_parameters: best.num_parameters,
curved_turning,
curved_evidence_margin,
})
}
#[derive(Debug, Clone)]
pub struct HybridSplitSelection {
pub atoms: Vec<HybridAtomChoice>,
pub total_negative_log_evidence: f64,
pub total_parameters: usize,
pub curved_atom_count: usize,
}
impl HybridSplitSelection {
pub fn linear_atom_count(&self) -> usize {
self.atoms.len() - self.curved_atom_count
}
pub fn is_pure_linear(&self) -> bool {
self.curved_atom_count == 0 && !self.atoms.is_empty()
}
pub fn is_pure_curved(&self) -> bool {
self.curved_atom_count == self.atoms.len() && !self.atoms.is_empty()
}
}
pub fn select_hybrid_split(
slots: &[Vec<HybridAtomCandidate>],
) -> Result<HybridSplitSelection, String> {
let mut atoms = Vec::with_capacity(slots.len());
let mut total_nle = 0.0_f64;
let mut total_parameters = 0usize;
let mut curved_atom_count = 0usize;
for (i, slot) in slots.iter().enumerate() {
let choice = select_hybrid_atom(slot)
.ok_or_else(|| format!("hybrid split slot {i} has no candidate parameterizations"))?;
if !choice.negative_log_evidence.is_finite() {
return Err(format!(
"hybrid split slot {i} selected a non-finite evidence ({})",
choice.negative_log_evidence
));
}
if !choice.param.is_linear() {
curved_atom_count += 1;
}
total_nle += choice.negative_log_evidence;
total_parameters += choice.num_parameters;
atoms.push(choice);
}
Ok(HybridSplitSelection {
atoms,
total_negative_log_evidence: total_nle,
total_parameters,
curved_atom_count,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::arrow_schur::ArrowFactorSlab;
use ndarray::array;
fn dense_inverse(h: &Array2<f64>) -> Array2<f64> {
let p = h.nrows();
let mut aug = Array2::<f64>::zeros((p, 2 * p));
for i in 0..p {
for j in 0..p {
aug[[i, j]] = h[[i, j]];
}
aug[[i, p + i]] = 1.0;
}
for col in 0..p {
let mut pivot = col;
for row in (col + 1)..p {
if aug[[row, col]].abs() > aug[[pivot, col]].abs() {
pivot = row;
}
}
if pivot != col {
for j in 0..(2 * p) {
aug.swap([col, j], [pivot, j]);
}
}
let d = aug[[col, col]];
for j in 0..(2 * p) {
aug[[col, j]] /= d;
}
for row in 0..p {
if row == col {
continue;
}
let f = aug[[row, col]];
if f != 0.0 {
for j in 0..(2 * p) {
aug[[row, j]] -= f * aug[[col, j]];
}
}
}
}
let mut inv = Array2::<f64>::zeros((p, p));
for i in 0..p {
for j in 0..p {
inv[[i, j]] = aug[[i, p + j]];
}
}
inv
}
#[test]
fn coupling_components_block_diagonal_is_all_singletons_by_block() {
let mut h = Array2::<f64>::eye(4);
h[[0, 1]] = 0.3;
h[[1, 0]] = 0.3;
h[[2, 3]] = 0.7;
h[[3, 2]] = 0.7;
let labels = coupling_components(h.view());
assert_eq!(labels[0], labels[1]);
assert_eq!(labels[2], labels[3]);
assert_ne!(labels[0], labels[2]);
let mut uniq = labels.clone();
uniq.sort_unstable();
uniq.dedup();
assert_eq!(uniq.len(), 2);
}
#[test]
fn coupling_components_fully_coupled_is_one_component() {
let mut h = Array2::<f64>::eye(3);
for i in 0..3 {
for j in 0..3 {
if i != j {
h[[i, j]] = 0.1;
}
}
}
let labels = coupling_components(h.view());
assert!(labels.iter().all(|&l| l == labels[0]));
}
#[test]
fn coupling_components_transitive_chain_merges() {
let mut h = Array2::<f64>::eye(3);
h[[0, 1]] = 0.5;
h[[1, 0]] = 0.5;
h[[1, 2]] = 0.5;
h[[2, 1]] = 0.5;
let labels = coupling_components(h.view());
assert_eq!(labels[0], labels[1]);
assert_eq!(labels[1], labels[2]);
}
#[test]
fn compare_reml_fits_delta_and_bayes_factor_never_contradict_winner_gh1465() {
let cand = |name: &str, score: f64, edf: f64| RemlCandidate {
index: 0,
name: name.to_string(),
score,
edf: Some(edf),
log_lik: Some(0.0),
family: Some("gaussian".to_string()),
n_obs: Some(100),
};
let candidates = vec![
cand("m1", 53.748, 50.0),
cand("m2", 41.605, 51.0),
cand("m3", 120.011, 65.0),
];
let cmp = compare_reml_fits(candidates).expect("comparison");
assert_eq!(cmp.winner, "m1", "AIC winner");
for row in &cmp.ranking {
assert!(
row.delta >= 0.0,
"ranking delta for {} must be >= 0, got {}",
row.name,
row.delta
);
assert!(
row.bayes_factor >= 1.0 - 1e-12,
"ranking bayes_factor for {} must be >= 1, got {}",
row.name,
row.bayes_factor
);
}
let winner_row = cmp.ranking.iter().find(|r| r.name == "m1").unwrap();
assert!(winner_row.delta.abs() < 1e-12, "winner delta == 0");
assert!(
(winner_row.bayes_factor - 1.0).abs() < 1e-9,
"winner bayes_factor == 1"
);
for row in &cmp.score_table {
assert!(
row.delta_reml >= 0.0,
"score-table delta_reml for {} must be >= 0, got {}",
row.name,
row.delta_reml
);
assert!(
row.bayes_factor_best_over_model >= 1.0 - 1e-12,
"score-table bayes_factor for {} must be >= 1, got {}",
row.name,
row.bayes_factor_best_over_model
);
}
let m2 = cmp.score_table.iter().find(|r| r.name == "m2").unwrap();
assert!(
m2.delta_reml.abs() < 1e-12,
"the minimum-raw-REML row has delta_reml 0"
);
}
#[test]
fn cone_of_influence_empty_support_is_empty() {
let labels = vec![0usize, 0, 1, 1];
assert!(cone_of_influence(&labels, &[]).is_empty());
}
#[test]
fn cone_of_influence_returns_full_component() {
let labels = vec![0usize, 0, 1, 1];
assert_eq!(cone_of_influence(&labels, &[0]), vec![0, 1]);
assert_eq!(cone_of_influence(&labels, &[1, 2]), vec![0, 1, 2, 3]);
}
#[test]
fn coned_matches_full_solve_on_fully_coupled_hessian() {
let h = Array2::from_shape_vec((3, 3), vec![4.0, 1.0, 0.5, 1.0, 3.0, 0.8, 0.5, 0.8, 2.5])
.unwrap();
let inv = dense_inverse(&h);
let mut dg = Array2::<f64>::zeros((3, 2));
dg[[0, 0]] = 1.3;
dg[[2, 1]] = -0.7;
let supports = vec![0..1usize, 2..3usize];
let eye: Array2<f64> = Array2::eye(3);
let op = crate::sensitivity::FitSensitivity::from_projected(&eye, &inv);
let full = op.mode_response(dg.view()).unwrap();
let coned = op
.mode_response_coned(h.view(), dg.view(), &supports)
.unwrap();
for i in 0..3 {
for a in 0..2 {
assert!(
(full[[i, a]] - coned[[i, a]]).abs() < 1e-12,
"fully-coupled mismatch at ({i},{a}): {} vs {}",
full[[i, a]],
coned[[i, a]]
);
}
}
}
#[test]
fn coned_confines_to_component_on_decoupled_hessian() {
let mut h = Array2::<f64>::zeros((4, 4));
h[[0, 0]] = 4.0;
h[[1, 1]] = 3.0;
h[[0, 1]] = 1.0;
h[[1, 0]] = 1.0;
h[[2, 2]] = 2.0;
h[[3, 3]] = 5.0;
h[[2, 3]] = 0.6;
h[[3, 2]] = 0.6;
let inv = dense_inverse(&h);
let mut dg = Array2::<f64>::zeros((4, 1));
dg[[0, 0]] = 0.9;
dg[[1, 0]] = -0.4;
let support_range = 0..2usize;
let supports = std::slice::from_ref(&support_range);
let eye: Array2<f64> = Array2::eye(4);
let coned = crate::sensitivity::FitSensitivity::from_projected(&eye, &inv)
.mode_response_coned(h.view(), dg.view(), supports)
.unwrap();
let q = dg.column(0).to_owned();
let exact = inv.dot(&q).mapv(|v| -v);
for i in 0..4 {
assert!(
(coned[[i, 0]] - exact[[i]]).abs() < 1e-12,
"decoupled mismatch at {i}: {} vs {}",
coned[[i, 0]],
exact[[i]]
);
}
assert_eq!(coned[[2, 0]], 0.0);
assert_eq!(coned[[3, 0]], 0.0);
}
#[test]
fn coned_skips_inactive_column_with_empty_support() {
let h = Array2::<f64>::eye(2);
let dg = Array2::<f64>::zeros((2, 1));
let empty_support = 0..0usize;
let supports = std::slice::from_ref(&empty_support);
let eye: Array2<f64> = Array2::eye(2);
let nan_inv = Array2::<f64>::from_elem((2, 2), f64::NAN);
let coned = crate::sensitivity::FitSensitivity::from_projected(&eye, &nan_inv)
.mode_response_coned(h.view(), dg.view(), supports)
.unwrap();
assert_eq!(coned[[0, 0]], 0.0);
assert_eq!(coned[[1, 0]], 0.0);
}
fn make_minimal_cache() -> ArrowFactorCache {
let l_huu = Array2::from_shape_vec((1, 1), vec![std::f64::consts::SQRT_2]).unwrap();
let l_schur = Array2::from_shape_vec((1, 1), vec![(1.875_f64).sqrt()]).unwrap();
let htbeta = Array2::from_shape_vec((1, 1), vec![0.5]).unwrap();
let mut cache = ArrowFactorCache {
htt_factors: ArrowFactorSlab::from_blocks(vec![l_huu]),
htt_factors_undamped: crate::arrow_schur::ArrowUndampedFactors::SameAsDamped,
schur_factor: Some(l_schur),
schur_factor_is_undamped: true,
beta_schur_deflation: None,
joint_hessian_log_det: None,
solver_mode: crate::arrow_schur::ArrowSolverMode::Direct,
ridge_t: 0.0,
ridge_beta: 0.0,
htbeta: crate::arrow_schur::ArrowHtbetaCache::Dense {
blocks: std::sync::Arc::from(vec![htbeta]),
estimated_bytes: std::mem::size_of::<f64>(),
},
d: 1,
row_dims: std::sync::Arc::from(vec![1usize]),
row_offsets: std::sync::Arc::from(vec![0usize, 1usize]),
k: 1,
manifold_mode_fingerprint: 0,
row_hessian_fingerprint: 0,
pcg_diagnostics: crate::arrow_schur::ArrowPcgDiagnostics::default(),
gauge_deflated_directions: 0,
deflated_row_directions: std::sync::Arc::from(Vec::new()),
deflation_row_spectra: std::sync::Arc::from(Vec::new()),
beta_gauge_quotient: None,
};
cache.joint_hessian_log_det = cache.compute_undamped_arrow_log_det();
cache
}
#[test]
fn laplace_evidence_returns_finite_for_minimal_cache() {
let cache = make_minimal_cache();
let v = laplace_evidence(
EvidenceLogDetSource::FactoredArrow {
cache: &cache,
fallback_hvp: None,
},
0.0,
0.0,
2.0,
1.0,
);
assert!(v.is_finite());
let expected =
0.5 * (2.0_f64.ln() + 1.875_f64.ln()) - 0.5 * (2.0 * std::f64::consts::PI).ln();
assert!((v - expected).abs() < 1e-12);
}
fn k0_direct_cache_no_schur(latent_diag: f64) -> ArrowFactorCache {
let l_huu = Array2::from_shape_vec((1, 1), vec![latent_diag.sqrt()]).unwrap();
let mut cache = ArrowFactorCache {
htt_factors: ArrowFactorSlab::from_blocks(vec![l_huu]),
htt_factors_undamped: crate::arrow_schur::ArrowUndampedFactors::SameAsDamped,
schur_factor: None,
schur_factor_is_undamped: true,
beta_schur_deflation: None,
joint_hessian_log_det: None,
solver_mode: crate::arrow_schur::ArrowSolverMode::Direct,
ridge_t: 0.0,
ridge_beta: 0.0,
htbeta: crate::arrow_schur::ArrowHtbetaCache::Disabled { estimated_bytes: 0 },
d: 1,
row_dims: std::sync::Arc::from(vec![1usize]),
row_offsets: std::sync::Arc::from(vec![0usize, 1usize]),
k: 0,
manifold_mode_fingerprint: 0,
row_hessian_fingerprint: 0,
pcg_diagnostics: crate::arrow_schur::ArrowPcgDiagnostics::default(),
gauge_deflated_directions: 0,
deflated_row_directions: std::sync::Arc::from(Vec::new()),
deflation_row_spectra: std::sync::Arc::from(Vec::new()),
beta_gauge_quotient: None,
};
cache.joint_hessian_log_det = cache.compute_undamped_arrow_log_det();
cache
}
#[test]
fn arrow_log_det_some_for_k0_direct_cache_without_schur() {
let cache = k0_direct_cache_no_schur(3.0);
let log_det = arrow_log_det_from_cache(&cache)
.expect("k==0 Direct cache must yield Some(per-row sum), not None (#1132)");
assert!(
(log_det - 3.0_f64.ln()).abs() < 1e-12,
"log_det = {log_det}"
);
let cached = cache
.compute_undamped_arrow_log_det()
.expect("compute_undamped_arrow_log_det must be Some for k==0");
assert!((cached - 3.0_f64.ln()).abs() < 1e-12, "cached = {cached}");
}
#[test]
fn arrow_log_det_none_for_kpos_cache_without_schur() {
let mut cache = k0_direct_cache_no_schur(3.0);
cache.k = 1;
cache.solver_mode = crate::arrow_schur::ArrowSolverMode::InexactPCG;
cache.joint_hessian_log_det = None;
assert!(arrow_log_det_from_cache(&cache).is_none());
assert!(cache.compute_undamped_arrow_log_det().is_none());
}
#[test]
fn laplace_evidence_nan_when_authoritative_logdet_missing() {
let mut cache = make_minimal_cache();
cache.ridge_t = 1e-3;
cache.joint_hessian_log_det = None;
assert!(
laplace_evidence(
EvidenceLogDetSource::FactoredArrow {
cache: &cache,
fallback_hvp: None,
},
0.0,
0.0,
2.0,
1.0,
)
.is_nan()
);
}
#[test]
fn laplace_evidence_uses_hvp_fallback_without_authoritative_logdet() {
let mut cache = make_minimal_cache();
cache.schur_factor = None;
cache.joint_hessian_log_det = None;
let hvp = |x: &[f64]| -> Vec<f64> { vec![2.0 * x[0], 1.875 * x[1]] };
let v = laplace_evidence(
EvidenceLogDetSource::FactoredArrow {
cache: &cache,
fallback_hvp: Some(EvidenceHvpLogDet {
dim: 2,
apply: &hvp,
}),
},
0.0,
0.0,
2.0,
1.0,
);
let expected =
0.5 * (2.0_f64.ln() + 1.875_f64.ln()) - 0.5 * (2.0 * std::f64::consts::PI).ln();
assert!((v - expected).abs() < 1e-12);
}
#[test]
fn ift_du_dbeta_has_expected_shape() {
let cache = make_minimal_cache();
let du_db = ift_du_dbeta(&cache);
assert_eq!(du_db.shape(), &[1, 1]);
assert!((du_db[[0, 0]] - (-0.25)).abs() < 1e-12);
}
#[test]
fn ift_dbeta_drho_returns_some_for_direct_cache() {
let cache = make_minimal_cache();
let q = Array2::from_shape_vec((1, 1), vec![1.0]).unwrap();
let out = ift_dbeta_drho(&cache, q.view()).unwrap();
assert_eq!(out.shape(), &[1, 1]);
assert!((out[[0, 0]] + 1.0 / 1.875).abs() < 1e-12);
}
#[test]
fn topology_select_picks_lowest_negative_log_evidence() {
let candidates = vec![
TopologyCandidate {
kind: TopologyKind::Flat,
negative_log_evidence: 10.0,
effective_dim: 4.0,
n_obs: 100,
converged: true,
exclusion_reason: None,
},
TopologyCandidate {
kind: TopologyKind::Sphere,
negative_log_evidence: 8.0,
effective_dim: 5.0,
n_obs: 100,
converged: true,
exclusion_reason: None,
},
TopologyCandidate {
kind: TopologyKind::Torus,
negative_log_evidence: f64::NAN,
effective_dim: 6.0,
n_obs: 100,
converged: false,
exclusion_reason: Some("torus periods missing".to_string()),
},
];
let sel = select_topology(&candidates, TopologySelectOptions::default());
assert_eq!(sel.winner, TopologyKind::Sphere);
assert!(!sel.tie);
}
#[test]
fn topology_select_tie_breaks_to_simpler() {
let candidates = vec![
TopologyCandidate {
kind: TopologyKind::Sphere,
negative_log_evidence: 5.0,
effective_dim: 5.0,
n_obs: 100,
converged: true,
exclusion_reason: None,
},
TopologyCandidate {
kind: TopologyKind::Flat,
negative_log_evidence: 5.0 + 1e-6,
effective_dim: 4.0,
n_obs: 100,
converged: true,
exclusion_reason: None,
},
];
let sel = select_topology(&candidates, TopologySelectOptions::default());
assert_eq!(sel.winner, TopologyKind::Flat);
assert!(sel.tie);
}
fn gaussian_logpdf(y: f64, mean: f64, sd: f64) -> f64 {
let z = (y - mean) / sd;
-0.5 * (2.0 * std::f64::consts::PI).ln() - sd.ln() - 0.5 * z * z
}
#[test]
fn stacking_single_candidate_gets_full_weight() {
let log_density = Array2::from_shape_vec((3, 1), vec![-1.0, -2.0, -0.5]).unwrap();
let out = solve_stacking_weights(log_density.view(), StackingConfig::default()).unwrap();
assert!((out.weights[0] - 1.0).abs() < 1e-12);
assert_eq!(out.weights.len(), 1);
}
#[test]
fn stacking_dominant_candidate_attracts_nearly_all_weight() {
let mut log_density = Array2::<f64>::zeros((50, 2));
for i in 0..50 {
log_density[[i, 0]] = -0.1;
log_density[[i, 1]] = -5.0;
}
let out = solve_stacking_weights(log_density.view(), StackingConfig::default()).unwrap();
assert!(out.weights[0] > 0.99, "w0 = {}", out.weights[0]);
assert!(out.weights[1] < 0.01, "w1 = {}", out.weights[1]);
}
#[test]
fn stacking_complementary_candidates_share_weight() {
let n = 40;
let mut log_density = Array2::<f64>::zeros((n, 2));
for i in 0..n {
if i < n / 2 {
log_density[[i, 0]] = gaussian_logpdf(0.0, 0.0, 0.5);
log_density[[i, 1]] = gaussian_logpdf(0.0, 1.5, 0.5);
} else {
log_density[[i, 0]] = gaussian_logpdf(0.0, 1.5, 0.5);
log_density[[i, 1]] = gaussian_logpdf(0.0, 0.0, 0.5);
}
}
let out = solve_stacking_weights(log_density.view(), StackingConfig::default()).unwrap();
assert!(
out.weights[0] > 0.2 && out.weights[0] < 0.8,
"w0 = {}",
out.weights[0]
);
assert!((out.weights.sum() - 1.0).abs() < 1e-9);
}
#[test]
fn stacking_weights_stay_on_the_simplex() {
let log_density = Array2::from_shape_vec(
(3, 3),
vec![-1.0, -2.0, -3.0, -2.5, -1.0, -2.0, -3.0, -2.0, -1.0],
)
.unwrap();
let out = solve_stacking_weights(log_density.view(), StackingConfig::default()).unwrap();
assert!((out.weights.sum() - 1.0).abs() < 1e-9);
assert!(out.weights.iter().all(|&w| w >= -1e-12));
}
#[test]
fn stacking_solution_satisfies_the_simplex_kkt_certificate() {
let log_density = Array2::from_shape_vec(
(5, 2),
vec![-0.2, -3.0, -3.0, -0.2, -0.5, -1.5, -1.5, -0.5, -0.1, -2.0],
)
.unwrap();
let config = StackingConfig::default();
let out = solve_stacking_weights(log_density.view(), config).unwrap();
assert!(out.certificate.residual() <= config.kkt_tol);
let n = log_density.nrows();
for k in 0..2 {
let mut g = 0.0_f64;
for i in 0..n {
let mix: f64 = (0..2)
.map(|c| out.weights[c] * log_density[[i, c]].exp())
.sum();
g += log_density[[i, k]].exp() / mix;
}
g /= n as f64;
assert!(
g <= 1.0 + config.kkt_tol,
"stationarity violated for candidate {k}: g = {g}"
);
assert!(
out.weights[k] * (g - 1.0).abs() <= config.kkt_tol * (1.0 + 1e-6),
"complementary slackness violated for candidate {k}: w = {}, g = {g}",
out.weights[k]
);
}
}
#[test]
fn stacking_exhaustion_without_certificate_is_an_error_not_weights() {
let log_density = Array2::from_shape_vec(
(6, 3),
vec![
0.0, -2.0, -4.0, -0.4, -0.1, -3.0, -2.0, 0.0, -0.3, -3.0, -1.0, 0.0, -0.2, -2.0,
-0.5, -1.0, -0.3, -2.0,
],
)
.unwrap();
let config = StackingConfig {
max_iter: 1,
..StackingConfig::default()
};
let err = solve_stacking_weights(log_density.view(), config).unwrap_err();
let checkpoint = match err {
StackingError::DidNotConverge {
certificate,
checkpoint,
..
} => {
assert!(certificate.residual() > config.kkt_tol);
assert_eq!(checkpoint.completed_iterations, 1);
checkpoint
}
other => panic!("expected typed stacking exhaustion, got {other}"),
};
let encoded = serde_json::to_string(&checkpoint).unwrap();
let checkpoint: StackingCheckpoint = serde_json::from_str(&encoded).unwrap();
let mut other_density = log_density.clone();
other_density[[0, 0]] += 0.25;
assert!(matches!(
resume_stacking_weights(other_density.view(), StackingConfig::default(), &checkpoint,),
Err(StackingError::InvalidInput { .. })
));
let resumed =
resume_stacking_weights(log_density.view(), StackingConfig::default(), &checkpoint)
.unwrap();
let uninterrupted =
solve_stacking_weights(log_density.view(), StackingConfig::default()).unwrap();
for (resumed, uninterrupted) in resumed.weights.iter().zip(uninterrupted.weights.iter()) {
assert!((resumed - uninterrupted).abs() <= 1.0e-10);
}
}
#[test]
fn stacking_near_tied_boundary_uses_newton_not_millions_of_em_steps() {
let log_density =
Array2::from_shape_fn(
(64, 2),
|(_, candidate)| {
if candidate == 0 { 0.0 } else { -1.0e-6 }
},
);
let out = solve_stacking_weights(log_density.view(), StackingConfig::default()).unwrap();
assert!(out.weights[0] >= 1.0 - StackingConfig::default().kkt_tol);
assert!(out.iterations < 8, "iterations = {}", out.iterations);
}
#[test]
fn stacking_dead_candidate_column_gets_zero_weight() {
let log_density = Array2::from_shape_vec(
(3, 2),
vec![
-1.0,
f64::NEG_INFINITY,
-2.0,
f64::NEG_INFINITY,
-0.5,
f64::NEG_INFINITY,
],
)
.unwrap();
let out = solve_stacking_weights(log_density.view(), StackingConfig::default()).unwrap();
assert_eq!(out.weights[1], 0.0);
assert!((out.weights[0] - 1.0).abs() < 1e-12);
}
#[test]
fn stacking_rejects_invalid_and_unscorable_rows() {
let log_density = Array2::from_shape_vec(
(3, 2),
vec![-1.0, -2.0, f64::NAN, f64::NEG_INFINITY, -2.0, -1.0],
)
.unwrap();
assert!(matches!(
solve_stacking_weights(log_density.view(), StackingConfig::default()),
Err(StackingError::InvalidInput { .. })
));
let unscorable = Array2::from_shape_vec(
(2, 2),
vec![-1.0, -2.0, f64::NEG_INFINITY, f64::NEG_INFINITY],
)
.unwrap();
assert!(matches!(
solve_stacking_weights(unscorable.view(), StackingConfig::default()),
Err(StackingError::InvalidInput { .. })
));
}
fn two_cluster_mixture_data() -> Array2<f64> {
Array2::from_shape_vec(
(12, 1),
vec![
-2.2, -2.0, -1.9, -2.1, -1.8, -2.05, 1.8, 2.0, 2.2, 1.9, 2.1, 2.05,
],
)
.unwrap()
}
#[test]
fn gaussian_mixture_monotonicity_resolves_composite_map_noise_2264() {
let objective_scale = 1.0;
let composite_resolution = f64::EPSILON.sqrt() * objective_scale;
let uncertainty = gaussian_mixture_monotonicity_uncertainty(objective_scale, 0.0, 0.0);
assert_eq!(uncertainty, composite_resolution);
let noise_scale_decrease = -0.5 * composite_resolution;
assert!(noise_scale_decrease >= -uncertainty);
let resolved_decrease = -2.0 * composite_resolution;
assert!(resolved_decrease < -uncertainty);
let larger_reduction_bound = 2.0 * composite_resolution;
assert_eq!(
gaussian_mixture_monotonicity_uncertainty(objective_scale, larger_reduction_bound, 0.0),
larger_reduction_bound,
);
}
#[test]
fn gaussian_mixture_issue_scale_negative_step_is_within_computed_uncertainty_2264() {
let objective_scale = 1.0;
let recorded_step = -1.4e-13;
let uncertainty = gaussian_mixture_monotonicity_uncertainty(objective_scale, 0.0, 0.0);
let certificate = GaussianMixtureCertificate {
mean_log_likelihood: -objective_scale,
mean_log_likelihood_gain: recorded_step,
monotonicity_uncertainty: uncertainty,
objective_residual: recorded_step.abs() / objective_scale,
objective_tolerance: f64::EPSILON.sqrt(),
parameter_residual: 0.0,
parameter_tolerance: f64::EPSILON.sqrt(),
};
assert_eq!(
certificate.monotonicity_uncertainty,
f64::EPSILON.sqrt() * objective_scale,
"reported uncertainty must be the computed composite-map resolution"
);
assert!(
certificate.mean_log_likelihood_gain >= -certificate.monotonicity_uncertainty,
"the recorded noise-scale decrease must not be a monotonicity violation"
);
}
#[test]
fn gaussian_mixture_below_roundoff_positive_gain_can_certify_2264() {
let objective_scale = 1.0;
let recorded_gain = 6.6e-15;
let recorded_reduction_bound = 1.5e-14;
let objective_tolerance = f64::EPSILON.sqrt();
let parameter_tolerance = f64::EPSILON.sqrt();
let uncertainty = gaussian_mixture_monotonicity_uncertainty(
objective_scale,
recorded_reduction_bound,
0.0,
);
let certificate = GaussianMixtureCertificate {
mean_log_likelihood: -objective_scale,
mean_log_likelihood_gain: recorded_gain,
monotonicity_uncertainty: uncertainty,
objective_residual: recorded_gain / objective_scale,
objective_tolerance,
parameter_residual: 0.5 * parameter_tolerance,
parameter_tolerance,
};
assert_eq!(
certificate.monotonicity_uncertainty,
(f64::EPSILON.sqrt() * objective_scale).max(recorded_reduction_bound),
"reported uncertainty must come from the composite-map and reduction bounds"
);
assert!(certificate.mean_log_likelihood_gain >= -certificate.monotonicity_uncertainty);
assert!(certificate.objective_residual <= certificate.objective_tolerance);
assert!(certificate.parameter_residual <= certificate.parameter_tolerance);
}
#[test]
fn gaussian_mixture_parameter_map_uses_component_measure_geometry_2324() {
let tolerance = f64::EPSILON.sqrt();
let raw_coordinate_step: f64 = 4.6e-8;
assert!(raw_coordinate_step > tolerance);
let weights = array![0.25, 0.75];
let previous_means = array![[0.0], [2.0]];
let next_means = array![[raw_coordinate_step], [2.0]];
let covariance = vec![array![[1.0]], array![[1.0]]];
let residual = mixture_parameter_residual(
&weights,
&previous_means,
&covariance,
&weights,
&next_means,
&covariance,
);
assert_eq!(residual, weights[0] * raw_coordinate_step);
assert!(residual <= tolerance);
let next_weights = array![
weights[0] + raw_coordinate_step,
weights[1] - raw_coordinate_step
];
let mass_residual = mixture_parameter_residual(
&weights,
&previous_means,
&covariance,
&next_weights,
&previous_means,
&covariance,
);
assert!(mass_residual > tolerance);
}
#[test]
fn gaussian_mixture_fit_certificate_describes_the_exact_returned_iterate() {
let data = two_cluster_mixture_data();
let config = GaussianMixtureConfig::default();
let fit = fit_gaussian_mixture(data.view(), 2, config).unwrap();
let certificate = fit.certificate();
assert!(certificate.objective_residual <= certificate.objective_tolerance);
assert!(certificate.parameter_residual <= certificate.parameter_tolerance);
let checkpoint = GaussianMixtureCheckpoint {
weights: fit.weights.clone(),
means: fit.means.clone(),
covariances: fit.covariances.clone(),
mean_log_likelihood: certificate.mean_log_likelihood,
completed_iterations: fit.iterations,
data_fingerprint: mixture_data_fingerprint(data.view()),
covariance_floor: config.covariance_floor,
};
let current = mixture_e_step(
data.view(),
&checkpoint.weights,
&checkpoint.means,
&checkpoint.covariances,
)
.unwrap();
let (weights, means, covariances) = mixture_m_step(
data.view(),
current.responsibilities.view(),
config.covariance_floor,
)
.unwrap();
let residual = mixture_parameter_residual(
&checkpoint.weights,
&checkpoint.means,
&checkpoint.covariances,
&weights,
&means,
&covariances,
);
let next = mixture_e_step(data.view(), &weights, &means, &covariances).unwrap();
assert!(residual <= config.parameter_tol);
assert_eq!(certificate.mean_log_likelihood, current.mean_log_likelihood);
assert_eq!(
certificate.mean_log_likelihood_gain,
next.mean_log_likelihood - current.mean_log_likelihood
);
assert_eq!(
certificate.monotonicity_uncertainty,
gaussian_mixture_monotonicity_uncertainty(
current
.mean_log_likelihood
.abs()
.max(next.mean_log_likelihood.abs())
.max(1.0),
current.mean_log_likelihood_roundoff,
next.mean_log_likelihood_roundoff,
)
);
assert_eq!(certificate.parameter_residual, residual);
assert!(
(next.mean_log_likelihood - current.mean_log_likelihood).abs()
/ current
.mean_log_likelihood
.abs()
.max(next.mean_log_likelihood.abs())
.max(1.0)
<= config.loglik_tol
);
}
#[test]
fn gaussian_mixture_exhaustion_is_typed_and_resumable() {
let data = two_cluster_mixture_data();
let short = GaussianMixtureConfig {
max_iter: 1,
..GaussianMixtureConfig::default()
};
let err = fit_gaussian_mixture(data.view(), 2, short).unwrap_err();
let checkpoint = match err {
GaussianMixtureError::DidNotConverge {
certificate,
checkpoint,
..
} => {
assert!(
certificate.objective_residual > short.loglik_tol
|| certificate.parameter_residual > short.parameter_tol
);
assert_eq!(checkpoint.completed_iterations, 1);
let at_checkpoint = mixture_e_step(
data.view(),
&checkpoint.weights,
&checkpoint.means,
&checkpoint.covariances,
)
.unwrap();
assert_eq!(
certificate.mean_log_likelihood, at_checkpoint.mean_log_likelihood,
"exhaustion evidence and checkpoint must describe one iterate"
);
checkpoint
}
other => panic!("expected typed EM exhaustion, got {other}"),
};
let encoded = serde_json::to_string(&checkpoint).unwrap();
let checkpoint: GaussianMixtureCheckpoint = serde_json::from_str(&encoded).unwrap();
let mut other_data = data.clone();
other_data[[0, 0]] += 0.01;
assert!(matches!(
resume_gaussian_mixture(
other_data.view(),
GaussianMixtureConfig::default(),
checkpoint.clone(),
),
Err(GaussianMixtureError::InvalidInput { .. })
));
let resumed =
resume_gaussian_mixture(data.view(), GaussianMixtureConfig::default(), checkpoint)
.unwrap();
let uninterrupted =
fit_gaussian_mixture(data.view(), 2, GaussianMixtureConfig::default()).unwrap();
for (resumed, uninterrupted) in resumed.weights.iter().zip(uninterrupted.weights.iter()) {
assert!((resumed - uninterrupted).abs() <= 1.0e-10);
}
assert!(resumed.bic().is_finite());
}
#[test]
fn gaussian_mixture_bic_is_finite_with_an_active_covariance_floor() {
let per_cluster = 45usize;
let mut data = Array2::<f64>::zeros((2 * per_cluster, 2));
for sample in 0..per_cluster {
let phase = std::f64::consts::TAU * sample as f64 / per_cluster as f64;
data[[2 * sample, 0]] = -2.0;
data[[2 * sample, 1]] = 0.08 * phase.sin();
data[[2 * sample + 1, 0]] = 2.0 + 0.12 * phase.cos();
data[[2 * sample + 1, 1]] = 0.08 * phase.sin();
}
let fit = fit_gaussian_mixture(data.view(), 2, GaussianMixtureConfig::default())
.expect("the covariance floor defines a valid constrained mixture fit");
let bic = fit.bic();
assert!(bic.is_finite());
assert_eq!(
bic,
-fit.loglik + 0.5 * fit.num_free_parameters() as f64 * (data.nrows() as f64).ln()
);
}
fn seven_clusters_on_a_circle_2262() -> Array2<f64> {
let clusters = 7usize;
let per_cluster = 32usize;
let mut data = Array2::<f64>::zeros((clusters * per_cluster, 2));
for cluster in 0..clusters {
let angle = std::f64::consts::TAU * cluster as f64 / clusters as f64;
let (sin_angle, cos_angle) = angle.sin_cos();
for sample in 0..per_cluster {
let phase = std::f64::consts::TAU * sample as f64 / per_cluster as f64;
let local_radius = 0.035 * (1.0 + 0.3 * (3.0 * phase).cos());
let radial_noise = local_radius * phase.cos();
let tangent_noise = local_radius * phase.sin();
let radius = 2.0 + radial_noise;
let row = cluster * per_cluster + sample;
data[[row, 0]] = 0.4 + radius * cos_angle - tangent_noise * sin_angle;
data[[row, 1]] = -0.3 + radius * sin_angle + tangent_noise * cos_angle;
}
}
data
}
fn two_noisy_circles_for_union() -> Array2<f64> {
let rows_per_circle = 96usize;
let mut data = Array2::<f64>::zeros((2 * rows_per_circle, 2));
for (circle, (center, radius)) in [([-4.0_f64, 0.3_f64], 1.2_f64), ([4.0, -0.2], 0.9)]
.into_iter()
.enumerate()
{
for sample in 0..rows_per_circle {
let angle = std::f64::consts::TAU * sample as f64 / rows_per_circle as f64;
let noisy_radius =
radius + 0.045 * (3.0 * angle).cos() + 0.018 * (5.0 * angle).sin();
let row = circle * rows_per_circle + sample;
data[[row, 0]] = center[0] + noisy_radius * angle.cos();
data[[row, 1]] = center[1] + noisy_radius * angle.sin();
}
}
data
}
#[test]
fn circular_gaussian_density_avoids_extreme_scale_intermediate_overflow() {
let noise_variance = f64::MAX / 2.0;
let fit =
CircularGaussianFit2d::from_parameters([0.0, 0.0], 1.1e154, noise_variance).unwrap();
let center_log_density = fit.log_density(0.0, 0.0);
let off_center_log_density = fit.log_density(1.7e154, 0.0);
assert!(center_log_density.is_finite());
assert!(off_center_log_density.is_finite());
let expected_center = -std::f64::consts::TAU.ln()
- noise_variance.ln()
- 0.5 * (fit.radius() / noise_variance.sqrt()).powi(2);
assert_eq!(center_log_density, expected_center);
}
#[test]
fn union_circles_use_the_shared_normalized_cartesian_density() {
let data = two_noisy_circles_for_union();
let config = GaussianMixtureConfig::default();
let density_fit =
fit_union_density(data.view(), UnionStructure::CircleCircle, config).unwrap();
let union = fit_union_structure(data.view(), UnionStructure::CircleCircle, config).unwrap();
assert_eq!(
union.total_parameters,
2 * CircularGaussianFit2d::NUM_FREE_PARAMETERS + 1
);
let component_weight_sum: f64 = union
.components
.iter()
.map(|component| component.mixing_weight)
.sum();
assert!((component_weight_sum - 1.0).abs() <= 8.0 * f64::EPSILON);
let mut fitted_centers = Array2::<f64>::zeros((density_fit.components.len(), 2));
for (index, component) in density_fit.components.iter().enumerate() {
let UnionDensityModel::Circle(fit) = &component.model else {
panic!("circle+circle union produced a non-circle density");
};
let center = fit.center();
fitted_centers[[index, 0]] = center[0];
fitted_centers[[index, 1]] = center[1];
let at_center = fit.log_density(center[0], center[1]);
let expected = -std::f64::consts::TAU.ln()
- fit.noise_variance().ln()
- 0.5 * (fit.radius() / fit.noise_variance().sqrt()).powi(2);
assert!(at_center.is_finite());
assert!((at_center - expected).abs() < 1.0e-12 * (1.0 + expected.abs()));
}
let training_log_density = union_per_point_log_density(
data.view(),
data.view(),
UnionStructure::CircleCircle,
config,
)
.unwrap();
let direct_log_likelihood = pairwise_sum(
training_log_density
.as_slice()
.expect("owned score vector is contiguous"),
);
assert!(
(union.log_likelihood - direct_log_likelihood).abs()
<= 1.0e-12 * (1.0 + direct_log_likelihood.abs())
);
let expected_bic = -direct_log_likelihood
+ 0.5 * union.total_parameters as f64 * (data.nrows() as f64).ln();
assert!((union.bic - expected_bic).abs() <= 1.0e-12 * (1.0 + expected_bic.abs()));
let held_out = union_per_point_log_density(
data.view(),
fitted_centers.view(),
UnionStructure::CircleCircle,
config,
)
.unwrap();
assert!(held_out.iter().all(|value| value.is_finite()));
}
fn circle_and_point_union_data() -> (Array2<f64>, Vec<Vec<usize>>) {
let circle_rows = 32usize;
let point_rows = 12usize;
let mut data = Array2::<f64>::zeros((circle_rows + point_rows, 2));
for row in 0..circle_rows {
let angle = std::f64::consts::TAU * row as f64 / circle_rows as f64;
let radius = 1.0 + 0.025 * (3.0 * angle).cos();
data[[row, 0]] = -4.0 + radius * angle.cos();
data[[row, 1]] = 0.2 + radius * angle.sin();
}
for offset in 0..point_rows {
let phase = offset as f64;
let row = circle_rows + offset;
data[[row, 0]] = 4.0 + 0.055 * (1.7 * phase).cos() + 0.018 * (0.4 * phase).sin();
data[[row, 1]] = -0.3 + 0.052 * (1.3 * phase).sin() - 0.015 * (0.9 * phase).cos();
}
(
data,
vec![
(0..circle_rows).collect(),
(circle_rows..circle_rows + point_rows).collect(),
],
)
}
#[test]
fn heterogeneous_union_role_assignment_is_group_label_invariant() {
let (data, groups) = circle_and_point_union_data();
let config = GaussianMixtureConfig::default();
let forward = fit_union_density_from_groups(
data.view(),
UnionStructure::CirclePointCluster,
&groups,
config,
)
.unwrap();
let reversed_groups = vec![groups[1].clone(), groups[0].clone()];
let reversed = fit_union_density_from_groups(
data.view(),
UnionStructure::CirclePointCluster,
&reversed_groups,
config,
)
.unwrap();
assert_eq!(forward.components[0].kind, UnionComponentKind::Circle);
assert_eq!(forward.components[1].kind, UnionComponentKind::PointCluster);
assert_eq!(
reversed.components[0].kind,
UnionComponentKind::PointCluster
);
assert_eq!(reversed.components[1].kind, UnionComponentKind::Circle);
assert_eq!(forward.total_parameters, 4 + 3 + 1);
assert_eq!(reversed.total_parameters, forward.total_parameters);
assert!(
(forward.log_likelihood - reversed.log_likelihood).abs()
<= 1.0e-12 * (1.0 + forward.log_likelihood.abs())
);
assert!((forward.bic - reversed.bic).abs() <= 1.0e-12 * (1.0 + forward.bic.abs()));
}
#[test]
fn point_cluster_is_isotropic_and_line_remains_full_covariance() {
let (mut data, mut groups) = circle_and_point_union_data();
for row in 0..groups[0].len() {
let coordinate = (row as f64 - 15.5) / 4.0;
data[[row, 0]] = -4.0 + coordinate;
data[[row, 1]] = 0.2 + 0.018 * coordinate + 0.006 * (1.9 * row as f64).sin();
}
let fit = fit_union_density_from_groups(
data.view(),
UnionStructure::LineCluster,
&groups,
GaussianMixtureConfig::default(),
)
.unwrap();
assert_eq!(fit.components[0].kind, UnionComponentKind::Line);
assert_eq!(fit.components[0].num_parameters, 5);
assert_eq!(fit.components[1].kind, UnionComponentKind::PointCluster);
assert_eq!(fit.components[1].num_parameters, 3);
assert_eq!(fit.total_parameters, 5 + 3 + 1);
let UnionDensityModel::Gaussian(point) = &fit.components[1].model else {
panic!("point cluster did not produce a Gaussian density");
};
assert_eq!(point.precision[[0, 1]], 0.0);
assert_eq!(point.precision[[1, 0]], 0.0);
assert_eq!(point.precision[[0, 0]], point.precision[[1, 1]]);
let total_weight: f64 = fit
.components
.iter()
.map(|component| component.mixing_weight)
.sum();
assert!((total_weight - 1.0).abs() <= 8.0 * f64::EPSILON);
groups.reverse();
let reversed = fit_union_density_from_groups(
data.view(),
UnionStructure::LineCluster,
&groups,
GaussianMixtureConfig::default(),
)
.unwrap();
assert_eq!(
reversed.components[0].kind,
UnionComponentKind::PointCluster
);
assert_eq!(reversed.components[1].kind, UnionComponentKind::Line);
assert!((fit.bic - reversed.bic).abs() <= 1.0e-12 * (1.0 + fit.bic.abs()));
}
#[test]
fn isotropic_union_density_uses_the_same_fractional_mean_chart_as_its_mle() {
let translated = ndarray::array![[1.0e16], [1.0e16 + 2.0], [1.0e16 + 2.0]];
let fit = fit_isotropic_gaussian_component(translated.view(), 1.0e-12).unwrap();
let variance = fit.precision[[0, 0]].recip();
assert!((variance - 8.0 / 9.0).abs() <= 32.0 * f64::EPSILON);
let residuals = [-4.0 / 3.0, 2.0 / 3.0, 2.0 / 3.0];
let expected_log_norm = -0.5 * ((2.0 * std::f64::consts::PI).ln() + variance.ln());
for (row, residual) in residuals.into_iter().enumerate() {
let expected = expected_log_norm - 0.5 * residual * residual / variance;
let actual = fit.log_density(translated.row(row));
assert!(
(actual - expected).abs() <= 32.0 * f64::EPSILON * (1.0 + expected.abs()),
"row {row}: density chart disagrees with fitted MLE residual: actual={actual}, expected={expected}"
);
}
let subnormal = f64::from_bits(1);
let constant = ndarray::array![[subnormal], [subnormal], [subnormal]];
let constant_fit = fit_isotropic_gaussian_component(constant.view(), 1.0).unwrap();
assert_eq!(constant_fit.residual(constant.row(0)), vec![0.0]);
assert_eq!(
constant_fit.log_density(constant.row(0)),
constant_fit.log_norm
);
}
#[test]
fn union_ladder_fails_closed_when_one_declared_structure_fails() {
let mut data = Array2::<f64>::zeros((8, 2));
for row in 0..5 {
let angle = std::f64::consts::TAU * row as f64 / 5.0;
data[[row, 0]] = -5.0 + angle.cos();
data[[row, 1]] = angle.sin();
}
data[[5, 0]] = 5.00;
data[[5, 1]] = 0.00;
data[[6, 0]] = 5.08;
data[[6, 1]] = 0.02;
data[[7, 0]] = 4.97;
data[[7, 1]] = 0.07;
let error = fit_union_ladder(data.view(), GaussianMixtureConfig::default()).unwrap_err();
assert!(error.contains("every declared structure must fit"));
assert!(error.contains(UnionStructure::CircleCircle.as_str()));
assert!(error.contains("needs at least 5 rows"));
}
#[test]
fn ring_of_clusters_fit_is_stationary_and_complexity_priced_2262() {
let data = seven_clusters_on_a_circle_2262();
let config = GaussianMixtureConfig::default();
let fit = fit_ring_gaussian_mixture(data.view(), 7, config).unwrap();
let certificate = fit.certificate();
assert!(certificate.objective_residual <= certificate.objective_tolerance);
assert!(certificate.parameter_residual <= certificate.parameter_tolerance);
assert_eq!(fit.num_free_parameters(), 17);
assert!((fit.center()[0] - 0.4).abs() < 0.05);
assert!((fit.center()[1] + 0.3).abs() < 0.05);
assert!((fit.radius() - 2.0).abs() < 0.05);
assert!(fit.variance().is_finite() && fit.variance() > 0.0);
assert!(
fit.per_point_log_density(data.view())
.unwrap()
.iter()
.all(|value| value.is_finite())
);
assert!(fit.bic().is_finite());
let free = fit_gaussian_mixture(data.view(), 7, config).unwrap();
assert_eq!(free.num_free_parameters(), 41);
assert!(fit.num_free_parameters() < free.num_free_parameters());
}
#[test]
fn ring_certificate_uses_identifiable_component_means() {
let y = 0.91_f64.sqrt();
let weights = Array1::from_vec(vec![0.2, 0.3, 0.5]);
let previous = RingMixtureState {
weights: weights.clone(),
center: Array1::from_vec(vec![0.0, 0.0]),
radius: 1.0,
directions: Array2::from_shape_vec((3, 2), vec![0.3, y, 0.3, -y, 0.3, y]).unwrap(),
variance: 0.25,
mean_log_likelihood: -1.0,
completed_iterations: 10,
};
let next = RingMixtureState {
weights,
center: Array1::from_vec(vec![0.6, 0.0]),
radius: 1.0,
directions: Array2::from_shape_vec((3, 2), vec![-0.3, y, -0.3, -y, -0.3, y]).unwrap(),
variance: 0.25,
mean_log_likelihood: -1.0,
completed_iterations: 11,
};
assert!(relative_parameter_step(previous.center[0], next.center[0]) > 0.5);
assert_eq!(ring_identifiable_parameter_residual(&previous, &next), 0.0);
}
#[test]
fn ring_parameter_map_weights_component_motion_by_predictive_mass_2324() {
let raw_coordinate_step: f64 = 4.6e-8;
let (next_y, next_x) = raw_coordinate_step.sin_cos();
let previous = RingMixtureState {
weights: array![0.25, 0.75],
center: array![0.0, 0.0],
radius: 1.0,
directions: array![[1.0, 0.0], [0.0, 1.0]],
variance: 1.0,
mean_log_likelihood: -1.0,
completed_iterations: 10,
};
let next = RingMixtureState {
weights: previous.weights.clone(),
center: previous.center.clone(),
radius: previous.radius,
directions: array![[next_x, next_y], [0.0, 1.0]],
variance: previous.variance,
mean_log_likelihood: -1.0,
completed_iterations: 11,
};
let raw_mean_step = (next_x - 1.0).hypot(next_y);
assert!(raw_mean_step > f64::EPSILON.sqrt());
assert!(ring_identifiable_parameter_residual(&previous, &next) <= f64::EPSILON.sqrt());
}
#[test]
fn stacked_mean_is_weighted_combination() {
let weights = Array1::from_vec(vec![0.25, 0.75]);
let means = vec![
Array1::from_vec(vec![1.0, 2.0, 3.0]),
Array1::from_vec(vec![5.0, 6.0, 7.0]),
];
let out = stacked_predictive_mean(&weights, &means).unwrap();
assert!((out[0] - (0.25 * 1.0 + 0.75 * 5.0)).abs() < 1e-12);
assert!((out[2] - (0.25 * 3.0 + 0.75 * 7.0)).abs() < 1e-12);
}
#[test]
fn stacked_mean_rejects_shape_mismatch() {
let weights = Array1::from_vec(vec![0.5, 0.5]);
let means = vec![
Array1::from_vec(vec![1.0, 2.0]),
Array1::from_vec(vec![3.0]),
];
assert!(stacked_predictive_mean(&weights, &means).is_err());
}
fn hybrid_slot(
linear_nle: f64,
p_linear: usize,
latent_dim: usize,
p_curved: usize,
theta: f64,
curved_loglik_gain: f64,
) -> Vec<HybridAtomCandidate> {
let param_price =
0.5 * (p_curved as f64 - p_linear as f64) * (2.0 * std::f64::consts::PI).ln();
let curved_nle = linear_nle - curved_loglik_gain + param_price;
vec![
HybridAtomCandidate::linear(linear_nle, p_linear),
HybridAtomCandidate::curved(latent_dim, curved_nle, p_curved, Some(theta)),
]
}
#[test]
fn hybrid_dominance_floor_selects_linear_when_turning_is_zero() {
let slot = hybrid_slot(100.0, 2, 1, 5, 0.0, 0.0);
let choice = select_hybrid_atom(&slot).unwrap();
assert!(choice.param.is_linear());
assert_eq!(choice.param, HybridAtomParam::Linear);
assert!(choice.curved_turning.unwrap() <= HYBRID_LINEAR_TURNING_FLOOR);
}
#[test]
fn hybrid_selects_curved_when_turning_pays_for_itself() {
let slot = hybrid_slot(100.0, 2, 1, 5, 2.0 * std::f64::consts::PI, 30.0);
let choice = select_hybrid_atom(&slot).unwrap();
assert_eq!(choice.param, HybridAtomParam::Curved { latent_dim: 1 });
assert!(choice.curved_evidence_margin > 0.0);
}
#[test]
fn hybrid_keeps_linear_when_curvature_doesnt_pay_its_price() {
let slot = hybrid_slot(100.0, 2, 1, 5, 0.05, 0.1);
let choice = select_hybrid_atom(&slot).unwrap();
assert!(choice.param.is_linear());
assert!(choice.curved_evidence_margin <= 0.0);
}
#[test]
fn hybrid_tie_breaks_to_the_cheaper_linear_atom() {
let theta = 0.5; let nle = 42.0;
let slot = vec![
HybridAtomCandidate::linear(nle, 2),
HybridAtomCandidate::curved(1, nle, 5, Some(theta)),
];
let choice = select_hybrid_atom(&slot).unwrap();
assert!(choice.param.is_linear());
assert_eq!(choice.num_parameters, 2);
}
#[test]
fn hybrid_split_reduces_to_pure_linear_when_all_features_are_straight() {
let slots: Vec<Vec<HybridAtomCandidate>> = (0..6)
.map(|i| hybrid_slot(50.0 + i as f64, 2, 1, 5, 0.0, 0.0))
.collect();
let split = select_hybrid_split(&slots).unwrap();
assert!(split.is_pure_linear());
assert_eq!(split.curved_atom_count, 0);
assert_eq!(split.linear_atom_count(), 6);
let pure_linear: f64 = (0..6).map(|i| 50.0 + i as f64).sum();
assert!((split.total_negative_log_evidence - pure_linear).abs() < 1e-12);
}
#[test]
fn hybrid_split_reduces_to_pure_curved_when_every_feature_curves() {
let slots: Vec<Vec<HybridAtomCandidate>> = (0..5)
.map(|i| hybrid_slot(80.0 + i as f64, 2, 1, 5, 2.0 * std::f64::consts::PI, 40.0))
.collect();
let split = select_hybrid_split(&slots).unwrap();
assert!(split.is_pure_curved());
assert_eq!(split.curved_atom_count, 5);
assert_eq!(split.linear_atom_count(), 0);
}
#[test]
fn hybrid_split_on_mixed_dictionary_picks_curved_for_circles_linear_for_directions() {
let mut slots: Vec<Vec<HybridAtomCandidate>> = Vec::new();
let mut pure_linear_baseline = 0.0_f64;
for i in 0..3 {
let linear_nle = 120.0 + 3.0 * i as f64;
pure_linear_baseline += linear_nle;
slots.push(hybrid_slot(
linear_nle,
2,
1,
5,
2.0 * std::f64::consts::PI,
35.0,
));
}
for i in 0..4 {
let linear_nle = 90.0 + 2.0 * i as f64;
pure_linear_baseline += linear_nle;
slots.push(hybrid_slot(linear_nle, 2, 1, 5, 0.0, 0.0));
}
let split = select_hybrid_split(&slots).unwrap();
for (idx, choice) in split.atoms.iter().enumerate() {
if idx < 3 {
assert_eq!(
choice.param,
HybridAtomParam::Curved { latent_dim: 1 },
"circle slot {idx} should select curved"
);
} else {
assert!(
choice.param.is_linear(),
"direction slot {idx} should select linear"
);
}
}
assert_eq!(split.curved_atom_count, 3);
assert_eq!(split.linear_atom_count(), 4);
assert!(
split.total_negative_log_evidence <= pure_linear_baseline + 1e-9,
"hybrid NLE {} must be <= summed linear-candidate NLE {}",
split.total_negative_log_evidence,
pure_linear_baseline
);
assert!(split.total_negative_log_evidence < pure_linear_baseline);
}
#[test]
fn hybrid_split_rejects_empty_slot() {
let slots = vec![hybrid_slot(10.0, 2, 1, 5, 0.0, 0.0), Vec::new()];
assert!(select_hybrid_split(&slots).is_err());
}
fn cand(name: &str, score: f64, edf: f64, log_lik: f64) -> RemlCandidate {
RemlCandidate {
index: 0,
name: name.to_string(),
score,
edf: Some(edf),
log_lik: Some(log_lik),
family: None,
n_obs: None,
}
}
#[test]
fn ranking_score_is_conditional_aic_when_loglik_and_edf_present() {
let c = cand("m", 999.0, 6.748, -32.0866);
let expected = -2.0 * -32.0866 + 2.0 * 6.748;
assert!((c.ranking_score() - expected).abs() < 1e-9);
}
#[test]
fn ranking_score_falls_back_to_evidence_without_loglik() {
let c = RemlCandidate {
index: 0,
name: "m".to_string(),
score: 151.28,
edf: Some(6.0),
log_lik: None,
family: None,
n_obs: None,
};
assert_eq!(c.ranking_score(), 151.28);
}
#[test]
fn compare_models_rejects_pure_noise_smooth_despite_lower_evidence() {
let small = cand("small", 180.526, 6.748, -32.0866);
let big = cand("big", 177.404, 14.250, -32.1212);
assert!(big.score < small.score);
let cmp = compare_reml_fits(vec![small, big]).expect("compare");
assert_eq!(
cmp.winner, "small",
"compare_models must Occam-penalise the pure-noise smooth and pick the smaller model"
);
let small_row = cmp
.score_table
.iter()
.find(|r| r.name == "small")
.expect("small row");
let big_row = cmp
.score_table
.iter()
.find(|r| r.name == "big")
.expect("big row");
assert!((small_row.reml_score - 180.526).abs() < 1e-9);
assert!((big_row.reml_score - 177.404).abs() < 1e-9);
}
#[test]
fn ranking_bayes_factor_is_akaike_evidence_ratio_not_its_square() {
let delta_aic = 27.68_f64;
let winner = cand("winner", 100.0, 0.0, 0.0);
let loser = cand("loser", 110.0, 0.0, -delta_aic / 2.0);
let cmp = compare_reml_fits(vec![winner, loser]).expect("compare");
assert_eq!(cmp.winner, "winner");
let loser_row = cmp
.ranking
.iter()
.find(|r| r.name == "loser")
.expect("loser ranking row");
assert!((loser_row.delta - delta_aic).abs() < 1e-9);
let expected = (0.5 * delta_aic).exp();
assert!(
(loser_row.bayes_factor / expected - 1.0).abs() < 1e-9,
"ranking bayes_factor {} should be exp(½ΔAIC)={}, not exp(ΔAIC)={}",
loser_row.bayes_factor,
expected,
delta_aic.exp()
);
assert!(loser_row.bayes_factor < delta_aic.exp() * 0.5);
let loser_score_row = cmp
.score_table
.iter()
.find(|r| r.name == "loser")
.expect("loser score row");
let expected_reml_bf = 10.0_f64.exp();
assert!(
(loser_score_row.bayes_factor_best_over_model / expected_reml_bf - 1.0).abs() < 1e-9,
"raw-REML bayes_factor_best_over_model must stay exp(Δreml)=exp(10), got {}",
loser_score_row.bayes_factor_best_over_model
);
}
#[test]
fn compare_models_keeps_power_for_a_relevant_smooth() {
let small = cand("small", 1025.067, 6.75, -368.985);
let big = cand("big", 199.509, 14.25, -33.165);
let cmp = compare_reml_fits(vec![small, big]).expect("compare");
assert_eq!(
cmp.winner, "big",
"compare_models must retain power: the relevant smooth's model must win"
);
}
#[test]
fn compare_models_rejects_mismatched_observation_counts() {
let with_n = |name: &str, n: usize| RemlCandidate {
index: 0,
name: name.to_string(),
score: 100.0,
edf: Some(5.0),
log_lik: Some(-40.0),
family: Some("gaussian".to_string()),
n_obs: Some(n),
};
let err = compare_reml_fits(vec![with_n("big", 500), with_n("small", 100)])
.expect_err("cross-n comparison must be rejected");
assert!(
err.contains("number of observations") && err.contains("500") && err.contains("100"),
"n-guard error should name the incomparable counts, got: {err}"
);
compare_reml_fits(vec![with_n("a", 250), with_n("b", 250)])
.expect("same-n comparison must succeed");
let without_n = RemlCandidate {
index: 0,
name: "legacy".to_string(),
score: 90.0,
edf: Some(4.0),
log_lik: Some(-35.0),
family: Some("gaussian".to_string()),
n_obs: None,
};
compare_reml_fits(vec![with_n("counted", 500), without_n])
.expect("an unconstrained (None) count must not trip the guard");
}
}