antecedent_validate/stability/
env_holdout.rs1#![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#[derive(Clone, Debug)]
18pub struct EnvironmentHoldoutReport {
19 pub discovery_links: Arc<[LaggedLink]>,
21 pub holdout_links: Arc<[LaggedLink]>,
23 pub shared_frequency: f64,
25 pub jaccard: f64,
27}
28
29#[derive(Clone, Debug)]
31pub struct EnvironmentHoldout {
32 pub jpcmci: JpcmciPlus,
34 pub split: EnvHoldoutSplit,
36}
37
38impl EnvironmentHoldout {
39 #[must_use]
41 pub fn new(jpcmci: JpcmciPlus, split: EnvHoldoutSplit) -> Self {
42 Self { jpcmci, split }
43 }
44
45 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}