use std::sync::Arc;
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub enum HessianFactorization {
Cholesky,
Ldlt,
Analytic,
Mcmc,
}
#[derive(Clone, Debug, PartialEq)]
pub struct InferenceDiagnostics {
pub converged: bool,
pub iterations: u32,
pub grad_inf_norm: f64,
pub hessian_condition: f64,
pub factorization: HessianFactorization,
pub separation_warning: bool,
pub notes: Vec<Arc<str>>,
pub backend_id: Arc<str>,
pub n_chains: Option<u32>,
pub n_warmup: Option<u32>,
pub ess_bulk_min: Option<f64>,
pub rhat_max: Option<f64>,
pub n_divergences: Option<u32>,
}
impl InferenceDiagnostics {
#[must_use]
pub fn analytic(backend_id: impl Into<Arc<str>>) -> Self {
Self {
converged: true,
iterations: 0,
grad_inf_norm: 0.0,
hessian_condition: 1.0,
factorization: HessianFactorization::Analytic,
separation_warning: false,
notes: Vec::new(),
backend_id: backend_id.into(),
n_chains: None,
n_warmup: None,
ess_bulk_min: None,
rhat_max: None,
n_divergences: None,
}
}
#[must_use]
pub fn allows_posterior(&self) -> bool {
match self.factorization {
HessianFactorization::Analytic => true,
HessianFactorization::Mcmc => {
let rhat_ok = self.rhat_max.is_some_and(|r| r.is_finite() && r < 1.05);
let ess_ok = self.ess_bulk_min.is_some_and(|e| e.is_finite() && e > 10.0);
let div_ok = self.n_divergences.is_some();
self.converged && rhat_ok && ess_ok && div_ok
}
HessianFactorization::Cholesky | HessianFactorization::Ldlt => {
self.converged
&& self.grad_inf_norm.is_finite()
&& self.hessian_condition.is_finite()
&& self.hessian_condition > 0.0
}
}
}
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct PriorSensitivitySummary {
pub prior_scales: Arc<[f64]>,
pub alphas: Arc<[f64]>,
pub effect_means: Arc<[f64]>,
pub effect_sds: Arc<[f64]>,
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct ConflictSummary {
pub source_ids: Arc<[Arc<str>]>,
pub alphas_requested: Arc<[f64]>,
pub alphas_applied: Arc<[f64]>,
pub p_values: Arc<[Option<f64>]>,
pub kl_values: Arc<[Option<f64>]>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn laplace_requires_convergence() {
let mut d = InferenceDiagnostics {
converged: false,
iterations: 10,
grad_inf_norm: 1.0,
hessian_condition: 10.0,
factorization: HessianFactorization::Cholesky,
separation_warning: false,
notes: Vec::new(),
backend_id: Arc::from("laplace"),
n_chains: None,
n_warmup: None,
ess_bulk_min: None,
rhat_max: None,
n_divergences: None,
};
assert!(!d.allows_posterior());
d.converged = true;
assert!(d.allows_posterior());
}
#[test]
fn mcmc_requires_rhat_and_ess() {
let mut d = InferenceDiagnostics {
converged: true,
iterations: 100,
grad_inf_norm: 0.0,
hessian_condition: f64::NAN,
factorization: HessianFactorization::Mcmc,
separation_warning: false,
notes: Vec::new(),
backend_id: Arc::from("hmc"),
n_chains: Some(4),
n_warmup: Some(50),
ess_bulk_min: Some(5.0),
rhat_max: Some(1.2),
n_divergences: Some(0),
};
assert!(!d.allows_posterior());
d.ess_bulk_min = Some(50.0);
d.rhat_max = Some(1.01);
assert!(d.allows_posterior());
}
}