Skip to main content

rusty_alto/
sorted_language.rs

1//! Lazy k-best language iteration for explicit weighted tree automata.
2
3use crate::{Explicit, FxHashSet, StateId, Symbol, TopDownTa};
4use fixedbitset::FixedBitSet;
5use packed_term_arena::tree::{Tree, TreeArena};
6use std::{cmp::Ordering, collections::BinaryHeap, mem};
7
8/// A weighted tree generated from an automaton language.
9///
10/// The tree root refers to the arena owned by the
11/// [`SortedLanguageIterator`] that produced it. This is intentionally a lean
12/// handle: advancing the iterator may invalidate previously returned tree
13/// handles in future implementations. Clone a tree out through
14/// [`SortedLanguageIterator::clone_tree`] before advancing if it must be kept.
15#[derive(Clone, Copy, Debug, PartialEq)]
16pub struct WeightedTree {
17    tree: Tree,
18    weight: f64,
19}
20
21impl WeightedTree {
22    /// Return this tree's root in the producing iterator's arena.
23    pub fn tree(&self) -> Tree {
24        self.tree
25    }
26
27    /// Return the tree weight.
28    pub fn weight(&self) -> f64 {
29        self.weight
30    }
31}
32
33/// Lazily enumerates trees accepted by an [`Explicit`] automaton in descending
34/// weight order.
35///
36/// This is the Rust port of Alto's sorted language iterator. It keeps one
37/// stream per state and one stream per top-down rule. Rule streams contain
38/// unevaluated child-rank tuples and only ask child streams for their k-best
39/// trees when the tuple is needed. Recursive states are guarded by a
40/// per-expansion visiting set, so productive recursive languages can be
41/// enumerated without eagerly materializing the language.
42///
43/// The ordering assumes the usual k-best monotonicity condition: combining a
44/// rule with lower-ranked child trees must not increase the item score. For
45/// the built-in evaluator this is the natural condition for non-negative
46/// multiplicative weights on productive automata.
47pub struct SortedLanguageIterator<'a> {
48    accepting: Vec<StateId>,
49    state_streams: Vec<Option<StateStream>>,
50    rule_streams: Vec<RuleStream>,
51    arena: TreeArena<Symbol>,
52    visiting: FixedBitSet,
53    _automaton: &'a Explicit,
54}
55
56impl Explicit {
57    /// Iterate over accepted trees in descending rule-weight product order.
58    pub fn sorted_language(&self) -> SortedLanguageIterator<'_> {
59        SortedLanguageIterator::new(self)
60    }
61}
62
63impl<'a> SortedLanguageIterator<'a> {
64    /// Create a lazy sorted language iterator for an explicit automaton.
65    pub fn new(automaton: &'a Explicit) -> Self {
66        let mut accepting = Vec::new();
67        automaton.initial_states(&mut |q| accepting.push(q));
68
69        let mut state_streams = Vec::with_capacity(automaton.num_states() as usize);
70        state_streams.resize_with(automaton.num_states() as usize, || None);
71
72        Self {
73            accepting,
74            state_streams,
75            rule_streams: Vec::new(),
76            arena: TreeArena::new(),
77            visiting: FixedBitSet::with_capacity(automaton.num_states() as usize),
78            _automaton: automaton,
79        }
80    }
81
82    /// Return the arena that contains trees produced by this iterator.
83    ///
84    /// [`WeightedTree::tree`] values returned by this iterator are roots in
85    /// this arena. The arena is owned by the iterator so generated subtrees can
86    /// be reused without reference counting or copying.
87    pub fn arena(&self) -> &TreeArena<Symbol> {
88        &self.arena
89    }
90
91    /// Clone a generated tree into a fresh arena.
92    ///
93    /// Use this before advancing the iterator if the tree must be retained
94    /// independently of the iterator.
95    pub fn clone_tree(&self, root: Tree) -> (TreeArena<Symbol>, Tree) {
96        let mut target = TreeArena::new();
97        let root = self.arena.copy_into(root, &mut target);
98        (target, root)
99    }
100
101    fn ensure_state_stream(&mut self, state: StateId) {
102        let idx = state.index();
103        if self.state_streams[idx].is_some() {
104            return;
105        }
106
107        let mut rule_streams = Vec::new();
108        for rule in self._automaton.rules_topdown(state) {
109            let stream_idx = self.rule_streams.len();
110            self.rule_streams.push(RuleStream::new(rule.into()));
111            rule_streams.push(stream_idx);
112        }
113
114        self.state_streams[idx] = Some(StateStream {
115            known: Vec::new(),
116            rule_streams,
117            next_item: 0,
118        });
119    }
120
121    fn state_stream(&self, state: StateId) -> &StateStream {
122        self.state_streams[state.index()]
123            .as_ref()
124            .expect("state stream must be initialized")
125    }
126
127    fn state_stream_mut(&mut self, state: StateId) -> &mut StateStream {
128        self.state_streams[state.index()]
129            .as_mut()
130            .expect("state stream must be initialized")
131    }
132
133    fn state_item(&mut self, state: StateId, k: usize) -> Option<EvaluatedItem> {
134        self.ensure_state_stream(state);
135
136        if let Some(item) = self.state_stream(state).known.get(k) {
137            return Some(item.clone());
138        }
139
140        if k != self.state_stream(state).known.len() || self.visiting.contains(state.index()) {
141            return None;
142        }
143
144        self.visiting.set(state.index(), true);
145        let best_stream = self.best_rule_stream_for_state(state);
146        let best = best_stream.and_then(|stream| self.rule_pop(stream));
147        self.visiting.set(state.index(), false);
148
149        if let Some(item) = best {
150            self.state_stream_mut(state).known.push(item.clone());
151            Some(item)
152        } else {
153            None
154        }
155    }
156
157    fn state_pop_next(&mut self, state: StateId) -> Option<EvaluatedItem> {
158        self.ensure_state_stream(state);
159        let next = self.state_stream(state).next_item;
160        let item = self.state_item(state, next)?;
161        self.state_stream_mut(state).next_item += 1;
162        Some(item)
163    }
164
165    fn state_is_finished(&mut self, state: StateId) -> bool {
166        self.ensure_state_stream(state);
167        if self.visiting.contains(state.index()) {
168            return false;
169        }
170        let streams = self.state_stream(state).rule_streams.clone();
171        streams
172            .into_iter()
173            .all(|stream| self.rule_is_finished(stream))
174    }
175
176    fn best_rule_stream_for_state(&mut self, state: StateId) -> Option<usize> {
177        let streams = self.state_stream(state).rule_streams.clone();
178        streams
179            .into_iter()
180            .filter_map(|stream| self.rule_peek_weight(stream).map(|weight| (stream, weight)))
181            .max_by(|a, b| compare_weight(a.1, b.1))
182            .map(|(stream, _)| stream)
183    }
184
185    fn rule_peek_weight(&mut self, stream: usize) -> Option<f64> {
186        self.evaluate_unevaluated(stream);
187        self.rule_streams[stream]
188            .evaluated
189            .peek()
190            .map(|entry| entry.item.item_weight)
191    }
192
193    fn rule_pop(&mut self, stream: usize) -> Option<EvaluatedItem> {
194        self.evaluate_unevaluated(stream);
195        let item = self.rule_streams[stream].evaluated.pop()?.item;
196        let popped_tuple = item.item.clone();
197        let tree = self.arena.add_node(item.symbol, item.children);
198        let evaluated = EvaluatedItem {
199            tree,
200            tree_weight: item.tree_weight,
201            item_weight: item.item_weight,
202        };
203        self.rule_streams[stream]
204            .pending_variations
205            .push(popped_tuple);
206        Some(evaluated)
207    }
208
209    fn rule_is_finished(&mut self, stream: usize) -> bool {
210        self.evaluate_unevaluated(stream);
211        self.rule_streams[stream].evaluated.is_empty()
212            && self.rule_streams[stream].unevaluated.is_empty()
213    }
214
215    fn evaluate_unevaluated(&mut self, stream: usize) {
216        self.expand_pending_variations(stream);
217        let items = mem::take(&mut self.rule_streams[stream].unevaluated);
218        if items.is_empty() {
219            return;
220        }
221
222        let rule = self.rule_streams[stream].rule.clone();
223        let mut leftovers = Vec::new();
224        let mut evaluated = Vec::new();
225
226        for item in items {
227            if item.rule_position > 0 {
228                continue;
229            }
230
231            let mut children = Vec::with_capacity(rule.children.len());
232            let mut child_weight = 1.0;
233            let mut available = true;
234            let mut keep = true;
235
236            for (&child_state, &rank) in rule.children.iter().zip(&item.child_positions) {
237                if let Some(child_item) = self.state_item(child_state, rank) {
238                    child_weight *= child_item.tree_weight;
239                    children.push(child_item.tree);
240                } else {
241                    available = false;
242                    if self.state_is_finished(child_state) {
243                        keep = false;
244                    }
245                    break;
246                }
247            }
248
249            if available {
250                let tree_weight = rule.weight * child_weight;
251                let eval = ScoredItem {
252                    item,
253                    symbol: rule.symbol,
254                    children,
255                    tree_weight,
256                    item_weight: tree_weight,
257                };
258                evaluated.push(eval);
259            } else if keep {
260                leftovers.push(item);
261            }
262        }
263
264        let rule_stream = &mut self.rule_streams[stream];
265        rule_stream.unevaluated.extend(leftovers);
266        for item in evaluated {
267            let seq = rule_stream.next_seq;
268            rule_stream.next_seq += 1;
269            rule_stream.evaluated.push(HeapItem { item, seq });
270        }
271    }
272
273    fn expand_pending_variations(&mut self, stream: usize) {
274        let pending = mem::take(&mut self.rule_streams[stream].pending_variations);
275        if pending.is_empty() {
276            return;
277        }
278
279        let rule_stream = &mut self.rule_streams[stream];
280        for item in pending {
281            for variation in item.variations() {
282                if rule_stream.discovered.insert(variation.clone()) {
283                    rule_stream.unevaluated.push(variation);
284                }
285            }
286        }
287    }
288}
289
290impl Iterator for SortedLanguageIterator<'_> {
291    type Item = WeightedTree;
292
293    fn next(&mut self) -> Option<Self::Item> {
294        let best_state = self
295            .accepting
296            .clone()
297            .into_iter()
298            .filter_map(|state| {
299                self.ensure_state_stream(state);
300                let next = self.state_stream(state).next_item;
301                self.state_item(state, next)
302                    .map(|item| (state, item.item_weight))
303            })
304            .max_by(|a, b| compare_weight(a.1, b.1))
305            .map(|(state, _)| state)?;
306
307        let item = self.state_pop_next(best_state)?;
308        Some(WeightedTree {
309            tree: item.tree,
310            weight: item.tree_weight,
311        })
312    }
313}
314
315#[derive(Clone, Debug)]
316struct OwnedRule {
317    symbol: Symbol,
318    children: Vec<StateId>,
319    weight: f64,
320}
321
322impl From<crate::Rule<'_>> for OwnedRule {
323    fn from(rule: crate::Rule<'_>) -> Self {
324        Self {
325            symbol: rule.symbol,
326            children: rule.children.to_vec(),
327            weight: rule.weight,
328        }
329    }
330}
331
332#[derive(Clone, Debug)]
333struct StateStream {
334    known: Vec<EvaluatedItem>,
335    rule_streams: Vec<usize>,
336    next_item: usize,
337}
338
339#[derive(Clone, Debug)]
340struct RuleStream {
341    rule: OwnedRule,
342    evaluated: BinaryHeap<HeapItem>,
343    unevaluated: Vec<UnevaluatedItem>,
344    pending_variations: Vec<UnevaluatedItem>,
345    discovered: FxHashSet<UnevaluatedItem>,
346    next_seq: usize,
347}
348
349impl RuleStream {
350    fn new(rule: OwnedRule) -> Self {
351        let zero = UnevaluatedItem {
352            rule_position: 0,
353            child_positions: vec![0; rule.children.len()],
354        };
355        let mut discovered = FxHashSet::default();
356        discovered.insert(zero.clone());
357
358        Self {
359            rule,
360            evaluated: BinaryHeap::new(),
361            unevaluated: vec![zero],
362            pending_variations: Vec::new(),
363            discovered,
364            next_seq: 0,
365        }
366    }
367}
368
369#[derive(Clone, Debug)]
370struct EvaluatedItem {
371    tree: Tree,
372    tree_weight: f64,
373    item_weight: f64,
374}
375
376#[derive(Clone, Debug)]
377struct ScoredItem {
378    item: UnevaluatedItem,
379    symbol: Symbol,
380    children: Vec<Tree>,
381    tree_weight: f64,
382    item_weight: f64,
383}
384
385#[derive(Clone, Debug, PartialEq, Eq, Hash)]
386struct UnevaluatedItem {
387    rule_position: usize,
388    child_positions: Vec<usize>,
389}
390
391impl UnevaluatedItem {
392    fn variations(&self) -> impl Iterator<Item = UnevaluatedItem> + '_ {
393        (0..=self.child_positions.len()).map(|pos| {
394            let mut item = self.clone();
395            if pos == 0 {
396                item.rule_position += 1;
397            } else {
398                item.child_positions[pos - 1] += 1;
399            }
400            item
401        })
402    }
403}
404
405#[derive(Clone, Debug)]
406struct HeapItem {
407    item: ScoredItem,
408    seq: usize,
409}
410
411impl PartialEq for HeapItem {
412    fn eq(&self, other: &Self) -> bool {
413        self.item.item_weight.total_cmp(&other.item.item_weight) == Ordering::Equal
414            && self.seq == other.seq
415    }
416}
417
418impl Eq for HeapItem {}
419
420impl PartialOrd for HeapItem {
421    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
422        Some(self.cmp(other))
423    }
424}
425
426impl Ord for HeapItem {
427    fn cmp(&self, other: &Self) -> Ordering {
428        compare_weight(self.item.item_weight, other.item.item_weight)
429            .then_with(|| other.seq.cmp(&self.seq))
430    }
431}
432
433fn compare_weight(left: f64, right: f64) -> Ordering {
434    left.total_cmp(&right)
435}
436
437#[cfg(test)]
438mod tests {
439    use super::*;
440    use crate::{ExplicitBuilder, Signature};
441
442    fn symbols(names: &[(&str, usize)]) -> Signature {
443        let mut sig = Signature::new();
444        for &(name, arity) in names {
445            sig.intern(name.to_owned(), arity).unwrap();
446        }
447        sig
448    }
449
450    fn show(arena: &TreeArena<Symbol>, tree: WeightedTree, signature: &Signature) -> String {
451        fn rec(arena: &TreeArena<Symbol>, node: Tree, signature: &Signature, out: &mut String) {
452            out.push_str(signature.resolve(*arena.get_label(node)));
453            if !arena.get_children(node).is_empty() {
454                out.push('(');
455                for (idx, &child) in arena.get_children(node).iter().enumerate() {
456                    if idx > 0 {
457                        out.push(',');
458                    }
459                    rec(arena, child, signature, out);
460                }
461                out.push(')');
462            }
463        }
464
465        let mut out = String::new();
466        rec(arena, tree.tree(), signature, &mut out);
467        out
468    }
469
470    #[test]
471    fn enumerates_nonrecursive_language_by_descending_weight() {
472        let sig = symbols(&[("b", 0), ("f", 1), ("g", 1)]);
473        let b = sig.get("b").unwrap();
474        let f = sig.get("f").unwrap();
475        let g = sig.get("g").unwrap();
476
477        let mut builder = ExplicitBuilder::new();
478        let qb = builder.new_state();
479        let qa = builder.new_state();
480        builder.add_weighted_rule(b, vec![], qb, 0.5);
481        builder.add_weighted_rule(f, vec![qb], qa, 0.7);
482        builder.add_weighted_rule(g, vec![qb], qa, 0.3);
483        builder.add_accepting(qa);
484        let automaton = builder.build();
485
486        let mut it = automaton.sorted_language();
487        let first = it.next().unwrap();
488        assert_eq!(show(it.arena(), first, &sig), "f(b)");
489        assert_eq!(first.weight(), 0.35);
490
491        let second = it.next().unwrap();
492        assert_eq!(show(it.arena(), second, &sig), "g(b)");
493        assert_eq!(second.weight(), 0.15);
494
495        assert!(it.next().is_none());
496    }
497
498    #[test]
499    fn handles_recursive_productive_language_lazily() {
500        let sig = symbols(&[("b", 0), ("f", 1)]);
501        let b = sig.get("b").unwrap();
502        let f = sig.get("f").unwrap();
503
504        let mut builder = ExplicitBuilder::new();
505        let q = builder.new_state();
506        builder.add_weighted_rule(b, vec![], q, 0.5);
507        builder.add_weighted_rule(f, vec![q], q, 0.5);
508        builder.add_accepting(q);
509        let automaton = builder.build();
510
511        let mut it = automaton.sorted_language();
512        let first = it.next().unwrap();
513        let second = it.next().unwrap();
514        let third = it.next().unwrap();
515
516        assert_eq!(show(it.arena(), first, &sig), "b");
517        assert_eq!(first.weight(), 0.5);
518        assert_eq!(show(it.arena(), second, &sig), "f(b)");
519        assert_eq!(second.weight(), 0.25);
520        assert_eq!(show(it.arena(), third, &sig), "f(f(b))");
521        assert_eq!(third.weight(), 0.125);
522    }
523
524    #[test]
525    fn merges_multiple_accepting_state_streams() {
526        let sig = symbols(&[("b", 0), ("f", 1), ("g", 1)]);
527        let b = sig.get("b").unwrap();
528        let f = sig.get("f").unwrap();
529        let g = sig.get("g").unwrap();
530
531        let mut builder = ExplicitBuilder::new();
532        let qb = builder.new_state();
533        let qa = builder.new_state();
534        builder.add_weighted_rule(b, vec![], qb, 0.5);
535        builder.add_weighted_rule(f, vec![qb], qb, 0.5);
536        builder.add_weighted_rule(g, vec![qb], qa, 0.4);
537        builder.add_accepting(qb);
538        builder.add_accepting(qa);
539        let automaton = builder.build();
540
541        let mut it = automaton.sorted_language();
542        let mut got = Vec::new();
543        for _ in 0..5 {
544            let tree = it.next().unwrap();
545            got.push((show(it.arena(), tree, &sig), tree.weight()));
546        }
547
548        assert_eq!(
549            got,
550            vec![
551                ("b".to_owned(), 0.5),
552                ("f(b)".to_owned(), 0.25),
553                ("g(b)".to_owned(), 0.2),
554                ("f(f(b))".to_owned(), 0.125),
555                ("g(f(b))".to_owned(), 0.1),
556            ]
557        );
558    }
559
560    #[test]
561    fn empty_language_yields_no_items() {
562        let sig = symbols(&[("g", 2)]);
563        let g = sig.get("g").unwrap();
564
565        let mut builder = ExplicitBuilder::new();
566        let q = builder.new_state();
567        let q1 = builder.new_state();
568        let q2 = builder.new_state();
569        builder.add_weighted_rule(g, vec![q1, q2], q, 1.0);
570        builder.add_accepting(q);
571        let automaton = builder.build();
572
573        assert!(automaton.sorted_language().next().is_none());
574    }
575
576    #[test]
577    fn clones_weighted_tree_to_independent_arena() {
578        let sig = symbols(&[("b", 0), ("f", 1)]);
579        let b = sig.get("b").unwrap();
580        let f = sig.get("f").unwrap();
581
582        let mut builder = ExplicitBuilder::new();
583        let qb = builder.new_state();
584        let qa = builder.new_state();
585        builder.add_rule(b, vec![], qb);
586        builder.add_rule(f, vec![qb], qa);
587        builder.add_accepting(qa);
588        let automaton = builder.build();
589
590        let mut it = automaton.sorted_language();
591        let tree = it.next().unwrap();
592        let (arena, root) = it.clone_tree(tree.tree());
593        assert_eq!(arena.get_label(root), &f);
594        let child = arena.get_children(root)[0];
595        assert_eq!(arena.get_label(child), &b);
596    }
597
598    #[test]
599    fn mirrors_alto_gontrum_recursive_regression() {
600        let sig = symbols(&[("r1", 2), ("r2", 0), ("r3", 1), ("r4", 2), ("r5", 0)]);
601        let r1 = sig.get("r1").unwrap();
602        let r2 = sig.get("r2").unwrap();
603        let r3 = sig.get("r3").unwrap();
604        let r4 = sig.get("r4").unwrap();
605        let r5 = sig.get("r5").unwrap();
606
607        let mut builder = ExplicitBuilder::new();
608        let s = builder.new_state();
609        let a = builder.new_state();
610        let b = builder.new_state();
611        builder.add_weighted_rule(r1, vec![a, b], s, 1.0);
612        builder.add_weighted_rule(r2, vec![], a, 1.0);
613        builder.add_weighted_rule(r3, vec![a], a, 0.0);
614        builder.add_weighted_rule(r4, vec![b, b], b, 0.7);
615        builder.add_weighted_rule(r5, vec![], b, 0.3);
616        builder.add_accepting(s);
617        let automaton = builder.build();
618
619        let mut it = automaton.sorted_language();
620        let first = it.next().unwrap();
621        assert_eq!(show(it.arena(), first, &sig), "r1(r2,r5)");
622        let second = it.next().unwrap();
623        assert_eq!(show(it.arena(), second, &sig), "r1(r2,r4(r5,r5))");
624    }
625}