#![allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)]
use std::sync::Arc;
use antecedent_core::VariableId;
use antecedent_core::{TemporalIndexer, TemporalNodeKey};
use crate::dag::Dag;
use crate::error::GraphError;
use crate::temporal::TemporalDag;
use crate::types::{DenseNodeId, NodeRef};
#[derive(Clone, Debug)]
pub struct UnfoldedTemporalGraph {
pub dag: Dag,
pub indexer: TemporalIndexer,
}
#[derive(Clone, Debug)]
pub struct LazyUnfoldedTemporalGraph {
pub template: TemporalDag,
pub indexer: TemporalIndexer,
}
impl TemporalDag {
pub fn unfold_lazy(
&self,
indexer: TemporalIndexer,
) -> Result<LazyUnfoldedTemporalGraph, GraphError> {
for (i, _) in self.nodes().iter().enumerate() {
let id = DenseNodeId::try_from_usize(i)?;
let _ = self
.temporal_key(id)
.ok_or(GraphError::InvalidEndpoints { message: "unfold requires lagged nodes" })?;
}
Ok(LazyUnfoldedTemporalGraph { template: self.clone(), indexer })
}
pub fn unfold(&self, indexer: TemporalIndexer) -> Result<UnfoldedTemporalGraph, GraphError> {
self.unfold_lazy(indexer)?.materialize()
}
}
impl LazyUnfoldedTemporalGraph {
pub fn has_edge(&self, from: TemporalNodeKey, to: TemporalNodeKey) -> Result<bool, GraphError> {
let _ = self.indexer.dense_id(from).map_err(|_| GraphError::InvalidEndpoints {
message: "unfold endpoint outside window",
})?;
let _ = self.indexer.dense_id(to).map_err(|_| GraphError::InvalidEndpoints {
message: "unfold endpoint outside window",
})?;
for (from_i, _) in self.template.nodes().iter().enumerate() {
let from_id = DenseNodeId::try_from_usize(from_i)?;
let Some(from_key) = self.template.temporal_key(from_id) else {
continue;
};
for &to_id in self.template.children(from_id) {
let Some(to_key) = self.template.temporal_key(to_id) else {
continue;
};
if edge_matches(from_key, to_key, from, to) {
return Ok(true);
}
}
}
Ok(false)
}
pub fn materialize(&self) -> Result<UnfoldedTemporalGraph, GraphError> {
let n = self.indexer.dense_len();
let n_u32 = u32::try_from(n).map_err(|_| GraphError::TooManyNodes)?;
let mut dag = Dag::with_variables(n_u32);
for (from_i, _) in self.template.nodes().iter().enumerate() {
let from = DenseNodeId::try_from_usize(from_i)?;
let from_key = self
.template
.temporal_key(from)
.ok_or(GraphError::InvalidEndpoints { message: "unfold requires lagged nodes" })?;
for &to in self.template.children(from) {
let to_key =
self.template.temporal_key(to).ok_or(GraphError::InvalidEndpoints {
message: "unfold requires lagged nodes",
})?;
insert_replicated_edges(&mut dag, &self.indexer, from_key, to_key)?;
}
}
Ok(UnfoldedTemporalGraph { dag, indexer: self.indexer.clone() })
}
}
fn edge_matches(
template_from: TemporalNodeKey,
template_to: TemporalNodeKey,
concrete_from: TemporalNodeKey,
concrete_to: TemporalNodeKey,
) -> bool {
template_from.variable == concrete_from.variable
&& template_to.variable == concrete_to.variable
&& concrete_from.offset.wrapping_sub(template_from.offset)
== concrete_to.offset.wrapping_sub(template_to.offset)
}
fn insert_replicated_edges(
dag: &mut Dag,
indexer: &TemporalIndexer,
from_key: TemporalNodeKey,
to_key: TemporalNodeKey,
) -> Result<(), GraphError> {
let min_off = -(indexer.history() as i32);
let max_off = (indexer.horizon() as i32) - 1;
let lo = min_off - from_key.offset.min(to_key.offset);
let hi = max_off - from_key.offset.max(to_key.offset);
for delta in lo..=hi {
let a = TemporalNodeKey {
variable: from_key.variable,
offset: from_key.offset.saturating_add(delta),
};
let b = TemporalNodeKey {
variable: to_key.variable,
offset: to_key.offset.saturating_add(delta),
};
if a.offset < min_off || a.offset > max_off || b.offset < min_off || b.offset > max_off {
continue;
}
let from_dense = indexer.dense_id(a).map_err(|_| GraphError::InvalidEndpoints {
message: "unfold endpoint outside window",
})?;
let to_dense = indexer.dense_id(b).map_err(|_| GraphError::InvalidEndpoints {
message: "unfold endpoint outside window",
})?;
let from = DenseNodeId::from_raw(from_dense);
let to = DenseNodeId::from_raw(to_dense);
if from == to {
continue;
}
match dag.insert_directed(from, to) {
Ok(()) | Err(GraphError::DuplicateEdge { .. }) => {}
Err(e) => return Err(e),
}
}
Ok(())
}
#[derive(Clone, Debug)]
pub struct TemporalGraphReview {
pub graph: TemporalDag,
pub pending_edges: Arc<[(TemporalNodeKey, TemporalNodeKey)]>,
pub algorithm: Arc<str>,
}
impl TemporalGraphReview {
#[must_use]
pub fn from_graph(graph: TemporalDag, algorithm: impl Into<Arc<str>>) -> Self {
let mut pending = Vec::new();
for (i, _) in graph.nodes().iter().enumerate() {
let from = DenseNodeId::try_from_usize(i).expect("node fit");
let Some(from_key) = graph.temporal_key(from) else {
continue;
};
for &to in graph.children(from) {
if let Some(to_key) = graph.temporal_key(to) {
pending.push((from_key, to_key));
}
}
}
Self { graph, pending_edges: Arc::from(pending), algorithm: algorithm.into() }
}
#[must_use]
pub fn accept_edge(mut self, from: TemporalNodeKey, to: TemporalNodeKey) -> Self {
let pending: Vec<_> =
self.pending_edges.iter().copied().filter(|e| *e != (from, to)).collect();
self.pending_edges = Arc::from(pending);
self
}
#[must_use]
pub fn is_complete(&self) -> bool {
self.pending_edges.is_empty()
}
}
#[derive(Clone, Debug)]
pub struct TemporalCpdagReview {
pub graph: crate::TemporalCpdag,
pub pending_edges: Arc<[(TemporalNodeKey, TemporalNodeKey)]>,
pub pending_undirected: Arc<[(TemporalNodeKey, TemporalNodeKey)]>,
pub algorithm: Arc<str>,
}
impl TemporalCpdagReview {
#[must_use]
pub fn from_cpdag(graph: crate::TemporalCpdag, algorithm: impl Into<Arc<str>>) -> Self {
let mut pending = Vec::new();
let mut undirected = Vec::new();
for e in graph.edges() {
if let Some((from, to)) = e.parent_child() {
if let (Some(fk), Some(tk)) = (graph.temporal_key(from), graph.temporal_key(to)) {
pending.push((fk, tk));
}
} else if e.is_undirected() {
if let (Some(ak), Some(bk)) = (graph.temporal_key(e.a), graph.temporal_key(e.b)) {
if (ak.variable, ak.offset) <= (bk.variable, bk.offset) {
undirected.push((ak, bk));
} else {
undirected.push((bk, ak));
}
}
}
}
Self {
graph,
pending_edges: Arc::from(pending),
pending_undirected: Arc::from(undirected),
algorithm: algorithm.into(),
}
}
#[must_use]
pub fn accept_edge(mut self, from: TemporalNodeKey, to: TemporalNodeKey) -> Self {
let pending: Vec<_> =
self.pending_edges.iter().copied().filter(|e| *e != (from, to)).collect();
self.pending_edges = Arc::from(pending);
self
}
pub fn orient_edge(
mut self,
from: TemporalNodeKey,
to: TemporalNodeKey,
) -> Result<Self, GraphError> {
let from_id = self.resolve_key(from)?;
let to_id = self.resolve_key(to)?;
self.graph.orient_undirected(from_id, to_id)?;
let undirected: Vec<_> = self
.pending_undirected
.iter()
.copied()
.filter(|&(a, b)| (a, b) != (from, to) && (a, b) != (to, from))
.collect();
self.pending_undirected = Arc::from(undirected);
if !self.pending_edges.iter().any(|e| *e == (from, to)) {
let mut pending = self.pending_edges.to_vec();
pending.push((from, to));
self.pending_edges = Arc::from(pending);
}
Ok(self)
}
#[must_use]
pub fn is_complete(&self) -> bool {
self.pending_edges.is_empty() && self.pending_undirected.is_empty()
}
pub fn try_into_temporal_dag(self) -> Result<TemporalDag, GraphError> {
if !self.is_complete() {
return Err(GraphError::InvalidEndpoints {
message: "TemporalCpdagReview is incomplete; accept directed and orient undirected edges first",
});
}
self.graph.try_into_temporal_dag()
}
fn resolve_key(&self, key: TemporalNodeKey) -> Result<DenseNodeId, GraphError> {
for i in 0..self.graph.node_count() {
let id = DenseNodeId::try_from_usize(i)?;
if self.graph.temporal_key(id) == Some(key) {
return Ok(id);
}
}
Err(GraphError::UnknownNode { id: key.variable.raw() })
}
}
pub fn ensure_lagged(
graph: &mut TemporalDag,
variable: VariableId,
lag: antecedent_core::Lag,
) -> Result<DenseNodeId, GraphError> {
for (i, n) in graph.nodes().iter().enumerate() {
if let NodeRef::Lagged { variable: v, lag: l } = n {
if *v == variable && *l == lag {
return DenseNodeId::try_from_usize(i);
}
}
}
graph.add_lagged(variable, lag)
}
#[cfg(test)]
#[allow(clippy::many_single_char_names)]
mod tests {
use antecedent_core::TemporalIndexer;
use antecedent_core::{Lag, VariableId};
use super::*;
use crate::dsep::DSeparationWorkspace;
#[test]
fn lazy_has_edge_matches_materialize() {
let mut g = TemporalDag::empty();
let past = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
let now = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
g.insert_directed(past, now).unwrap();
let indexer = TemporalIndexer::new(2, 1, 2).unwrap();
let lazy = g.unfold_lazy(indexer.clone()).unwrap();
let from = TemporalNodeKey { variable: VariableId::from_raw(0), offset: -1 };
let to = TemporalNodeKey { variable: VariableId::from_raw(1), offset: 0 };
assert!(lazy.has_edge(from, to).unwrap());
let unfolded = lazy.materialize().unwrap();
assert_eq!(unfolded.dag.node_count(), 6);
}
#[test]
fn unfold_replicates_lagged_edge() {
let mut g = TemporalDag::empty();
let past = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
let now = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
g.insert_directed(past, now).unwrap();
let indexer = TemporalIndexer::new(2, 1, 2).unwrap();
let unfolded = g.unfold(indexer).unwrap();
assert_eq!(unfolded.dag.node_count(), 6);
let mut edge_count = 0usize;
for i in 0..unfolded.dag.node_count() {
edge_count += unfolded.dag.children(DenseNodeId::from_raw(i as u32)).len();
}
assert!(edge_count >= 1);
}
#[test]
fn unfold_dsep_on_chain() {
let mut g = TemporalDag::empty();
let x1 = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
let y0 = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
let z0 = g.add_lagged(VariableId::from_raw(2), Lag::CONTEMPORANEOUS).unwrap();
g.insert_directed(x1, y0).unwrap();
g.insert_directed(y0, z0).unwrap();
let indexer = TemporalIndexer::new(3, 1, 1).unwrap();
let unfolded = g.unfold(indexer).unwrap();
let mut ws = DSeparationWorkspace::default();
let y = DenseNodeId::from_raw(
unfolded
.indexer
.dense_id(TemporalNodeKey { variable: VariableId::from_raw(1), offset: 0 })
.unwrap(),
);
let z = DenseNodeId::from_raw(
unfolded
.indexer
.dense_id(TemporalNodeKey { variable: VariableId::from_raw(2), offset: 0 })
.unwrap(),
);
let x = DenseNodeId::from_raw(
unfolded
.indexer
.dense_id(TemporalNodeKey { variable: VariableId::from_raw(0), offset: -1 })
.unwrap(),
);
assert!(!unfolded.dag.is_d_separated(y, z, &[], &mut ws).unwrap());
assert!(unfolded.dag.is_d_separated(x, z, &[y], &mut ws).unwrap());
assert!(!unfolded.dag.is_d_separated(x, z, &[], &mut ws).unwrap());
}
#[test]
fn unfold_replicates_edge_with_both_endpoints_lagged() {
let mut g = TemporalDag::empty();
let x2 = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(2)).unwrap();
let y1 = g.add_lagged(VariableId::from_raw(1), Lag::from_raw(1)).unwrap();
g.insert_directed(x2, y1).unwrap();
let indexer = TemporalIndexer::new(2, 1, 1).unwrap();
let lazy = g.unfold_lazy(indexer.clone()).unwrap();
let from = TemporalNodeKey { variable: VariableId::from_raw(0), offset: -1 };
let to = TemporalNodeKey { variable: VariableId::from_raw(1), offset: 0 };
assert!(lazy.has_edge(from, to).unwrap());
let unfolded = lazy.materialize().unwrap();
let edges: Vec<_> = unfolded.dag.edges().collect();
assert_eq!(edges.len(), 1);
let (f, t) = edges[0].parent_child().unwrap();
assert_eq!(f.raw(), indexer.dense_id(from).unwrap());
assert_eq!(t.raw(), indexer.dense_id(to).unwrap());
for &va in &[0u32, 1] {
for oa in -1..=0 {
for &vb in &[0u32, 1] {
for ob in -1..=0 {
let a = TemporalNodeKey { variable: VariableId::from_raw(va), offset: oa };
let b = TemporalNodeKey { variable: VariableId::from_raw(vb), offset: ob };
let dense_a = DenseNodeId::from_raw(indexer.dense_id(a).unwrap());
let dense_b = DenseNodeId::from_raw(indexer.dense_id(b).unwrap());
let eager = unfolded.dag.children(dense_a).contains(&dense_b);
assert_eq!(lazy.has_edge(a, b).unwrap(), eager, "{a:?} -> {b:?}");
}
}
}
}
}
#[test]
fn review_accept_clears_pending() {
let mut g = TemporalDag::empty();
let a = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
let b = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
g.insert_directed(a, b).unwrap();
let review = TemporalGraphReview::from_graph(g, "pcmci");
assert!(!review.is_complete());
let from = TemporalNodeKey { variable: VariableId::from_raw(0), offset: -1 };
let to = TemporalNodeKey { variable: VariableId::from_raw(1), offset: 0 };
let review = review.accept_edge(from, to);
assert!(review.is_complete());
}
fn random_small_template(rng: &mut antecedent_core::CausalRng) -> (TemporalDag, u32, u32, u32) {
let n_vars = 2 + (rng.next_u64() % 2) as u32; let max_lag = 1 + (rng.next_u64() % 2) as u32; let history = max_lag;
let horizon = 1 + (rng.next_u64() % 2) as u32;
let mut g = TemporalDag::empty();
let mut ids = Vec::new();
for v in 0..n_vars {
for lag in 0..=max_lag {
let id = g.add_lagged(VariableId::from_raw(v), Lag::from_raw(lag)).unwrap();
ids.push((v, lag, id));
}
}
for &(va, la, a) in &ids {
for &(vb, lb, b) in &ids {
if a == b {
continue;
}
let time_ok = la > lb || (la == lb && va < vb);
if !time_ok {
continue;
}
if rng.next_u64() % 3 == 0 {
let _ = g.insert_directed(a, b);
}
}
}
(g, n_vars, history, horizon)
}
fn dag_from_lazy_scan(lazy: &LazyUnfoldedTemporalGraph) -> Dag {
let n = lazy.indexer.dense_len();
let mut dag = Dag::with_variables(u32::try_from(n).unwrap());
let min_off = -(lazy.indexer.history() as i32);
let max_off = (lazy.indexer.horizon() as i32) - 1;
let n_vars = lazy.indexer.variable_count();
for va in 0..n_vars {
for oa in min_off..=max_off {
for vb in 0..n_vars {
for ob in min_off..=max_off {
let from =
TemporalNodeKey { variable: VariableId::from_raw(va), offset: oa };
let to = TemporalNodeKey { variable: VariableId::from_raw(vb), offset: ob };
if !lazy.has_edge(from, to).unwrap() {
continue;
}
let a = DenseNodeId::from_raw(lazy.indexer.dense_id(from).unwrap());
let b = DenseNodeId::from_raw(lazy.indexer.dense_id(to).unwrap());
if a != b {
let _ = dag.insert_directed(a, b);
}
}
}
}
}
dag
}
#[test]
fn property_lazy_unfold_matches_materialize_on_small_templates() {
use antecedent_core::CausalRng;
let mut rng = CausalRng::from_seed(91);
for _ in 0..50 {
let (g, n_vars, history, horizon) = random_small_template(&mut rng);
let indexer = TemporalIndexer::new(n_vars, history, horizon).unwrap();
let lazy = g.unfold_lazy(indexer.clone()).unwrap();
let unfolded = lazy.materialize().unwrap();
let min_off = -(history as i32);
let max_off = (horizon as i32) - 1;
for va in 0..n_vars {
for oa in min_off..=max_off {
for vb in 0..n_vars {
for ob in min_off..=max_off {
let from =
TemporalNodeKey { variable: VariableId::from_raw(va), offset: oa };
let to =
TemporalNodeKey { variable: VariableId::from_raw(vb), offset: ob };
let dense_a = DenseNodeId::from_raw(indexer.dense_id(from).unwrap());
let dense_b = DenseNodeId::from_raw(indexer.dense_id(to).unwrap());
let eager = unfolded.dag.children(dense_a).contains(&dense_b);
assert_eq!(
lazy.has_edge(from, to).unwrap(),
eager,
"lazy≠materialize {from:?}->{to:?}"
);
}
}
}
}
}
}
#[test]
fn property_unfolded_dsep_lazy_scan_matches_materialize() {
use antecedent_core::CausalRng;
let mut rng = CausalRng::from_seed(113);
let mut ws = DSeparationWorkspace::default();
for _ in 0..30 {
let (g, n_vars, history, horizon) = random_small_template(&mut rng);
let indexer = TemporalIndexer::new(n_vars, history, horizon).unwrap();
let lazy = g.unfold_lazy(indexer).unwrap();
let unfolded = lazy.materialize().unwrap();
let scanned = dag_from_lazy_scan(&lazy);
let n = unfolded.dag.node_count() as u32;
assert_eq!(scanned.node_count(), unfolded.dag.node_count());
for i in 0..n {
let u = DenseNodeId::from_raw(i);
let mut a = scanned.children(u).to_vec();
let mut b = unfolded.dag.children(u).to_vec();
a.sort_by_key(|x| x.raw());
b.sort_by_key(|x| x.raw());
assert_eq!(a, b, "lazy-scan adjacency ≠ materialize at {i}");
}
for _ in 0..10 {
let x = DenseNodeId::from_raw(rng.next_u64() as u32 % n);
let mut y = DenseNodeId::from_raw(rng.next_u64() as u32 % n);
while y == x {
y = DenseNodeId::from_raw(rng.next_u64() as u32 % n);
}
let mut z = Vec::new();
for i in 0..n {
let v = DenseNodeId::from_raw(i);
if v == x || v == y {
continue;
}
if rng.next_u64() % 3 == 0 {
z.push(v);
}
}
let lazy_sep = scanned.is_d_separated(x, y, &z, &mut ws).unwrap();
let mat_sep = unfolded.dag.is_d_separated(x, y, &z, &mut ws).unwrap();
assert_eq!(
lazy_sep,
mat_sep,
"unfolded d-sep mismatch x={} y={}",
x.raw(),
y.raw()
);
}
}
}
}