1use crate::{Explicit, ProbabilityScorer, StateId, Symbol, TopDownTa, WeightScorer};
4use packed_term_arena::tree::{Tree, TreeArena};
5use smallvec::SmallVec;
6
7#[derive(Debug)]
9pub struct ViterbiTree {
10 arena: TreeArena<Symbol>,
11 root: Tree,
12 weight: f64,
13 score: f64,
14}
15
16impl ViterbiTree {
17 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 pub fn arena(&self) -> &TreeArena<Symbol> {
34 &self.arena
35 }
36
37 pub fn root(&self) -> Tree {
39 self.root
40 }
41
42 pub fn weight(&self) -> f64 {
44 self.weight
45 }
46
47 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 pub fn viterbi(&self) -> Option<ViterbiTree> {
74 self.viterbi_with(&ProbabilityScorer)
75 }
76
77 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 #[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}