use std::sync::Arc;
use antecedent_core::ExecutionContext;
use antecedent_estimate::BayesianBackendKind;
use antecedent_estimate::{HydrateMapping, PreparedBayesianProblem, hydrate_prior};
use antecedent_io::PosteriorQuantityWire;
use antecedent_io::PriorMapping;
use antecedent_io::{decode_posterior_artifact, extract_prior_source_meta, read_and_migrate};
use antecedent_prob::{
BayesLikelihood, ComposedPrior, ConflictSummary, ExternalPriorSource, PosteriorQuantityKind,
PriorSet,
};
use antecedent_validate::{ConflictPolicy, PriorPredictiveCheck, compose_with_conflict_policy};
use crate::error::CausalError;
use crate::io::decode_causal_posterior_bytes;
#[derive(Clone, Debug, PartialEq)]
pub enum InferenceMode {
Frequentist,
Bayesian(BayesianConfig),
}
impl Default for InferenceMode {
fn default() -> Self {
Self::Frequentist
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct ExternalComposeSpec {
pub sources: Arc<[ExternalPriorSource]>,
pub composed: ComposedPrior,
pub conflict_policy: Option<ConflictPolicy>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct BayesianConfig {
pub backend: BayesianBackendKind,
pub likelihood: BayesLikelihood,
pub n_draws: usize,
pub prior_scale: f64,
pub prior: Option<PriorSet>,
pub prior_artifact: Option<Arc<[u8]>>,
pub prior_mapping: Option<PriorMapping>,
pub external_compose: Option<Box<ExternalComposeSpec>>,
}
impl BayesianConfig {
#[must_use]
pub fn laplace() -> Self {
Self {
backend: BayesianBackendKind::Laplace,
likelihood: BayesLikelihood::GaussianIdentity,
n_draws: 1000,
prior_scale: 10.0,
prior: None,
prior_artifact: None,
prior_mapping: None,
external_compose: None,
}
}
#[must_use]
pub fn conjugate() -> Self {
Self {
backend: BayesianBackendKind::ConjugateGaussian,
likelihood: BayesLikelihood::GaussianIdentity,
n_draws: 1000,
prior_scale: 10.0,
prior: None,
prior_artifact: None,
prior_mapping: None,
external_compose: None,
}
}
#[must_use]
pub fn hmc() -> Self {
Self {
backend: BayesianBackendKind::Hmc,
likelihood: BayesLikelihood::GaussianIdentity,
n_draws: 4_000,
prior_scale: 10.0,
prior: None,
prior_artifact: None,
prior_mapping: None,
external_compose: None,
}
}
#[must_use]
pub fn prior_scale(mut self, scale: f64) -> Self {
self.prior_scale = scale;
self
}
#[must_use]
pub fn n_draws(mut self, n: usize) -> Self {
self.n_draws = n;
self
}
#[must_use]
pub fn prior(mut self, prior: PriorSet) -> Self {
self.prior = Some(prior);
self.external_compose = None;
self
}
#[must_use]
pub fn prior_from_artifact(
mut self,
bytes: impl Into<Arc<[u8]>>,
mapping: Option<PriorMapping>,
) -> Self {
self.prior_artifact = Some(bytes.into());
self.prior_mapping = mapping;
self.prior = None;
self.external_compose = None;
self
}
#[must_use]
pub fn prior_from_composed(
mut self,
sources: impl Into<Arc<[ExternalPriorSource]>>,
composed: ComposedPrior,
conflict: Option<ConflictPolicy>,
) -> Self {
let sources = sources.into();
self.prior = Some(composed.prior.clone());
self.prior_artifact = None;
self.prior_mapping = None;
self.external_compose =
Some(Box::new(ExternalComposeSpec { sources, composed, conflict_policy: conflict }));
self
}
}
#[must_use]
pub fn hydrate_mapping_from_io(mapping: &PriorMapping) -> HydrateMapping {
match mapping {
PriorMapping::IdenticalCoefficientSubspace => HydrateMapping::IdenticalCoefficientSubspace,
PriorMapping::EffectFunctional { source_quantity } => {
HydrateMapping::EffectFunctional { source_quantity: source_quantity.clone() }
}
PriorMapping::NamedParameters { pairs } => {
HydrateMapping::NamedParameters { pairs: pairs.clone() }
}
}
}
fn wire_quantities_to_kinds(wire: &[PosteriorQuantityWire]) -> Vec<PosteriorQuantityKind> {
wire.iter()
.map(|q| match q {
PosteriorQuantityWire::Coefficient { index, name } => {
PosteriorQuantityKind::Coefficient {
index: *index as usize,
name: name.as_ref().map(|s| Arc::<str>::from(s.as_str())),
}
}
PosteriorQuantityWire::ResidualVariance => PosteriorQuantityKind::ResidualVariance,
PosteriorQuantityWire::Effect { name } => {
PosteriorQuantityKind::Effect { name: Arc::from(name.as_str()) }
}
PosteriorQuantityWire::Scalar { name } => {
PosteriorQuantityKind::Scalar { name: Arc::from(name.as_str()) }
}
})
.collect()
}
fn coef_names_for_problem(prep: &PreparedBayesianProblem) -> Arc<[Arc<str>]> {
if let Some(names) = &prep.coef_names {
return Arc::clone(names);
}
let names: Vec<Arc<str>> =
(0..prep.design.ncols).map(|i| Arc::<str>::from(format!("coef_{i}"))).collect();
Arc::from(names)
}
fn designs_compatible(
source_coef_names: &[Option<&str>],
target_ncols: usize,
target_names: &[Arc<str>],
) -> bool {
if source_coef_names.len() != target_ncols {
return false;
}
if source_coef_names.iter().all(Option::is_none) {
return true;
}
if source_coef_names.len() != target_names.len() {
return false;
}
source_coef_names
.iter()
.zip(target_names.iter())
.all(|(src, tgt)| src.is_some_and(|name| name == tgt.as_ref()))
}
fn default_hydrate_mapping(
bytes: &[u8],
prep: &PreparedBayesianProblem,
) -> Result<HydrateMapping, CausalError> {
let artifact = read_and_migrate(bytes)?;
if let Some(meta) = extract_prior_source_meta(&artifact)? {
if let Some(mapping) = &meta.declared_mapping {
return Ok(hydrate_mapping_from_io(mapping));
}
}
let (wire, _) = decode_posterior_artifact(&artifact)?;
let source_coef_names: Vec<Option<&str>> = wire
.quantities
.iter()
.filter_map(|q| match q {
PosteriorQuantityWire::Coefficient { name, .. } => Some(name.as_deref()),
_ => None,
})
.collect();
let target_names = coef_names_for_problem(prep);
if designs_compatible(&source_coef_names, prep.design.ncols, &target_names) {
return Ok(HydrateMapping::IdenticalCoefficientSubspace);
}
let source_quantity = wire.quantities.iter().find_map(|q| match q {
PosteriorQuantityWire::Effect { name } => Some(name.clone()),
_ => None,
});
match source_quantity {
Some(source_quantity) => Ok(HydrateMapping::EffectFunctional { source_quantity }),
None => Err(CausalError::Compile {
message: "prior artifact mapping required: designs differ and no Effect \
quantity is available for EffectFunctional default"
.into(),
}),
}
}
pub fn resolve_bayesian_prior(
cfg: &BayesianConfig,
prep: &PreparedBayesianProblem,
) -> Result<Option<PriorSet>, CausalError> {
let (prior, _) = resolve_bayesian_prior_with_conflict(cfg, prep, None)?;
Ok(prior)
}
pub fn resolve_bayesian_prior_with_conflict(
cfg: &BayesianConfig,
prep: &PreparedBayesianProblem,
ctx: Option<&ExecutionContext>,
) -> Result<(Option<PriorSet>, Option<ConflictSummary>), CausalError> {
if let Some(ext) = &cfg.external_compose {
if let (Some(policy), Some(ctx)) = (&ext.conflict_policy, ctx) {
let baseline = PriorSet::weakly_informative(prep.design.ncols);
let ppc = PriorPredictiveCheck {
n_sims: 200,
seed: ctx.rng.master_seed(),
..PriorPredictiveCheck::new()
};
let (composed, summary) =
compose_with_conflict_policy(&ext.sources, &baseline, policy, prep, ctx, &ppc)
.map_err(CausalError::from)?;
return Ok((Some(composed.prior), Some(summary)));
}
return Ok((Some(ext.composed.prior.clone()), None));
}
if let Some(p) = &cfg.prior {
return Ok((Some(p.clone()), None));
}
let Some(bytes) = cfg.prior_artifact.as_ref() else {
return Ok((None, None));
};
let mapping = match cfg.prior_mapping.as_ref() {
Some(m) => hydrate_mapping_from_io(m),
None => default_hydrate_mapping(bytes, prep)?,
};
let names = coef_names_for_problem(prep);
let baseline = PriorSet::weakly_informative(prep.design.ncols);
let treatment_col = prep.design.treatment_column();
Ok((
Some(hydrate_prior_from_posterior_bytes(
bytes,
&mapping,
&baseline,
&names,
treatment_col,
)?),
None,
))
}
pub fn hydrate_prior_from_posterior_bytes(
bytes: &[u8],
mapping: &HydrateMapping,
baseline: &PriorSet,
target_coef_names: &[Arc<str>],
treatment_col: Option<usize>,
) -> Result<PriorSet, CausalError> {
let (wire, _) = decode_causal_posterior_bytes(bytes)?;
let quantities = wire_quantities_to_kinds(&wire.quantities);
hydrate_prior(
mapping,
&quantities,
&wire.mean,
&wire.sd,
baseline,
target_coef_names,
treatment_col,
)
.map_err(CausalError::from)
}