use crate::{BottomUpTa, Explicit, ProbabilityScorer, StateId, TopDownTa, WeightScorer};
use fixedbitset::FixedBitSet;
use std::collections::BinaryHeap;
#[derive(Clone, Copy, PartialEq)]
struct OrdF64(f64);
impl Eq for OrdF64 {}
impl Ord for OrdF64 {
fn cmp(&self, o: &Self) -> std::cmp::Ordering {
self.0.total_cmp(&o.0)
}
}
impl PartialOrd for OrdF64 {
fn partial_cmp(&self, o: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(o))
}
}
pub trait IntersectionHeuristic<R: BottomUpTa> {
fn outside_estimate(&self, left: StateId, right: &R::State) -> f64;
#[inline]
fn admits(&self, _left: StateId, _right: &R::State) -> bool {
true
}
#[inline]
fn estimate_if_admitted(&self, left: StateId, right: &R::State) -> Option<f64> {
self.admits(left, right)
.then(|| self.outside_estimate(left, right))
}
#[inline]
fn memoize_admission(&self) -> bool {
false
}
#[inline]
fn estimate_after_admission(&self, left: StateId, right: &R::State) -> f64 {
self.outside_estimate(left, right)
}
}
pub struct MinHeuristic<A, B> {
a: A,
b: B,
}
impl<A, B> MinHeuristic<A, B> {
pub fn new(a: A, b: B) -> Self {
Self { a, b }
}
}
impl<R, A, B> IntersectionHeuristic<R> for MinHeuristic<A, B>
where
R: BottomUpTa,
A: IntersectionHeuristic<R>,
B: IntersectionHeuristic<R>,
{
#[inline]
fn outside_estimate(&self, left: StateId, right: &R::State) -> f64 {
self.a
.outside_estimate(left, right)
.min(self.b.outside_estimate(left, right))
}
#[inline]
fn admits(&self, left: StateId, right: &R::State) -> bool {
self.a.admits(left, right) && self.b.admits(left, right)
}
#[inline]
fn estimate_if_admitted(&self, left: StateId, right: &R::State) -> Option<f64> {
let a = self.a.estimate_if_admitted(left, right)?;
let b = self.b.estimate_if_admitted(left, right)?;
Some(a.min(b))
}
#[inline]
fn memoize_admission(&self) -> bool {
self.a.memoize_admission() || self.b.memoize_admission()
}
#[inline]
fn estimate_after_admission(&self, left: StateId, right: &R::State) -> f64 {
self.a
.estimate_after_admission(left, right)
.min(self.b.estimate_after_admission(left, right))
}
}
pub struct ZeroHeuristic;
impl<R: BottomUpTa> IntersectionHeuristic<R> for ZeroHeuristic {
#[inline]
fn outside_estimate(&self, _left: StateId, _right: &R::State) -> f64 {
1.0
}
}
pub struct ScoredZeroHeuristic {
one: f64,
}
impl ScoredZeroHeuristic {
pub fn new<S: WeightScorer>(scorer: &S) -> Self {
Self { one: scorer.one() }
}
}
impl<R: BottomUpTa> IntersectionHeuristic<R> for ScoredZeroHeuristic {
#[inline]
fn outside_estimate(&self, _left: StateId, _right: &R::State) -> f64 {
self.one
}
}
pub struct OutsideHeuristic {
out: Vec<f64>,
zero: f64,
}
impl OutsideHeuristic {
pub fn from_grammar(grammar: &Explicit) -> Self {
Self::from_grammar_with(grammar, &ProbabilityScorer)
}
pub fn from_grammar_with<S: WeightScorer>(grammar: &Explicit, scorer: &S) -> Self {
let n = grammar.num_states() as usize;
let mut inside = vec![scorer.zero(); n];
let mut fin_in = FixedBitSet::with_capacity(n);
let rules: Vec<_> = grammar.rules().collect();
let mut by_child: Vec<Vec<usize>> = vec![Vec::new(); n];
for (idx, rule) in rules.iter().enumerate() {
let mut seen_in_rule = FixedBitSet::with_capacity(n);
for &child in rule.children {
if !seen_in_rule.contains(child.index()) {
seen_in_rule.set(child.index(), true);
by_child[child.index()].push(idx);
}
}
}
let mut pending: Vec<usize> = rules
.iter()
.map(|r| {
let mut seen = FixedBitSet::with_capacity(n);
r.children
.iter()
.filter(|&&c| {
if seen.contains(c.index()) {
false
} else {
seen.set(c.index(), true);
true
}
})
.count()
})
.collect();
let mut partial: Vec<f64> = rules.iter().map(|r| scorer.rule_score(r.weight)).collect();
let mut heap: BinaryHeap<(OrdF64, u32)> = BinaryHeap::new();
for (idx, rule) in rules.iter().enumerate() {
if rule.children.is_empty() {
let ri = rule.result.index();
let rule_score = scorer.rule_score(rule.weight);
if scorer.better(rule_score, inside[ri]) {
inside[ri] = rule_score;
heap.push((OrdF64(rule_score), ri as u32));
}
let _ = idx; }
}
while let Some((OrdF64(w), si)) = heap.pop() {
let si = si as usize;
if fin_in.contains(si) {
continue;
}
if w != inside[si] {
continue;
}
fin_in.set(si, true);
for &rule_idx in &by_child[si] {
let rule = &rules[rule_idx];
partial[rule_idx] = scorer.times(partial[rule_idx], inside[si]);
pending[rule_idx] -= 1;
if pending[rule_idx] == 0 {
let ri = rule.result.index();
let cand = partial[rule_idx];
if scorer.better(cand, inside[ri]) {
inside[ri] = cand;
heap.push((OrdF64(cand), ri as u32));
}
}
}
}
let mut outside = vec![scorer.zero(); n];
let mut fin_out = FixedBitSet::with_capacity(n);
let mut out_heap: BinaryHeap<(OrdF64, u32)> = BinaryHeap::new();
grammar.initial_states(&mut |state| {
if !state.is_stuck() && state.index() < n {
let si = state.index();
let one = scorer.one();
if scorer.better(one, outside[si]) {
outside[si] = one;
out_heap.push((OrdF64(one), si as u32));
}
}
});
while let Some((OrdF64(w), si)) = out_heap.pop() {
let si = si as usize;
if fin_out.contains(si) {
continue;
}
if w != outside[si] {
continue;
}
fin_out.set(si, true);
let state = StateId(si as u32);
for rule in grammar.rules_topdown(state) {
if rule.children.is_empty() {
continue;
}
let nc = rule.children.len();
let mut prefix = vec![scorer.one(); nc + 1];
for i in 0..nc {
prefix[i + 1] = scorer.times(prefix[i], inside[rule.children[i].index()]);
}
let mut suffix = vec![scorer.one(); nc + 1];
for i in (0..nc).rev() {
suffix[i] = scorer.times(suffix[i + 1], inside[rule.children[i].index()]);
}
for p in 0..nc {
let child_p = rule.children[p];
if child_p.is_stuck() {
continue;
}
let ci = child_p.index();
if fin_out.contains(ci) {
continue;
}
let sibling_product = scorer.times(prefix[p], suffix[p + 1]);
let new_out = scorer.times(
scorer.times(w, scorer.rule_score(rule.weight)),
sibling_product,
);
if scorer.better(new_out, outside[ci]) {
outside[ci] = new_out;
out_heap.push((OrdF64(new_out), ci as u32));
}
}
}
}
OutsideHeuristic {
out: outside,
zero: scorer.zero(),
}
}
}
impl<R: BottomUpTa> IntersectionHeuristic<R> for OutsideHeuristic {
#[inline]
fn outside_estimate(&self, left: StateId, _right: &R::State) -> f64 {
self.out.get(left.index()).copied().unwrap_or(self.zero)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{Explicit, ExplicitBuilder, Symbol};
fn build_grammar() -> (Explicit, [StateId; 5]) {
let mut b = ExplicitBuilder::new();
let s0 = b.new_state(); let s1 = b.new_state(); let s2 = b.new_state(); let s3 = b.new_state(); let s4 = b.new_state();
let a = Symbol(0);
let f = Symbol(1);
let g = Symbol(2);
b.add_weighted_rule(a, vec![], s0, 0.5); b.add_weighted_rule(a, vec![], s2, 0.3); b.add_weighted_rule(f, vec![s0], s1, 0.8); b.add_weighted_rule(g, vec![s1], s4, 0.9);
b.add_accepting(s4);
let grammar = b.build();
(grammar, [s0, s1, s2, s3, s4])
}
#[test]
fn inside_weights_are_correct() {
let (grammar, [s0, s1, s2, s3, s4]) = build_grammar();
let h = OutsideHeuristic::from_grammar(&grammar);
let inside = compute_inside(&grammar);
assert!(
(inside[s0.index()] - 0.5).abs() < 1e-10,
"IN[s0] = {}",
inside[s0.index()]
);
assert!(
(inside[s1.index()] - 0.4).abs() < 1e-10,
"IN[s1] = {}",
inside[s1.index()]
);
assert!(
(inside[s2.index()] - 0.3).abs() < 1e-10,
"IN[s2] = {}",
inside[s2.index()]
);
assert!(
inside[s3.index()] == 0.0,
"IN[s3] should be 0, got {}",
inside[s3.index()]
);
assert!(
(inside[s4.index()] - 0.36).abs() < 1e-10,
"IN[s4] = {}",
inside[s4.index()]
);
let _ = h;
}
#[test]
fn outside_weights_are_correct() {
let (grammar, [s0, s1, s2, s3, s4]) = build_grammar();
let h = OutsideHeuristic::from_grammar(&grammar);
assert!(
(h.out[s4.index()] - 1.0).abs() < 1e-10,
"OUT[s4] = {}",
h.out[s4.index()]
);
assert!(
(h.out[s1.index()] - 0.9).abs() < 1e-10,
"OUT[s1] = {}",
h.out[s1.index()]
);
assert!(
(h.out[s0.index()] - 0.72).abs() < 1e-10,
"OUT[s0] = {}",
h.out[s0.index()]
);
assert!(
h.out[s2.index()] == 0.0,
"OUT[s2] should be 0, got {}",
h.out[s2.index()]
);
assert!(
h.out[s3.index()] == 0.0,
"OUT[s3] should be 0, got {}",
h.out[s3.index()]
);
}
#[test]
fn outside_estimate_returns_out_value() {
let (grammar, [s0, _s1, _s2, _s3, s4]) = build_grammar();
let h = OutsideHeuristic::from_grammar(&grammar);
let dummy_right = StateId(0);
let est_s0: f64 = <OutsideHeuristic as IntersectionHeuristic<Explicit>>::outside_estimate(
&h,
s0,
&dummy_right,
);
let est_s4: f64 = <OutsideHeuristic as IntersectionHeuristic<Explicit>>::outside_estimate(
&h,
s4,
&dummy_right,
);
assert!((est_s0 - 0.72).abs() < 1e-10, "estimate for s0 = {est_s0}");
assert!((est_s4 - 1.0).abs() < 1e-10, "estimate for s4 = {est_s4}");
}
#[test]
fn zero_heuristic_always_returns_one() {
let h = ZeroHeuristic;
let dummy: StateId = StateId(0);
let est: f64 =
<ZeroHeuristic as IntersectionHeuristic<Explicit>>::outside_estimate(&h, dummy, &dummy);
assert_eq!(est, 1.0);
}
#[test]
fn outside_estimate_out_of_range_returns_zero() {
let (grammar, _states) = build_grammar();
let h = OutsideHeuristic::from_grammar(&grammar);
let far = StateId(9999);
let dummy_right = StateId(0);
let est: f64 = <OutsideHeuristic as IntersectionHeuristic<Explicit>>::outside_estimate(
&h,
far,
&dummy_right,
);
assert_eq!(est, 0.0);
}
fn compute_inside(grammar: &Explicit) -> Vec<f64> {
let n = grammar.num_states() as usize;
let mut inside = vec![0.0f64; n];
let mut fin_in = FixedBitSet::with_capacity(n);
let rules: Vec<_> = grammar.rules().collect();
let mut by_child: Vec<Vec<usize>> = vec![Vec::new(); n];
for (idx, rule) in rules.iter().enumerate() {
let mut seen = FixedBitSet::with_capacity(n);
for &child in rule.children {
if !seen.contains(child.index()) {
seen.set(child.index(), true);
by_child[child.index()].push(idx);
}
}
}
let mut pending: Vec<usize> = rules
.iter()
.map(|r| {
let mut seen = FixedBitSet::with_capacity(n);
r.children
.iter()
.filter(|&&c| {
if seen.contains(c.index()) {
false
} else {
seen.set(c.index(), true);
true
}
})
.count()
})
.collect();
let mut partial: Vec<f64> = rules.iter().map(|r| r.weight).collect();
let mut heap: BinaryHeap<(OrdF64, u32)> = BinaryHeap::new();
for rule in rules.iter() {
if rule.children.is_empty() {
let ri = rule.result.index();
if rule.weight > inside[ri] {
inside[ri] = rule.weight;
heap.push((OrdF64(rule.weight), ri as u32));
}
}
}
while let Some((OrdF64(w), si)) = heap.pop() {
let si = si as usize;
if fin_in.contains(si) {
continue;
}
if w < inside[si] - 1e-15 * inside[si].max(1e-15) {
continue;
}
fin_in.set(si, true);
for &rule_idx in &by_child[si] {
partial[rule_idx] *= inside[si];
pending[rule_idx] -= 1;
if pending[rule_idx] == 0 {
let ri = rules[rule_idx].result.index();
let cand = partial[rule_idx];
if cand > inside[ri] {
inside[ri] = cand;
heap.push((OrdF64(cand), ri as u32));
}
}
}
}
inside
}
}