Skip to main content

antecedent_estimate/
bayesian.rs

1//! Bayesian mechanisms, g-computation, and posterior functional evaluation .
2//!
3//! SPDX-License-Identifier: MIT OR Apache-2.0
4
5#![allow(
6    clippy::cast_precision_loss,
7    clippy::cast_possible_truncation,
8    clippy::needless_range_loop,
9    clippy::too_many_arguments,
10    clippy::too_many_lines,
11    clippy::needless_pass_by_value,
12    clippy::doc_markdown,
13    clippy::many_single_char_names
14)]
15
16use std::sync::Arc;
17
18use antecedent_core::IdentificationStatus;
19use antecedent_core::{
20    Assumption, AssumptionRecord, AssumptionScope, AssumptionSet, AssumptionSource,
21    AssumptionStatus, AverageEffectQuery, ExecutionContext, PriorAssumption, TargetPopulation,
22    VariableId,
23};
24use antecedent_data::{TableView, TabularData};
25use antecedent_expr::IdentifiedEstimand;
26use antecedent_prob::{
27    BayesDesignRef, BayesFitOptions, BayesLikelihood, ConflictSummary, ConjugateGaussianBackend,
28    EffectBatch, EffectPrior, GaussianCoefficientPrior, HmcGlmBackend, HmcOptions,
29    InferenceBackend, InferenceDiagnostics, LaplaceGlmBackend, LaplaceWorkspace, PosteriorBatch,
30    PosteriorDraws, PosteriorEvalWorkspace, PosteriorQuantityKind, PosteriorSchema,
31    PosteriorSummary, PriorSensitivitySummary, PriorSet, PriorSpec, sample_gaussian_mvn,
32};
33use antecedent_stats::{CompiledDesign, DesignColumnRole, GlmFamily};
34
35use crate::adjustment::{PreparedEstimationProblem, intervention_f64};
36use crate::error::EstimationError;
37use crate::overlap::OverlapPolicy;
38use crate::util::require_explicit_override;
39
40/// Causal posterior over an identified functional.
41#[derive(Clone, Debug)]
42pub struct CausalPosterior {
43    /// Columnar effect (and optional coefficient) draws.
44    pub draws: PosteriorDraws,
45    /// Summary of `draws`.
46    pub summaries: PosteriorSummary,
47    /// Identification status — priors never upgrade this.
48    pub identification: IdentificationStatus,
49    /// Optional prior-sensitivity grid.
50    pub prior_sensitivity: Option<PriorSensitivitySummary>,
51    /// Optional external-prior conflict shrink summary.
52    pub conflict_summary: Option<ConflictSummary>,
53    /// Inference diagnostics.
54    pub diagnostics: InferenceDiagnostics,
55    /// Assumptions including prior restrictions.
56    pub assumptions: AssumptionSet,
57    /// Unidentified graph mass retained when aggregating envelopes (0 if single graph).
58    pub unidentified_mass: f64,
59    /// Adaptive draw early-stop (Laplace / conjugate Gaussian redraw path).
60    pub early_stopped: bool,
61}
62
63impl CausalPosterior {
64    /// Primary effect column index (first `Effect` quantity), if any.
65    #[must_use]
66    pub fn effect_column(&self) -> Option<usize> {
67        self.draws
68            .schema
69            .quantities
70            .iter()
71            .position(|q| matches!(q, PosteriorQuantityKind::Effect { .. }))
72    }
73
74    /// Empirical P(effect < threshold) for the primary effect column.
75    ///
76    /// # Errors
77    ///
78    /// Missing effect column.
79    pub fn probability_below(&self, threshold: f64) -> Result<f64, EstimationError> {
80        let q = self
81            .effect_column()
82            .ok_or_else(|| EstimationError::stats_msg("CausalPosterior has no effect column"))?;
83        self.draws.probability_below(q, threshold).map_err(EstimationError::from)
84    }
85}
86
87/// Minimum coefficient prior variance when hydrating from a posterior (numerical floor).
88const HYDRATE_VAR_FLOOR: f64 = 1e-12;
89
90/// Build a Gaussian coefficient [`PriorSet`] from posterior quantity summaries.
91///
92/// Uses coefficient-column posterior means and SDs (index-aligned). Effect /
93/// residual columns are ignored. When `expected_n_coef` is `Some`, it must match
94/// the number of coefficient columns.
95///
96/// # Errors
97///
98/// No coefficient columns, non-finite summaries, non-contiguous indices, or
99/// dimension mismatch vs `expected_n_coef`.
100pub fn hydrate_prior_from_quantity_summaries(
101    quantities: &[PosteriorQuantityKind],
102    mean: &[f64],
103    sd: &[f64],
104    expected_n_coef: Option<usize>,
105) -> Result<PriorSet, EstimationError> {
106    if mean.len() != quantities.len() || sd.len() != quantities.len() {
107        return Err(EstimationError::stats_msg(
108            "hydrate_prior: mean/sd length must match quantities",
109        ));
110    }
111    let mut coef_cols: Vec<(usize, usize)> = quantities
112        .iter()
113        .enumerate()
114        .filter_map(|(col, q)| match q {
115            PosteriorQuantityKind::Coefficient { index, .. } => Some((*index, col)),
116            _ => None,
117        })
118        .collect();
119    coef_cols.sort_by_key(|(index, _)| *index);
120    let n_coef = coef_cols.len();
121    if n_coef == 0 {
122        return Err(EstimationError::stats_msg(
123            "hydrate_prior_from_posterior: no coefficient columns in posterior",
124        ));
125    }
126    if let Some(expected) = expected_n_coef {
127        if n_coef != expected {
128            return Err(EstimationError::stats_msg(format!(
129                "posterior coefficient dimension {n_coef} != expected n_coef {expected}"
130            )));
131        }
132    }
133    for (i, (index, _)) in coef_cols.iter().enumerate() {
134        if *index != i {
135            return Err(EstimationError::stats_msg(format!(
136                "posterior coefficient indices are not contiguous (expected {i}, got {index})"
137            )));
138        }
139    }
140    let mut means = Vec::with_capacity(n_coef);
141    let mut variance = Vec::with_capacity(n_coef);
142    for (_, col) in &coef_cols {
143        let m = mean[*col];
144        let s = sd[*col];
145        if !m.is_finite() || !s.is_finite() {
146            return Err(EstimationError::stats_msg(
147                "posterior coefficient summary is non-finite; cannot hydrate prior",
148            ));
149        }
150        means.push(m);
151        variance.push((s * s).max(HYDRATE_VAR_FLOOR));
152    }
153    let coef = GaussianCoefficientPrior { mean: Arc::from(means), variance: Arc::from(variance) };
154    coef.validate().map_err(EstimationError::from)?;
155    Ok(PriorSet {
156        specs: vec![PriorSpec::GaussianCoefficients(coef)],
157        contrast: None,
158        categorical: Vec::new(),
159        restrictions: Vec::new(),
160    })
161}
162
163/// Build a Gaussian coefficient [`PriorSet`] from a fitted posterior (sequential Bayes).
164///
165/// # Errors
166///
167/// See [`hydrate_prior_from_quantity_summaries`].
168pub fn hydrate_prior_from_posterior(
169    posterior: &CausalPosterior,
170    expected_n_coef: Option<usize>,
171) -> Result<PriorSet, EstimationError> {
172    hydrate_prior_from_quantity_summaries(
173        &posterior.draws.schema.quantities,
174        &posterior.summaries.mean,
175        &posterior.summaries.sd,
176        expected_n_coef,
177    )
178}
179
180/// Bridge from a banked posterior into a target design's coefficient prior.
181///
182/// Mirrors [`antecedent_io::PriorMapping`] without depending on `antecedent-io` (avoids a
183/// cycle). Convert at the facade.
184#[derive(Clone, Debug, PartialEq, Eq, Hash)]
185pub enum HydrateMapping {
186    /// Identical coefficient subspace (P1-C sequential Bayes).
187    IdenticalCoefficientSubspace,
188    /// Effect-functional transfer via a named source quantity (e.g. `"ate"`).
189    EffectFunctional {
190        /// Source effect / quantity name.
191        source_quantity: String,
192    },
193    /// Explicit source→target quantity name pairs.
194    NamedParameters {
195        /// `(source_name, target_name)` pairs.
196        pairs: Vec<(String, String)>,
197    },
198}
199
200/// Build a coefficient [`PriorSet`] under a declared [`HydrateMapping`].
201///
202/// - [`HydrateMapping::IdenticalCoefficientSubspace`]: full coef hydrate; hard-errors
203///   when source coef count ≠ baseline length.
204/// - [`HydrateMapping::EffectFunctional`]: maps source effect moments onto the
205///   treatment coefficient (identity-link ATE bridge); other dims keep `baseline`.
206/// - [`HydrateMapping::NamedParameters`]: maps named source moments onto named
207///   target coefficients; unmapped dims keep `baseline`.
208///
209/// Records `external_effect_prior` / `external_named_prior` on
210/// [`PriorSet::restrictions`].
211///
212/// # Errors
213///
214/// Dimension mismatch, missing effect column, unknown names, or invalid baseline.
215pub fn hydrate_prior(
216    mapping: &HydrateMapping,
217    quantities: &[PosteriorQuantityKind],
218    mean: &[f64],
219    sd: &[f64],
220    baseline: &PriorSet,
221    target_coef_names: &[Arc<str>],
222    treatment_col: Option<usize>,
223) -> Result<PriorSet, EstimationError> {
224    if mean.len() != quantities.len() || sd.len() != quantities.len() {
225        return Err(EstimationError::stats_msg(
226            "hydrate_prior: mean/sd length must match quantities",
227        ));
228    }
229    let n_target = target_coef_names.len();
230    let base_coef = baseline.gaussian_coefficients().ok_or_else(|| {
231        EstimationError::stats_msg("hydrate_prior: baseline missing GaussianCoefficients")
232    })?;
233    if base_coef.len() != n_target {
234        return Err(EstimationError::stats_msg(format!(
235            "hydrate_prior: baseline n_coef {} != target_coef_names {}",
236            base_coef.len(),
237            n_target
238        )));
239    }
240
241    match mapping {
242        HydrateMapping::IdenticalCoefficientSubspace => {
243            let mut prior =
244                hydrate_prior_from_quantity_summaries(quantities, mean, sd, Some(n_target))?;
245            // Preserve residual specs from baseline when present.
246            merge_baseline_residuals(&mut prior, baseline);
247            Ok(prior)
248        }
249        HydrateMapping::EffectFunctional { source_quantity } => {
250            let t_col = treatment_col.ok_or_else(|| {
251                EstimationError::stats_msg("hydrate_prior: EffectFunctional requires treatment_col")
252            })?;
253            if t_col >= n_target {
254                return Err(EstimationError::stats_msg(format!(
255                    "hydrate_prior: treatment_col {t_col} out of range for {n_target} coefs"
256                )));
257            }
258            let (m, s) = quantity_moments(quantities, mean, sd, source_quantity.as_str())?;
259            let effect = EffectPrior::new(m, s.max(HYDRATE_VAR_FLOOR.sqrt()))
260                .map_err(EstimationError::from)?;
261            let mut means: Vec<f64> = base_coef.mean.to_vec();
262            let mut vars: Vec<f64> = base_coef.variance.to_vec();
263            means[t_col] = effect.mean;
264            vars[t_col] = (effect.sd * effect.sd).max(HYDRATE_VAR_FLOOR);
265            let coef =
266                GaussianCoefficientPrior { mean: Arc::from(means), variance: Arc::from(vars) };
267            coef.validate().map_err(EstimationError::from)?;
268            let mut prior = PriorSet {
269                specs: vec![PriorSpec::GaussianCoefficients(coef)],
270                contrast: baseline.contrast,
271                categorical: baseline.categorical.clone(),
272                restrictions: vec![PriorAssumption {
273                    id: Arc::from("external_effect_prior"),
274                    description: Arc::from(format!(
275                        "external effect-functional prior from quantity `{source_quantity}` onto treatment coefficient"
276                    )),
277                }],
278            };
279            merge_baseline_residuals(&mut prior, baseline);
280            Ok(prior)
281        }
282        HydrateMapping::NamedParameters { pairs } => {
283            if pairs.is_empty() {
284                return Err(EstimationError::stats_msg(
285                    "hydrate_prior: NamedParameters requires at least one pair",
286                ));
287            }
288            let mut means: Vec<f64> = base_coef.mean.to_vec();
289            let mut vars: Vec<f64> = base_coef.variance.to_vec();
290            let name_index: std::collections::HashMap<&str, usize> =
291                target_coef_names.iter().enumerate().map(|(i, n)| (n.as_ref(), i)).collect();
292            for (src, tgt) in pairs {
293                let (m, s) = quantity_moments(quantities, mean, sd, src)?;
294                let Some(&idx) = name_index.get(tgt.as_str()) else {
295                    return Err(EstimationError::stats_msg(format!(
296                        "hydrate_prior: unknown target coefficient name `{tgt}`"
297                    )));
298                };
299                means[idx] = m;
300                vars[idx] = (s * s).max(HYDRATE_VAR_FLOOR);
301            }
302            let coef =
303                GaussianCoefficientPrior { mean: Arc::from(means), variance: Arc::from(vars) };
304            coef.validate().map_err(EstimationError::from)?;
305            let pair_desc =
306                pairs.iter().map(|(a, b)| format!("{a}->{b}")).collect::<Vec<_>>().join(", ");
307            let mut prior = PriorSet {
308                specs: vec![PriorSpec::GaussianCoefficients(coef)],
309                contrast: baseline.contrast,
310                categorical: baseline.categorical.clone(),
311                restrictions: vec![PriorAssumption {
312                    id: Arc::from("external_named_prior"),
313                    description: Arc::from(format!("external named-parameter prior ({pair_desc})")),
314                }],
315            };
316            merge_baseline_residuals(&mut prior, baseline);
317            Ok(prior)
318        }
319    }
320}
321
322fn merge_baseline_residuals(prior: &mut PriorSet, baseline: &PriorSet) {
323    for spec in &baseline.specs {
324        match spec {
325            PriorSpec::ResidualInvGamma(_) | PriorSpec::KnownResidualVariance(_) => {
326                if !prior.specs.iter().any(|s| {
327                    matches!(
328                        s,
329                        PriorSpec::ResidualInvGamma(_) | PriorSpec::KnownResidualVariance(_)
330                    )
331                }) {
332                    prior.specs.push(spec.clone());
333                }
334            }
335            PriorSpec::GaussianCoefficients(_) => {}
336        }
337    }
338}
339
340fn quantity_moments(
341    quantities: &[PosteriorQuantityKind],
342    mean: &[f64],
343    sd: &[f64],
344    name: &str,
345) -> Result<(f64, f64), EstimationError> {
346    for (i, q) in quantities.iter().enumerate() {
347        let q_name = match q {
348            PosteriorQuantityKind::Effect { name: n }
349            | PosteriorQuantityKind::Scalar { name: n } => Some(n.as_ref()),
350            PosteriorQuantityKind::Coefficient { name: n, .. } => {
351                n.as_ref().map(std::convert::AsRef::as_ref)
352            }
353            PosteriorQuantityKind::ResidualVariance => Some("residual_variance"),
354        };
355        if q_name == Some(name) {
356            let m = mean[i];
357            let s = sd[i];
358            if !m.is_finite() || !s.is_finite() {
359                return Err(EstimationError::stats_msg(format!(
360                    "hydrate_prior: non-finite summary for quantity `{name}`"
361                )));
362            }
363            return Ok((m, s.max(HYDRATE_VAR_FLOOR.sqrt())));
364        }
365    }
366    Err(EstimationError::stats_msg(format!("hydrate_prior: missing quantity `{name}`")))
367}
368
369/// Bayesian linear / GLM mechanism fit (coefficient posterior).
370#[derive(Clone, Debug)]
371pub struct BayesianGlmMechanism {
372    /// Fitted coefficient draws (columnar).
373    pub coefficient_draws: PosteriorDraws,
374    /// MAP / posterior mode coefficients.
375    pub map: Vec<f64>,
376    /// Likelihood used.
377    pub likelihood: BayesLikelihood,
378    /// Diagnostics.
379    pub diagnostics: InferenceDiagnostics,
380    /// Compiled design retained for g-computation.
381    pub design: CompiledDesign,
382    /// Treatment column index in the design.
383    pub treatment_col: usize,
384    /// Active / control levels.
385    pub active: f64,
386    /// Control level.
387    pub control: f64,
388}
389
390/// Which inference backend to use for Bayesian g-computation.
391#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
392pub enum BayesianBackendKind {
393    /// Analytic conjugate Gaussian (identity link only).
394    ConjugateGaussian,
395    /// Native Laplace GLM.
396    Laplace,
397    /// Native HMC GLM (multi-chain; ESS / R-hat gated).
398    Hmc,
399}
400
401/// Bayesian g-computation ATE estimator.
402#[derive(Clone, Debug)]
403pub struct BayesianGComputationAte {
404    /// Backend kind.
405    pub backend: BayesianBackendKind,
406    /// Likelihood (Laplace); conjugate forces GaussianIdentity.
407    pub likelihood: BayesLikelihood,
408    /// Draw count.
409    pub n_draws: usize,
410    /// RNG seed.
411    pub seed: u64,
412    /// Overlap policy (must be ExplicitOverride).
413    pub overlap: OverlapPolicy,
414    /// Prior scale for isotropic Gaussian coefficients (weakly informative default 10).
415    pub prior_scale: f64,
416    /// Optional explicit coefficient prior (e.g. hydrated from a previous posterior).
417    /// When set, overrides isotropic [`Self::prior_scale`].
418    pub prior: Option<PriorSet>,
419}
420
421impl Default for BayesianGComputationAte {
422    fn default() -> Self {
423        Self::new()
424    }
425}
426
427impl BayesianGComputationAte {
428    /// Laplace Gaussian defaults.
429    #[must_use]
430    pub fn new() -> Self {
431        Self {
432            backend: BayesianBackendKind::Laplace,
433            likelihood: BayesLikelihood::GaussianIdentity,
434            n_draws: 1000,
435            seed: 0,
436            overlap: OverlapPolicy::ExplicitOverride,
437            prior_scale: 10.0,
438            prior: None,
439        }
440    }
441
442    /// Conjugate Gaussian linear path.
443    #[must_use]
444    pub fn conjugate() -> Self {
445        Self {
446            backend: BayesianBackendKind::ConjugateGaussian,
447            likelihood: BayesLikelihood::GaussianIdentity,
448            ..Self::new()
449        }
450    }
451
452    /// Prepare from data + identified estimand (same IR as frequentist adjustment).
453    ///
454    /// # Errors
455    ///
456    /// Overlap / estimand / data failures.
457    pub fn prepare(
458        &self,
459        data: &TabularData,
460        estimand: &IdentifiedEstimand,
461        query: &AverageEffectQuery,
462    ) -> Result<PreparedBayesianProblem, EstimationError> {
463        require_explicit_override(
464            self.overlap,
465            "BayesianGComputationAte requires ExplicitOverride overlap policy",
466        )?;
467        if !matches!(
468            estimand.method_kind().ok(),
469            Some(
470                antecedent_expr::EstimandMethod::BackdoorAdjustment
471                    | antecedent_expr::EstimandMethod::BackdoorEfficient
472            )
473        ) {
474            return Err(EstimationError::IncompatibleEstimand {
475                message: "BayesianGComputationAte expects backdoor.adjustment/efficient",
476            });
477        }
478        query.validate()?;
479        if !query.effect_modifiers.is_empty() {
480            return Err(EstimationError::unsupported(
481                "Bayesian g-comp does not support effect modifiers",
482            ));
483        }
484        if query.target_population != TargetPopulation::AllObserved {
485            return Err(EstimationError::unsupported(
486                "Bayesian g-comp only supports TargetPopulation::AllObserved",
487            ));
488        }
489        let active = intervention_f64(&query.active)?;
490        let control = intervention_f64(&query.control)?;
491        if (active - control).abs() < f64::EPSILON {
492            return Err(EstimationError::unsupported(
493                "active and control treatment levels must differ",
494            ));
495        }
496
497        let treatment = query.treatment;
498        let outcome = query.outcome;
499        let mut ids = Vec::with_capacity(2 + estimand.adjustment_set.len());
500        ids.push(treatment);
501        ids.push(outcome);
502        ids.extend_from_slice(&estimand.adjustment_set);
503        let row_mask = data.complete_case_mask(&ids).map_err(EstimationError::from)?;
504        let t = data.float64_masked(treatment, &row_mask).map_err(EstimationError::from)?;
505        let y = data.float64_masked(outcome, &row_mask).map_err(EstimationError::from)?;
506        let mut covs: Vec<(VariableId, Vec<f64>)> = Vec::new();
507        for &z in estimand.adjustment_set.iter() {
508            covs.push((z, data.float64_masked(z, &row_mask).map_err(EstimationError::from)?));
509        }
510        let cov_refs: Vec<(VariableId, &[f64])> =
511            covs.iter().map(|(id, v)| (*id, v.as_slice())).collect();
512        let selected_rows: Vec<usize> =
513            row_mask.iter().enumerate().filter_map(|(i, keep)| keep.then_some(i)).collect();
514        let design = CompiledDesign::linear_adjustment(&t, &cov_refs, &y, &selected_rows)
515            .map_err(EstimationError::from)?;
516        let schema = data.schema();
517        let treatment_name = schema.get(treatment).map(|v| v.name.as_ref()).unwrap_or("treatment");
518        let coef_names = coefficient_names_from_design(&design, treatment_name, |id| {
519            schema.get(id).ok().map(|v| Arc::clone(&v.name))
520        });
521        Ok(PreparedBayesianProblem {
522            design,
523            method: Arc::clone(&estimand.method),
524            adjustment_set: Arc::clone(&estimand.adjustment_set),
525            active,
526            control,
527            overlap: self.overlap,
528            coef_names: Some(coef_names),
529        })
530    }
531
532    /// Adapt a frequentist prepared design (e.g. lag-aligned temporal) for Bayesian fit.
533    ///
534    /// Used by the temporal pulse/sustained path: prepare via
535    /// [`crate::TemporalLinearAdjustment`], then fit with this estimator.
536    #[must_use]
537    pub fn from_prepared_estimation(prep: &PreparedEstimationProblem) -> PreparedBayesianProblem {
538        PreparedBayesianProblem {
539            design: prep.design.clone(),
540            method: Arc::clone(&prep.method),
541            adjustment_set: Arc::clone(&prep.adjustment_set),
542            active: prep.active,
543            control: prep.control,
544            overlap: prep.overlap,
545            coef_names: None,
546        }
547    }
548
549    /// Fit mechanism + evaluate ATE g-computation posterior.
550    ///
551    /// `identification` is recorded as-is; informative priors never change it.
552    ///
553    /// # Errors
554    ///
555    /// Backend / evaluation failures.
556    pub fn fit(
557        &self,
558        problem: &PreparedBayesianProblem,
559        identification: IdentificationStatus,
560        workspace: &mut BayesianGCompWorkspace,
561        ctx: &ExecutionContext,
562    ) -> Result<CausalPosterior, EstimationError> {
563        let sequential = self.prior.is_some();
564        let prior = if let Some(p) = &self.prior {
565            if let Some(coef) = p.gaussian_coefficients() {
566                if coef.len() != problem.design.ncols {
567                    return Err(EstimationError::stats_msg(format!(
568                        "sequential prior coefficient dimension {} != design ncols {}",
569                        coef.len(),
570                        problem.design.ncols
571                    )));
572                }
573            } else {
574                return Err(EstimationError::stats_msg(
575                    "sequential prior missing GaussianCoefficients entry",
576                ));
577            }
578            p.clone()
579        } else {
580            PriorSet {
581                specs: vec![PriorSpec::GaussianCoefficients(
582                    antecedent_prob::GaussianCoefficientPrior::isotropic(
583                        problem.design.ncols,
584                        self.prior_scale,
585                    ),
586                )],
587                contrast: None,
588                categorical: Vec::new(),
589                restrictions: Vec::new(),
590            }
591        };
592        let mut assumptions = AssumptionSet::new();
593        let source = if sequential {
594            AssumptionSource::Artifact
595        } else {
596            AssumptionSource::AlgorithmDefault { algorithm: Arc::from("bayesian_gcomp") }
597        };
598        for spec in &prior.specs {
599            let mut pa = spec.as_assumption();
600            if sequential {
601                pa.description = Arc::from(format!(
602                    "{} (sequential prior from posterior artifact)",
603                    pa.description
604                ));
605            }
606            assumptions.push(AssumptionRecord {
607                assumption: Assumption::PriorRestriction(pa),
608                source: source.clone(),
609                scope: AssumptionScope::Estimation,
610                status: AssumptionStatus::Untestable,
611            });
612        }
613        for pa in &prior.restrictions {
614            assumptions.push(AssumptionRecord {
615                assumption: Assumption::PriorRestriction(pa.clone()),
616                source: AssumptionSource::Artifact,
617                scope: AssumptionScope::Estimation,
618                status: AssumptionStatus::Untestable,
619            });
620        }
621
622        let likelihood = match self.backend {
623            BayesianBackendKind::ConjugateGaussian => BayesLikelihood::GaussianIdentity,
624            BayesianBackendKind::Laplace | BayesianBackendKind::Hmc => self.likelihood,
625        };
626        let max_draws = self.n_draws.max(1);
627        let adaptive = ctx.adaptive_draws;
628        let laplace_adaptive = adaptive.enabled
629            && matches!(self.backend, BayesianBackendKind::Laplace)
630            && max_draws > adaptive.min_draws.max(2);
631        let initial_draws =
632            if laplace_adaptive { adaptive.min_draws.max(2).min(max_draws) } else { max_draws };
633        let opts = BayesFitOptions {
634            n_draws: initial_draws,
635            seed: self.seed,
636            ..BayesFitOptions::default()
637        };
638        let design_ref = BayesDesignRef {
639            x_colmajor: &problem.design.matrix,
640            nrows: problem.design.nrows,
641            ncols: problem.design.ncols,
642            y: &problem.design.outcome,
643            weights: None,
644            offsets: None,
645        };
646
647        let mut fit = match self.backend {
648            BayesianBackendKind::ConjugateGaussian => ConjugateGaussianBackend.fit(
649                likelihood,
650                design_ref,
651                &prior,
652                &opts,
653                &mut workspace.laplace,
654                ctx,
655            ),
656            BayesianBackendKind::Laplace => LaplaceGlmBackend.fit(
657                likelihood,
658                design_ref,
659                &prior,
660                &opts,
661                &mut workspace.laplace,
662                ctx,
663            ),
664            BayesianBackendKind::Hmc => HmcGlmBackend::new()
665                .with_options(HmcOptions {
666                    n_chains: 2,
667                    n_warmup: (self.n_draws / 2).max(50),
668                    ..HmcOptions::default()
669                })
670                .fit(likelihood, design_ref, &prior, &opts, &mut workspace.laplace, ctx),
671        }
672        .map_err(prob_err)?;
673
674        if !fit.diagnostics.allows_posterior() {
675            return Err(EstimationError::stats_msg("Bayesian fit refused without diagnostics"));
676        }
677
678        let t_col = problem
679            .design
680            .treatment_column()
681            .ok_or_else(|| EstimationError::stats_msg("missing treatment column"))?;
682
683        let glm_family = likelihood_to_glm_family(likelihood);
684        let evaluator = GCompAteEvaluator {
685            family: glm_family,
686            treatment_col: t_col,
687            active: problem.active,
688            control: problem.control,
689            nrows: problem.design.nrows,
690            ncols: problem.design.ncols,
691            matrix: Arc::clone(&problem.design.matrix),
692        };
693        let compiled = evaluator.compile()?;
694
695        let mut early_stopped = false;
696        let mut n_draws = fit.draws.n_draws;
697        // Adaptive MVN sampling uses the β-block covariance only; drop residual-variance
698        // columns so batch merges match `PosteriorSchema::coefficients`.
699        let mut coef_draws = coefficient_only_draws(&fit.draws)?;
700
701        if laplace_adaptive {
702            let cov = fit.cov.as_ref().ok_or_else(|| {
703                EstimationError::stats_msg("Laplace adaptive draws require posterior covariance")
704            })?;
705            let map = fit.map.clone();
706            let batch = 32usize;
707            let mut effect_acc: Vec<f64> = Vec::with_capacity(max_draws);
708            let mut width_prev: Option<f64> = None;
709
710            // Evaluate initial block.
711            {
712                workspace.eval.prepare(n_draws, problem.design.ncols);
713                let mut effect_out = EffectBatch::default();
714                effect_out.prepare(n_draws);
715                let batch_view = coef_draws.batch(0, n_draws).map_err(EstimationError::from)?;
716                evaluator.evaluate_batch(
717                    &compiled,
718                    batch_view,
719                    &mut effect_out,
720                    &mut workspace.eval,
721                    ctx,
722                )?;
723                effect_acc.extend_from_slice(&effect_out.values[..n_draws]);
724            }
725
726            loop {
727                let width = quantile_width_95(&effect_acc);
728                let ess = effect_acc.len() as f64; // independent MVN draws
729                if effect_acc.len() >= adaptive.min_draws.max(2) {
730                    let width_ok = width_prev.is_some_and(|prev| {
731                        let rel = (width - prev).abs() / prev.abs().max(1e-12);
732                        rel < adaptive.quantile_width_rel_epsilon
733                    });
734                    if width_ok || ess >= adaptive.ess_target {
735                        early_stopped = n_draws < max_draws;
736                        break;
737                    }
738                }
739                width_prev = Some(width);
740                if n_draws >= max_draws {
741                    break;
742                }
743                let next = (n_draws + batch).min(max_draws);
744                let add = next - n_draws;
745                let extra = sample_gaussian_mvn(
746                    &map,
747                    cov,
748                    add,
749                    self.seed.wrapping_add(n_draws as u64),
750                    &mut workspace.laplace,
751                )
752                .map_err(EstimationError::from)?;
753                let extra_draws = PosteriorDraws::from_column_major(
754                    PosteriorSchema::coefficients(problem.design.ncols),
755                    add,
756                    extra,
757                )
758                .map_err(EstimationError::from)?;
759                workspace.eval.prepare(add, problem.design.ncols);
760                let mut effect_out = EffectBatch::default();
761                effect_out.prepare(add);
762                let batch_view = extra_draws.batch(0, add).map_err(EstimationError::from)?;
763                evaluator.evaluate_batch(
764                    &compiled,
765                    batch_view,
766                    &mut effect_out,
767                    &mut workspace.eval,
768                    ctx,
769                )?;
770                effect_acc.extend_from_slice(&effect_out.values[..add]);
771                coef_draws = merge_coefficient_draws(&coef_draws, &extra_draws)?;
772                n_draws = next;
773                fit.draws = coef_draws.clone();
774            }
775
776            // Rebuild combined posterior from accumulated effects + final coef draws.
777            let mechanism_draws = coef_draws;
778            let mut quantities = mechanism_draws.schema.quantities.to_vec();
779            quantities.retain(|q| !matches!(q, PosteriorQuantityKind::ResidualVariance));
780            let effect_idx = quantities.len();
781            quantities.push(PosteriorQuantityKind::Effect { name: Arc::from("ate") });
782            let n_q = quantities.len();
783            let mut values = vec![0.0; n_draws * n_q];
784            for (qi, q) in mechanism_draws.schema.quantities.iter().enumerate() {
785                if matches!(q, PosteriorQuantityKind::ResidualVariance) {
786                    continue;
787                }
788                let dest = quantities.iter().position(|qq| qq == q).ok_or_else(|| {
789                    EstimationError::stats_msg(format!(
790                        "posterior quantity missing from schema: {q:?}"
791                    ))
792                })?;
793                let coef_col = mechanism_draws.column(qi).map_err(EstimationError::from)?;
794                values[dest * n_draws..(dest + 1) * n_draws].copy_from_slice(coef_col);
795            }
796            values[effect_idx * n_draws..(effect_idx + 1) * n_draws]
797                .copy_from_slice(&effect_acc[..n_draws]);
798            if let Some(names) = problem.coef_names.as_ref() {
799                apply_coefficient_names(&mut quantities, names);
800            }
801            let draws = PosteriorDraws::from_column_major(
802                PosteriorSchema { quantities: Arc::from(quantities) },
803                n_draws,
804                values,
805            )
806            .map_err(EstimationError::from)?;
807            let summaries = draws.summarize();
808            return Ok(CausalPosterior {
809                draws,
810                summaries,
811                identification,
812                prior_sensitivity: None,
813                conflict_summary: None,
814                diagnostics: fit.diagnostics,
815                assumptions,
816                unidentified_mass: 0.0,
817                early_stopped,
818            });
819        }
820
821        let mechanism = BayesianGlmMechanism {
822            coefficient_draws: coef_draws,
823            map: fit.map,
824            likelihood,
825            diagnostics: fit.diagnostics.clone(),
826            design: problem.design.clone(),
827            treatment_col: t_col,
828            active: problem.active,
829            control: problem.control,
830        };
831
832        workspace.eval.prepare(n_draws, problem.design.ncols);
833        let mut effect_out = EffectBatch::default();
834        effect_out.prepare(n_draws);
835        let batch = mechanism.coefficient_draws.batch(0, n_draws).map_err(EstimationError::from)?;
836        evaluator.evaluate_batch(&compiled, batch, &mut effect_out, &mut workspace.eval, ctx)?;
837
838        let mut quantities = mechanism.coefficient_draws.schema.quantities.to_vec();
839        // Drop residual variance column from combined effect artifact if present — keep coefs + effect.
840        quantities.retain(|q| !matches!(q, PosteriorQuantityKind::ResidualVariance));
841        let effect_idx = quantities.len();
842        quantities.push(PosteriorQuantityKind::Effect { name: Arc::from("ate") });
843        let n_q = quantities.len();
844        let mut values = vec![0.0; n_draws * n_q];
845        for (qi, q) in mechanism.coefficient_draws.schema.quantities.iter().enumerate() {
846            if matches!(q, PosteriorQuantityKind::ResidualVariance) {
847                continue;
848            }
849            let dest = quantities.iter().position(|qq| qq == q).ok_or_else(|| {
850                EstimationError::stats_msg(format!("posterior quantity missing from schema: {q:?}"))
851            })?;
852            let col = mechanism.coefficient_draws.column(qi).map_err(EstimationError::from)?;
853            values[dest * n_draws..(dest + 1) * n_draws].copy_from_slice(col);
854        }
855        values[effect_idx * n_draws..(effect_idx + 1) * n_draws]
856            .copy_from_slice(&effect_out.values[..n_draws]);
857
858        if let Some(names) = problem.coef_names.as_ref() {
859            apply_coefficient_names(&mut quantities, names);
860        }
861
862        let draws = PosteriorDraws::from_column_major(
863            PosteriorSchema { quantities: Arc::from(quantities) },
864            n_draws,
865            values,
866        )
867        .map_err(EstimationError::from)?;
868        let summaries = draws.summarize();
869
870        let _ = mechanism;
871        Ok(CausalPosterior {
872            draws,
873            summaries,
874            identification,
875            prior_sensitivity: None,
876            conflict_summary: None,
877            diagnostics: fit.diagnostics,
878            assumptions,
879            unidentified_mass: 0.0,
880            early_stopped: false,
881        })
882    }
883}
884
885/// Bayesian g-computation on a lag-aligned temporal design.
886///
887/// Prepare with [`crate::TemporalLinearAdjustment::prepare`], convert via
888/// [`BayesianGComputationAte::from_prepared_estimation`], then [`BayesianGComputationAte::fit`].
889/// This type documents the temporal entry point; fitting delegates to [`BayesianGComputationAte`].
890#[derive(Clone, Debug, Default)]
891pub struct BayesianTemporalGcomp {
892    /// Shared Bayesian estimator configuration.
893    pub inner: BayesianGComputationAte,
894}
895
896impl BayesianTemporalGcomp {
897    /// Laplace Gaussian defaults.
898    #[must_use]
899    pub fn new() -> Self {
900        Self { inner: BayesianGComputationAte::new() }
901    }
902
903    /// Conjugate Gaussian linear path.
904    #[must_use]
905    pub fn conjugate() -> Self {
906        Self { inner: BayesianGComputationAte::conjugate() }
907    }
908
909    /// Convert a temporal prepared design for Bayesian fit.
910    #[must_use]
911    pub fn from_prepared_estimation(prep: &PreparedEstimationProblem) -> PreparedBayesianProblem {
912        BayesianGComputationAte::from_prepared_estimation(prep)
913    }
914
915    /// Fit on a prepared Bayesian problem (typically from a temporal design).
916    ///
917    /// # Errors
918    ///
919    /// Backend / evaluation failures.
920    pub fn fit(
921        &self,
922        problem: &PreparedBayesianProblem,
923        identification: IdentificationStatus,
924        workspace: &mut BayesianGCompWorkspace,
925        ctx: &ExecutionContext,
926    ) -> Result<CausalPosterior, EstimationError> {
927        self.inner.fit(problem, identification, workspace, ctx)
928    }
929}
930
931/// Durable coefficient names from a design + schema name resolver.
932///
933/// Convention: `intercept`, `coef_{treatment}`, `coef_{covariate}`.
934#[must_use]
935pub fn coefficient_names_from_design(
936    design: &CompiledDesign,
937    treatment_name: &str,
938    covariate_name: impl Fn(VariableId) -> Option<Arc<str>>,
939) -> Arc<[Arc<str>]> {
940    let names: Vec<Arc<str>> = design
941        .columns
942        .iter()
943        .map(|col| match col.role {
944            DesignColumnRole::Intercept => Arc::from("intercept"),
945            DesignColumnRole::Treatment => Arc::from(format!("coef_{treatment_name}")),
946            DesignColumnRole::Covariate(id) => covariate_name(id).map_or_else(
947                || Arc::from(format!("coef_var_{}", id.raw())),
948                |n| Arc::from(format!("coef_{n}")),
949            ),
950        })
951        .collect();
952    Arc::from(names)
953}
954
955/// Apply durable names onto coefficient quantities (in place).
956fn apply_coefficient_names(quantities: &mut [PosteriorQuantityKind], names: &[Arc<str>]) {
957    for q in quantities {
958        if let PosteriorQuantityKind::Coefficient { index, name } = q {
959            if let Some(n) = names.get(*index) {
960                *name = Some(Arc::clone(n));
961            }
962        }
963    }
964}
965
966/// Prepared Bayesian g-comp problem.
967#[derive(Clone, Debug)]
968pub struct PreparedBayesianProblem {
969    /// Design.
970    pub design: CompiledDesign,
971    /// Estimand method.
972    pub method: Arc<str>,
973    /// Adjustment set.
974    pub adjustment_set: Arc<[VariableId]>,
975    /// Active treatment.
976    pub active: f64,
977    /// Control treatment.
978    pub control: f64,
979    /// Overlap.
980    pub overlap: OverlapPolicy,
981    /// Optional durable coefficient names aligned to design columns.
982    pub coef_names: Option<Arc<[Arc<str>]>>,
983}
984
985/// Workspace for Bayesian g-comp.
986#[derive(Clone, Debug, Default)]
987pub struct BayesianGCompWorkspace {
988    /// Laplace / conjugate workspace.
989    pub laplace: LaplaceWorkspace,
990    /// Posterior functional eval scratch.
991    pub eval: PosteriorEvalWorkspace,
992}
993
994/// Trait for batched posterior functional evaluation.
995pub trait PosteriorFunctionalEvaluator {
996    /// Compiled plan type.
997    type Compiled;
998
999    /// Compile against a posterior schema.
1000    ///
1001    /// # Errors
1002    ///
1003    /// Incompatible schema.
1004    fn compile(&self) -> Result<Self::Compiled, EstimationError>;
1005
1006    /// Evaluate a batch of coefficient draws into effects.
1007    ///
1008    /// # Errors
1009    ///
1010    /// Shape / numerical failures.
1011    fn evaluate_batch(
1012        &self,
1013        compiled: &Self::Compiled,
1014        posterior: PosteriorBatch<'_>,
1015        output: &mut EffectBatch,
1016        workspace: &mut PosteriorEvalWorkspace,
1017        ctx: &ExecutionContext,
1018    ) -> Result<(), EstimationError>;
1019}
1020
1021/// Compiled g-comp ATE evaluator (finite-difference mean contrast).
1022#[derive(Clone, Debug)]
1023pub struct GCompAteEvaluator {
1024    /// Mean family.
1025    pub family: GlmFamily,
1026    /// Treatment column.
1027    pub treatment_col: usize,
1028    /// Active level.
1029    pub active: f64,
1030    /// Control level.
1031    pub control: f64,
1032    /// Rows.
1033    pub nrows: usize,
1034    /// Cols.
1035    pub ncols: usize,
1036    /// Design matrix (column-major).
1037    pub matrix: Arc<[f64]>,
1038}
1039
1040/// Empty compiled marker (evaluator is self-contained).
1041#[derive(Clone, Copy, Debug, Default)]
1042pub struct CompiledGCompAte;
1043
1044impl PosteriorFunctionalEvaluator for GCompAteEvaluator {
1045    type Compiled = CompiledGCompAte;
1046
1047    fn compile(&self) -> Result<Self::Compiled, EstimationError> {
1048        if self.treatment_col >= self.ncols {
1049            return Err(EstimationError::stats_msg("treatment column out of range"));
1050        }
1051        Ok(CompiledGCompAte)
1052    }
1053
1054    fn evaluate_batch(
1055        &self,
1056        _compiled: &Self::Compiled,
1057        posterior: PosteriorBatch<'_>,
1058        output: &mut EffectBatch,
1059        workspace: &mut PosteriorEvalWorkspace,
1060        _ctx: &ExecutionContext,
1061    ) -> Result<(), EstimationError> {
1062        let n_draws = posterior.len;
1063        workspace.prepare(n_draws, self.ncols);
1064        output.prepare(n_draws);
1065
1066        // Coefficient columns 0..ncols from the batch (ignore extra quantities).
1067        let mut coef_cols: Vec<&[f64]> = Vec::with_capacity(self.ncols);
1068        for c in 0..self.ncols {
1069            let col = posterior.column(c).map_err(EstimationError::from)?;
1070            coef_cols.push(col);
1071        }
1072
1073        for d in 0..n_draws {
1074            for c in 0..self.ncols {
1075                workspace.row[c] = coef_cols[c][d];
1076            }
1077            let beta = &workspace.row[..self.ncols];
1078            let mut sum = 0.0;
1079            for r in 0..self.nrows {
1080                let mu_a = predict_row(
1081                    self.family,
1082                    &self.matrix,
1083                    self.nrows,
1084                    self.ncols,
1085                    self.treatment_col,
1086                    beta,
1087                    r,
1088                    self.active,
1089                );
1090                let mu_c = predict_row(
1091                    self.family,
1092                    &self.matrix,
1093                    self.nrows,
1094                    self.ncols,
1095                    self.treatment_col,
1096                    beta,
1097                    r,
1098                    self.control,
1099                );
1100                sum += mu_a - mu_c;
1101            }
1102            output.values[d] = sum / self.nrows as f64;
1103        }
1104        Ok(())
1105    }
1106}
1107
1108fn predict_row(
1109    family: GlmFamily,
1110    matrix: &[f64],
1111    nrows: usize,
1112    ncols: usize,
1113    t_col: usize,
1114    beta: &[f64],
1115    row: usize,
1116    t_value: f64,
1117) -> f64 {
1118    let mut eta = 0.0;
1119    for c in 0..ncols {
1120        let x = if c == t_col { t_value } else { matrix[c * nrows + row] };
1121        eta += x * beta[c];
1122    }
1123    family.mean_from_eta(eta)
1124}
1125
1126fn likelihood_to_glm_family(l: BayesLikelihood) -> GlmFamily {
1127    match l {
1128        BayesLikelihood::GaussianIdentity => GlmFamily::GaussianIdentity,
1129        BayesLikelihood::BernoulliLogit => GlmFamily::BinomialLogit,
1130        BayesLikelihood::BernoulliProbit => GlmFamily::BinomialProbit,
1131        BayesLikelihood::PoissonLog => GlmFamily::PoissonLog,
1132    }
1133}
1134
1135fn prob_err(e: antecedent_prob::ProbError) -> EstimationError {
1136    EstimationError::from(e)
1137}
1138
1139/// 95% quantile width of a scalar draw vector.
1140fn quantile_width_95(values: &[f64]) -> f64 {
1141    if values.len() < 2 {
1142        return f64::NAN;
1143    }
1144    // Reuse posterior summarization for consistent quantiles.
1145    let schema = PosteriorSchema {
1146        quantities: Arc::from([PosteriorQuantityKind::Effect { name: Arc::from("w") }]),
1147    };
1148    let Ok(draws) = PosteriorDraws::from_column_major(schema, values.len(), values.to_vec()) else {
1149        return f64::NAN;
1150    };
1151    let s = draws.summarize();
1152    s.q975[0] - s.q025[0]
1153}
1154
1155/// Concatenate two coefficient-only posterior draw tables (same schema).
1156fn merge_coefficient_draws(
1157    a: &PosteriorDraws,
1158    b: &PosteriorDraws,
1159) -> Result<PosteriorDraws, EstimationError> {
1160    if a.schema != b.schema {
1161        return Err(EstimationError::stats_msg("merge_coefficient_draws: schema mismatch"));
1162    }
1163    let n_q = a.schema.quantities.len();
1164    let n = a.n_draws + b.n_draws;
1165    let mut values = vec![0.0; n * n_q];
1166    for q in 0..n_q {
1167        let col_a = a.column(q).map_err(EstimationError::from)?;
1168        let col_b = b.column(q).map_err(EstimationError::from)?;
1169        values[q * n..q * n + a.n_draws].copy_from_slice(col_a);
1170        values[q * n + a.n_draws..(q + 1) * n].copy_from_slice(col_b);
1171    }
1172    PosteriorDraws::from_column_major(a.schema.clone(), n, values).map_err(EstimationError::from)
1173}
1174
1175/// Keep coefficient columns only (drop residual-variance / other non-β quantities).
1176fn coefficient_only_draws(draws: &PosteriorDraws) -> Result<PosteriorDraws, EstimationError> {
1177    let coef_idx: Vec<usize> = draws
1178        .schema
1179        .quantities
1180        .iter()
1181        .enumerate()
1182        .filter_map(|(i, q)| matches!(q, PosteriorQuantityKind::Coefficient { .. }).then_some(i))
1183        .collect();
1184    if coef_idx.is_empty() {
1185        return Err(EstimationError::stats_msg(
1186            "coefficient_only_draws: no coefficient quantities",
1187        ));
1188    }
1189    if coef_idx.len() == draws.schema.quantities.len() {
1190        return Ok(draws.clone());
1191    }
1192    let n = draws.n_draws;
1193    let n_q = coef_idx.len();
1194    let mut quantities = Vec::with_capacity(n_q);
1195    let mut values = vec![0.0; n * n_q];
1196    for (dest, &src) in coef_idx.iter().enumerate() {
1197        quantities.push(draws.schema.quantities[src].clone());
1198        let col = draws.column(src).map_err(EstimationError::from)?;
1199        values[dest * n..(dest + 1) * n].copy_from_slice(col);
1200    }
1201    PosteriorDraws::from_column_major(
1202        PosteriorSchema { quantities: Arc::from(quantities) },
1203        n,
1204        values,
1205    )
1206    .map_err(EstimationError::from)
1207}
1208
1209/// Build a non-identified posterior artifact that still records priors (exit criterion #2).
1210///
1211/// Samples prior-predictive draws for a scalar effect mean (isotropic Gaussian / weakly
1212/// informative scale from `prior`) so Bayesian envelopes can surface uncertainty without
1213/// inventing identification. Status remains [`IdentificationStatus::NotIdentified`].
1214#[must_use]
1215pub fn nonidentified_with_prior(
1216    prior: &PriorSet,
1217    diagnostics: InferenceDiagnostics,
1218    n_draws: usize,
1219    seed: u64,
1220) -> CausalPosterior {
1221    let mut assumptions = AssumptionSet::new();
1222    for spec in &prior.specs {
1223        assumptions.push(AssumptionRecord {
1224            assumption: Assumption::PriorRestriction(spec.as_assumption()),
1225            source: AssumptionSource::UserDeclared,
1226            scope: AssumptionScope::Estimation,
1227            status: AssumptionStatus::Untestable,
1228        });
1229    }
1230    let schema = PosteriorSchema {
1231        quantities: Arc::from([PosteriorQuantityKind::Effect { name: Arc::from("ate") }]),
1232    };
1233    let (mean, scale) = prior_predictive_effect_params(prior);
1234    let n = n_draws.max(1);
1235    let mut values = vec![0.0; n];
1236    let mut rng = ExecutionContext::for_tests(seed).rng.stream(0xBA7E_u64);
1237    for v in &mut values {
1238        *v = mean + scale * antecedent_kernels::standard_normal(&mut rng);
1239    }
1240    let draws = PosteriorDraws::from_column_major(schema, n, Arc::<[f64]>::from(values))
1241        .unwrap_or_else(|_| PosteriorDraws {
1242            schema: PosteriorSchema {
1243                quantities: Arc::from([PosteriorQuantityKind::Effect { name: Arc::from("ate") }]),
1244            },
1245            n_draws: 0,
1246            values: Arc::from([]),
1247        });
1248    let summaries = draws.summarize();
1249    CausalPosterior {
1250        draws,
1251        summaries,
1252        identification: IdentificationStatus::NotIdentified,
1253        prior_sensitivity: None,
1254        conflict_summary: None,
1255        diagnostics,
1256        assumptions,
1257        unidentified_mass: 1.0,
1258        early_stopped: false,
1259    }
1260}
1261
1262fn prior_predictive_effect_params(prior: &PriorSet) -> (f64, f64) {
1263    if let Some(g) = prior.gaussian_coefficients() {
1264        let mean = g.mean.first().copied().unwrap_or(0.0);
1265        let var = g.variance.first().copied().unwrap_or(100.0).max(1e-12);
1266        return (mean, var.sqrt());
1267    }
1268    (0.0, 10.0)
1269}
1270
1271#[cfg(test)]
1272mod tests {
1273    use super::*;
1274    use antecedent_core::{
1275        CausalSchemaBuilder, MeasurementSpec, RoleHint, SmallRoleSet, ValueType, VariableId,
1276    };
1277    use antecedent_data::column::{Float64Column, ValidityBitmap};
1278    use antecedent_data::{OwnedColumn, OwnedColumnarStorage, TabularData};
1279    use antecedent_expr::{ExprId, IdentifiedEstimand};
1280    use antecedent_prob::InferenceDiagnostics;
1281
1282    fn linear_scm_table(n: usize) -> (TabularData, VariableId, VariableId, VariableId) {
1283        let mut b = CausalSchemaBuilder::new();
1284        b.add_variable(
1285            "Z",
1286            ValueType::Continuous,
1287            SmallRoleSet::from_hint(RoleHint::Context),
1288            None,
1289            None,
1290            MeasurementSpec::default(),
1291        )
1292        .unwrap();
1293        b.add_variable(
1294            "T",
1295            ValueType::Continuous,
1296            SmallRoleSet::from_hint(RoleHint::TreatmentCandidate),
1297            None,
1298            None,
1299            MeasurementSpec::default(),
1300        )
1301        .unwrap();
1302        b.add_variable(
1303            "Y",
1304            ValueType::Continuous,
1305            SmallRoleSet::from_hint(RoleHint::OutcomeCandidate),
1306            None,
1307            None,
1308            MeasurementSpec::default(),
1309        )
1310        .unwrap();
1311        let schema = b.build().unwrap();
1312        let z = VariableId::from_raw(0);
1313        let t = VariableId::from_raw(1);
1314        let y = VariableId::from_raw(2);
1315        let mut zv = vec![0.0; n];
1316        let mut tv = vec![0.0; n];
1317        let mut yv = vec![0.0; n];
1318        for i in 0..n {
1319            zv[i] = (i as f64) * 0.1;
1320            tv[i] = if i % 2 == 0 { 1.0 } else { 0.0 };
1321            yv[i] = 2.0 * tv[i] + 0.5 * zv[i];
1322        }
1323        let validity = ValidityBitmap::all_valid(n);
1324        let cols = vec![
1325            OwnedColumn::Float64(Float64Column::new(z, Arc::from(zv), validity.clone()).unwrap()),
1326            OwnedColumn::Float64(Float64Column::new(t, Arc::from(tv), validity.clone()).unwrap()),
1327            OwnedColumn::Float64(Float64Column::new(y, Arc::from(yv), validity).unwrap()),
1328        ];
1329        let storage = OwnedColumnarStorage::try_new(schema, cols, None, None).unwrap();
1330        (TabularData::new(storage), t, y, z)
1331    }
1332
1333    #[test]
1334    fn bayesian_and_frequentist_share_ate() {
1335        let n = 80;
1336        let (data, t, y, z) = linear_scm_table(n);
1337        let estimand = IdentifiedEstimand::backdoor(
1338            "backdoor.adjustment",
1339            Arc::from(vec![z]),
1340            ExprId::from_raw(0),
1341        );
1342        let query = AverageEffectQuery::binary_ate(t, y);
1343
1344        let freq = crate::adjustment::LinearAdjustmentAte {
1345            bootstrap_replicates: 0,
1346            ..crate::adjustment::LinearAdjustmentAte::new()
1347        };
1348        let prep = freq.prepare(&data, &estimand, &query).unwrap();
1349        let mut ws = crate::adjustment::EstimationWorkspace::default();
1350        let freq_est = freq
1351            .fit(&prep, &mut ws, &ExecutionContext::for_tests(1), AssumptionSet::new())
1352            .unwrap();
1353
1354        let bayes = BayesianGComputationAte {
1355            backend: BayesianBackendKind::ConjugateGaussian,
1356            n_draws: 400,
1357            seed: 5,
1358            prior_scale: 100.0,
1359            ..BayesianGComputationAte::new()
1360        };
1361        let bprep = bayes.prepare(&data, &estimand, &query).unwrap();
1362        let mut bws = BayesianGCompWorkspace::default();
1363        let post = bayes
1364            .fit(
1365                &bprep,
1366                IdentificationStatus::NonparametricallyIdentified,
1367                &mut bws,
1368                &ExecutionContext::for_tests(1),
1369            )
1370            .unwrap();
1371        let eq = post.effect_column().unwrap();
1372        let mean = post.summaries.mean[eq];
1373        assert!((freq_est.ate - 2.0).abs() < 1e-6, "frequentist ate={}", freq_est.ate);
1374        assert!((mean - freq_est.ate).abs() < 0.05, "bayes={mean} freq={}", freq_est.ate);
1375        assert_eq!(post.identification, IdentificationStatus::NonparametricallyIdentified);
1376        let coef_names: Vec<_> = post
1377            .draws
1378            .schema
1379            .quantities
1380            .iter()
1381            .filter_map(|q| match q {
1382                PosteriorQuantityKind::Coefficient { name, .. } => name.as_ref().map(AsRef::as_ref),
1383                _ => None,
1384            })
1385            .collect();
1386        assert!(coef_names.contains(&"intercept"), "{coef_names:?}");
1387        assert!(coef_names.iter().any(|n| n.starts_with("coef_")), "{coef_names:?}");
1388    }
1389
1390    #[test]
1391    fn prior_does_not_create_identification() {
1392        let prior = PriorSet::weakly_informative(3);
1393        let post = nonidentified_with_prior(&prior, InferenceDiagnostics::analytic("none"), 64, 1);
1394        assert_eq!(post.identification, IdentificationStatus::NotIdentified);
1395        assert!(!post.assumptions.is_empty());
1396        assert!((post.unidentified_mass - 1.0).abs() < 1e-12);
1397        assert!(post.draws.n_draws > 0, "prior-predictive draws required");
1398    }
1399
1400    #[test]
1401    fn temporal_prepared_design_conjugate_recovers_pulse() {
1402        use antecedent_core::{
1403            CausalSchemaBuilder, Lag, MeasurementSpec, RoleHint, SmallRoleSet, TemporalEffectQuery,
1404            TemporalPolicy, ValueType,
1405        };
1406        use antecedent_data::{
1407            Float64Column, OwnedColumn, OwnedColumnarStorage, SamplingRegularity, TimeIndex,
1408            TimeSeriesData, ValidityBitmap,
1409        };
1410        use antecedent_graph::{TemporalDag, ensure_lagged};
1411        use antecedent_identify::TemporalBackdoorIdentifier;
1412
1413        use crate::temporal_adjustment::TemporalLinearAdjustment;
1414
1415        let n = 300usize;
1416        let mut b = CausalSchemaBuilder::new();
1417        b.add_variable(
1418            "x",
1419            ValueType::Continuous,
1420            SmallRoleSet::from_hint(RoleHint::TreatmentCandidate),
1421            None,
1422            None,
1423            MeasurementSpec::default(),
1424        )
1425        .unwrap();
1426        b.add_variable(
1427            "y",
1428            ValueType::Continuous,
1429            SmallRoleSet::from_hint(RoleHint::OutcomeCandidate),
1430            None,
1431            None,
1432            MeasurementSpec::default(),
1433        )
1434        .unwrap();
1435        let schema = b.build().unwrap();
1436        let mut x = vec![0.0; n];
1437        let mut y = vec![0.0; n];
1438        for t in 1..n {
1439            x[t] = ((t as f64) * 0.07).sin();
1440            y[t] = 0.8 * x[t - 1];
1441        }
1442        let cols = vec![
1443            OwnedColumn::Float64(
1444                Float64Column::new(
1445                    VariableId::from_raw(0),
1446                    Arc::from(x),
1447                    ValidityBitmap::all_valid(n),
1448                )
1449                .unwrap(),
1450            ),
1451            OwnedColumn::Float64(
1452                Float64Column::new(
1453                    VariableId::from_raw(1),
1454                    Arc::from(y),
1455                    ValidityBitmap::all_valid(n),
1456                )
1457                .unwrap(),
1458            ),
1459        ];
1460        let storage = OwnedColumnarStorage::try_new(schema, cols, None, None).unwrap();
1461        let data = TimeSeriesData::try_new(
1462            storage,
1463            TimeIndex { regularity: SamplingRegularity::Regular { interval_ns: 1 }, length: n },
1464        )
1465        .unwrap();
1466        let mut g = TemporalDag::empty();
1467        let x1 = ensure_lagged(&mut g, VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
1468        let y0 = ensure_lagged(&mut g, VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
1469        g.insert_directed(x1, y0).unwrap();
1470
1471        let q = TemporalEffectQuery::pulse(VariableId::from_raw(0), VariableId::from_raw(1), 1.0)
1472            .with_policy(TemporalPolicy::pulse(-1))
1473            .with_horizon_steps(1)
1474            .with_max_history_lag(Some(1));
1475        let id_res = TemporalBackdoorIdentifier::new().identify_temporal(&g, &q).unwrap();
1476        let estimand = id_res.result.estimands.first().unwrap();
1477        let temporal = TemporalLinearAdjustment::new();
1478        let prep = temporal
1479            .prepare(
1480                &data,
1481                estimand,
1482                &q,
1483                &id_res.indexer,
1484                None,
1485                &ExecutionContext::for_tests(1).kernel_policy,
1486            )
1487            .unwrap();
1488        let bayes = BayesianTemporalGcomp {
1489            inner: BayesianGComputationAte {
1490                backend: BayesianBackendKind::ConjugateGaussian,
1491                n_draws: 200,
1492                seed: 7,
1493                prior_scale: 100.0,
1494                ..BayesianGComputationAte::new()
1495            },
1496        };
1497        let bprep = BayesianTemporalGcomp::from_prepared_estimation(&prep);
1498        let mut ws = BayesianGCompWorkspace::default();
1499        let post = bayes
1500            .fit(
1501                &bprep,
1502                IdentificationStatus::NonparametricallyIdentified,
1503                &mut ws,
1504                &ExecutionContext::for_tests(1),
1505            )
1506            .unwrap();
1507        let eq = post.effect_column().unwrap();
1508        let mean = post.summaries.mean[eq];
1509        assert!((mean - 0.8).abs() < 0.05, "bayesian temporal pulse mean={mean}");
1510        assert!(post.probability_below(0.0).unwrap().is_finite());
1511    }
1512
1513    #[test]
1514    fn hydrate_prior_from_posterior_and_refit() {
1515        let n = 60;
1516        let (data, t, y, z) = linear_scm_table(n);
1517        let estimand = IdentifiedEstimand::backdoor(
1518            "backdoor.adjustment",
1519            Arc::from(vec![z]),
1520            ExprId::from_raw(0),
1521        );
1522        let query = AverageEffectQuery::binary_ate(t, y);
1523        let bayes = BayesianGComputationAte {
1524            backend: BayesianBackendKind::ConjugateGaussian,
1525            n_draws: 200,
1526            seed: 3,
1527            prior_scale: 10.0,
1528            ..BayesianGComputationAte::new()
1529        };
1530        let prep = bayes.prepare(&data, &estimand, &query).unwrap();
1531        let mut ws = BayesianGCompWorkspace::default();
1532        let post = bayes
1533            .fit(
1534                &prep,
1535                IdentificationStatus::NonparametricallyIdentified,
1536                &mut ws,
1537                &ExecutionContext::for_tests(1),
1538            )
1539            .unwrap();
1540        let prior = hydrate_prior_from_posterior(&post, Some(prep.design.ncols)).unwrap();
1541        assert_eq!(prior.gaussian_coefficients().unwrap().len(), prep.design.ncols);
1542        assert!(hydrate_prior_from_posterior(&post, Some(prep.design.ncols + 1)).is_err());
1543
1544        let sequential = BayesianGComputationAte { prior: Some(prior), ..bayes };
1545        let post2 = sequential
1546            .fit(
1547                &prep,
1548                IdentificationStatus::NonparametricallyIdentified,
1549                &mut ws,
1550                &ExecutionContext::for_tests(1),
1551            )
1552            .unwrap();
1553        assert!(post2.assumptions.entries.iter().any(|a| {
1554            matches!(a.source, AssumptionSource::Artifact)
1555                && matches!(&a.assumption, Assumption::PriorRestriction(pa) if pa.description.contains("sequential"))
1556        }));
1557        let eq = post2.effect_column().unwrap();
1558        assert!(post2.summaries.mean[eq].is_finite());
1559    }
1560
1561    #[test]
1562    fn hydrate_effect_functional_maps_treatment_coef() {
1563        let quantities = vec![
1564            PosteriorQuantityKind::Coefficient { index: 0, name: Some(Arc::from("intercept")) },
1565            PosteriorQuantityKind::Coefficient { index: 1, name: Some(Arc::from("coef_t")) },
1566            PosteriorQuantityKind::Effect { name: Arc::from("ate") },
1567        ];
1568        let mean = vec![0.1, 0.5, 2.0];
1569        let sd = vec![1.0, 1.0, 0.4];
1570        let names: Vec<Arc<str>> =
1571            vec![Arc::from("intercept"), Arc::from("coef_t"), Arc::from("coef_z")];
1572        let baseline = PriorSet::weakly_informative(3);
1573        let prior = hydrate_prior(
1574            &HydrateMapping::EffectFunctional { source_quantity: "ate".into() },
1575            &quantities,
1576            &mean,
1577            &sd,
1578            &baseline,
1579            &names,
1580            Some(1),
1581        )
1582        .unwrap();
1583        let coef = prior.gaussian_coefficients().unwrap();
1584        assert!((coef.mean[1] - 2.0).abs() < 1e-12);
1585        assert!((coef.variance[1] - 0.16).abs() < 1e-12);
1586        // Unmapped dims keep baseline (isotropic scale 10 → var 100).
1587        assert!((coef.mean[0] - 0.0).abs() < 1e-12);
1588        assert!((coef.variance[0] - 100.0).abs() < 1e-12);
1589        assert!((coef.variance[2] - 100.0).abs() < 1e-12);
1590        assert!(prior.restrictions.iter().any(|r| r.id.as_ref() == "external_effect_prior"));
1591    }
1592
1593    #[test]
1594    fn hydrate_mapping_hard_errors() {
1595        let quantities = vec![
1596            PosteriorQuantityKind::Coefficient { index: 0, name: Some(Arc::from("intercept")) },
1597            PosteriorQuantityKind::Coefficient { index: 1, name: Some(Arc::from("coef_t")) },
1598            PosteriorQuantityKind::Effect { name: Arc::from("ate") },
1599        ];
1600        let mean = vec![0.0, 1.0, 2.0];
1601        let sd = vec![1.0, 1.0, 0.5];
1602        let names2: Vec<Arc<str>> = vec![Arc::from("intercept"), Arc::from("coef_t")];
1603        let baseline2 = PriorSet::weakly_informative(2);
1604        // Identical with wrong expected dim via target names of different length than source coefs.
1605        let names3: Vec<Arc<str>> =
1606            vec![Arc::from("intercept"), Arc::from("coef_t"), Arc::from("coef_w")];
1607        let baseline3 = PriorSet::weakly_informative(3);
1608        assert!(
1609            hydrate_prior(
1610                &HydrateMapping::IdenticalCoefficientSubspace,
1611                &quantities,
1612                &mean,
1613                &sd,
1614                &baseline3,
1615                &names3,
1616                None,
1617            )
1618            .is_err()
1619        );
1620
1621        assert!(
1622            hydrate_prior(
1623                &HydrateMapping::EffectFunctional { source_quantity: "missing".into() },
1624                &quantities,
1625                &mean,
1626                &sd,
1627                &baseline2,
1628                &names2,
1629                Some(1),
1630            )
1631            .is_err()
1632        );
1633
1634        assert!(
1635            hydrate_prior(
1636                &HydrateMapping::NamedParameters {
1637                    pairs: vec![("ate".into(), "no_such_coef".into())],
1638                },
1639                &quantities,
1640                &mean,
1641                &sd,
1642                &baseline2,
1643                &names2,
1644                None,
1645            )
1646            .is_err()
1647        );
1648
1649        assert!(
1650            hydrate_prior(
1651                &HydrateMapping::NamedParameters {
1652                    pairs: vec![("no_src".into(), "coef_t".into())],
1653                },
1654                &quantities,
1655                &mean,
1656                &sd,
1657                &baseline2,
1658                &names2,
1659                None,
1660            )
1661            .is_err()
1662        );
1663    }
1664
1665    #[test]
1666    fn hydrate_named_parameters_overwrites_target() {
1667        let quantities = vec![
1668            PosteriorQuantityKind::Coefficient { index: 0, name: Some(Arc::from("intercept")) },
1669            PosteriorQuantityKind::Coefficient { index: 1, name: Some(Arc::from("coef_t")) },
1670            PosteriorQuantityKind::Effect { name: Arc::from("ate") },
1671        ];
1672        let mean = vec![0.0, 0.0, 1.5];
1673        let sd = vec![1.0, 1.0, 0.2];
1674        let names: Vec<Arc<str>> = vec![Arc::from("intercept"), Arc::from("coef_t")];
1675        let baseline = PriorSet::weakly_informative(2);
1676        let prior = hydrate_prior(
1677            &HydrateMapping::NamedParameters { pairs: vec![("ate".into(), "coef_t".into())] },
1678            &quantities,
1679            &mean,
1680            &sd,
1681            &baseline,
1682            &names,
1683            None,
1684        )
1685        .unwrap();
1686        let coef = prior.gaussian_coefficients().unwrap();
1687        assert!((coef.mean[1] - 1.5).abs() < 1e-12);
1688        assert!(prior.restrictions.iter().any(|r| r.id.as_ref() == "external_named_prior"));
1689    }
1690}