use crate::causal_discovery::brcd::brcd_error::{BrcdError, BrcdErrorEnum};
use deep_causality_rand::Rng;
use deep_causality_topology::{EdgeKind, MixedGraph};
use std::collections::{BTreeMap, BTreeSet, VecDeque};
pub const MEC_ENUM_BOUND: usize = 100_000;
pub fn mec_size<N>(graph: &MixedGraph<N>) -> Result<usize, BrcdError> {
validate_cpdag(graph)?;
let mut size: usize = 1;
for component in chain_components(graph) {
let amos = enumerate_amos(&component)?;
size = checked_class_product(size, amos.len())?;
}
Ok(size)
}
pub fn representative_dag<N: Clone>(graph: &MixedGraph<N>) -> Result<MixedGraph<N>, BrcdError> {
build_member(graph, |amo_count| {
debug_assert!(amo_count > 0);
0
})
}
pub fn mec_sample_dag<N: Clone, R: Rng>(
graph: &MixedGraph<N>,
rng: &mut R,
) -> Result<MixedGraph<N>, BrcdError> {
build_member(graph, |amo_count| rng.random_range(0..amo_count))
}
fn validate_cpdag<N>(graph: &MixedGraph<N>) -> Result<(), BrcdError> {
for edge in graph.edges().values() {
match edge.kind() {
EdgeKind::Directed | EdgeKind::Undirected => {}
_ => return Err(BrcdError(BrcdErrorEnum::NotACpdag)),
}
}
if graph.has_cycle() {
return Err(BrcdError(BrcdErrorEnum::NotAcyclic));
}
Ok(())
}
fn checked_class_product(acc: usize, factor: usize) -> Result<usize, BrcdError> {
match acc.checked_mul(factor) {
Some(p) if p <= MEC_ENUM_BOUND => Ok(p),
_ => Err(BrcdError(BrcdErrorEnum::ClassTooLarge {
bound: MEC_ENUM_BOUND,
})),
}
}
struct Component {
adj: BTreeMap<usize, Vec<usize>>,
}
fn chain_components<N>(graph: &MixedGraph<N>) -> Vec<Component> {
let mut undirected_vertices: BTreeSet<usize> = BTreeSet::new();
for &(a, b) in graph.undirected_edges().iter() {
undirected_vertices.insert(a);
undirected_vertices.insert(b);
}
let mut seen: BTreeSet<usize> = BTreeSet::new();
let mut components = Vec::new();
for &start in undirected_vertices.iter() {
if seen.contains(&start) {
continue;
}
let mut members: BTreeSet<usize> = BTreeSet::new();
let mut queue: VecDeque<usize> = VecDeque::new();
queue.push_back(start);
seen.insert(start);
members.insert(start);
while let Some(v) = queue.pop_front() {
for nb in graph.undirected_neighbors(v) {
if seen.insert(nb) {
members.insert(nb);
queue.push_back(nb);
}
}
}
let mut adj: BTreeMap<usize, Vec<usize>> = BTreeMap::new();
for &v in members.iter() {
adj.insert(v, graph.undirected_neighbors(v));
}
components.push(Component { adj });
}
components
}
fn enumerate_amos(component: &Component) -> Result<Vec<Vec<(usize, usize)>>, BrcdError> {
let adj = &component.adj;
let total_edges: usize = adj.values().map(Vec::len).sum::<usize>() / 2;
let mut labels: BTreeMap<usize, i64> = adj.keys().map(|&v| (v, 0i64)).collect();
let mut visited: BTreeSet<usize> = BTreeSet::new();
let mut unique: BTreeSet<Vec<(usize, usize)>> = BTreeSet::new();
let mut out: Vec<Vec<(usize, usize)>> = Vec::new();
mcs_enum(
adj,
total_edges,
&mut visited,
&mut labels,
&[],
&mut unique,
&mut out,
)?;
Ok(out)
}
#[allow(clippy::too_many_arguments)]
fn mcs_enum(
adj: &BTreeMap<usize, Vec<usize>>,
total_edges: usize,
visited: &mut BTreeSet<usize>,
labels: &mut BTreeMap<usize, i64>,
oriented: &[(usize, usize)],
unique: &mut BTreeSet<Vec<(usize, usize)>>,
out: &mut Vec<Vec<(usize, usize)>>,
) -> Result<(), BrcdError> {
if visited.len() == adj.len() {
if oriented.len() == total_edges {
let mut key = oriented.to_vec();
key.sort_unstable();
if unique.insert(key.clone()) {
if out.len() >= MEC_ENUM_BOUND {
return Err(BrcdError(BrcdErrorEnum::ClassTooLarge {
bound: MEC_ENUM_BOUND,
}));
}
out.push(key);
}
}
return Ok(());
}
let max_label = adj
.keys()
.filter(|v| !visited.contains(v))
.map(|v| labels[v])
.max()
.expect("non-empty unvisited set");
let candidates: Vec<usize> = adj
.keys()
.copied()
.filter(|v| !visited.contains(v) && labels[v] == max_label)
.collect();
for v in candidates {
visited.insert(v);
let saved_labels = labels.clone();
let mut next_oriented = oriented.to_vec();
for &neighbor in &adj[&v] {
if !visited.contains(&neighbor) {
*labels.get_mut(&neighbor).expect("labelled vertex") += 1;
} else if !next_oriented.contains(&(neighbor, v))
&& !next_oriented.contains(&(v, neighbor))
&& !creates_invalid_collider(v, neighbor, &next_oriented, adj)
{
next_oriented.push((neighbor, v));
}
}
mcs_enum(
adj,
total_edges,
visited,
labels,
&next_oriented,
unique,
out,
)?;
visited.remove(&v);
*labels = saved_labels;
}
Ok(())
}
fn creates_invalid_collider(
v: usize,
neighbor: usize,
oriented: &[(usize, usize)],
adj: &BTreeMap<usize, Vec<usize>>,
) -> bool {
for &other in &adj[&v] {
if other != neighbor
&& !adj[&neighbor].contains(&other)
&& oriented.contains(&(other, v))
&& !oriented.contains(&(v, other))
{
return true;
}
}
false
}
fn build_member<N: Clone>(
graph: &MixedGraph<N>,
mut choose: impl FnMut(usize) -> usize,
) -> Result<MixedGraph<N>, BrcdError> {
validate_cpdag(graph)?;
let mut dag = graph.clone();
let mut running: usize = 1;
for component in chain_components(graph) {
let amos = enumerate_amos(&component)?;
if amos.is_empty() {
return Err(BrcdError(BrcdErrorEnum::NotACpdag));
}
running = checked_class_product(running, amos.len())?;
let pick = choose(amos.len());
for &(parent, child) in &amos[pick] {
dag.orient(parent, child)
.expect("AMO edge corresponds to an undirected edge of the clone");
}
}
if dag.has_cycle() {
return Err(BrcdError(BrcdErrorEnum::NotAcyclic));
}
Ok(dag)
}