1#![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#[derive(Clone, Debug)]
42pub struct CausalPosterior {
43 pub draws: PosteriorDraws,
45 pub summaries: PosteriorSummary,
47 pub identification: IdentificationStatus,
49 pub prior_sensitivity: Option<PriorSensitivitySummary>,
51 pub conflict_summary: Option<ConflictSummary>,
53 pub diagnostics: InferenceDiagnostics,
55 pub assumptions: AssumptionSet,
57 pub unidentified_mass: f64,
59 pub early_stopped: bool,
61}
62
63impl CausalPosterior {
64 #[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 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
87const HYDRATE_VAR_FLOOR: f64 = 1e-12;
89
90pub 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
163pub 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#[derive(Clone, Debug, PartialEq, Eq, Hash)]
185pub enum HydrateMapping {
186 IdenticalCoefficientSubspace,
188 EffectFunctional {
190 source_quantity: String,
192 },
193 NamedParameters {
195 pairs: Vec<(String, String)>,
197 },
198}
199
200pub 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 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#[derive(Clone, Debug)]
371pub struct BayesianGlmMechanism {
372 pub coefficient_draws: PosteriorDraws,
374 pub map: Vec<f64>,
376 pub likelihood: BayesLikelihood,
378 pub diagnostics: InferenceDiagnostics,
380 pub design: CompiledDesign,
382 pub treatment_col: usize,
384 pub active: f64,
386 pub control: f64,
388}
389
390#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
392pub enum BayesianBackendKind {
393 ConjugateGaussian,
395 Laplace,
397 Hmc,
399}
400
401#[derive(Clone, Debug)]
403pub struct BayesianGComputationAte {
404 pub backend: BayesianBackendKind,
406 pub likelihood: BayesLikelihood,
408 pub n_draws: usize,
410 pub seed: u64,
412 pub overlap: OverlapPolicy,
414 pub prior_scale: f64,
416 pub prior: Option<PriorSet>,
419}
420
421impl Default for BayesianGComputationAte {
422 fn default() -> Self {
423 Self::new()
424 }
425}
426
427impl BayesianGComputationAte {
428 #[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 #[must_use]
444 pub fn conjugate() -> Self {
445 Self {
446 backend: BayesianBackendKind::ConjugateGaussian,
447 likelihood: BayesLikelihood::GaussianIdentity,
448 ..Self::new()
449 }
450 }
451
452 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 #[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 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 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 {
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; 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 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 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#[derive(Clone, Debug, Default)]
891pub struct BayesianTemporalGcomp {
892 pub inner: BayesianGComputationAte,
894}
895
896impl BayesianTemporalGcomp {
897 #[must_use]
899 pub fn new() -> Self {
900 Self { inner: BayesianGComputationAte::new() }
901 }
902
903 #[must_use]
905 pub fn conjugate() -> Self {
906 Self { inner: BayesianGComputationAte::conjugate() }
907 }
908
909 #[must_use]
911 pub fn from_prepared_estimation(prep: &PreparedEstimationProblem) -> PreparedBayesianProblem {
912 BayesianGComputationAte::from_prepared_estimation(prep)
913 }
914
915 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#[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
955fn 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#[derive(Clone, Debug)]
968pub struct PreparedBayesianProblem {
969 pub design: CompiledDesign,
971 pub method: Arc<str>,
973 pub adjustment_set: Arc<[VariableId]>,
975 pub active: f64,
977 pub control: f64,
979 pub overlap: OverlapPolicy,
981 pub coef_names: Option<Arc<[Arc<str>]>>,
983}
984
985#[derive(Clone, Debug, Default)]
987pub struct BayesianGCompWorkspace {
988 pub laplace: LaplaceWorkspace,
990 pub eval: PosteriorEvalWorkspace,
992}
993
994pub trait PosteriorFunctionalEvaluator {
996 type Compiled;
998
999 fn compile(&self) -> Result<Self::Compiled, EstimationError>;
1005
1006 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#[derive(Clone, Debug)]
1023pub struct GCompAteEvaluator {
1024 pub family: GlmFamily,
1026 pub treatment_col: usize,
1028 pub active: f64,
1030 pub control: f64,
1032 pub nrows: usize,
1034 pub ncols: usize,
1036 pub matrix: Arc<[f64]>,
1038}
1039
1040#[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 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
1139fn quantile_width_95(values: &[f64]) -> f64 {
1141 if values.len() < 2 {
1142 return f64::NAN;
1143 }
1144 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
1155fn 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
1175fn 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#[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 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 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}