Skip to main content

antecedent_graph/
cpdag_completion.rs

1//! Streamed / bounded CPDAG MEC completion sampling.
2//!
3//! Yields DAG members of the Markov equivalence class of a static [`Cpdag`]:
4//! orientations of undirected (Tail–Tail) edges that remain acyclic and do not
5//! introduce new unshielded colliders. Completions are never retained without
6//! bound (`max_completions` caps **valid** yields). Conflict (`x-x`) edges are
7//! refused at construction.
8//!
9//! SPDX-License-Identifier: MIT OR Apache-2.0
10
11use crate::cpdag::Cpdag;
12use crate::dag::Dag;
13use crate::error::GraphError;
14use crate::types::DenseNodeId;
15
16/// One DAG completion of a CPDAG (MEC member).
17#[derive(Clone, Debug)]
18pub struct CpdagCompletion {
19    /// Completed DAG.
20    pub graph: Dag,
21    /// Index of this completion in the stream (0-based among valid yields).
22    pub index: usize,
23}
24
25/// Streams CPDAG → DAG completions with a hard cap (no unbounded retain).
26#[derive(Clone, Debug)]
27pub struct CpdagCompletionSampler {
28    base: Cpdag,
29    /// Undirected edges `(a, b)` with `a.raw() < b.raw()`.
30    undirected: Vec<(DenseNodeId, DenseNodeId)>,
31    max_completions: usize,
32    next_index: usize,
33    /// Bitmask: bit i = 0 → orient a→b, bit i = 1 → orient b→a.
34    assign: u64,
35}
36
37impl CpdagCompletionSampler {
38    /// Build a sampler that yields at most `max_completions` **valid** MEC DAGs.
39    ///
40    /// # Errors
41    ///
42    /// Conflict edges present, or more than 63 undirected edges (mask capacity).
43    pub fn new(cpdag: Cpdag, max_completions: usize) -> Result<Self, GraphError> {
44        if cpdag.conflict_edge_count() > 0 {
45            return Err(GraphError::InvalidEndpoints {
46                message: "CpdagCompletionSampler refuses conflict (x-x) edges",
47            });
48        }
49        let mut undirected = Vec::new();
50        for e in cpdag.edges() {
51            if e.is_undirected() {
52                let (a, b) = if e.a.raw() <= e.b.raw() { (e.a, e.b) } else { (e.b, e.a) };
53                undirected.push((a, b));
54            }
55        }
56        undirected.sort_by_key(|(a, b)| (a.raw(), b.raw()));
57        undirected.dedup();
58        if undirected.len() > 63 {
59            return Err(GraphError::InvalidEndpoints {
60                message: "too many undirected edges for CpdagCompletionSampler mask",
61            });
62        }
63        Ok(Self { base: cpdag, undirected, max_completions, next_index: 0, assign: 0 })
64    }
65
66    /// Hard cap on yielded valid completions.
67    #[must_use]
68    pub fn max_completions(&self) -> usize {
69        self.max_completions
70    }
71
72    /// Number of undirected edges being oriented.
73    #[must_use]
74    pub fn n_undirected(&self) -> usize {
75        self.undirected.len()
76    }
77
78    fn build_completion(&self, mask: u64) -> Option<Dag> {
79        let mut g = self.base.clone();
80        for (i, &(a, b)) in self.undirected.iter().enumerate() {
81            let reverse = ((mask >> i) & 1) == 1;
82            let (from, to) = if reverse { (b, a) } else { (a, b) };
83            if g.orient_undirected(from, to).is_err() {
84                return None;
85            }
86        }
87        let dag = g.try_into_dag().ok()?;
88        if is_mec_member(&self.base, &dag) { Some(dag) } else { None }
89    }
90}
91
92/// Whether `dag` is a Markov-equivalence member of `cpdag` (same skeleton, same
93/// unshielded colliders, all compelled directed edges of the CPDAG present).
94#[must_use]
95pub fn is_mec_member(cpdag: &Cpdag, dag: &Dag) -> bool {
96    if cpdag.node_count() != dag.node_count() {
97        return false;
98    }
99    // Compelled directed edges must appear.
100    for e in cpdag.edges() {
101        if let Some((from, to)) = e.parent_child() {
102            if !dag.children(from).contains(&to) {
103                return false;
104            }
105        } else if e.is_undirected() {
106            let a = e.a;
107            let b = e.b;
108            let ab = dag.children(a).contains(&b);
109            let ba = dag.children(b).contains(&a);
110            if ab == ba {
111                // missing or both — not a simple orientation
112                return false;
113            }
114        } else if e.is_conflict() {
115            return false;
116        }
117    }
118    // Skeleton: every DAG edge must exist in the CPDAG (any mark).
119    for e in dag.edges() {
120        if let Some((from, to)) = e.parent_child() {
121            if !cpdag.has_edge(from, to) {
122                return false;
123            }
124        }
125    }
126    // Unshielded colliders must match.
127    let cpdag_colliders = unshielded_colliders_cpdag(cpdag);
128    let dag_colliders = unshielded_colliders_dag(dag);
129    cpdag_colliders == dag_colliders
130}
131
132fn unshielded_colliders_cpdag(g: &Cpdag) -> Vec<(u32, u32, u32)> {
133    let n = g.node_count();
134    let mut out = Vec::new();
135    for center_idx in 0..n {
136        let Ok(center_raw) = u32::try_from(center_idx) else {
137            break;
138        };
139        let center = DenseNodeId::from_raw(center_raw);
140        let parents = g.parents(center);
141        for left_i in 0..parents.len() {
142            for right_i in (left_i + 1)..parents.len() {
143                let left_parent = parents[left_i];
144                let right_parent = parents[right_i];
145                if !g.has_edge(left_parent, right_parent) {
146                    let (lo, hi) = if left_parent.raw() <= right_parent.raw() {
147                        (left_parent.raw(), right_parent.raw())
148                    } else {
149                        (right_parent.raw(), left_parent.raw())
150                    };
151                    out.push((lo, center.raw(), hi));
152                }
153            }
154        }
155    }
156    out.sort_unstable();
157    out.dedup();
158    out
159}
160
161fn unshielded_colliders_dag(g: &Dag) -> Vec<(u32, u32, u32)> {
162    let n = g.node_count();
163    let mut out = Vec::new();
164    for center_idx in 0..n {
165        let Ok(center_raw) = u32::try_from(center_idx) else {
166            break;
167        };
168        let center = DenseNodeId::from_raw(center_raw);
169        let parents = g.parents(center);
170        for left_i in 0..parents.len() {
171            for right_i in (left_i + 1)..parents.len() {
172                let left_parent = parents[left_i];
173                let right_parent = parents[right_i];
174                let adjacent = g.children(left_parent).contains(&right_parent)
175                    || g.children(right_parent).contains(&left_parent);
176                if !adjacent {
177                    let (lo, hi) = if left_parent.raw() <= right_parent.raw() {
178                        (left_parent.raw(), right_parent.raw())
179                    } else {
180                        (right_parent.raw(), left_parent.raw())
181                    };
182                    out.push((lo, center.raw(), hi));
183                }
184            }
185        }
186    }
187    out.sort_unstable();
188    out.dedup();
189    out
190}
191
192impl Iterator for CpdagCompletionSampler {
193    type Item = CpdagCompletion;
194
195    fn next(&mut self) -> Option<Self::Item> {
196        if self.next_index >= self.max_completions {
197            return None;
198        }
199        let n = self.undirected.len();
200        let total = if n == 0 { 1u64 } else { 1u64 << n };
201        while self.assign < total {
202            let mask = self.assign;
203            self.assign += 1;
204            if let Some(graph) = self.build_completion(mask) {
205                let index = self.next_index;
206                self.next_index += 1;
207                return Some(CpdagCompletion { graph, index });
208            }
209        }
210        None
211    }
212}
213
214#[cfg(test)]
215mod tests {
216    use super::*;
217    use crate::cpdag::Cpdag;
218
219    #[test]
220    fn fully_oriented_yields_one() {
221        let mut g = Cpdag::with_variables(2);
222        g.insert_directed(DenseNodeId::from_raw(0), DenseNodeId::from_raw(1)).unwrap();
223        let collected: Vec<_> = CpdagCompletionSampler::new(g, 10).unwrap().collect();
224        assert_eq!(collected.len(), 1);
225        assert!(
226            collected[0]
227                .graph
228                .children(DenseNodeId::from_raw(0))
229                .contains(&DenseNodeId::from_raw(1))
230        );
231    }
232
233    #[test]
234    fn chain_undirected_has_three_mec_dags() {
235        // A — B — C: MEC has A→B→C, A←B←C, and A←B→C — not A→B←C (new v-structure).
236        let mut g = Cpdag::with_variables(3);
237        let a = DenseNodeId::from_raw(0);
238        let b = DenseNodeId::from_raw(1);
239        let c = DenseNodeId::from_raw(2);
240        g.insert_undirected(a, b).unwrap();
241        g.insert_undirected(b, c).unwrap();
242        let collected: Vec<_> = CpdagCompletionSampler::new(g.clone(), 16).unwrap().collect();
243        assert_eq!(g.undirected_edge_count(), 2);
244        assert_eq!(collected.len(), 3, "expected 3 MEC DAGs, got {}", collected.len());
245        for c in &collected {
246            assert!(is_mec_member(&g, &c.graph));
247        }
248    }
249
250    #[test]
251    fn respects_max_completions_bound() {
252        let mut g = Cpdag::with_variables(3);
253        g.insert_undirected(DenseNodeId::from_raw(0), DenseNodeId::from_raw(1)).unwrap();
254        g.insert_undirected(DenseNodeId::from_raw(1), DenseNodeId::from_raw(2)).unwrap();
255        let collected: Vec<_> = CpdagCompletionSampler::new(g, 2).unwrap().collect();
256        assert_eq!(collected.len(), 2);
257    }
258
259    #[test]
260    fn refuses_conflict_edges() {
261        let mut g = Cpdag::with_variables(2);
262        g.insert_undirected(DenseNodeId::from_raw(0), DenseNodeId::from_raw(1)).unwrap();
263        g.mark_conflict(DenseNodeId::from_raw(0), DenseNodeId::from_raw(1)).unwrap();
264        assert!(CpdagCompletionSampler::new(g, 4).is_err());
265    }
266
267    #[test]
268    fn compelled_collider_preserved() {
269        // A → B ← C with A—C absent: classic v-structure CPDAG (A—C may be absent).
270        let mut g = Cpdag::with_variables(3);
271        let a = DenseNodeId::from_raw(0);
272        let b = DenseNodeId::from_raw(1);
273        let c = DenseNodeId::from_raw(2);
274        g.insert_directed(a, b).unwrap();
275        g.insert_directed(c, b).unwrap();
276        let collected: Vec<_> = CpdagCompletionSampler::new(g.clone(), 8).unwrap().collect();
277        assert_eq!(collected.len(), 1);
278        assert!(is_mec_member(&g, &collected[0].graph));
279    }
280}