Skip to main content

antecedent_validate/stability/
pcmci_grid.rs

1//! PCMCI link-frequency stability and parameter grids.
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::BTreeMap;
8use std::sync::Arc;
9
10use antecedent_core::{ExecutionContext, Lag, VariableId};
11use antecedent_data::{ResamplingPlan, TableView, TimeSeriesData, resample_timeseries};
12use antecedent_discovery::{DiscoveryWorkspace, LaggedLink, Pcmci, ci_from_name};
13
14use crate::error::ValidationError;
15
16/// Stability frequency for one lagged link.
17#[derive(Clone, Copy, Debug, PartialEq)]
18pub struct LinkStability {
19    /// Link.
20    pub link: LaggedLink,
21    /// Fraction of bootstrap replicates / grid cells retaining the link.
22    pub frequency: f64,
23}
24
25/// Report from discovery stability or parameter-sensitivity grids.
26///
27/// For [`BlockBootstrapStability`], `replicates` is the bootstrap count and
28/// `block_size` is the moving-block length. For parameter sweeps
29/// ([`AlphaThresholdSensitivity`], [`LagWindowSensitivity`], [`CiTestSensitivity`]),
30/// `replicates` is the number of grid cells and `block_size` is `0`.
31#[derive(Clone, Debug)]
32pub struct DiscoveryStabilityReport {
33    /// Per-link frequencies (links seen in ≥1 replicate / grid cell).
34    pub frequencies: Arc<[LinkStability]>,
35    /// Replicates run (bootstrap) or grid cell count (parameter sweeps).
36    pub replicates: u32,
37    /// Block size used (`0` for parameter sweeps).
38    pub block_size: usize,
39}
40
41/// Block-bootstrap stability around a [`Pcmci`] configuration.
42#[derive(Clone, Debug)]
43pub struct BlockBootstrapStability {
44    /// PCMCI configuration to re-run.
45    pub pcmci: Pcmci,
46    /// Bootstrap replicates.
47    pub replicates: u32,
48    /// Block length.
49    pub block_size: usize,
50}
51
52impl Default for BlockBootstrapStability {
53    fn default() -> Self {
54        Self::new()
55    }
56}
57
58impl BlockBootstrapStability {
59    /// Defaults: 20 replicates, block size 20.
60    #[must_use]
61    pub fn new() -> Self {
62        Self { pcmci: Pcmci::new().with_fdr(false), replicates: 20, block_size: 20 }
63    }
64
65    /// Run stability assessment.
66    ///
67    /// # Errors
68    ///
69    /// Data or discovery failures.
70    pub fn run(
71        &self,
72        data: &TimeSeriesData,
73        variables: &[VariableId],
74        workspace: &mut DiscoveryWorkspace,
75        ctx: &ExecutionContext,
76    ) -> Result<DiscoveryStabilityReport, ValidationError> {
77        if self.replicates == 0 || self.block_size == 0 {
78            return Err(ValidationError::NotApplicable {
79                message: "stability requires positive replicates and block_size",
80            });
81        }
82        if self.block_size > data.row_count() {
83            return Err(ValidationError::NotApplicable {
84                message: "block_size exceeds series length",
85            });
86        }
87        let mut counts: BTreeMap<LaggedLink, u32> = BTreeMap::new();
88        let mut rng = ctx.rng.stream(0x57AB_u64);
89        let mut index_scratch = Vec::new();
90        for _ in 0..self.replicates {
91            let boot = resample_timeseries(
92                data,
93                ResamplingPlan::MovingBlock { length: self.block_size },
94                &mut rng,
95                &mut index_scratch,
96            )
97            .map_err(ValidationError::from)?;
98            let result =
99                self.pcmci.run(&boot, variables, workspace, ctx).map_err(ValidationError::from)?;
100            for s in result.evidence.links.iter() {
101                *counts.entry(s.link).or_insert(0) += 1;
102            }
103        }
104        Ok(report_from_counts(counts, self.replicates, self.block_size))
105    }
106}
107
108/// Alpha-threshold sensitivity: re-run PCMCI across an `alpha` grid on the same data.
109#[derive(Clone, Debug)]
110pub struct AlphaThresholdSensitivity {
111    /// Base PCMCI configuration (FDR, CI, `max_lag`, …).
112    pub pcmci: Pcmci,
113    /// Significance levels to sweep.
114    pub alphas: Arc<[f64]>,
115}
116
117impl AlphaThresholdSensitivity {
118    /// Build with a base config and alpha grid.
119    #[must_use]
120    pub fn new(pcmci: Pcmci, alphas: impl Into<Arc<[f64]>>) -> Self {
121        Self { pcmci, alphas: alphas.into() }
122    }
123
124    /// Run the alpha grid.
125    ///
126    /// # Errors
127    ///
128    /// Empty/invalid grid, or discovery failures.
129    pub fn run(
130        &self,
131        data: &TimeSeriesData,
132        variables: &[VariableId],
133        workspace: &mut DiscoveryWorkspace,
134        ctx: &ExecutionContext,
135    ) -> Result<DiscoveryStabilityReport, ValidationError> {
136        if self.alphas.is_empty() {
137            return Err(ValidationError::NotApplicable {
138                message: "alpha sensitivity requires a non-empty alphas grid",
139            });
140        }
141        if self.alphas.iter().any(|&a| !(a > 0.0 && a <= 1.0)) {
142            return Err(ValidationError::NotApplicable {
143                message: "alpha sensitivity requires alphas in (0, 1]",
144            });
145        }
146        let configs = self.alphas.iter().map(|&alpha| {
147            let mut constraints = self.pcmci.engine().constraints.clone();
148            constraints.alpha = alpha;
149            self.pcmci.clone().with_constraints(constraints)
150        });
151        run_param_grid(configs, data, variables, workspace, ctx)
152    }
153}
154
155/// Lag-window sensitivity: re-run PCMCI across a `max_lag` grid on the same data.
156#[derive(Clone, Debug)]
157pub struct LagWindowSensitivity {
158    /// Base PCMCI configuration.
159    pub pcmci: Pcmci,
160    /// Maximum lags to sweep.
161    pub max_lags: Arc<[u32]>,
162}
163
164impl LagWindowSensitivity {
165    /// Build with a base config and max-lag grid.
166    #[must_use]
167    pub fn new(pcmci: Pcmci, max_lags: impl Into<Arc<[u32]>>) -> Self {
168        Self { pcmci, max_lags: max_lags.into() }
169    }
170
171    /// Run the lag-window grid.
172    ///
173    /// # Errors
174    ///
175    /// Empty/invalid grid, or discovery failures.
176    pub fn run(
177        &self,
178        data: &TimeSeriesData,
179        variables: &[VariableId],
180        workspace: &mut DiscoveryWorkspace,
181        ctx: &ExecutionContext,
182    ) -> Result<DiscoveryStabilityReport, ValidationError> {
183        if self.max_lags.is_empty() {
184            return Err(ValidationError::NotApplicable {
185                message: "lag-window sensitivity requires a non-empty max_lags grid",
186            });
187        }
188        let min_lag = self.pcmci.engine().constraints.temporal.min_lag.raw();
189        if self.max_lags.iter().any(|&m| m < min_lag) {
190            return Err(ValidationError::NotApplicable {
191                message: "lag-window sensitivity requires max_lag ≥ constraints.min_lag",
192            });
193        }
194        let configs = self.max_lags.iter().map(|&max_lag| {
195            let mut constraints = self.pcmci.engine().constraints.clone();
196            constraints.temporal.max_lag = Lag::from_raw(max_lag);
197            self.pcmci.clone().with_constraints(constraints)
198        });
199        run_param_grid(configs, data, variables, workspace, ctx)
200    }
201}
202
203/// CI-test sensitivity: re-run PCMCI across named CI tests on the same data.
204#[derive(Clone, Debug)]
205pub struct CiTestSensitivity {
206    /// Base PCMCI configuration (constraints / FDR fixed).
207    pub pcmci: Pcmci,
208    /// CI test names resolved via [`ci_from_name`].
209    pub ci_names: Arc<[Arc<str>]>,
210}
211
212impl CiTestSensitivity {
213    /// Build with a base config and CI name grid.
214    #[must_use]
215    pub fn new(pcmci: Pcmci, ci_names: impl Into<Arc<[Arc<str>]>>) -> Self {
216        Self { pcmci, ci_names: ci_names.into() }
217    }
218
219    /// Run the CI-test grid.
220    ///
221    /// # Errors
222    ///
223    /// Empty grid, unknown CI name, or discovery failures.
224    pub fn run(
225        &self,
226        data: &TimeSeriesData,
227        variables: &[VariableId],
228        workspace: &mut DiscoveryWorkspace,
229        ctx: &ExecutionContext,
230    ) -> Result<DiscoveryStabilityReport, ValidationError> {
231        if self.ci_names.is_empty() {
232            return Err(ValidationError::NotApplicable {
233                message: "CI-test sensitivity requires a non-empty ci_names grid",
234            });
235        }
236        let mut configs = Vec::with_capacity(self.ci_names.len());
237        for name in self.ci_names.iter() {
238            let ci = ci_from_name(name).map_err(|_e| ValidationError::NotApplicable {
239                message: "CI-test sensitivity: unknown or unsupported CI name",
240            })?;
241            configs.push(self.pcmci.clone().with_ci(ci));
242        }
243        run_param_grid(configs, data, variables, workspace, ctx)
244    }
245}
246
247pub(crate) fn run_param_grid(
248    configs: impl IntoIterator<Item = Pcmci>,
249    data: &TimeSeriesData,
250    variables: &[VariableId],
251    workspace: &mut DiscoveryWorkspace,
252    ctx: &ExecutionContext,
253) -> Result<DiscoveryStabilityReport, ValidationError> {
254    let mut counts: BTreeMap<LaggedLink, u32> = BTreeMap::new();
255    let mut cells = 0u32;
256    for pcmci in configs {
257        cells = cells.saturating_add(1);
258        let result = pcmci.run(data, variables, workspace, ctx).map_err(ValidationError::from)?;
259        for s in result.evidence.links.iter() {
260            *counts.entry(s.link).or_insert(0) += 1;
261        }
262    }
263    if cells == 0 {
264        return Err(ValidationError::NotApplicable {
265            message: "parameter sensitivity grid produced zero cells",
266        });
267    }
268    Ok(report_from_counts(counts, cells, 0))
269}
270
271pub(crate) fn report_from_counts(
272    counts: BTreeMap<LaggedLink, u32>,
273    replicates: u32,
274    block_size: usize,
275) -> DiscoveryStabilityReport {
276    let mut frequencies = Vec::with_capacity(counts.len());
277    for (link, c) in counts {
278        frequencies.push(LinkStability { link, frequency: f64::from(c) / f64::from(replicates) });
279    }
280    frequencies
281        .sort_by(|a, b| b.frequency.partial_cmp(&a.frequency).unwrap_or(std::cmp::Ordering::Equal));
282    DiscoveryStabilityReport { frequencies: Arc::from(frequencies), replicates, block_size }
283}
284
285#[cfg(test)]
286#[allow(clippy::cast_precision_loss, clippy::many_single_char_names)]
287mod tests {
288    use antecedent_core::{
289        CausalSchemaBuilder, ExecutionContext, Lag, MeasurementSpec, RoleHint, SmallRoleSet,
290        ValueType, VariableId,
291    };
292    use antecedent_data::{
293        Float64Column, OwnedColumn, OwnedColumnarStorage, SamplingRegularity, TimeIndex,
294        TimeSeriesData, ValidityBitmap,
295    };
296    use antecedent_discovery::{DiscoveryConstraints, DiscoveryWorkspace, TemporalConstraints};
297
298    use super::*;
299
300    fn linked_series() -> (TimeSeriesData, Vec<VariableId>) {
301        let n = 300usize;
302        let mut b = CausalSchemaBuilder::new();
303        b.add_variable(
304            "x",
305            ValueType::Continuous,
306            SmallRoleSet::from_hint(RoleHint::Context),
307            None,
308            None,
309            MeasurementSpec::default(),
310        )
311        .unwrap();
312        b.add_variable(
313            "y",
314            ValueType::Continuous,
315            SmallRoleSet::from_hint(RoleHint::Context),
316            None,
317            None,
318            MeasurementSpec::default(),
319        )
320        .unwrap();
321        let schema = b.build().unwrap();
322        let mut x = vec![0.0; n];
323        let mut y = vec![0.0; n];
324        for t in 1..n {
325            x[t] = ((t as f64) * 0.02).sin();
326            y[t] = 0.9 * x[t - 1];
327        }
328        let cols = vec![
329            OwnedColumn::Float64(
330                Float64Column::new(
331                    VariableId::from_raw(0),
332                    Arc::from(x),
333                    ValidityBitmap::all_valid(n),
334                )
335                .unwrap(),
336            ),
337            OwnedColumn::Float64(
338                Float64Column::new(
339                    VariableId::from_raw(1),
340                    Arc::from(y),
341                    ValidityBitmap::all_valid(n),
342                )
343                .unwrap(),
344            ),
345        ];
346        let storage = OwnedColumnarStorage::try_new(schema, cols, None, None).unwrap();
347        let data = TimeSeriesData::try_new(
348            storage,
349            TimeIndex { regularity: SamplingRegularity::Regular { interval_ns: 1 }, length: n },
350        )
351        .unwrap();
352        (data, vec![VariableId::from_raw(0), VariableId::from_raw(1)])
353    }
354
355    fn base_pcmci() -> Pcmci {
356        Pcmci::new().with_fdr(false).with_constraints(DiscoveryConstraints {
357            temporal: TemporalConstraints { max_lag: Lag::from_raw(2), min_lag: Lag::from_raw(1) },
358            max_cond_size: 1,
359            alpha: 0.05,
360            ..DiscoveryConstraints::default()
361        })
362    }
363
364    fn true_link_freq(report: &DiscoveryStabilityReport) -> f64 {
365        report
366            .frequencies
367            .iter()
368            .find(|f| {
369                f.link.source == VariableId::from_raw(0)
370                    && f.link.target == VariableId::from_raw(1)
371                    && f.link.source_lag.raw() == 1
372            })
373            .map_or(0.0, |f| f.frequency)
374    }
375
376    #[test]
377    fn true_link_is_stable() {
378        let (data, vars) = linked_series();
379        let mut stab = BlockBootstrapStability::new();
380        stab.replicates = 8;
381        stab.block_size = 25;
382        stab.pcmci = base_pcmci();
383        let mut ws = DiscoveryWorkspace::default();
384        let ctx = ExecutionContext::for_tests(5);
385        let report = stab.run(&data, &vars, &mut ws, &ctx).unwrap();
386        assert!(
387            true_link_freq(&report) > 0.0,
388            "expected true link to appear; report={:?}",
389            report.frequencies
390        );
391    }
392
393    #[test]
394    fn alpha_threshold_retains_true_link() {
395        let (data, vars) = linked_series();
396        let sens = AlphaThresholdSensitivity::new(base_pcmci(), Arc::from([0.05f64, 0.1, 0.2]));
397        let mut ws = DiscoveryWorkspace::default();
398        let ctx = ExecutionContext::for_tests(5);
399        let report = sens.run(&data, &vars, &mut ws, &ctx).unwrap();
400        assert_eq!(report.replicates, 3);
401        assert_eq!(report.block_size, 0);
402        assert!(true_link_freq(&report) > 0.0);
403    }
404
405    #[test]
406    fn lag_window_retains_true_link() {
407        let (data, vars) = linked_series();
408        let sens = LagWindowSensitivity::new(base_pcmci(), Arc::from([1u32, 2, 3]));
409        let mut ws = DiscoveryWorkspace::default();
410        let ctx = ExecutionContext::for_tests(5);
411        let report = sens.run(&data, &vars, &mut ws, &ctx).unwrap();
412        assert_eq!(report.replicates, 3);
413        assert!(true_link_freq(&report) > 0.0);
414    }
415
416    #[test]
417    fn ci_test_retains_true_link() {
418        let (data, vars) = linked_series();
419        let names: Arc<[Arc<str>]> =
420            Arc::from([Arc::<str>::from("parcorr"), Arc::<str>::from("robust_parcorr")]);
421        let sens = CiTestSensitivity::new(base_pcmci(), names);
422        let mut ws = DiscoveryWorkspace::default();
423        let ctx = ExecutionContext::for_tests(5);
424        let report = sens.run(&data, &vars, &mut ws, &ctx).unwrap();
425        assert_eq!(report.replicates, 2);
426        assert!(true_link_freq(&report) > 0.0);
427    }
428
429    #[test]
430    fn empty_grids_not_applicable() {
431        let (data, vars) = linked_series();
432        let mut ws = DiscoveryWorkspace::default();
433        let ctx = ExecutionContext::for_tests(5);
434        assert!(matches!(
435            AlphaThresholdSensitivity::new(base_pcmci(), Arc::from([]) as Arc<[f64]>)
436                .run(&data, &vars, &mut ws, &ctx),
437            Err(ValidationError::NotApplicable { .. })
438        ));
439        assert!(matches!(
440            LagWindowSensitivity::new(base_pcmci(), Arc::from([]) as Arc<[u32]>)
441                .run(&data, &vars, &mut ws, &ctx),
442            Err(ValidationError::NotApplicable { .. })
443        ));
444        assert!(matches!(
445            CiTestSensitivity::new(base_pcmci(), Arc::from([]) as Arc<[Arc<str>]>)
446                .run(&data, &vars, &mut ws, &ctx),
447            Err(ValidationError::NotApplicable { .. })
448        ));
449    }
450}