use crate::brcd::BrcdError;
use crate::brcd::brcd_boss_score::FamilyScorer;
use deep_causality_algebra::RealField;
use std::cmp::Ordering;
struct GstNode<T> {
add: Option<usize>,
grow_score: T,
shrink_score: T,
branches: Option<Vec<GstNode<T>>>,
remove: Option<Vec<usize>>,
}
impl<T: RealField> GstNode<T> {
fn leaf(add: usize, score: T) -> Self {
Self {
add: Some(add),
grow_score: score,
shrink_score: score,
branches: None,
remove: None,
}
}
fn grow<S: FamilyScorer<T>>(
&mut self,
vertex: usize,
available: &[usize],
parents: &mut Vec<usize>,
scorer: &S,
) -> Result<(), BrcdError> {
let mut branches = Vec::new();
for &add in available {
parents.push(add);
let score = scorer.score(vertex, parents)?;
parents.pop();
if score > self.grow_score {
branches.push(GstNode::leaf(add, score));
}
}
branches.sort_by(|a, b| {
a.grow_score
.partial_cmp(&b.grow_score)
.unwrap_or(Ordering::Equal)
});
self.branches = Some(branches);
Ok(())
}
fn shrink<S: FamilyScorer<T>>(
&mut self,
vertex: usize,
parents: &mut Vec<usize>,
scorer: &S,
) -> Result<(), BrcdError> {
let mut removed = Vec::new();
loop {
let mut best: Option<usize> = None;
let candidates = parents.clone();
for r in candidates {
let pos = parents.iter().position(|&x| x == r).expect("present");
parents.remove(pos);
let score = scorer.score(vertex, parents)?;
parents.insert(pos, r);
if score > self.shrink_score {
self.shrink_score = score;
best = Some(r);
}
}
match best {
None => break,
Some(b) => {
removed.push(b);
parents.retain(|&x| x != b);
}
}
}
self.remove = Some(removed);
Ok(())
}
fn trace<S: FamilyScorer<T>>(
&mut self,
vertex: usize,
prefix: &[usize],
available: &mut Vec<usize>,
parents: &mut Vec<usize>,
scorer: &S,
) -> Result<T, BrcdError> {
if self.branches.is_none() {
self.grow(vertex, available, parents, scorer)?;
}
let branches = self.branches.as_mut().expect("grown above");
for branch in branches.iter_mut() {
let add = branch.add.expect("non-root branch");
available.retain(|&x| x != add);
if prefix.contains(&add) {
parents.push(add);
return branch.trace(vertex, prefix, available, parents, scorer);
}
}
if self.remove.is_none() {
self.shrink(vertex, parents, scorer)?;
return Ok(self.shrink_score);
}
if let Some(removed) = &self.remove {
for &r in removed {
parents.retain(|&x| x != r);
}
}
Ok(self.shrink_score)
}
}
pub struct Gst<T> {
vertex: usize,
num_vars: usize,
forbidden: Vec<usize>,
root: GstNode<T>,
}
impl<T: RealField> Gst<T> {
pub fn new<S: FamilyScorer<T>>(vertex: usize, scorer: &S) -> Result<Self, BrcdError> {
let empty = scorer.score(vertex, &[])?;
Ok(Self {
vertex,
num_vars: scorer.num_vars(),
forbidden: vec![vertex],
root: GstNode {
add: None,
grow_score: empty,
shrink_score: empty,
branches: None,
remove: None,
},
})
}
pub fn vertex(&self) -> usize {
self.vertex
}
pub fn trace<S: FamilyScorer<T>>(
&mut self,
prefix: &[usize],
scorer: &S,
) -> Result<(Vec<usize>, T), BrcdError> {
let mut available: Vec<usize> = (0..self.num_vars)
.filter(|i| !self.forbidden.contains(i))
.collect();
let mut parents = Vec::new();
let score = self
.root
.trace(self.vertex, prefix, &mut available, &mut parents, scorer)?;
parents.sort_unstable();
Ok((parents, score))
}
}