use crate::dag::Dag;
use crate::dsep::{DSeparationWorkspace, SeparationResult};
use crate::error::GraphError;
use crate::types::DenseNodeId;
use crate::workspace::{BitSet, GraphWorkspace};
#[derive(Clone, Debug)]
pub struct GraphOverlay {
hide_incoming: BitSet,
hide_outgoing: BitSet,
}
impl GraphOverlay {
#[must_use]
pub fn observational(n: usize) -> Self {
Self { hide_incoming: BitSet::with_len(n), hide_outgoing: BitSet::with_len(n) }
}
#[must_use]
pub fn do_intervention(n: usize, intervened: &[DenseNodeId]) -> Self {
let mut overlay = Self::observational(n);
for &v in intervened {
if v.as_usize() < n {
overlay.hide_incoming.insert(v);
}
}
overlay
}
#[must_use]
pub fn remove_outgoing(n: usize, nodes: &[DenseNodeId]) -> Self {
let mut overlay = Self::observational(n);
for &v in nodes {
if v.as_usize() < n {
overlay.hide_outgoing.insert(v);
}
}
overlay
}
#[must_use]
pub fn is_observational(&self) -> bool {
!self.hide_incoming.any() && !self.hide_outgoing.any()
}
#[must_use]
pub fn edge_visible(&self, from: DenseNodeId, to: DenseNodeId) -> bool {
!self.hide_outgoing.contains(from) && !self.hide_incoming.contains(to)
}
}
#[derive(Clone, Copy, Debug)]
pub struct DagView<'a> {
dag: &'a Dag,
overlay: &'a GraphOverlay,
}
impl<'a> DagView<'a> {
#[must_use]
pub fn node_count(&self) -> usize {
self.dag.node_count()
}
#[must_use]
pub fn dag(&self) -> &'a Dag {
self.dag
}
#[must_use]
pub fn overlay(&self) -> &'a GraphOverlay {
self.overlay
}
pub fn parents_into(&self, id: DenseNodeId, out: &mut Vec<DenseNodeId>) {
out.clear();
if id.as_usize() >= self.dag.node_count() {
return;
}
for &p in self.dag.parents(id) {
if self.overlay.edge_visible(p, id) {
out.push(p);
}
}
}
pub fn children_into(&self, id: DenseNodeId, out: &mut Vec<DenseNodeId>) {
out.clear();
if id.as_usize() >= self.dag.node_count() {
return;
}
for &c in self.dag.children(id) {
if self.overlay.edge_visible(id, c) {
out.push(c);
}
}
}
pub fn ancestors_of(&self, nodes: &[DenseNodeId], out: &mut BitSet, ws: &mut GraphWorkspace) {
self.dag.ancestors_of_with(nodes, out, ws, Some(self.overlay));
}
pub fn descendants_of(&self, nodes: &[DenseNodeId], out: &mut BitSet, ws: &mut GraphWorkspace) {
self.dag.descendants_of_with(nodes, out, ws, Some(self.overlay));
}
pub fn is_d_separated(
&self,
x: DenseNodeId,
y: DenseNodeId,
z: &[DenseNodeId],
ws: &mut DSeparationWorkspace,
) -> Result<bool, GraphError> {
self.dag.is_d_separated_with(x, y, z, ws, Some(self.overlay))
}
pub fn d_separation(
&self,
x: DenseNodeId,
y: DenseNodeId,
z: &[DenseNodeId],
ws: &mut DSeparationWorkspace,
) -> Result<SeparationResult, GraphError> {
self.dag.d_separation_with(x, y, z, ws, Some(self.overlay))
}
pub fn materialize(&self) -> Result<Dag, GraphError> {
let n = u32::try_from(self.dag.node_count()).map_err(|_| GraphError::TooManyNodes)?;
let mut out = Dag::with_variables(n);
for i in 0..self.dag.node_count() {
let from = DenseNodeId::from_raw(u32::try_from(i).expect("node fit"));
for &to in self.dag.children(from) {
if self.overlay.edge_visible(from, to) {
out.insert_directed_unchecked(from, to);
}
}
}
Ok(out)
}
}
impl Dag {
#[must_use]
pub fn view<'a>(&'a self, overlay: &'a GraphOverlay) -> DagView<'a> {
DagView { dag: self, overlay }
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dsep::DSeparationWorkspace;
fn chain3() -> Dag {
let mut g = Dag::with_variables(3);
g.insert_directed(DenseNodeId::from_raw(0), DenseNodeId::from_raw(1)).unwrap();
g.insert_directed(DenseNodeId::from_raw(1), DenseNodeId::from_raw(2)).unwrap();
g
}
#[test]
fn do_intervention_view_matches_materialize() {
let g = chain3();
let t = DenseNodeId::from_raw(1);
let overlay = GraphOverlay::do_intervention(g.node_count(), &[t]);
let view = g.view(&overlay);
let mut parents = Vec::new();
view.parents_into(t, &mut parents);
assert!(parents.is_empty());
let mut children = Vec::new();
view.children_into(DenseNodeId::from_raw(0), &mut children);
assert!(children.is_empty());
view.children_into(t, &mut children);
assert_eq!(children, vec![DenseNodeId::from_raw(2)]);
let m = view.materialize().unwrap();
assert!(m.parents(t).is_empty());
assert!(m.children(DenseNodeId::from_raw(0)).is_empty());
assert_eq!(m.children(t).len(), 1);
}
#[test]
fn remove_outgoing_hides_treatment_children() {
let g = chain3();
let t = DenseNodeId::from_raw(1);
let overlay = GraphOverlay::remove_outgoing(g.node_count(), &[t]);
let view = g.view(&overlay);
let mut children = Vec::new();
view.children_into(t, &mut children);
assert!(children.is_empty());
let mut parents = Vec::new();
view.parents_into(t, &mut parents);
assert_eq!(parents, vec![DenseNodeId::from_raw(0)]);
}
#[test]
fn view_dsep_matches_materialized_mutilate() {
let mut g = Dag::with_variables(3);
let a = DenseNodeId::from_raw(0);
let t = DenseNodeId::from_raw(1);
let y = DenseNodeId::from_raw(2);
g.insert_directed(a, t).unwrap();
g.insert_directed(t, y).unwrap();
g.insert_directed(a, y).unwrap();
let overlay = GraphOverlay::do_intervention(g.node_count(), &[t]);
let view = g.view(&overlay);
let materialized = view.materialize().unwrap();
let mut ws = DSeparationWorkspace::default();
let view_sep = view.is_d_separated(t, y, &[a], &mut ws).unwrap();
let mat_sep = materialized.is_d_separated(t, y, &[a], &mut ws).unwrap();
assert_eq!(view_sep, mat_sep);
assert!(!view_sep);
}
#[test]
fn property_overlay_dsep_matches_mutilate_on_random_dags() {
use antecedent_core::CausalRng;
let mut rng = CausalRng::from_seed(17);
let mut ws = DSeparationWorkspace::default();
for _ in 0..40 {
let node_count = 4 + u32::try_from(rng.next_u64() % 3).unwrap_or(0); let mut graph = Dag::with_variables(node_count);
let mut order: Vec<u32> = (0..node_count).collect();
let n_usize = usize::try_from(node_count).unwrap_or(0);
for i in (1..n_usize).rev() {
let bound = u64::try_from(i + 1).unwrap_or(1);
let j = usize::try_from(rng.next_u64() % bound).unwrap_or(0);
order.swap(i, j);
}
for i in 0..n_usize {
for j in (i + 1)..n_usize {
if rng.next_u64() % 3 == 0 {
let _ = graph.insert_directed(
DenseNodeId::from_raw(order[i]),
DenseNodeId::from_raw(order[j]),
);
}
}
}
let treat_cap = n_usize.clamp(1, 3);
let treat_bound = u64::try_from(treat_cap).unwrap_or(1);
let n_treated = 1 + usize::try_from(rng.next_u64() % treat_bound).unwrap_or(0);
let mut treated = Vec::new();
while treated.len() < n_treated {
let raw = u32::try_from(rng.next_u64() % u64::from(node_count)).unwrap_or(0);
let treatment = DenseNodeId::from_raw(raw);
if !treated.contains(&treatment) {
treated.push(treatment);
}
}
let overlay = GraphOverlay::do_intervention(graph.node_count(), &treated);
let view = graph.view(&overlay);
let mutilated = graph.mutilate(&treated).unwrap();
let mat = view.materialize().unwrap();
for i in 0..node_count {
let u = DenseNodeId::from_raw(i);
assert_eq!(mat.children(u), mutilated.children(u));
}
for _ in 0..12 {
let x_raw = u32::try_from(rng.next_u64() % u64::from(node_count)).unwrap_or(0);
let source = DenseNodeId::from_raw(x_raw);
let mut y_raw = u32::try_from(rng.next_u64() % u64::from(node_count)).unwrap_or(0);
let mut target = DenseNodeId::from_raw(y_raw);
while target == source {
y_raw = u32::try_from(rng.next_u64() % u64::from(node_count)).unwrap_or(0);
target = DenseNodeId::from_raw(y_raw);
}
let mut conditioning = Vec::new();
for i in 0..node_count {
let node = DenseNodeId::from_raw(i);
if node == source || node == target {
continue;
}
if rng.next_u64() % 2 == 0 {
conditioning.push(node);
}
}
let view_sep = view.is_d_separated(source, target, &conditioning, &mut ws).unwrap();
let mut_sep =
mutilated.is_d_separated(source, target, &conditioning, &mut ws).unwrap();
assert_eq!(
view_sep,
mut_sep,
"overlay≠mutilate d-sep x={} y={} z={:?} T={:?}",
source.raw(),
target.raw(),
conditioning.iter().map(|v| v.raw()).collect::<Vec<_>>(),
treated.iter().map(|v| v.raw()).collect::<Vec<_>>()
);
}
}
}
}