Skip to main content

antecedent_validate/stability/
env_holdout.rs

1//! Environment holdout validation via J-PCMCI+.
2//!
3//! SPDX-License-Identifier: MIT OR Apache-2.0
4
5#![allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
6
7use std::collections::BTreeSet;
8use std::sync::Arc;
9
10use antecedent_core::{ExecutionContext, VariableId};
11use antecedent_data::{EnvHoldoutSplit, MultiEnvironmentData};
12use antecedent_discovery::{DiscoveryWorkspace, JpcmciPlus, LaggedLink};
13
14use crate::error::ValidationError;
15
16/// Report comparing discovery vs holdout environment graphs.
17#[derive(Clone, Debug)]
18pub struct EnvironmentHoldoutReport {
19    /// Links discovered on training environments.
20    pub discovery_links: Arc<[LaggedLink]>,
21    /// Links discovered on holdout environments.
22    pub holdout_links: Arc<[LaggedLink]>,
23    /// Fraction of discovery links also present on holdout.
24    pub shared_frequency: f64,
25    /// Jaccard index of the two link sets.
26    pub jaccard: f64,
27}
28
29/// Environment-holdout discovery agreement under [`JpcmciPlus`].
30#[derive(Clone, Debug)]
31pub struct EnvironmentHoldout {
32    /// J-PCMCI+ configuration.
33    pub jpcmci: JpcmciPlus,
34    /// Discovery vs estimation environment indexes.
35    pub split: EnvHoldoutSplit,
36}
37
38impl EnvironmentHoldout {
39    /// Build with a J-PCMCI+ config and holdout split.
40    #[must_use]
41    pub fn new(jpcmci: JpcmciPlus, split: EnvHoldoutSplit) -> Self {
42        Self { jpcmci, split }
43    }
44
45    /// Discover independently on train and holdout env subsets; report link overlap.
46    ///
47    /// # Errors
48    ///
49    /// Split indexes out of range, empty subsets, or discovery failures.
50    pub fn run(
51        &self,
52        data: &MultiEnvironmentData,
53        variables: &[VariableId],
54        workspace: &mut DiscoveryWorkspace,
55        ctx: &ExecutionContext,
56    ) -> Result<EnvironmentHoldoutReport, ValidationError> {
57        let train = subset_envs(data, &self.split.discovery_envs)?;
58        let holdout = subset_envs(data, &self.split.estimation_envs)?;
59        let train_res =
60            self.jpcmci.run(&train, variables, workspace, ctx).map_err(ValidationError::from)?;
61        let hold_res =
62            self.jpcmci.run(&holdout, variables, workspace, ctx).map_err(ValidationError::from)?;
63        let train_set: BTreeSet<LaggedLink> =
64            train_res.evidence.links.iter().map(|s| s.link).collect();
65        let hold_set: BTreeSet<LaggedLink> =
66            hold_res.evidence.links.iter().map(|s| s.link).collect();
67        let shared = train_set.intersection(&hold_set).count();
68        let union = train_set.union(&hold_set).count();
69        let shared_frequency =
70            if train_set.is_empty() { 1.0 } else { shared as f64 / train_set.len() as f64 };
71        let jaccard = if union == 0 { 1.0 } else { shared as f64 / union as f64 };
72        Ok(EnvironmentHoldoutReport {
73            discovery_links: Arc::from(train_set.into_iter().collect::<Vec<_>>()),
74            holdout_links: Arc::from(hold_set.into_iter().collect::<Vec<_>>()),
75            shared_frequency,
76            jaccard,
77        })
78    }
79}
80
81fn subset_envs(
82    data: &MultiEnvironmentData,
83    idxs: &[usize],
84) -> Result<MultiEnvironmentData, ValidationError> {
85    if idxs.is_empty() {
86        return Err(ValidationError::NotApplicable {
87            message: "environment holdout subset is empty",
88        });
89    }
90    let mut envs = Vec::with_capacity(idxs.len());
91    for &i in idxs {
92        let env = data.environment(i).map_err(ValidationError::from)?;
93        envs.push(env.clone());
94    }
95    MultiEnvironmentData::try_new(envs).map_err(ValidationError::from)
96}
97
98#[cfg(test)]
99#[allow(clippy::cast_precision_loss)]
100mod tests {
101    use antecedent_core::{
102        CausalSchemaBuilder, ExecutionContext, Lag, MeasurementSpec, RoleHint, SmallRoleSet,
103        ValueType, VariableId,
104    };
105    use antecedent_data::{
106        Float64Column, OwnedColumn, OwnedColumnarStorage, SamplingRegularity, TimeIndex,
107        TimeSeriesData, ValidityBitmap,
108    };
109    use antecedent_discovery::{DiscoveryConstraints, DiscoveryWorkspace, TemporalConstraints};
110    use std::sync::Arc;
111
112    use super::*;
113
114    fn shared_lag_env(n: usize, seed: f64) -> TimeSeriesData {
115        let mut b = CausalSchemaBuilder::new();
116        for name in ["x", "y"] {
117            b.add_variable(
118                name,
119                ValueType::Continuous,
120                SmallRoleSet::from_hint(RoleHint::Context),
121                None,
122                None,
123                MeasurementSpec::default(),
124            )
125            .unwrap();
126        }
127        let schema = b.build().unwrap();
128        let mut x = vec![0.0; n];
129        let mut y = vec![0.0; n];
130        for t in 1..n {
131            x[t] = 0.4 * x[t - 1] + ((t as f64) * 0.02 + seed).sin() * 0.1;
132            y[t] = 0.75 * x[t - 1] + 0.2 * y[t - 1];
133        }
134        let cols = vec![
135            OwnedColumn::Float64(
136                Float64Column::new(
137                    VariableId::from_raw(0),
138                    Arc::from(x),
139                    ValidityBitmap::all_valid(n),
140                )
141                .unwrap(),
142            ),
143            OwnedColumn::Float64(
144                Float64Column::new(
145                    VariableId::from_raw(1),
146                    Arc::from(y),
147                    ValidityBitmap::all_valid(n),
148                )
149                .unwrap(),
150            ),
151        ];
152        let storage = OwnedColumnarStorage::try_new(schema, cols, None, None).unwrap();
153        TimeSeriesData::try_new(
154            storage,
155            TimeIndex { regularity: SamplingRegularity::Regular { interval_ns: 1 }, length: n },
156        )
157        .unwrap()
158    }
159
160    #[test]
161    fn env_holdout_runs_two_envs() {
162        let multi =
163            MultiEnvironmentData::try_new([shared_lag_env(180, 0.0), shared_lag_env(180, 1.0)])
164                .unwrap();
165        let split = EnvHoldoutSplit::try_prefix(2, 1).unwrap();
166        let constraints = DiscoveryConstraints {
167            temporal: TemporalConstraints { max_lag: Lag::from_raw(1), min_lag: Lag::from_raw(1) },
168            max_cond_size: 1,
169            alpha: 0.15,
170            ..Default::default()
171        };
172        let hold = EnvironmentHoldout::new(
173            JpcmciPlus::new().with_fdr(false).with_constraints(constraints),
174            split,
175        );
176        let mut ws = DiscoveryWorkspace::default();
177        let ctx = ExecutionContext::for_tests(4);
178        let vars = [VariableId::from_raw(0), VariableId::from_raw(1)];
179        let report = hold.run(&multi, &vars, &mut ws, &ctx).unwrap();
180        assert!(report.jaccard >= 0.0 && report.jaccard <= 1.0);
181        assert!(report.shared_frequency >= 0.0 && report.shared_frequency <= 1.0);
182    }
183}