use std::sync::Arc;
use antecedent_core::{PriorAssumption, VariableId};
use crate::error::ProbError;
const EFFECT_PRIOR_SD_FLOOR: f64 = 1e-12;
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub enum ContrastCoding {
Treatment,
Sum,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct EffectPrior {
pub mean: f64,
pub sd: f64,
}
impl EffectPrior {
pub fn new(mean: f64, sd: f64) -> Result<Self, ProbError> {
let p = Self { mean, sd };
p.validate()?;
Ok(p)
}
pub fn from_effect_draws(draws: &[f64]) -> Result<Self, ProbError> {
if draws.is_empty() {
return Err(ProbError::InvalidPrior { message: "from_effect_draws: empty draws" });
}
let n = draws.len() as f64;
let mean = draws.iter().sum::<f64>() / n;
if !mean.is_finite() {
return Err(ProbError::InvalidPrior { message: "from_effect_draws: non-finite mean" });
}
let sd = if draws.len() == 1 {
EFFECT_PRIOR_SD_FLOOR
} else {
let var = draws
.iter()
.map(|&x| {
let d = x - mean;
d * d
})
.sum::<f64>()
/ (n - 1.0);
var.sqrt().max(EFFECT_PRIOR_SD_FLOOR)
};
if !sd.is_finite() {
return Err(ProbError::InvalidPrior { message: "from_effect_draws: non-finite sd" });
}
Self::new(mean, sd)
}
pub fn validate(self) -> Result<(), ProbError> {
if !self.mean.is_finite() {
return Err(ProbError::InvalidPrior { message: "effect prior mean must be finite" });
}
if !(self.sd > 0.0) || !self.sd.is_finite() {
return Err(ProbError::InvalidPrior {
message: "effect prior sd must be finite and > 0",
});
}
Ok(())
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct GaussianCoefficientPrior {
pub mean: Arc<[f64]>,
pub variance: Arc<[f64]>,
}
impl GaussianCoefficientPrior {
#[must_use]
pub fn isotropic(n_coef: usize, scale: f64) -> Self {
let var = scale * scale;
Self { mean: Arc::from(vec![0.0; n_coef]), variance: Arc::from(vec![var; n_coef]) }
}
pub fn shared(n_coef: usize, mean: f64, variance: f64) -> Result<Self, ProbError> {
if n_coef == 0 {
return Err(ProbError::InvalidPrior { message: "n_coef must be > 0" });
}
if !(variance > 0.0) {
return Err(ProbError::InvalidPrior { message: "variance must be > 0" });
}
Ok(Self {
mean: Arc::from(vec![mean; n_coef]),
variance: Arc::from(vec![variance; n_coef]),
})
}
#[must_use]
pub fn len(&self) -> usize {
self.mean.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.mean.is_empty()
}
#[must_use]
pub fn precision(&self) -> Vec<f64> {
self.variance.iter().map(|&v| 1.0 / v).collect()
}
pub fn validate(&self) -> Result<(), ProbError> {
if self.mean.len() != self.variance.len() {
return Err(ProbError::InvalidPrior { message: "mean and variance length mismatch" });
}
if self.mean.is_empty() {
return Err(ProbError::InvalidPrior { message: "empty coefficient prior" });
}
for &v in self.variance.iter() {
if !(v > 0.0) && v.is_finite() {
return Err(ProbError::InvalidPrior { message: "variance must be > 0" });
}
if !v.is_finite() {
return Err(ProbError::InvalidPrior { message: "variance must be finite" });
}
}
Ok(())
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct InvGammaPrior {
pub shape: f64,
pub scale: f64,
}
impl InvGammaPrior {
#[must_use]
pub const fn weakly_informative() -> Self {
Self { shape: 1e-3, scale: 1e-3 }
}
pub fn validate(self) -> Result<(), ProbError> {
if !(self.shape > 0.0) || !(self.scale > 0.0) {
return Err(ProbError::InvalidPrior {
message: "InvGamma shape and scale must be > 0",
});
}
Ok(())
}
}
#[derive(Clone, Debug, PartialEq)]
pub enum PriorSpec {
GaussianCoefficients(GaussianCoefficientPrior),
ResidualInvGamma(InvGammaPrior),
KnownResidualVariance(f64),
}
impl PriorSpec {
#[must_use]
pub fn as_assumption(&self) -> PriorAssumption {
match self {
Self::GaussianCoefficients(_) => PriorAssumption {
id: Arc::from("gaussian_coefficients"),
description: Arc::from("Gaussian prior on regression coefficients"),
},
Self::ResidualInvGamma(_) => PriorAssumption {
id: Arc::from("residual_inv_gamma"),
description: Arc::from("Inverse-Gamma prior on residual variance"),
},
Self::KnownResidualVariance(_) => PriorAssumption {
id: Arc::from("known_residual_variance"),
description: Arc::from("Known residual variance (no prior uncertainty)"),
},
}
}
pub fn validate(&self) -> Result<(), ProbError> {
match self {
Self::GaussianCoefficients(p) => p.validate(),
Self::ResidualInvGamma(p) => p.validate(),
Self::KnownResidualVariance(v) => {
if !(*v > 0.0) || !v.is_finite() {
return Err(ProbError::InvalidPrior {
message: "known residual variance must be finite and > 0",
});
}
Ok(())
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum GaussianVarianceModel {
Known {
sigma2: f64,
},
InvGamma {
shape: f64,
scale: f64,
},
}
impl GaussianVarianceModel {
pub fn from_prior_set(prior: &PriorSet) -> Result<Self, ProbError> {
let mut known: Option<f64> = None;
let mut inv_gamma: Option<InvGammaPrior> = None;
for spec in &prior.specs {
match spec {
PriorSpec::KnownResidualVariance(v) => {
if known.is_some() || inv_gamma.is_some() {
return Err(ProbError::InvalidPrior {
message: "PriorSet must contain at most one residual variance specification",
});
}
known = Some(*v);
}
PriorSpec::ResidualInvGamma(p) => {
if known.is_some() || inv_gamma.is_some() {
return Err(ProbError::InvalidPrior {
message: "PriorSet must contain at most one residual variance specification",
});
}
inv_gamma = Some(*p);
}
PriorSpec::GaussianCoefficients(_) => {}
}
}
if let Some(sigma2) = known {
if !(sigma2 > 0.0) || !sigma2.is_finite() {
return Err(ProbError::InvalidPrior {
message: "known residual variance must be finite and > 0",
});
}
return Ok(Self::Known { sigma2 });
}
let ig = inv_gamma.unwrap_or_else(InvGammaPrior::weakly_informative);
ig.validate()?;
Ok(Self::InvGamma { shape: ig.shape, scale: ig.scale })
}
#[must_use]
pub const fn state_dim(self, ncols: usize) -> usize {
match self {
Self::Known { .. } => ncols,
Self::InvGamma { .. } => ncols.saturating_add(1),
}
}
#[must_use]
pub const fn include_sigma2(self) -> bool {
matches!(self, Self::InvGamma { .. })
}
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct PriorSet {
pub specs: Vec<PriorSpec>,
pub contrast: Option<ContrastCoding>,
pub categorical: Vec<VariableId>,
pub restrictions: Vec<PriorAssumption>,
}
impl PriorSet {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn weakly_informative(n_coef: usize) -> Self {
Self {
specs: vec![
PriorSpec::GaussianCoefficients(GaussianCoefficientPrior::isotropic(n_coef, 10.0)),
PriorSpec::ResidualInvGamma(InvGammaPrior::weakly_informative()),
],
contrast: None,
categorical: Vec::new(),
restrictions: Vec::new(),
}
}
pub fn push(&mut self, spec: PriorSpec) {
self.specs.push(spec);
}
pub fn validate_contrasts(&self) -> Result<(), ProbError> {
if !self.categorical.is_empty() && self.contrast.is_none() {
return Err(ProbError::InvalidPrior {
message: "categorical predictors require explicit contrast coding",
});
}
Ok(())
}
pub fn validate(&self) -> Result<(), ProbError> {
for s in &self.specs {
s.validate()?;
}
self.validate_contrasts()?;
let mut n_residual = 0usize;
for s in &self.specs {
match s {
PriorSpec::ResidualInvGamma(_) | PriorSpec::KnownResidualVariance(_) => {
n_residual = n_residual.saturating_add(1);
}
PriorSpec::GaussianCoefficients(_) => {}
}
}
if n_residual > 1 {
return Err(ProbError::InvalidPrior {
message: "PriorSet must contain at most one residual variance specification",
});
}
Ok(())
}
#[must_use]
pub fn gaussian_coefficients(&self) -> Option<&GaussianCoefficientPrior> {
self.specs.iter().find_map(|s| match s {
PriorSpec::GaussianCoefficients(p) => Some(p),
_ => None,
})
}
#[must_use]
pub fn residual_inv_gamma(&self) -> Option<InvGammaPrior> {
self.specs.iter().find_map(|s| match s {
PriorSpec::ResidualInvGamma(p) => Some(*p),
_ => None,
})
}
#[must_use]
pub fn known_residual_variance(&self) -> Option<f64> {
self.specs.iter().find_map(|s| match s {
PriorSpec::KnownResidualVariance(v) => Some(*v),
_ => None,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn weakly_informative_validates() {
let p = PriorSet::weakly_informative(3);
p.validate().unwrap();
assert_eq!(p.gaussian_coefficients().unwrap().len(), 3);
}
#[test]
fn categorical_requires_contrast() {
let mut p = PriorSet::weakly_informative(2);
p.categorical.push(VariableId::from_raw(0));
assert!(p.validate().is_err());
p.contrast = Some(ContrastCoding::Treatment);
p.validate().unwrap();
}
#[test]
fn effect_prior_from_draws_moments() {
let draws = [1.0, 3.0, 5.0];
let p = EffectPrior::from_effect_draws(&draws).unwrap();
assert!((p.mean - 3.0).abs() < 1e-12);
assert!((p.sd - 2.0).abs() < 1e-12);
}
#[test]
fn effect_prior_rejects_empty_and_nonfinite() {
assert!(EffectPrior::from_effect_draws(&[]).is_err());
assert!(EffectPrior::new(f64::NAN, 1.0).is_err());
assert!(EffectPrior::new(0.0, 0.0).is_err());
assert!(EffectPrior::new(0.0, -1.0).is_err());
}
#[test]
fn effect_prior_single_draw_floors_sd() {
let p = EffectPrior::from_effect_draws(&[2.5]).unwrap();
assert!((p.mean - 2.5).abs() < 1e-12);
assert!(p.sd > 0.0);
}
#[test]
fn residual_specs_must_be_unique() {
let mut p = PriorSet::new();
p.push(PriorSpec::GaussianCoefficients(GaussianCoefficientPrior::isotropic(1, 1.0)));
p.push(PriorSpec::KnownResidualVariance(1.0));
p.push(PriorSpec::ResidualInvGamma(InvGammaPrior::weakly_informative()));
assert!(p.validate().is_err());
assert!(GaussianVarianceModel::from_prior_set(&p).is_err());
}
#[test]
fn variance_model_defaults_to_weak_inv_gamma() {
let mut p = PriorSet::new();
p.push(PriorSpec::GaussianCoefficients(GaussianCoefficientPrior::isotropic(2, 1.0)));
p.validate().unwrap();
let model = GaussianVarianceModel::from_prior_set(&p).unwrap();
let weak = InvGammaPrior::weakly_informative();
assert_eq!(model, GaussianVarianceModel::InvGamma { shape: weak.shape, scale: weak.scale });
assert_eq!(model.state_dim(2), 3);
assert!(model.include_sigma2());
}
#[test]
fn variance_model_known() {
let mut p = PriorSet::new();
p.push(PriorSpec::KnownResidualVariance(2.5));
let model = GaussianVarianceModel::from_prior_set(&p).unwrap();
assert_eq!(model, GaussianVarianceModel::Known { sigma2: 2.5 });
assert_eq!(model.state_dim(4), 4);
assert!(!model.include_sigma2());
}
}