use crate::brcd::brcd_meek::meek_complete;
use crate::brcd::{BrcdError, BrcdErrorEnum};
use deep_causality_tensor::CausalTensor;
use deep_causality_topology::MixedGraph;
use std::collections::BTreeSet;
pub fn dag_to_cpdag(parents: &[Vec<usize>]) -> Result<MixedGraph<()>, BrcdError> {
let n = parents.len();
if n == 0 {
return Err(BrcdError(BrcdErrorEnum::EmptyData));
}
for (c, ps) in parents.iter().enumerate() {
for &p in ps {
if p >= n {
return Err(BrcdError(BrcdErrorEnum::NodeOutOfBounds));
}
if p == c {
return Err(BrcdError(BrcdErrorEnum::NotAcyclic));
}
}
}
if !is_dag(parents) {
return Err(BrcdError(BrcdErrorEnum::NotAcyclic));
}
let mut adjacency: BTreeSet<(usize, usize)> = BTreeSet::new();
for (c, ps) in parents.iter().enumerate() {
for &p in ps {
adjacency.insert((p.min(c), p.max(c)));
}
}
let is_adjacent = |a: usize, b: usize| adjacency.contains(&(a.min(b), a.max(b)));
let mut compelled: BTreeSet<(usize, usize)> = BTreeSet::new();
for (c, ps) in parents.iter().enumerate() {
for i in 0..ps.len() {
for j in (i + 1)..ps.len() {
let (a, b) = (ps[i], ps[j]);
if !is_adjacent(a, b) {
compelled.insert((a, c));
compelled.insert((b, c));
}
}
}
}
let data = CausalTensor::new(vec![(); n], vec![n])
.map_err(|_| BrcdError(BrcdErrorEnum::DimensionMismatch))?;
let mut graph = MixedGraph::<()>::new(n, data, 0)
.map_err(|_| BrcdError(BrcdErrorEnum::DimensionMismatch))?;
for (c, ps) in parents.iter().enumerate() {
for &p in ps {
if compelled.contains(&(p, c)) {
graph
.add_arc(p, c)
.map_err(|_| BrcdError(BrcdErrorEnum::DimensionMismatch))?;
} else {
graph
.add_undirected(p, c)
.map_err(|_| BrcdError(BrcdErrorEnum::DimensionMismatch))?;
}
}
}
meek_complete(&mut graph);
Ok(graph)
}
fn is_dag(parents: &[Vec<usize>]) -> bool {
let n = parents.len();
let mut indegree = vec![0usize; n];
let mut children: Vec<Vec<usize>> = vec![Vec::new(); n];
for (c, ps) in parents.iter().enumerate() {
indegree[c] = ps.len();
for &p in ps {
children[p].push(c);
}
}
let mut queue: Vec<usize> = (0..n).filter(|&i| indegree[i] == 0).collect();
let mut processed = 0usize;
while let Some(v) = queue.pop() {
processed += 1;
for &c in &children[v] {
indegree[c] -= 1;
if indegree[c] == 0 {
queue.push(c);
}
}
}
processed == n
}