use crate::brcd::brcd_augment::incident_undirected_edges;
use crate::brcd::brcd_error::{BrcdError, BrcdErrorEnum};
use crate::brcd::brcd_validity::{baseline_parents, is_valid_configuration};
use deep_causality_algebra::RealField;
use deep_causality_topology::MixedGraph;
use std::collections::{BTreeMap, BTreeSet};
const MAX_MAPPRUNE_EDGES: usize = (usize::BITS - 2) as usize;
pub struct PrunedConfigs<N> {
pub configs: Vec<MixedGraph<N>>,
pub evals: usize,
}
struct Finder<'a, N, F> {
cpdag: &'a MixedGraph<N>,
targets: &'a [usize],
incident: &'a [(usize, usize)],
baseline: &'a BTreeMap<usize, BTreeSet<usize>>,
weight: F,
}
impl<T, N, F> Finder<'_, N, F>
where
T: RealField,
N: Clone,
F: Fn(&MixedGraph<N>) -> Result<T, BrcdError>,
{
fn orient(&self, bits: usize) -> MixedGraph<N> {
let mut g = self.cpdag.clone();
for (i, &(a, b)) in self.incident.iter().enumerate() {
if (bits >> i) & 1 == 0 {
g.orient(a, b)
.expect("incident edge is undirected in the clone");
} else {
g.orient(b, a)
.expect("incident edge is undirected in the clone");
}
}
g
}
fn eval(
&self,
bits: usize,
visited: &mut BTreeMap<usize, Option<T>>,
evals: &mut usize,
) -> Result<Option<T>, BrcdError> {
if let Some(&cached) = visited.get(&bits) {
return Ok(cached);
}
*evals += 1;
let mut g = self.orient(bits);
if !is_valid_configuration(&mut g, self.targets, self.baseline) {
visited.insert(bits, None);
return Ok(None);
}
let w = (self.weight)(&g)?;
visited.insert(bits, Some(w));
Ok(Some(w))
}
fn valid_start(
&self,
du: usize,
visited: &mut BTreeMap<usize, Option<T>>,
evals: &mut usize,
) -> Result<Option<(usize, T)>, BrcdError> {
let all_in = (1usize << du) - 1;
for bits in [0usize, all_in] {
if let Some(w) = self.eval(bits, visited, evals)? {
return Ok(Some((bits, w)));
}
}
for j in 0..du {
let bits = 1usize << j;
if let Some(w) = self.eval(bits, visited, evals)? {
return Ok(Some((bits, w)));
}
}
Ok(None)
}
}
pub fn find_map_configs<T, N, F>(
cpdag: &MixedGraph<N>,
targets: &[usize],
weight: F,
) -> Result<PrunedConfigs<N>, BrcdError>
where
T: RealField,
N: Clone,
F: Fn(&MixedGraph<N>) -> Result<T, BrcdError>,
{
let n = cpdag.num_vertices();
if targets.iter().any(|&t| t >= n) {
return Err(BrcdError(BrcdErrorEnum::NodeOutOfBounds));
}
let incident = incident_undirected_edges(cpdag, targets);
let du = incident.len();
if du > MAX_MAPPRUNE_EDGES {
return Err(BrcdError(BrcdErrorEnum::ConfigSpaceTooLarge { edges: du }));
}
let baseline = baseline_parents(cpdag, targets);
let finder = Finder {
cpdag,
targets,
incident: &incident,
baseline: &baseline,
weight,
};
let mut visited: BTreeMap<usize, Option<T>> = BTreeMap::new();
let mut evals = 0usize;
if du == 0 {
return match finder.eval(0, &mut visited, &mut evals)? {
Some(_) => Ok(PrunedConfigs {
configs: vec![finder.completed(0)],
evals,
}),
None => Ok(PrunedConfigs {
configs: Vec::new(),
evals,
}),
};
}
let Some((mut bits, mut cur)) = finder.valid_start(du, &mut visited, &mut evals)? else {
return Ok(PrunedConfigs {
configs: Vec::new(),
evals,
});
};
loop {
let mut best: Option<(usize, T)> = None;
for j in 0..du {
let cand = bits ^ (1usize << j);
if let Some(w) = finder.eval(cand, &mut visited, &mut evals)? {
let improves = w > cur;
let beats_best = best.as_ref().is_none_or(|(_, bw)| w > *bw);
if improves && beats_best {
best = Some((cand, w));
}
}
}
match best {
Some((cand, w)) => {
bits = cand;
cur = w;
}
None => break,
}
}
let mut ranked: Vec<(usize, T)> = visited
.into_iter()
.filter_map(|(bits, w)| w.map(|w| (bits, w)))
.collect();
ranked.sort_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
let configs = ranked
.iter()
.map(|(cfg_bits, _)| finder.completed(*cfg_bits))
.collect();
Ok(PrunedConfigs { configs, evals })
}
impl<T, N, F> Finder<'_, N, F>
where
T: RealField,
N: Clone,
F: Fn(&MixedGraph<N>) -> Result<T, BrcdError>,
{
fn completed(&self, bits: usize) -> MixedGraph<N> {
let mut g = self.orient(bits);
let _ = is_valid_configuration(&mut g, self.targets, self.baseline);
g
}
}