use crate::causal_discovery::brcd::brcd_error::{BrcdError, BrcdErrorEnum};
use crate::dag_sampling::chordal::mcs;
use crate::dag_sampling::clique_tree::CliqueTree;
use crate::dag_sampling::combinatorics;
use crate::dag_sampling::graph::Graph;
use crate::dag_sampling::index_set::IndexSet;
use crate::dag_sampling::lazy_tokens::LazyTokens;
use crate::dag_sampling::memoization::Memoization;
use crate::dag_sampling::utils::inverse_permutation;
use deep_causality_algebra::RealField;
use deep_causality_num::FromPrimitive;
use deep_causality_rand::Rng;
use deep_causality_topology::{EdgeKind, MixedGraph};
use std::collections::{BTreeMap, BTreeSet, VecDeque};
#[inline]
fn from_f64<T: FromPrimitive>(x: f64) -> T {
<T as FromPrimitive>::from_f64(x).expect("a finite f64 is representable in every RealField")
}
#[derive(Debug, Clone)]
struct WeightedChoice<T> {
cumulative: Vec<T>,
total: T,
}
impl<T: RealField + FromPrimitive> WeightedChoice<T> {
fn empty() -> Self {
WeightedChoice {
cumulative: Vec::new(),
total: T::zero(),
}
}
fn new(weights: &[T]) -> Self {
let mut cumulative = Vec::with_capacity(weights.len());
let mut running = T::zero();
for &w in weights {
running += w;
cumulative.push(running);
}
WeightedChoice {
cumulative,
total: running,
}
}
fn sample<R: Rng>(&self, rng: &mut R) -> usize {
let u: f64 = rng.random::<f64>();
let target = self.total * from_f64::<T>(u);
for (i, c) in self.cumulative.iter().enumerate() {
if target < *c {
return i;
}
}
self.cumulative.len() - 1
}
}
#[derive(Debug)]
struct ComponentSampler<T> {
n: usize,
clique_tree: CliqueTree,
separators: Vec<IndexSet>,
flowers: Vec<IndexSet>,
choices: Vec<WeightedChoice<T>>,
forbidden_sets: Vec<Vec<(usize, usize, usize)>>,
}
impl<T: RealField + FromPrimitive> ComponentSampler<T> {
fn init(g: &Graph) -> Self {
let clique_tree = CliqueTree::from(g);
let mut separators = clique_tree.separators();
separators.push(IndexSet::from_sorted(Vec::new()));
let mut flowers = clique_tree.flowers(&separators);
flowers.push(IndexSet::from_sorted((0..clique_tree.tree.n).collect()));
let num_subproblems = flowers.len();
let mut choices = vec![WeightedChoice::empty(); num_subproblems];
let forbidden_sets = clique_tree.forbidden_sets(&separators, &flowers);
let mut visited = LazyTokens::new(clique_tree.tree.n);
let mut considered = LazyTokens::new(clique_tree.tree.n);
let mut memoization = Memoization::<T>::new(clique_tree.tree.n, g.n);
Self::rec_count_init(
&mut choices,
num_subproblems - 1,
&mut visited,
&mut considered,
&mut memoization,
&clique_tree,
&separators,
&flowers,
&forbidden_sets,
);
ComponentSampler {
n: g.n,
clique_tree,
separators,
flowers,
choices,
forbidden_sets,
}
}
#[allow(clippy::too_many_arguments)]
fn rec_count_init(
choices: &mut [WeightedChoice<T>],
subproblem: usize,
visited: &mut LazyTokens,
considered: &mut LazyTokens,
memoization: &mut Memoization<T>,
clique_tree: &CliqueTree,
separators: &[IndexSet],
flowers: &[IndexSet],
forbidden_sets: &[Vec<(usize, usize, usize)>],
) -> T {
if let Some(res) = memoization.count[subproblem] {
return res;
}
let flower = &flowers[subproblem];
let separator = &separators[subproblem];
let mut sum = T::zero();
let mut amos_per_clique: Vec<T> = Vec::with_capacity(flower.len());
let mut pre: Vec<Option<T>> = vec![None; 2 * (clique_tree.tree.n - 1)];
let flower_cliques: Vec<usize> = flower.iter().copied().collect();
for clique_id in flower_cliques {
let mut forbidden_sizes = Vec::new();
forbidden_sizes.push(clique_tree.cliques[clique_id].len() - separator.len());
for &(u, v, size) in &forbidden_sets[clique_id] {
if !flower.contains(u) || !flower.contains(v) {
break;
}
if size > separator.len() {
forbidden_sizes.push(size - separator.len());
} else {
break;
}
}
let phi = combinatorics::rho(&forbidden_sizes, memoization);
visited.prepare();
considered.prepare();
visited.set(clique_id);
considered.set(clique_id);
let product = phi
* Self::rec_count_traversal(
choices,
flower,
clique_id,
&mut pre,
visited,
considered,
memoization,
clique_tree,
separators,
flowers,
forbidden_sets,
);
visited.restore();
considered.restore();
sum += product;
amos_per_clique.push(product);
}
memoization.count[subproblem] = Some(sum);
choices[subproblem] = WeightedChoice::new(&amos_per_clique);
sum
}
#[allow(clippy::too_many_arguments)]
fn rec_count_traversal(
choices: &mut [WeightedChoice<T>],
flower: &IndexSet,
i: usize,
pre: &mut Vec<Option<T>>,
visited: &mut LazyTokens,
considered: &mut LazyTokens,
memoization: &mut Memoization<T>,
clique_tree: &CliqueTree,
separators: &[IndexSet],
flowers: &[IndexSet],
forbidden_sets: &[Vec<(usize, usize, usize)>],
) -> T {
visited.set(i);
let mut product = T::one();
let neighbors: Vec<usize> = clique_tree.tree.neighbors(i).copied().collect();
for j in neighbors {
if !flower.contains(j) {
continue;
}
let edge_id = clique_tree.get_edge_id(i, j);
if !visited.check(j) && !considered.check(j) {
if let Some(pre_val) = pre[edge_id] {
visited.set(j);
product *= pre_val;
} else {
let next_flower_result = Self::rec_count_init(
choices,
edge_id,
visited,
considered,
memoization,
clique_tree,
separators,
flowers,
forbidden_sets,
);
for &new_clique_id in &flowers[edge_id] {
considered.set(new_clique_id);
}
let remaining_subtree_result = Self::rec_count_traversal(
choices,
flower,
j,
pre,
visited,
considered,
memoization,
clique_tree,
separators,
flowers,
forbidden_sets,
);
let combined = next_flower_result * remaining_subtree_result;
pre[edge_id] = Some(combined);
product *= combined;
}
} else if !visited.check(j) {
product *= Self::rec_count_traversal(
choices,
flower,
j,
pre,
visited,
considered,
memoization,
clique_tree,
separators,
flowers,
forbidden_sets,
);
}
}
product
}
fn sample_ordering<R: Rng>(&self, rng: &mut R) -> Vec<usize> {
let mut pos = vec![0usize; self.n];
let mut visited = LazyTokens::new(self.clique_tree.tree.n);
let mut considered = LazyTokens::new(self.clique_tree.tree.n);
self.rec_sample_ordering(
self.flowers.len() - 1,
&mut pos,
&mut visited,
&mut considered,
rng,
)
}
fn rec_sample_ordering<R: Rng>(
&self,
subproblem: usize,
pos: &mut [usize],
visited: &mut LazyTokens,
considered: &mut LazyTokens,
rng: &mut R,
) -> Vec<usize> {
let mut ordering = Vec::new();
let clique_local_idx = self.choices[subproblem].sample(rng);
let clique_id = self.flowers[subproblem].get(clique_local_idx);
let clique = &self.clique_tree.cliques[clique_id];
let flower = &self.flowers[subproblem];
let separator = &self.separators[subproblem];
let remaining_clique_vertices = clique.set_difference(separator);
let mut forbidden_prefixes = Vec::new();
for &(u, v, size) in &self.forbidden_sets[clique_id] {
if !flower.contains(u) || !flower.contains(v) {
break;
}
if size > separator.len() {
forbidden_prefixes.push(
self.separators[self.clique_tree.get_edge_id(u, v)].set_difference(separator),
);
} else {
break;
}
}
let mut clique_ordering = Self::draw_allowed_permutation(
&remaining_clique_vertices,
pos,
&forbidden_prefixes,
rng,
);
ordering.append(&mut clique_ordering);
let flower = &self.flowers[subproblem];
let mut queue = VecDeque::new();
queue.push_back(clique_id);
visited.prepare();
considered.prepare();
visited.set(clique_id);
considered.set(clique_id);
while let Some(u) = queue.pop_front() {
let neighbors: Vec<usize> = self.clique_tree.tree.neighbors(u).copied().collect();
for v in neighbors {
if !flower.contains(v) {
continue;
}
if !visited.check(v) {
queue.push_back(v);
visited.set(v);
}
if !considered.check(v) {
let new_flower_id = self.clique_tree.get_edge_id(u, v);
let mut sub =
self.rec_sample_ordering(new_flower_id, pos, visited, considered, rng);
ordering.append(&mut sub);
for &new_clique_id in &self.flowers[new_flower_id] {
considered.set(new_clique_id);
}
}
}
}
visited.restore();
considered.restore();
ordering
}
fn draw_allowed_permutation<R: Rng>(
clique: &IndexSet,
helper: &mut [usize],
forbidden_prefixes: &[IndexSet],
rng: &mut R,
) -> Vec<usize> {
for &u in clique {
helper[u] = clique.len();
}
for forbidden_prefix in forbidden_prefixes {
for &u in forbidden_prefix {
helper[u] = forbidden_prefix.len() - 1;
}
}
loop {
let mut perm = clique.to_vec();
shuffle(&mut perm, rng);
if Self::is_allowed(&perm, helper) {
return perm;
}
}
}
fn is_allowed(perm: &[usize], helper: &[usize]) -> bool {
let mut mx = 0;
for (i, &u) in perm.iter().enumerate() {
mx = mx.max(helper[u]);
if mx == i {
return false;
}
if mx >= perm.len() {
return true;
}
}
true
}
}
fn shuffle<R: Rng>(slice: &mut [usize], rng: &mut R) {
let len = slice.len();
if len <= 1 {
return;
}
for i in (1..len).rev() {
let j: usize = rng.random_range(0..(i + 1));
slice.swap(i, j);
}
}
pub fn sample_dag<T, N, R>(graph: &MixedGraph<N>, rng: &mut R) -> Result<MixedGraph<N>, BrcdError>
where
T: RealField + FromPrimitive,
N: Clone,
R: Rng,
{
validate_cpdag(graph)?;
let mut dag = graph.clone();
for component in chain_components(graph) {
let (g, local_to_global) = component.to_internal_graph();
let sampler = ComponentSampler::<T>::init(&g);
let ordering = sampler.sample_ordering(rng);
let mut order_pos = vec![0usize; g.n];
for (i, &local) in ordering.iter().enumerate() {
order_pos[local] = i;
}
for &(lu, lv) in &component.edges {
let (gu, gv) = (local_to_global[lu], local_to_global[lv]);
let (parent, child) = if order_pos[lu] < order_pos[lv] {
(gu, gv)
} else {
(gv, gu)
};
dag.orient(parent, child)
.map_err(|_| BrcdError(BrcdErrorEnum::NotACpdag))?;
}
}
if dag.has_cycle() {
return Err(BrcdError(BrcdErrorEnum::NotAcyclic));
}
Ok(dag)
}
pub fn representative_dag<N>(graph: &MixedGraph<N>) -> Result<MixedGraph<N>, BrcdError>
where
N: Clone,
{
validate_cpdag(graph)?;
let mut dag = graph.clone();
for component in chain_components(graph) {
let (g, local_to_global) = component.to_internal_graph();
let order = mcs(&g);
let rank = inverse_permutation(&order);
for &(lu, lv) in &component.edges {
let (parent, child) = if rank[lu] < rank[lv] {
(local_to_global[lu], local_to_global[lv])
} else {
(local_to_global[lv], local_to_global[lu])
};
dag.orient(parent, child)
.map_err(|_| BrcdError(BrcdErrorEnum::NotACpdag))?;
}
}
if dag.has_cycle() {
return Err(BrcdError(BrcdErrorEnum::NotAcyclic));
}
Ok(dag)
}
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(())
}
struct Component {
vertices: Vec<usize>,
edges: Vec<(usize, usize)>,
}
impl Component {
fn to_internal_graph(&self) -> (Graph, Vec<usize>) {
let edge_list: Vec<(usize, usize)> = self.edges.clone();
let g = Graph::from_edge_list(edge_list, self.vertices.len());
(g, self.vertices.clone())
}
}
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 vertices: Vec<usize> = members.iter().copied().collect();
let global_to_local: BTreeMap<usize, usize> =
vertices.iter().enumerate().map(|(i, &g)| (g, i)).collect();
let mut edges: Vec<(usize, usize)> = Vec::new();
for &v in &vertices {
for nb in graph.undirected_neighbors(v) {
if v < nb {
edges.push((global_to_local[&v], global_to_local[&nb]));
}
}
}
components.push(Component { vertices, edges });
}
components
}