Skip to main content

rusty_alto/
viterbi.rs

1//! One-best Viterbi extraction for explicit weighted tree automata.
2
3use crate::{Explicit, ProbabilityScorer, StateId, Symbol, TopDownTa, WeightScorer};
4use packed_term_arena::tree::{Tree, TreeArena};
5use smallvec::SmallVec;
6
7/// The highest-weighted tree found in an automaton language.
8#[derive(Debug)]
9pub struct ViterbiTree {
10    arena: TreeArena<Symbol>,
11    root: Tree,
12    weight: f64,
13    score: f64,
14}
15
16impl ViterbiTree {
17    /// Construct a `ViterbiTree` with an algorithm score and display weight.
18    pub(crate) fn new_with_score(
19        arena: TreeArena<Symbol>,
20        root: Tree,
21        score: f64,
22        weight: f64,
23    ) -> Self {
24        Self {
25            arena,
26            root,
27            weight,
28            score,
29        }
30    }
31
32    /// Return the arena containing the tree.
33    pub fn arena(&self) -> &TreeArena<Symbol> {
34        &self.arena
35    }
36
37    /// Return the root handle in [`Self::arena`].
38    pub fn root(&self) -> Tree {
39        self.root
40    }
41
42    /// Return the product of rule weights for this tree.
43    pub fn weight(&self) -> f64 {
44        self.weight
45    }
46
47    /// Return the score used to rank this tree.
48    ///
49    /// For [`Explicit::viterbi`] this equals [`Self::weight`]. For
50    /// [`Explicit::viterbi_with`] it is in the scorer's representation, e.g. a
51    /// log probability when using [`crate::LogProbabilityScorer`].
52    pub fn score(&self) -> f64 {
53        self.score
54    }
55}
56
57#[derive(Clone, Debug)]
58pub(crate) struct Backpointer {
59    pub(crate) symbol: Symbol,
60    pub(crate) children: SmallVec<[StateId; 2]>,
61    pub(crate) weight: f64,
62}
63
64impl Explicit {
65    /// Compute the highest-weighted accepted tree.
66    ///
67    /// This is a direct one-best dynamic program for acyclic parse charts. It
68    /// deliberately avoids the k-best sorted-language machinery when callers
69    /// only need the best derivation. Cyclic dependencies on the active DFS
70    /// path are ignored: for PCFG-style rule weights below one, traversing a
71    /// cycle cannot improve a finite derivation. Direct self-loops are skipped
72    /// for the same reason, matching Alto's convention.
73    pub fn viterbi(&self) -> Option<ViterbiTree> {
74        self.viterbi_with(&ProbabilityScorer)
75    }
76
77    /// Compute the highest-scoring accepted tree under `scorer`.
78    pub fn viterbi_with<S: WeightScorer>(&self, scorer: &S) -> Option<ViterbiTree> {
79        let mut marks = vec![0u8; self.num_states() as usize];
80        let mut best = vec![None::<Backpointer>; self.num_states() as usize];
81        let mut stack = Vec::new();
82        self.initial_states(&mut |state| {
83            visit_and_score(self, state, scorer, &mut marks, &mut best, &mut stack);
84        });
85        finish_best(self, scorer, &best)
86    }
87
88    /// Previous Viterbi implementation, retained only for performance comparisons.
89    #[cfg(feature = "viterbi-benchmark")]
90    #[doc(hidden)]
91    pub fn viterbi_old_benchmark(&self) -> Option<ViterbiTree> {
92        let scorer = ProbabilityScorer;
93        let mut order = Vec::new();
94        let mut marks = vec![0u8; self.num_states() as usize];
95        self.initial_states(&mut |state| visit_state_fast(self, state, &mut marks, &mut order));
96
97        let mut best = vec![None::<Backpointer>; self.num_states() as usize];
98        for state in order {
99            best[state.index()] = score_state(self, state, &scorer, &best);
100        }
101        finish_best(self, &scorer, &best)
102    }
103}
104
105fn visit_and_score<S: WeightScorer>(
106    auto: &Explicit,
107    start: StateId,
108    scorer: &S,
109    marks: &mut [u8],
110    best: &mut [Option<Backpointer>],
111    stack: &mut Vec<usize>,
112) {
113    if start.is_stuck() || start.index() >= marks.len() || marks[start.index()] != 0 {
114        return;
115    }
116
117    stack.clear();
118    stack.push(start.index() << 1);
119
120    while let Some(frame) = stack.pop() {
121        let state = StateId((frame >> 1) as u32);
122        if frame & 1 != 0 {
123            best[state.index()] = score_state(auto, state, scorer, best);
124            marks[state.index()] = 2;
125            continue;
126        }
127        if marks[state.index()] != 0 {
128            continue;
129        }
130
131        marks[state.index()] = 1;
132        stack.push((state.index() << 1) | 1);
133        for &rule_idx in auto.rule_indexes_topdown(state).iter().rev() {
134            let rule = auto.rule(rule_idx);
135            if rule.children.contains(&state) {
136                continue;
137            }
138            for &child in rule.children.iter().rev() {
139                if !child.is_stuck() && child.index() < marks.len() && marks[child.index()] == 0 {
140                    stack.push(child.index() << 1);
141                }
142            }
143        }
144    }
145}
146
147#[cfg(feature = "viterbi-benchmark")]
148fn visit_state_fast(auto: &Explicit, start: StateId, marks: &mut [u8], order: &mut Vec<StateId>) {
149    if start.is_stuck() || start.index() >= marks.len() || marks[start.index()] != 0 {
150        return;
151    }
152    let mut stack = vec![(start, false)];
153    marks[start.index()] = 1;
154    while let Some((state, exiting)) = stack.pop() {
155        if exiting {
156            marks[state.index()] = 2;
157            order.push(state);
158            continue;
159        }
160        stack.push((state, true));
161        for &rule_idx in auto.rule_indexes_topdown(state).iter().rev() {
162            let rule = auto.rule(rule_idx);
163            if rule.children.contains(&state) {
164                continue;
165            }
166            for &child in rule.children.iter().rev() {
167                if !child.is_stuck() && child.index() < marks.len() && marks[child.index()] == 0 {
168                    marks[child.index()] = 1;
169                    stack.push((child, false));
170                }
171            }
172        }
173    }
174}
175
176fn score_state<S: WeightScorer>(
177    auto: &Explicit,
178    state: StateId,
179    scorer: &S,
180    best: &[Option<Backpointer>],
181) -> Option<Backpointer> {
182    let mut best_here = None::<Backpointer>;
183    for rule in auto.rules_topdown(state) {
184        if rule.children.contains(&state) {
185            continue;
186        }
187
188        let mut weight = scorer.rule_score(rule.weight);
189        let mut all_children_available = true;
190        for &child in rule.children {
191            let Some(child_best) = best.get(child.index()).and_then(Option::as_ref) else {
192                all_children_available = false;
193                break;
194            };
195            weight = scorer.times(weight, child_best.weight);
196        }
197        if all_children_available
198            && best_here
199                .as_ref()
200                .is_none_or(|old| scorer.better(weight, old.weight))
201        {
202            best_here = Some(Backpointer {
203                symbol: rule.symbol,
204                children: rule.children.iter().copied().collect(),
205                weight,
206            });
207        }
208    }
209    best_here
210}
211
212fn finish_best<S: WeightScorer>(
213    auto: &Explicit,
214    scorer: &S,
215    best: &[Option<Backpointer>],
216) -> Option<ViterbiTree> {
217    let mut best_final = None::<(StateId, f64)>;
218    auto.initial_states(&mut |state| {
219        if let Some(backpointer) = best.get(state.index()).and_then(Option::as_ref)
220            && best_final
221                .is_none_or(|(_, old_weight)| scorer.better(backpointer.weight, old_weight))
222        {
223            best_final = Some((state, backpointer.weight));
224        }
225    });
226
227    let (state, score) = best_final?;
228    let mut arena = TreeArena::new();
229    let root = build_tree(state, best, &mut arena)?;
230    Some(ViterbiTree::new_with_score(
231        arena,
232        root,
233        score,
234        scorer.score_to_weight(score),
235    ))
236}
237
238pub(crate) fn build_tree(
239    state: StateId,
240    best: &[Option<Backpointer>],
241    arena: &mut TreeArena<Symbol>,
242) -> Option<Tree> {
243    let backpointer = best.get(state.index())?.as_ref()?;
244    let children = backpointer
245        .children
246        .iter()
247        .map(|&child| build_tree(child, best, arena))
248        .collect::<Option<Vec<_>>>()?;
249    Some(arena.add_node(backpointer.symbol, children))
250}
251
252pub(crate) fn build_tree_from_arena(
253    state: StateId,
254    backpointer_ids: &[Option<u32>],
255    backpointers: &[Backpointer],
256    arena: &mut TreeArena<Symbol>,
257) -> Option<Tree> {
258    let id = backpointer_ids.get(state.index())?.as_ref()?;
259    let backpointer = backpointers.get(*id as usize)?;
260    let children = backpointer
261        .children
262        .iter()
263        .map(|&child| build_tree_from_arena(child, backpointer_ids, backpointers, arena))
264        .collect::<Option<Vec<_>>>()?;
265    Some(arena.add_node(backpointer.symbol, children))
266}
267
268#[cfg(test)]
269mod tests {
270    use super::*;
271    use crate::ExplicitBuilder;
272
273    #[test]
274    fn chooses_highest_weighted_tree() {
275        let a = Symbol(0);
276        let b = Symbol(1);
277        let f = Symbol(2);
278
279        let mut builder = ExplicitBuilder::new();
280        let qa = builder.new_state();
281        let qb = builder.new_state();
282        let root = builder.new_state();
283        builder.add_weighted_rule(a, vec![], qa, 0.3);
284        builder.add_weighted_rule(b, vec![], qb, 0.8);
285        builder.add_weighted_rule(f, vec![qa], root, 0.9);
286        builder.add_weighted_rule(f, vec![qb], root, 0.4);
287        builder.add_accepting(root);
288        let automaton = builder.build();
289
290        let best = automaton.viterbi().unwrap();
291        assert!((best.weight() - 0.32).abs() < 1e-12);
292        assert_eq!(*best.arena().get_label(best.root()), f);
293        let child = best.arena().get_children(best.root())[0];
294        assert_eq!(*best.arena().get_label(child), b);
295    }
296
297    #[test]
298    fn returns_none_for_empty_language() {
299        let mut builder = ExplicitBuilder::new();
300        let root = builder.new_state();
301        builder.add_accepting(root);
302        let automaton = builder.build();
303
304        assert!(automaton.viterbi().is_none());
305    }
306
307    #[test]
308    fn preserves_binary_child_order() {
309        let a = Symbol(0);
310        let b = Symbol(1);
311        let f = Symbol(2);
312
313        let mut builder = ExplicitBuilder::new();
314        let qa = builder.new_state();
315        let qb = builder.new_state();
316        let root = builder.new_state();
317        builder.add_weighted_rule(a, vec![], qa, 0.5);
318        builder.add_weighted_rule(b, vec![], qb, 0.5);
319        builder.add_weighted_rule(f, vec![qa, qb], root, 0.5);
320        builder.add_accepting(root);
321        let automaton = builder.build();
322
323        let best = automaton.viterbi().unwrap();
324        assert!((best.weight() - 0.125).abs() < 1e-12);
325        let children = best.arena().get_children(best.root());
326        assert_eq!(children.len(), 2);
327        assert_eq!(*best.arena().get_label(children[0]), a);
328        assert_eq!(*best.arena().get_label(children[1]), b);
329    }
330
331    #[test]
332    fn skips_self_loop_rules_during_iterative_traversal() {
333        let a = Symbol(0);
334        let f = Symbol(1);
335
336        let mut builder = ExplicitBuilder::new();
337        let leaf = builder.new_state();
338        let root = builder.new_state();
339        builder.add_weighted_rule(a, vec![], leaf, 0.7);
340        builder.add_weighted_rule(f, vec![leaf], root, 0.5);
341        builder.add_weighted_rule(f, vec![root], root, 100.0);
342        builder.add_accepting(root);
343        let automaton = builder.build();
344
345        let best = automaton.viterbi().unwrap();
346        assert!((best.weight() - 0.35).abs() < 1e-12);
347        assert_eq!(*best.arena().get_label(best.root()), f);
348        let child = best.arena().get_children(best.root())[0];
349        assert_eq!(*best.arena().get_label(child), a);
350    }
351
352    #[test]
353    fn shared_dependency_is_scored_before_all_parents() {
354        let leaf_symbol = Symbol(0);
355        let unary_symbol = Symbol(1);
356        let root_symbol = Symbol(2);
357
358        let mut builder = ExplicitBuilder::new();
359        let shared = builder.new_state();
360        let left = builder.new_state();
361        let root = builder.new_state();
362        builder.add_weighted_rule(leaf_symbol, vec![], shared, 0.8);
363        builder.add_weighted_rule(unary_symbol, vec![shared], left, 0.7);
364        builder.add_weighted_rule(root_symbol, vec![left, shared], root, 0.6);
365        builder.add_accepting(root);
366        let automaton = builder.build();
367
368        let best = automaton.viterbi().expect("shared DAG has a derivation");
369        assert!((best.weight() - 0.8 * 0.7 * 0.8 * 0.6).abs() < 1e-12);
370    }
371
372    #[test]
373    fn unproductive_nontrivial_cycles_have_no_derivation() {
374        let f = Symbol(0);
375        let g = Symbol(1);
376        let mut builder = ExplicitBuilder::new();
377        let q0 = builder.new_state();
378        let q1 = builder.new_state();
379        builder.add_weighted_rule(f, vec![q1], q0, 0.5);
380        builder.add_weighted_rule(g, vec![q0], q1, 0.5);
381        builder.add_accepting(q0);
382        let automaton = builder.build();
383
384        assert!(automaton.viterbi().is_none());
385    }
386
387    #[test]
388    fn productive_nontrivial_cycle_uses_acyclic_exit() {
389        let leaf_symbol = Symbol(0);
390        let forward = Symbol(1);
391        let backward = Symbol(2);
392        let mut builder = ExplicitBuilder::new();
393        let q0 = builder.new_state();
394        let q1 = builder.new_state();
395        builder.add_weighted_rule(leaf_symbol, vec![], q1, 0.7);
396        builder.add_weighted_rule(forward, vec![q1], q0, 0.8);
397        builder.add_weighted_rule(backward, vec![q0], q1, 0.9);
398        builder.add_accepting(q0);
399        let automaton = builder.build();
400
401        let best = automaton.viterbi().expect("cycle has a productive exit");
402        assert!((best.weight() - 0.56).abs() < 1e-12);
403        assert_eq!(*best.arena().get_label(best.root()), forward);
404    }
405
406    #[test]
407    fn matches_sorted_language_on_shared_acyclic_automata() {
408        for width in 2..8 {
409            let mut builder = ExplicitBuilder::new();
410            let mut states = Vec::new();
411            for _ in 0..width {
412                states.push(builder.new_state());
413            }
414            builder.add_weighted_rule(Symbol(0), vec![], states[0], 0.91);
415            for i in 1..width {
416                builder.add_weighted_rule(
417                    Symbol((2 * i) as u32),
418                    vec![states[i - 1]],
419                    states[i],
420                    0.8 - i as f64 * 0.01,
421                );
422                builder.add_weighted_rule(
423                    Symbol((2 * i + 1) as u32),
424                    vec![states[i - 1], states[0]],
425                    states[i],
426                    0.7 - i as f64 * 0.01,
427                );
428            }
429            builder.add_accepting(states[width - 1]);
430            if width > 3 {
431                builder.add_accepting(states[width - 2]);
432            }
433            let automaton = builder.build();
434
435            let viterbi = automaton.viterbi().unwrap();
436            let sorted = automaton.sorted_language().next().unwrap();
437            assert!((viterbi.weight() - sorted.weight()).abs() < 1e-12);
438        }
439    }
440
441    #[test]
442    fn log_scorer_keeps_underflowed_derivation_orderable() {
443        let a = Symbol(0);
444        let f = Symbol(1);
445
446        let mut builder = ExplicitBuilder::new();
447        let mut states = Vec::new();
448        for _ in 0..220 {
449            states.push(builder.new_state());
450        }
451
452        builder.add_weighted_rule(a, vec![], states[0], 0.01);
453        for i in 1..states.len() {
454            builder.add_weighted_rule(f, vec![states[i - 1]], states[i], 0.01);
455        }
456        builder.add_accepting(*states.last().unwrap());
457        let automaton = builder.build();
458
459        let best_prob = automaton.viterbi().unwrap();
460        assert_eq!(best_prob.weight(), 0.0);
461
462        let scorer = crate::LogProbabilityScorer;
463        let best_log = automaton.viterbi_with(&scorer).unwrap();
464        assert!(best_log.score().is_finite());
465        assert_eq!(best_log.weight(), 0.0);
466    }
467}