antecedent_validate/stability/
false_positive.rs1#![allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
6
7use antecedent_core::{ExecutionContext, VariableId};
8use antecedent_data::{TimeSeriesData, surrogate_permute_columns, surrogate_phase_randomize};
9use antecedent_discovery::{DiscoveryWorkspace, Pcmci};
10
11use crate::error::ValidationError;
12
13#[derive(Clone, Copy, Debug, Eq, PartialEq)]
15pub enum NullTransform {
16 ColumnPermute,
18 PhaseRandomize,
20}
21
22#[derive(Clone, Debug)]
24pub struct FalsePositiveCheckReport {
25 pub method: NullTransform,
27 pub replicates: u32,
29 pub mean_edge_count: f64,
31 pub empirical_fpr: f64,
33 pub passed: bool,
35}
36
37#[derive(Clone, Debug)]
39pub struct FalsePositiveCheck {
40 pub pcmci: Pcmci,
42 pub transform: NullTransform,
44 pub replicates: u32,
46}
47
48impl FalsePositiveCheck {
49 #[must_use]
51 pub fn new(pcmci: Pcmci, transform: NullTransform, replicates: u32) -> Self {
52 Self { pcmci, transform, replicates }
53 }
54
55 pub fn run(
61 &self,
62 data: &TimeSeriesData,
63 variables: &[VariableId],
64 workspace: &mut DiscoveryWorkspace,
65 ctx: &ExecutionContext,
66 ) -> Result<FalsePositiveCheckReport, ValidationError> {
67 if self.replicates == 0 {
68 return Err(ValidationError::NotApplicable {
69 message: "false-positive check requires positive replicates",
70 });
71 }
72 let alpha = self.pcmci.engine().constraints.alpha;
73 let max_lag = self.pcmci.engine().constraints.temporal.max_lag.raw().max(1) as usize;
74 let family = (variables.len() * variables.len() * max_lag).max(1);
75 let mut rng = ctx.rng.stream(0xF41E_u64);
76 let mut total_edges = 0u64;
77 for _ in 0..self.replicates {
78 let null = match self.transform {
79 NullTransform::ColumnPermute => {
80 surrogate_permute_columns(data, &mut rng).map_err(ValidationError::from)?
81 }
82 NullTransform::PhaseRandomize => {
83 surrogate_phase_randomize(data, &mut rng).map_err(ValidationError::from)?
84 }
85 };
86 let result =
87 self.pcmci.run(&null, variables, workspace, ctx).map_err(ValidationError::from)?;
88 total_edges += result.evidence.links.len() as u64;
89 }
90 let mean_edge_count = total_edges as f64 / f64::from(self.replicates);
91 let empirical_fpr = mean_edge_count / family as f64;
92 let se = (alpha * (1.0 - alpha) / f64::from(self.replicates)).sqrt();
94 let passed = empirical_fpr <= alpha + (3.0 * se).max(0.05);
95 Ok(FalsePositiveCheckReport {
96 method: self.transform,
97 replicates: self.replicates,
98 mean_edge_count,
99 empirical_fpr,
100 passed,
101 })
102 }
103}
104
105#[cfg(test)]
106#[allow(clippy::cast_precision_loss)]
107mod tests {
108 use std::sync::Arc;
109
110 use antecedent_core::{
111 CausalSchemaBuilder, ExecutionContext, Lag, MeasurementSpec, RoleHint, SmallRoleSet,
112 ValueType, VariableId,
113 };
114 use antecedent_data::{
115 Float64Column, OwnedColumn, OwnedColumnarStorage, SamplingRegularity, TimeIndex,
116 ValidityBitmap,
117 };
118 use antecedent_discovery::{DiscoveryConstraints, TemporalConstraints};
119
120 use super::*;
121
122 fn linked_series() -> (TimeSeriesData, Vec<VariableId>) {
123 let n = 200usize;
124 let mut b = CausalSchemaBuilder::new();
125 for name in ["x", "y"] {
126 b.add_variable(
127 name,
128 ValueType::Continuous,
129 SmallRoleSet::from_hint(RoleHint::Context),
130 None,
131 None,
132 MeasurementSpec::default(),
133 )
134 .unwrap();
135 }
136 let schema = b.build().unwrap();
137 let mut x = vec![0.0; n];
138 let mut y = vec![0.0; n];
139 for t in 1..n {
140 x[t] = ((t as f64) * 0.02).sin();
141 y[t] = 0.9 * x[t - 1];
142 }
143 let cols = vec![
144 OwnedColumn::Float64(
145 Float64Column::new(
146 VariableId::from_raw(0),
147 Arc::from(x),
148 ValidityBitmap::all_valid(n),
149 )
150 .unwrap(),
151 ),
152 OwnedColumn::Float64(
153 Float64Column::new(
154 VariableId::from_raw(1),
155 Arc::from(y),
156 ValidityBitmap::all_valid(n),
157 )
158 .unwrap(),
159 ),
160 ];
161 let storage = OwnedColumnarStorage::try_new(schema, cols, None, None).unwrap();
162 let data = TimeSeriesData::try_new(
163 storage,
164 TimeIndex { regularity: SamplingRegularity::Regular { interval_ns: 1 }, length: n },
165 )
166 .unwrap();
167 (data, vec![VariableId::from_raw(0), VariableId::from_raw(1)])
168 }
169
170 #[test]
171 fn permute_null_reduces_edges() {
172 let (data, vars) = linked_series();
173 let constraints = DiscoveryConstraints {
174 temporal: TemporalConstraints { max_lag: Lag::from_raw(1), min_lag: Lag::from_raw(1) },
175 max_cond_size: 1,
176 alpha: 0.05,
177 ..Default::default()
178 };
179 let pcmci = Pcmci::new().with_fdr(false).with_constraints(constraints);
180 let mut ws = DiscoveryWorkspace::default();
181 let ctx = ExecutionContext::for_tests(8);
182 let before = pcmci.run(&data, &vars, &mut ws, &ctx).unwrap().evidence.links.len();
183 let check = FalsePositiveCheck::new(pcmci, NullTransform::ColumnPermute, 4);
184 let report = check.run(&data, &vars, &mut ws, &ctx).unwrap();
185 assert!(report.mean_edge_count <= before as f64 + 1.0);
186 assert_eq!(report.method, NullTransform::ColumnPermute);
187 }
188}