1#![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#[derive(Clone, Copy, Debug, PartialEq)]
18pub struct LinkStability {
19 pub link: LaggedLink,
21 pub frequency: f64,
23}
24
25#[derive(Clone, Debug)]
32pub struct DiscoveryStabilityReport {
33 pub frequencies: Arc<[LinkStability]>,
35 pub replicates: u32,
37 pub block_size: usize,
39}
40
41#[derive(Clone, Debug)]
43pub struct BlockBootstrapStability {
44 pub pcmci: Pcmci,
46 pub replicates: u32,
48 pub block_size: usize,
50}
51
52impl Default for BlockBootstrapStability {
53 fn default() -> Self {
54 Self::new()
55 }
56}
57
58impl BlockBootstrapStability {
59 #[must_use]
61 pub fn new() -> Self {
62 Self { pcmci: Pcmci::new().with_fdr(false), replicates: 20, block_size: 20 }
63 }
64
65 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#[derive(Clone, Debug)]
110pub struct AlphaThresholdSensitivity {
111 pub pcmci: Pcmci,
113 pub alphas: Arc<[f64]>,
115}
116
117impl AlphaThresholdSensitivity {
118 #[must_use]
120 pub fn new(pcmci: Pcmci, alphas: impl Into<Arc<[f64]>>) -> Self {
121 Self { pcmci, alphas: alphas.into() }
122 }
123
124 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#[derive(Clone, Debug)]
157pub struct LagWindowSensitivity {
158 pub pcmci: Pcmci,
160 pub max_lags: Arc<[u32]>,
162}
163
164impl LagWindowSensitivity {
165 #[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 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#[derive(Clone, Debug)]
205pub struct CiTestSensitivity {
206 pub pcmci: Pcmci,
208 pub ci_names: Arc<[Arc<str>]>,
210}
211
212impl CiTestSensitivity {
213 #[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 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}