#![allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
use std::collections::BTreeSet;
use std::sync::Arc;
use antecedent_core::{ExecutionContext, VariableId};
use antecedent_data::{EnvHoldoutSplit, MultiEnvironmentData};
use antecedent_discovery::{DiscoveryWorkspace, JpcmciPlus, LaggedLink};
use crate::error::ValidationError;
#[derive(Clone, Debug)]
pub struct EnvironmentHoldoutReport {
pub discovery_links: Arc<[LaggedLink]>,
pub holdout_links: Arc<[LaggedLink]>,
pub shared_frequency: f64,
pub jaccard: f64,
}
#[derive(Clone, Debug)]
pub struct EnvironmentHoldout {
pub jpcmci: JpcmciPlus,
pub split: EnvHoldoutSplit,
}
impl EnvironmentHoldout {
#[must_use]
pub fn new(jpcmci: JpcmciPlus, split: EnvHoldoutSplit) -> Self {
Self { jpcmci, split }
}
pub fn run(
&self,
data: &MultiEnvironmentData,
variables: &[VariableId],
workspace: &mut DiscoveryWorkspace,
ctx: &ExecutionContext,
) -> Result<EnvironmentHoldoutReport, ValidationError> {
let train = subset_envs(data, &self.split.discovery_envs)?;
let holdout = subset_envs(data, &self.split.estimation_envs)?;
let train_res =
self.jpcmci.run(&train, variables, workspace, ctx).map_err(ValidationError::from)?;
let hold_res =
self.jpcmci.run(&holdout, variables, workspace, ctx).map_err(ValidationError::from)?;
let train_set: BTreeSet<LaggedLink> =
train_res.evidence.links.iter().map(|s| s.link).collect();
let hold_set: BTreeSet<LaggedLink> =
hold_res.evidence.links.iter().map(|s| s.link).collect();
let shared = train_set.intersection(&hold_set).count();
let union = train_set.union(&hold_set).count();
let shared_frequency =
if train_set.is_empty() { 1.0 } else { shared as f64 / train_set.len() as f64 };
let jaccard = if union == 0 { 1.0 } else { shared as f64 / union as f64 };
Ok(EnvironmentHoldoutReport {
discovery_links: Arc::from(train_set.into_iter().collect::<Vec<_>>()),
holdout_links: Arc::from(hold_set.into_iter().collect::<Vec<_>>()),
shared_frequency,
jaccard,
})
}
}
fn subset_envs(
data: &MultiEnvironmentData,
idxs: &[usize],
) -> Result<MultiEnvironmentData, ValidationError> {
if idxs.is_empty() {
return Err(ValidationError::NotApplicable {
message: "environment holdout subset is empty",
});
}
let mut envs = Vec::with_capacity(idxs.len());
for &i in idxs {
let env = data.environment(i).map_err(ValidationError::from)?;
envs.push(env.clone());
}
MultiEnvironmentData::try_new(envs).map_err(ValidationError::from)
}
#[cfg(test)]
#[allow(clippy::cast_precision_loss)]
mod tests {
use antecedent_core::{
CausalSchemaBuilder, ExecutionContext, Lag, MeasurementSpec, RoleHint, SmallRoleSet,
ValueType, VariableId,
};
use antecedent_data::{
Float64Column, OwnedColumn, OwnedColumnarStorage, SamplingRegularity, TimeIndex,
TimeSeriesData, ValidityBitmap,
};
use antecedent_discovery::{DiscoveryConstraints, DiscoveryWorkspace, TemporalConstraints};
use std::sync::Arc;
use super::*;
fn shared_lag_env(n: usize, seed: f64) -> TimeSeriesData {
let mut b = CausalSchemaBuilder::new();
for name in ["x", "y"] {
b.add_variable(
name,
ValueType::Continuous,
SmallRoleSet::from_hint(RoleHint::Context),
None,
None,
MeasurementSpec::default(),
)
.unwrap();
}
let schema = b.build().unwrap();
let mut x = vec![0.0; n];
let mut y = vec![0.0; n];
for t in 1..n {
x[t] = 0.4 * x[t - 1] + ((t as f64) * 0.02 + seed).sin() * 0.1;
y[t] = 0.75 * x[t - 1] + 0.2 * y[t - 1];
}
let cols = vec![
OwnedColumn::Float64(
Float64Column::new(
VariableId::from_raw(0),
Arc::from(x),
ValidityBitmap::all_valid(n),
)
.unwrap(),
),
OwnedColumn::Float64(
Float64Column::new(
VariableId::from_raw(1),
Arc::from(y),
ValidityBitmap::all_valid(n),
)
.unwrap(),
),
];
let storage = OwnedColumnarStorage::try_new(schema, cols, None, None).unwrap();
TimeSeriesData::try_new(
storage,
TimeIndex { regularity: SamplingRegularity::Regular { interval_ns: 1 }, length: n },
)
.unwrap()
}
#[test]
fn env_holdout_runs_two_envs() {
let multi =
MultiEnvironmentData::try_new([shared_lag_env(180, 0.0), shared_lag_env(180, 1.0)])
.unwrap();
let split = EnvHoldoutSplit::try_prefix(2, 1).unwrap();
let constraints = DiscoveryConstraints {
temporal: TemporalConstraints { max_lag: Lag::from_raw(1), min_lag: Lag::from_raw(1) },
max_cond_size: 1,
alpha: 0.15,
..Default::default()
};
let hold = EnvironmentHoldout::new(
JpcmciPlus::new().with_fdr(false).with_constraints(constraints),
split,
);
let mut ws = DiscoveryWorkspace::default();
let ctx = ExecutionContext::for_tests(4);
let vars = [VariableId::from_raw(0), VariableId::from_raw(1)];
let report = hold.run(&multi, &vars, &mut ws, &ctx).unwrap();
assert!(report.jaccard >= 0.0 && report.jaccard <= 1.0);
assert!(report.shared_frequency >= 0.0 && report.shared_frequency <= 1.0);
}
}