extern crate alloc;
use alloc::vec::Vec;
use core::f64::consts::PI;
use nalgebra::Cholesky;
use rand::prelude::*;
use rand::SeedableRng;
use rand_distr::Gamma;
use rand_xoshiro::Xoshiro256PlusPlus;
use crate::math;
use crate::types::{Matrix9, Vector9};
pub const NU: f64 = 4.0;
pub const NU_L: f64 = 8.0;
pub const N_GIBBS: usize = 256;
pub const N_BURN: usize = 64;
pub const N_KEEP: usize = N_GIBBS - N_BURN;
const LAMBDA_MIN: f64 = 1e-10;
const LAMBDA_MAX: f64 = 1e10;
const KAPPA_MIN: f64 = 1e-10;
const KAPPA_MAX: f64 = 1e10;
const CONDITION_NUMBER_THRESHOLD: f64 = 1e4;
#[derive(Clone, Debug)]
pub struct GibbsResult {
pub delta_post: Vector9,
pub lambda_post: Matrix9,
pub leak_probability: f64,
pub effect_magnitude_ci: (f64, f64),
pub delta_draws: Vec<Vector9>,
pub lambda_mean: f64,
pub lambda_sd: f64,
pub lambda_cv: f64,
pub lambda_ess: f64,
pub lambda_mixing_ok: bool,
pub kappa_mean: f64,
pub kappa_sd: f64,
pub kappa_cv: f64,
pub kappa_ess: f64,
pub kappa_mixing_ok: bool,
}
pub struct GibbsSampler {
nu: f64,
sigma: f64,
l_r: Matrix9,
rng: Xoshiro256PlusPlus,
}
impl GibbsSampler {
pub fn new(l_r: &Matrix9, sigma: f64, seed: u64) -> Self {
Self {
nu: NU,
sigma,
l_r: *l_r,
rng: Xoshiro256PlusPlus::seed_from_u64(seed),
}
}
pub fn run(&mut self, delta_obs: &Vector9, sigma_n: &Matrix9, theta: f64) -> GibbsResult {
let sigma_n_reg = regularize_sigma_n(sigma_n);
let sigma_n_chol =
Cholesky::new(sigma_n_reg).expect("regularize_sigma_n should ensure SPD");
let sigma_n_inv = Self::invert_via_cholesky(&sigma_n_chol);
let r_inv = self.compute_r_inverse();
let mut retained_deltas: Vec<Vector9> = Vec::with_capacity(N_KEEP);
let mut retained_lambdas: Vec<f64> = Vec::with_capacity(N_KEEP);
let mut retained_kappas: Vec<f64> = Vec::with_capacity(N_KEEP);
let mut lambda = 1.0;
let mut kappa = 1.0;
for t in 0..N_GIBBS {
let delta = self.sample_delta_given_lambda_kappa(
&sigma_n_inv,
&r_inv,
delta_obs,
&sigma_n_chol,
lambda,
kappa,
);
lambda = self.sample_lambda_given_delta(&delta);
kappa = self.sample_kappa_given_delta(delta_obs, &delta, &sigma_n_chol);
if t >= N_BURN {
retained_deltas.push(delta);
retained_lambdas.push(lambda);
retained_kappas.push(kappa);
}
}
self.compute_summaries(
&retained_deltas,
&retained_lambdas,
&retained_kappas,
sigma_n,
theta,
)
}
fn sample_delta_given_lambda_kappa(
&mut self,
sigma_n_inv: &Matrix9,
r_inv: &Matrix9,
delta_obs: &Vector9,
sigma_n_chol: &Cholesky<f64, nalgebra::Const<9>>,
lambda: f64,
kappa: f64,
) -> Vector9 {
let scale_factor = lambda / (self.sigma * self.sigma);
let q = sigma_n_inv * kappa + r_inv * scale_factor;
let q_chol = match Cholesky::new(q) {
Some(c) => c,
None => {
let jittered = q + Matrix9::identity() * 1e-8;
Cholesky::new(jittered).expect("Q(λ, κ) must be SPD")
}
};
let sigma_n_inv_delta = sigma_n_chol.solve(delta_obs);
let kappa_sigma_n_inv_delta = sigma_n_inv_delta * kappa;
let mu = q_chol.solve(&kappa_sigma_n_inv_delta);
let z = self.sample_standard_normal_vector();
let l_q_inv_t_z = q_chol.l().solve_upper_triangular(&z).unwrap_or(z);
mu + l_q_inv_t_z
}
fn sample_lambda_given_delta(&mut self, delta: &Vector9) -> f64 {
let y = self.l_r.solve_lower_triangular(delta).unwrap_or(*delta);
let q = y.dot(&y);
let shape = (self.nu + 9.0) / 2.0; let rate = (self.nu + q / (self.sigma * self.sigma)) / 2.0;
let scale = 1.0 / rate;
let gamma = Gamma::new(shape, scale).unwrap();
let sample = gamma.sample(&mut self.rng);
sample.clamp(LAMBDA_MIN, LAMBDA_MAX)
}
fn sample_kappa_given_delta(
&mut self,
delta_obs: &Vector9,
delta: &Vector9,
sigma_n_chol: &Cholesky<f64, nalgebra::Const<9>>,
) -> f64 {
let residual = delta_obs - delta;
let y = sigma_n_chol.solve(&residual);
let s = residual.dot(&y);
let shape = (NU_L + 9.0) / 2.0; let rate = (NU_L + s) / 2.0;
let scale = 1.0 / rate;
let gamma = Gamma::new(shape, scale).unwrap();
let sample = gamma.sample(&mut self.rng);
sample.clamp(KAPPA_MIN, KAPPA_MAX)
}
fn compute_summaries(
&self,
retained_deltas: &[Vector9],
retained_lambdas: &[f64],
retained_kappas: &[f64],
_sigma_n: &Matrix9,
theta: f64,
) -> GibbsResult {
let n = retained_deltas.len() as f64;
let delta_post = {
let mut sum = Vector9::zeros();
for delta in retained_deltas {
sum += delta;
}
sum / n
};
let lambda_post = {
let mut cov = Matrix9::zeros();
for delta in retained_deltas {
let diff = delta - delta_post;
cov += diff * diff.transpose();
}
cov / (n - 1.0) };
let mut exceed_count = 0;
let mut max_effects: Vec<f64> = Vec::with_capacity(retained_deltas.len());
for delta in retained_deltas {
let max_effect = delta.iter().map(|x| x.abs()).fold(0.0_f64, f64::max);
max_effects.push(max_effect);
if max_effect > theta {
exceed_count += 1;
}
}
let leak_probability = exceed_count as f64 / n;
max_effects.sort_by(|a, b| a.partial_cmp(b).unwrap());
let ci_low = max_effects[(n * 0.025) as usize];
let ci_high = max_effects[((n * 0.975) as usize).min(max_effects.len() - 1)];
let lambda_mean = retained_lambdas.iter().sum::<f64>() / n;
let lambda_var = retained_lambdas
.iter()
.map(|&l| math::sq(l - lambda_mean))
.sum::<f64>()
/ (n - 1.0);
let lambda_sd = math::sqrt(lambda_var);
let lambda_cv = if lambda_mean > 0.0 {
lambda_sd / lambda_mean
} else {
0.0
};
let lambda_ess = compute_ess(retained_lambdas);
let lambda_mixing_ok = lambda_cv >= 0.1 && lambda_ess >= 20.0;
let kappa_mean = retained_kappas.iter().sum::<f64>() / n;
let kappa_var = retained_kappas
.iter()
.map(|&k| math::sq(k - kappa_mean))
.sum::<f64>()
/ (n - 1.0);
let kappa_sd = math::sqrt(kappa_var);
let kappa_cv = if kappa_mean > 0.0 {
kappa_sd / kappa_mean
} else {
0.0
};
let kappa_ess = compute_ess(retained_kappas);
let kappa_mixing_ok = kappa_cv >= 0.1 && kappa_ess >= 20.0;
let delta_draws = retained_deltas.to_vec();
GibbsResult {
delta_post,
lambda_post,
leak_probability,
effect_magnitude_ci: (ci_low, ci_high),
delta_draws,
lambda_mean,
lambda_sd,
lambda_cv,
lambda_ess,
lambda_mixing_ok,
kappa_mean,
kappa_sd,
kappa_cv,
kappa_ess,
kappa_mixing_ok,
}
}
fn sample_standard_normal_vector(&mut self) -> Vector9 {
let mut z = Vector9::zeros();
for i in 0..9 {
z[i] = self.sample_standard_normal();
}
z
}
fn sample_standard_normal(&mut self) -> f64 {
let u1: f64 = self.rng.random();
let u2: f64 = self.rng.random();
math::sqrt(-2.0 * math::ln(u1.max(1e-12))) * math::cos(2.0 * PI * u2)
}
fn invert_via_cholesky(chol: &Cholesky<f64, nalgebra::Const<9>>) -> Matrix9 {
let mut inv = Matrix9::zeros();
for j in 0..9 {
let mut e = Vector9::zeros();
e[j] = 1.0;
let col = chol.solve(&e);
for i in 0..9 {
inv[(i, j)] = col[i];
}
}
inv
}
fn compute_r_inverse(&self) -> Matrix9 {
let mut r_inv = Matrix9::zeros();
for j in 0..9 {
let mut e = Vector9::zeros();
e[j] = 1.0;
let y = self.l_r.solve_lower_triangular(&e).unwrap_or(e);
let x = self.l_r.transpose().solve_upper_triangular(&y).unwrap_or(y);
for i in 0..9 {
r_inv[(i, j)] = x[i];
}
}
r_inv
}
}
fn estimate_condition_number(m: &Matrix9) -> f64 {
if let Some(chol) = Cholesky::new(*m) {
let l = chol.l();
let diag: [f64; 9] = core::array::from_fn(|i| l[(i, i)].abs());
let max_l = diag.iter().cloned().fold(0.0_f64, f64::max);
let min_l = diag.iter().cloned().fold(f64::INFINITY, f64::min);
if min_l < 1e-12 {
return f64::INFINITY;
}
let cond_l = max_l / min_l;
return cond_l * cond_l;
}
f64::INFINITY
}
fn regularize_sigma_n(sigma_n: &Matrix9) -> Matrix9 {
let cond = estimate_condition_number(sigma_n);
if cond <= CONDITION_NUMBER_THRESHOLD {
if Cholesky::new(*sigma_n).is_some() {
return *sigma_n;
}
}
let diag_sigma = Matrix9::from_diagonal(&sigma_n.diagonal());
if cond > CONDITION_NUMBER_THRESHOLD * 1e2 || cond.is_infinite() {
return diag_sigma + Matrix9::identity() * 1e-6;
}
let log_excess = (cond / CONDITION_NUMBER_THRESHOLD).ln().max(0.0);
let lambda = (0.1 + 0.2 * log_excess).min(0.95);
let regularized = *sigma_n * (1.0 - lambda) + diag_sigma * lambda;
for &eps in &[1e-10, 1e-9, 1e-8, 1e-7, 1e-6, 1e-5] {
let jittered = regularized + Matrix9::identity() * eps;
if Cholesky::new(jittered).is_some() {
return jittered;
}
}
diag_sigma + Matrix9::identity() * 1e-6
}
fn compute_ess(chain: &[f64]) -> f64 {
let n = chain.len();
if n < 2 {
return n as f64;
}
let mean: f64 = chain.iter().sum::<f64>() / n as f64;
let var: f64 = chain.iter().map(|&x| math::sq(x - mean)).sum::<f64>() / n as f64;
if var < 1e-12 {
return n as f64; }
let mut sum_rho = 0.0;
for k in 1..=50.min(n / 2) {
let rho_k = autocorrelation(chain, k, mean, var);
if rho_k < 0.05 {
break;
}
sum_rho += rho_k;
}
n as f64 / (1.0 + 2.0 * sum_rho)
}
fn autocorrelation(chain: &[f64], k: usize, mean: f64, var: f64) -> f64 {
let n = chain.len();
if k >= n {
return 0.0;
}
let cov: f64 = (0..(n - k))
.map(|i| (chain[i] - mean) * (chain[i + k] - mean))
.sum::<f64>()
/ (n - k) as f64;
cov / var
}
pub fn run_gibbs_inference(
delta: &Vector9,
sigma_n: &Matrix9,
sigma_t: f64,
l_r: &Matrix9,
theta: f64,
seed: u64,
) -> GibbsResult {
let mut sampler = GibbsSampler::new(l_r, sigma_t, seed);
sampler.run(delta, sigma_n, theta)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_gibbs_determinism() {
let l_r = Matrix9::identity();
let sigma_n = Matrix9::identity() * 100.0;
let delta = Vector9::from_row_slice(&[10.0; 9]);
let sigma_t = 50.0;
let theta = 5.0;
let result1 = run_gibbs_inference(&delta, &sigma_n, sigma_t, &l_r, theta, 42);
let result2 = run_gibbs_inference(&delta, &sigma_n, sigma_t, &l_r, theta, 42);
assert!(
(result1.leak_probability - result2.leak_probability).abs() < 1e-10,
"Same seed should give same result"
);
}
#[test]
fn test_lambda_diagnostics() {
let l_r = Matrix9::identity();
let sigma_n = Matrix9::identity() * 100.0;
let delta = Vector9::from_row_slice(&[50.0; 9]);
let sigma_t = 50.0;
let theta = 10.0;
let result = run_gibbs_inference(&delta, &sigma_n, sigma_t, &l_r, theta, 42);
assert!(result.lambda_mean > 0.0);
assert!(result.lambda_sd > 0.0);
assert!(result.lambda_ess >= 1.0);
assert!(result.lambda_ess <= N_KEEP as f64);
}
#[test]
fn test_large_effect_detection() {
let l_r = Matrix9::identity();
let sigma_n = Matrix9::identity() * 100.0;
let delta = Vector9::from_row_slice(&[500.0; 9]); let sigma_t = 50.0;
let theta = 10.0;
let result = run_gibbs_inference(&delta, &sigma_n, sigma_t, &l_r, theta, 42);
assert!(
result.leak_probability > 0.95,
"Large effect should give high leak probability, got {}",
result.leak_probability
);
}
#[test]
fn test_no_effect_low_probability() {
let l_r = Matrix9::identity();
let sigma_n = Matrix9::identity() * 100.0;
let delta = Vector9::zeros(); let sigma_t = 50.0;
let theta = 100.0;
let result = run_gibbs_inference(&delta, &sigma_n, sigma_t, &l_r, theta, 42);
assert!(
result.leak_probability < 0.5,
"No effect should give low leak probability, got {}",
result.leak_probability
);
}
#[test]
fn test_ess_computation() {
let chain: Vec<f64> = (0..100).map(|i| (i as f64).sin()).collect();
let ess = compute_ess(&chain);
assert!(ess < 100.0);
assert!(ess > 0.0);
}
#[test]
fn test_kappa_diagnostics() {
let l_r = Matrix9::identity();
let sigma_n = Matrix9::identity() * 100.0;
let delta = Vector9::from_row_slice(&[50.0; 9]);
let sigma_t = 50.0;
let theta = 10.0;
let result = run_gibbs_inference(&delta, &sigma_n, sigma_t, &l_r, theta, 42);
assert!(result.kappa_mean > 0.0, "kappa_mean should be positive");
assert!(result.kappa_sd > 0.0, "kappa_sd should be positive");
assert!(
result.kappa_ess >= 1.0,
"kappa_ess should be >= 1, got {}",
result.kappa_ess
);
assert!(
result.kappa_ess <= N_KEEP as f64,
"kappa_ess should be <= N_KEEP"
);
assert!(result.kappa_cv >= 0.0, "kappa_cv should be non-negative");
}
#[test]
fn test_kappa_responds_to_residual_magnitude() {
let l_r = Matrix9::identity();
let sigma_n = Matrix9::identity(); let sigma_t = 50.0;
let theta = 10.0;
let delta_small = Vector9::from_row_slice(&[1.0; 9]);
let result_small = run_gibbs_inference(&delta_small, &sigma_n, sigma_t, &l_r, theta, 42);
let delta_large = Vector9::from_row_slice(&[100.0; 9]);
let result_large = run_gibbs_inference(&delta_large, &sigma_n, sigma_t, &l_r, theta, 42);
assert!(
result_large.kappa_mean < result_small.kappa_mean,
"Large residuals should give smaller kappa: {} vs {}",
result_large.kappa_mean,
result_small.kappa_mean
);
}
}