1#![allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)]
9
10use std::sync::Arc;
11
12use antecedent_core::VariableId;
13use antecedent_core::{TemporalIndexer, TemporalNodeKey};
14
15use crate::dag::Dag;
16use crate::error::GraphError;
17use crate::temporal::TemporalDag;
18use crate::types::{DenseNodeId, NodeRef};
19
20#[derive(Clone, Debug)]
22pub struct UnfoldedTemporalGraph {
23 pub dag: Dag,
25 pub indexer: TemporalIndexer,
27}
28
29#[derive(Clone, Debug)]
31pub struct LazyUnfoldedTemporalGraph {
32 pub template: TemporalDag,
34 pub indexer: TemporalIndexer,
36}
37
38impl TemporalDag {
39 pub fn unfold_lazy(
45 &self,
46 indexer: TemporalIndexer,
47 ) -> Result<LazyUnfoldedTemporalGraph, GraphError> {
48 for (i, _) in self.nodes().iter().enumerate() {
50 let id = DenseNodeId::try_from_usize(i)?;
51 let _ = self
52 .temporal_key(id)
53 .ok_or(GraphError::InvalidEndpoints { message: "unfold requires lagged nodes" })?;
54 }
55 Ok(LazyUnfoldedTemporalGraph { template: self.clone(), indexer })
56 }
57
58 pub fn unfold(&self, indexer: TemporalIndexer) -> Result<UnfoldedTemporalGraph, GraphError> {
64 self.unfold_lazy(indexer)?.materialize()
65 }
66}
67
68impl LazyUnfoldedTemporalGraph {
69 pub fn has_edge(&self, from: TemporalNodeKey, to: TemporalNodeKey) -> Result<bool, GraphError> {
75 let _ = self.indexer.dense_id(from).map_err(|_| GraphError::InvalidEndpoints {
76 message: "unfold endpoint outside window",
77 })?;
78 let _ = self.indexer.dense_id(to).map_err(|_| GraphError::InvalidEndpoints {
79 message: "unfold endpoint outside window",
80 })?;
81 for (from_i, _) in self.template.nodes().iter().enumerate() {
82 let from_id = DenseNodeId::try_from_usize(from_i)?;
83 let Some(from_key) = self.template.temporal_key(from_id) else {
84 continue;
85 };
86 for &to_id in self.template.children(from_id) {
87 let Some(to_key) = self.template.temporal_key(to_id) else {
88 continue;
89 };
90 if edge_matches(from_key, to_key, from, to) {
91 return Ok(true);
92 }
93 }
94 }
95 Ok(false)
96 }
97
98 pub fn materialize(&self) -> Result<UnfoldedTemporalGraph, GraphError> {
104 let n = self.indexer.dense_len();
105 let n_u32 = u32::try_from(n).map_err(|_| GraphError::TooManyNodes)?;
106 let mut dag = Dag::with_variables(n_u32);
107
108 for (from_i, _) in self.template.nodes().iter().enumerate() {
109 let from = DenseNodeId::try_from_usize(from_i)?;
110 let from_key = self
111 .template
112 .temporal_key(from)
113 .ok_or(GraphError::InvalidEndpoints { message: "unfold requires lagged nodes" })?;
114 for &to in self.template.children(from) {
115 let to_key =
116 self.template.temporal_key(to).ok_or(GraphError::InvalidEndpoints {
117 message: "unfold requires lagged nodes",
118 })?;
119 insert_replicated_edges(&mut dag, &self.indexer, from_key, to_key)?;
120 }
121 }
122
123 Ok(UnfoldedTemporalGraph { dag, indexer: self.indexer.clone() })
124 }
125}
126
127fn edge_matches(
128 template_from: TemporalNodeKey,
129 template_to: TemporalNodeKey,
130 concrete_from: TemporalNodeKey,
131 concrete_to: TemporalNodeKey,
132) -> bool {
133 template_from.variable == concrete_from.variable
134 && template_to.variable == concrete_to.variable
135 && concrete_from.offset.wrapping_sub(template_from.offset)
136 == concrete_to.offset.wrapping_sub(template_to.offset)
137}
138
139fn insert_replicated_edges(
140 dag: &mut Dag,
141 indexer: &TemporalIndexer,
142 from_key: TemporalNodeKey,
143 to_key: TemporalNodeKey,
144) -> Result<(), GraphError> {
145 let min_off = -(indexer.history() as i32);
146 let max_off = (indexer.horizon() as i32) - 1;
147 let lo = min_off - from_key.offset.min(to_key.offset);
150 let hi = max_off - from_key.offset.max(to_key.offset);
151 for delta in lo..=hi {
152 let a = TemporalNodeKey {
153 variable: from_key.variable,
154 offset: from_key.offset.saturating_add(delta),
155 };
156 let b = TemporalNodeKey {
157 variable: to_key.variable,
158 offset: to_key.offset.saturating_add(delta),
159 };
160 if a.offset < min_off || a.offset > max_off || b.offset < min_off || b.offset > max_off {
161 continue;
162 }
163 let from_dense = indexer.dense_id(a).map_err(|_| GraphError::InvalidEndpoints {
164 message: "unfold endpoint outside window",
165 })?;
166 let to_dense = indexer.dense_id(b).map_err(|_| GraphError::InvalidEndpoints {
167 message: "unfold endpoint outside window",
168 })?;
169 let from = DenseNodeId::from_raw(from_dense);
170 let to = DenseNodeId::from_raw(to_dense);
171 if from == to {
172 continue;
173 }
174 match dag.insert_directed(from, to) {
175 Ok(()) | Err(GraphError::DuplicateEdge { .. }) => {}
176 Err(e) => return Err(e),
177 }
178 }
179 Ok(())
180}
181
182#[derive(Clone, Debug)]
184pub struct TemporalGraphReview {
185 pub graph: TemporalDag,
187 pub pending_edges: Arc<[(TemporalNodeKey, TemporalNodeKey)]>,
189 pub algorithm: Arc<str>,
191}
192
193impl TemporalGraphReview {
194 #[must_use]
196 pub fn from_graph(graph: TemporalDag, algorithm: impl Into<Arc<str>>) -> Self {
197 let mut pending = Vec::new();
198 for (i, _) in graph.nodes().iter().enumerate() {
199 let from = DenseNodeId::try_from_usize(i).expect("node fit");
200 let Some(from_key) = graph.temporal_key(from) else {
201 continue;
202 };
203 for &to in graph.children(from) {
204 if let Some(to_key) = graph.temporal_key(to) {
205 pending.push((from_key, to_key));
206 }
207 }
208 }
209 Self { graph, pending_edges: Arc::from(pending), algorithm: algorithm.into() }
210 }
211
212 #[must_use]
214 pub fn accept_edge(mut self, from: TemporalNodeKey, to: TemporalNodeKey) -> Self {
215 let pending: Vec<_> =
216 self.pending_edges.iter().copied().filter(|e| *e != (from, to)).collect();
217 self.pending_edges = Arc::from(pending);
218 self
219 }
220
221 #[must_use]
223 pub fn is_complete(&self) -> bool {
224 self.pending_edges.is_empty()
225 }
226}
227
228#[derive(Clone, Debug)]
233pub struct TemporalCpdagReview {
234 pub graph: crate::TemporalCpdag,
236 pub pending_edges: Arc<[(TemporalNodeKey, TemporalNodeKey)]>,
238 pub pending_undirected: Arc<[(TemporalNodeKey, TemporalNodeKey)]>,
240 pub algorithm: Arc<str>,
242}
243
244impl TemporalCpdagReview {
245 #[must_use]
247 pub fn from_cpdag(graph: crate::TemporalCpdag, algorithm: impl Into<Arc<str>>) -> Self {
248 let mut pending = Vec::new();
249 let mut undirected = Vec::new();
250 for e in graph.edges() {
251 if let Some((from, to)) = e.parent_child() {
252 if let (Some(fk), Some(tk)) = (graph.temporal_key(from), graph.temporal_key(to)) {
253 pending.push((fk, tk));
254 }
255 } else if e.is_undirected() {
256 if let (Some(ak), Some(bk)) = (graph.temporal_key(e.a), graph.temporal_key(e.b)) {
257 if (ak.variable, ak.offset) <= (bk.variable, bk.offset) {
258 undirected.push((ak, bk));
259 } else {
260 undirected.push((bk, ak));
261 }
262 }
263 }
264 }
265 Self {
266 graph,
267 pending_edges: Arc::from(pending),
268 pending_undirected: Arc::from(undirected),
269 algorithm: algorithm.into(),
270 }
271 }
272
273 #[must_use]
275 pub fn accept_edge(mut self, from: TemporalNodeKey, to: TemporalNodeKey) -> Self {
276 let pending: Vec<_> =
277 self.pending_edges.iter().copied().filter(|e| *e != (from, to)).collect();
278 self.pending_edges = Arc::from(pending);
279 self
280 }
281
282 pub fn orient_edge(
288 mut self,
289 from: TemporalNodeKey,
290 to: TemporalNodeKey,
291 ) -> Result<Self, GraphError> {
292 let from_id = self.resolve_key(from)?;
293 let to_id = self.resolve_key(to)?;
294 self.graph.orient_undirected(from_id, to_id)?;
295 let undirected: Vec<_> = self
296 .pending_undirected
297 .iter()
298 .copied()
299 .filter(|&(a, b)| (a, b) != (from, to) && (a, b) != (to, from))
300 .collect();
301 self.pending_undirected = Arc::from(undirected);
302 if !self.pending_edges.iter().any(|e| *e == (from, to)) {
304 let mut pending = self.pending_edges.to_vec();
305 pending.push((from, to));
306 self.pending_edges = Arc::from(pending);
307 }
308 Ok(self)
309 }
310
311 #[must_use]
313 pub fn is_complete(&self) -> bool {
314 self.pending_edges.is_empty() && self.pending_undirected.is_empty()
315 }
316
317 pub fn try_into_temporal_dag(self) -> Result<TemporalDag, GraphError> {
323 if !self.is_complete() {
324 return Err(GraphError::InvalidEndpoints {
325 message: "TemporalCpdagReview is incomplete; accept directed and orient undirected edges first",
326 });
327 }
328 self.graph.try_into_temporal_dag()
329 }
330
331 fn resolve_key(&self, key: TemporalNodeKey) -> Result<DenseNodeId, GraphError> {
332 for i in 0..self.graph.node_count() {
333 let id = DenseNodeId::try_from_usize(i)?;
334 if self.graph.temporal_key(id) == Some(key) {
335 return Ok(id);
336 }
337 }
338 Err(GraphError::UnknownNode { id: key.variable.raw() })
339 }
340}
341
342pub fn ensure_lagged(
344 graph: &mut TemporalDag,
345 variable: VariableId,
346 lag: antecedent_core::Lag,
347) -> Result<DenseNodeId, GraphError> {
348 for (i, n) in graph.nodes().iter().enumerate() {
349 if let NodeRef::Lagged { variable: v, lag: l } = n {
350 if *v == variable && *l == lag {
351 return DenseNodeId::try_from_usize(i);
352 }
353 }
354 }
355 graph.add_lagged(variable, lag)
356}
357
358#[cfg(test)]
359#[allow(clippy::many_single_char_names)]
360mod tests {
361 use antecedent_core::TemporalIndexer;
362 use antecedent_core::{Lag, VariableId};
363
364 use super::*;
365 use crate::dsep::DSeparationWorkspace;
366
367 #[test]
368 fn lazy_has_edge_matches_materialize() {
369 let mut g = TemporalDag::empty();
370 let past = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
371 let now = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
372 g.insert_directed(past, now).unwrap();
373 let indexer = TemporalIndexer::new(2, 1, 2).unwrap();
374 let lazy = g.unfold_lazy(indexer.clone()).unwrap();
375 let from = TemporalNodeKey { variable: VariableId::from_raw(0), offset: -1 };
376 let to = TemporalNodeKey { variable: VariableId::from_raw(1), offset: 0 };
377 assert!(lazy.has_edge(from, to).unwrap());
378 let unfolded = lazy.materialize().unwrap();
379 assert_eq!(unfolded.dag.node_count(), 6);
380 }
381
382 #[test]
383 fn unfold_replicates_lagged_edge() {
384 let mut g = TemporalDag::empty();
385 let past = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
386 let now = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
387 g.insert_directed(past, now).unwrap();
388 let indexer = TemporalIndexer::new(2, 1, 2).unwrap();
389 let unfolded = g.unfold(indexer).unwrap();
390 assert_eq!(unfolded.dag.node_count(), 6);
391 let mut edge_count = 0usize;
392 for i in 0..unfolded.dag.node_count() {
393 edge_count += unfolded.dag.children(DenseNodeId::from_raw(i as u32)).len();
394 }
395 assert!(edge_count >= 1);
396 }
397
398 #[test]
399 fn unfold_dsep_on_chain() {
400 let mut g = TemporalDag::empty();
401 let x1 = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
402 let y0 = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
403 let z0 = g.add_lagged(VariableId::from_raw(2), Lag::CONTEMPORANEOUS).unwrap();
404 g.insert_directed(x1, y0).unwrap();
405 g.insert_directed(y0, z0).unwrap();
406 let indexer = TemporalIndexer::new(3, 1, 1).unwrap();
407 let unfolded = g.unfold(indexer).unwrap();
408 let mut ws = DSeparationWorkspace::default();
409 let y = DenseNodeId::from_raw(
410 unfolded
411 .indexer
412 .dense_id(TemporalNodeKey { variable: VariableId::from_raw(1), offset: 0 })
413 .unwrap(),
414 );
415 let z = DenseNodeId::from_raw(
416 unfolded
417 .indexer
418 .dense_id(TemporalNodeKey { variable: VariableId::from_raw(2), offset: 0 })
419 .unwrap(),
420 );
421 let x = DenseNodeId::from_raw(
422 unfolded
423 .indexer
424 .dense_id(TemporalNodeKey { variable: VariableId::from_raw(0), offset: -1 })
425 .unwrap(),
426 );
427 assert!(!unfolded.dag.is_d_separated(y, z, &[], &mut ws).unwrap());
430 assert!(unfolded.dag.is_d_separated(x, z, &[y], &mut ws).unwrap());
431 assert!(!unfolded.dag.is_d_separated(x, z, &[], &mut ws).unwrap());
432 }
433
434 #[test]
435 fn unfold_replicates_edge_with_both_endpoints_lagged() {
436 let mut g = TemporalDag::empty();
437 let x2 = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(2)).unwrap();
438 let y1 = g.add_lagged(VariableId::from_raw(1), Lag::from_raw(1)).unwrap();
439 g.insert_directed(x2, y1).unwrap();
440 let indexer = TemporalIndexer::new(2, 1, 1).unwrap();
442 let lazy = g.unfold_lazy(indexer.clone()).unwrap();
443 let from = TemporalNodeKey { variable: VariableId::from_raw(0), offset: -1 };
444 let to = TemporalNodeKey { variable: VariableId::from_raw(1), offset: 0 };
445 assert!(lazy.has_edge(from, to).unwrap());
446
447 let unfolded = lazy.materialize().unwrap();
448 let edges: Vec<_> = unfolded.dag.edges().collect();
449 assert_eq!(edges.len(), 1);
450 let (f, t) = edges[0].parent_child().unwrap();
451 assert_eq!(f.raw(), indexer.dense_id(from).unwrap());
452 assert_eq!(t.raw(), indexer.dense_id(to).unwrap());
453
454 for &va in &[0u32, 1] {
456 for oa in -1..=0 {
457 for &vb in &[0u32, 1] {
458 for ob in -1..=0 {
459 let a = TemporalNodeKey { variable: VariableId::from_raw(va), offset: oa };
460 let b = TemporalNodeKey { variable: VariableId::from_raw(vb), offset: ob };
461 let dense_a = DenseNodeId::from_raw(indexer.dense_id(a).unwrap());
462 let dense_b = DenseNodeId::from_raw(indexer.dense_id(b).unwrap());
463 let eager = unfolded.dag.children(dense_a).contains(&dense_b);
464 assert_eq!(lazy.has_edge(a, b).unwrap(), eager, "{a:?} -> {b:?}");
465 }
466 }
467 }
468 }
469 }
470
471 #[test]
472 fn review_accept_clears_pending() {
473 let mut g = TemporalDag::empty();
474 let a = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
475 let b = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
476 g.insert_directed(a, b).unwrap();
477 let review = TemporalGraphReview::from_graph(g, "pcmci");
478 assert!(!review.is_complete());
479 let from = TemporalNodeKey { variable: VariableId::from_raw(0), offset: -1 };
480 let to = TemporalNodeKey { variable: VariableId::from_raw(1), offset: 0 };
481 let review = review.accept_edge(from, to);
482 assert!(review.is_complete());
483 }
484
485 fn random_small_template(rng: &mut antecedent_core::CausalRng) -> (TemporalDag, u32, u32, u32) {
486 let n_vars = 2 + (rng.next_u64() % 2) as u32; let max_lag = 1 + (rng.next_u64() % 2) as u32; let history = max_lag;
489 let horizon = 1 + (rng.next_u64() % 2) as u32; let mut g = TemporalDag::empty();
492 let mut ids = Vec::new();
493 for v in 0..n_vars {
494 for lag in 0..=max_lag {
495 let id = g.add_lagged(VariableId::from_raw(v), Lag::from_raw(lag)).unwrap();
496 ids.push((v, lag, id));
497 }
498 }
499 for &(va, la, a) in &ids {
502 for &(vb, lb, b) in &ids {
503 if a == b {
504 continue;
505 }
506 let time_ok = la > lb || (la == lb && va < vb);
507 if !time_ok {
508 continue;
509 }
510 if rng.next_u64() % 3 == 0 {
511 let _ = g.insert_directed(a, b);
512 }
513 }
514 }
515 (g, n_vars, history, horizon)
516 }
517
518 fn dag_from_lazy_scan(lazy: &LazyUnfoldedTemporalGraph) -> Dag {
520 let n = lazy.indexer.dense_len();
521 let mut dag = Dag::with_variables(u32::try_from(n).unwrap());
522 let min_off = -(lazy.indexer.history() as i32);
523 let max_off = (lazy.indexer.horizon() as i32) - 1;
524 let n_vars = lazy.indexer.variable_count();
525 for va in 0..n_vars {
526 for oa in min_off..=max_off {
527 for vb in 0..n_vars {
528 for ob in min_off..=max_off {
529 let from =
530 TemporalNodeKey { variable: VariableId::from_raw(va), offset: oa };
531 let to = TemporalNodeKey { variable: VariableId::from_raw(vb), offset: ob };
532 if !lazy.has_edge(from, to).unwrap() {
533 continue;
534 }
535 let a = DenseNodeId::from_raw(lazy.indexer.dense_id(from).unwrap());
536 let b = DenseNodeId::from_raw(lazy.indexer.dense_id(to).unwrap());
537 if a != b {
538 let _ = dag.insert_directed(a, b);
539 }
540 }
541 }
542 }
543 }
544 dag
545 }
546
547 #[test]
549 fn property_lazy_unfold_matches_materialize_on_small_templates() {
550 use antecedent_core::CausalRng;
551
552 let mut rng = CausalRng::from_seed(91);
553 for _ in 0..50 {
554 let (g, n_vars, history, horizon) = random_small_template(&mut rng);
555 let indexer = TemporalIndexer::new(n_vars, history, horizon).unwrap();
556 let lazy = g.unfold_lazy(indexer.clone()).unwrap();
557 let unfolded = lazy.materialize().unwrap();
558 let min_off = -(history as i32);
559 let max_off = (horizon as i32) - 1;
560 for va in 0..n_vars {
561 for oa in min_off..=max_off {
562 for vb in 0..n_vars {
563 for ob in min_off..=max_off {
564 let from =
565 TemporalNodeKey { variable: VariableId::from_raw(va), offset: oa };
566 let to =
567 TemporalNodeKey { variable: VariableId::from_raw(vb), offset: ob };
568 let dense_a = DenseNodeId::from_raw(indexer.dense_id(from).unwrap());
569 let dense_b = DenseNodeId::from_raw(indexer.dense_id(to).unwrap());
570 let eager = unfolded.dag.children(dense_a).contains(&dense_b);
571 assert_eq!(
572 lazy.has_edge(from, to).unwrap(),
573 eager,
574 "lazy≠materialize {from:?}->{to:?}"
575 );
576 }
577 }
578 }
579 }
580 }
581 }
582
583 #[test]
585 fn property_unfolded_dsep_lazy_scan_matches_materialize() {
586 use antecedent_core::CausalRng;
587
588 let mut rng = CausalRng::from_seed(113);
589 let mut ws = DSeparationWorkspace::default();
590 for _ in 0..30 {
591 let (g, n_vars, history, horizon) = random_small_template(&mut rng);
592 let indexer = TemporalIndexer::new(n_vars, history, horizon).unwrap();
593 let lazy = g.unfold_lazy(indexer).unwrap();
594 let unfolded = lazy.materialize().unwrap();
595 let scanned = dag_from_lazy_scan(&lazy);
596 let n = unfolded.dag.node_count() as u32;
597 assert_eq!(scanned.node_count(), unfolded.dag.node_count());
598 for i in 0..n {
599 let u = DenseNodeId::from_raw(i);
600 let mut a = scanned.children(u).to_vec();
601 let mut b = unfolded.dag.children(u).to_vec();
602 a.sort_by_key(|x| x.raw());
603 b.sort_by_key(|x| x.raw());
604 assert_eq!(a, b, "lazy-scan adjacency ≠ materialize at {i}");
605 }
606 for _ in 0..10 {
607 let x = DenseNodeId::from_raw(rng.next_u64() as u32 % n);
608 let mut y = DenseNodeId::from_raw(rng.next_u64() as u32 % n);
609 while y == x {
610 y = DenseNodeId::from_raw(rng.next_u64() as u32 % n);
611 }
612 let mut z = Vec::new();
613 for i in 0..n {
614 let v = DenseNodeId::from_raw(i);
615 if v == x || v == y {
616 continue;
617 }
618 if rng.next_u64() % 3 == 0 {
619 z.push(v);
620 }
621 }
622 let lazy_sep = scanned.is_d_separated(x, y, &z, &mut ws).unwrap();
623 let mat_sep = unfolded.dag.is_d_separated(x, y, &z, &mut ws).unwrap();
624 assert_eq!(
625 lazy_sep,
626 mat_sep,
627 "unfolded d-sep mismatch x={} y={}",
628 x.raw(),
629 y.raw()
630 );
631 }
632 }
633 }
634}