Skip to main content

antecedent_validate/stability/
false_positive.rs

1//! False-positive checks via permute / phase-randomize surrogates.
2//!
3//! SPDX-License-Identifier: MIT OR Apache-2.0
4
5#![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/// Null transform applied to observed series before rediscovery.
14#[derive(Clone, Copy, Debug, Eq, PartialEq)]
15pub enum NullTransform {
16    /// Independently permute each column.
17    ColumnPermute,
18    /// Phase-randomize each column (preserve spectrum).
19    PhaseRandomize,
20}
21
22/// Report from [`FalsePositiveCheck`].
23#[derive(Clone, Debug)]
24pub struct FalsePositiveCheckReport {
25    /// Transform used.
26    pub method: NullTransform,
27    /// Surrogate replicates.
28    pub replicates: u32,
29    /// Mean retained edge count after nullification.
30    pub mean_edge_count: f64,
31    /// Empirical edge rate vs family size estimate.
32    pub empirical_fpr: f64,
33    /// Whether mean edge count is at/below the α-calibrated expectation band.
34    pub passed: bool,
35}
36
37/// Apply surrogate nulls to observed data and re-run PCMCI.
38#[derive(Clone, Debug)]
39pub struct FalsePositiveCheck {
40    /// PCMCI configuration.
41    pub pcmci: Pcmci,
42    /// Null transform.
43    pub transform: NullTransform,
44    /// Surrogate replicates.
45    pub replicates: u32,
46}
47
48impl FalsePositiveCheck {
49    /// Build a false-positive check.
50    #[must_use]
51    pub fn new(pcmci: Pcmci, transform: NullTransform, replicates: u32) -> Self {
52        Self { pcmci, transform, replicates }
53    }
54
55    /// Run surrogate false-positive assessment on observed `data`.
56    ///
57    /// # Errors
58    ///
59    /// Invalid config, surrogate, or discovery failures.
60    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        // Pass if empirical FPR is not far above α (allow 3√(α(1-α)/R) + 0.05 floor).
93        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}