use crate::brcd::BrcdError;
use crate::brcd::brcd_boss_gst::Gst;
use crate::brcd::brcd_boss_score::FamilyScorer;
use deep_causality_num::{FromPrimitive, RealField};
use deep_causality_rand::{Rng, Xoshiro256};
use std::cmp::Ordering;
use std::collections::BTreeSet;
const TOL: f64 = 1e-8;
const MAX_ROUNDS: usize = 2000;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct OrderSearchResult {
pub order: Vec<usize>,
pub parents: Vec<Vec<usize>>,
}
pub fn best_order_search<T, S>(scorer: &S, seed: u64) -> Result<OrderSearchResult, BrcdError>
where
T: RealField + FromPrimitive,
S: FamilyScorer<T>,
{
let p = scorer.num_vars();
let mut gsts: Vec<Gst<T>> = Vec::with_capacity(p);
for v in 0..p {
gsts.push(Gst::new(v, scorer)?);
}
let mut order: Vec<usize> = (0..p).collect();
if p <= 1 {
return finalize(order, &mut gsts, scorer);
}
let tol = from_f64::<T>(TOL);
let mut rng = Xoshiro256::from_seed(seed);
let mut visited: BTreeSet<Vec<usize>> = BTreeSet::new();
for _ in 0..MAX_ROUNDS {
if !visited.insert(order.clone()) {
break;
}
let current_total = total_score(&order, &mut gsts, scorer)?;
let mut variables = order.clone();
shuffle(&mut variables, &mut rng);
let mut improved = false;
for v in variables {
if better_mutation(v, &mut order, &mut gsts, scorer, tol)? {
improved = true;
}
}
let new_total = total_score(&order, &mut gsts, scorer)?;
if !improved || !new_total.is_finite() || new_total <= current_total + tol {
break;
}
}
finalize(order, &mut gsts, scorer)
}
fn total_score<T, S>(order: &[usize], gsts: &mut [Gst<T>], scorer: &S) -> Result<T, BrcdError>
where
T: RealField + FromPrimitive,
S: FamilyScorer<T>,
{
let mut total = T::zero();
let mut prefix: Vec<usize> = Vec::with_capacity(order.len());
for &w in order {
let (_, val) = gsts[w].trace(&prefix, scorer)?;
if !val.is_finite() {
return Ok(neg_inf::<T>());
}
total += val;
prefix.push(w);
}
Ok(total)
}
fn better_mutation<T, S>(
v: usize,
order: &mut Vec<usize>,
gsts: &mut [Gst<T>],
scorer: &S,
tol: T,
) -> Result<bool, BrcdError>
where
T: RealField + FromPrimitive,
S: FamilyScorer<T>,
{
let p = order.len();
let i = order
.iter()
.position(|&x| x == v)
.expect("v is in the order");
let mut scores = vec![neg_inf::<T>(); p + 1];
let mut prefix: Vec<usize> = Vec::with_capacity(p);
let mut accum = T::zero();
for (j, &w) in order.iter().enumerate() {
let (_, sv) = gsts[v].trace(&prefix, scorer)?;
if sv.is_finite() && accum.is_finite() {
scores[j] = sv + accum;
}
if v != w {
let (_, sw) = gsts[w].trace(&prefix, scorer)?;
accum = if sw.is_finite() && accum.is_finite() {
accum + sw
} else {
neg_inf::<T>()
};
prefix.push(w);
}
}
let (_, sv_end) = gsts[v].trace(&prefix, scorer)?;
if sv_end.is_finite() && accum.is_finite() {
scores[p] = sv_end + accum;
}
let best = argmax(&scores);
if scores[best].partial_cmp(&(scores[i] + tol)) != Some(Ordering::Greater) {
return Ok(false);
}
order.remove(i);
let insert_at = best - usize::from(best > i);
order.insert(insert_at, v);
Ok(true)
}
fn finalize<T, S>(
order: Vec<usize>,
gsts: &mut [Gst<T>],
scorer: &S,
) -> Result<OrderSearchResult, BrcdError>
where
T: RealField + FromPrimitive,
S: FamilyScorer<T>,
{
let p = order.len();
let mut parents = vec![Vec::new(); p];
let mut prefix: Vec<usize> = Vec::with_capacity(p);
for &w in &order {
let (pa, _) = gsts[w].trace(&prefix, scorer)?;
parents[w] = pa;
prefix.push(w);
}
Ok(OrderSearchResult { order, parents })
}
fn argmax<T: RealField>(scores: &[T]) -> usize {
let mut best = 0;
for (k, s) in scores.iter().enumerate().skip(1) {
if *s > scores[best] {
best = k;
}
}
best
}
fn shuffle(items: &mut [usize], rng: &mut Xoshiro256) {
for i in (1..items.len()).rev() {
let j: usize = rng.random_range(0..(i + 1));
items.swap(i, j);
}
}
fn neg_inf<T: RealField>() -> T {
T::zero().ln()
}
fn from_f64<T: FromPrimitive>(x: f64) -> T {
<T as FromPrimitive>::from_f64(x).expect("constant is representable in every RealField")
}