1use 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#[derive(Clone, Copy, Debug, PartialEq)]
16pub struct WeightedTree {
17 tree: Tree,
18 weight: f64,
19}
20
21impl WeightedTree {
22 pub fn tree(&self) -> Tree {
24 self.tree
25 }
26
27 pub fn weight(&self) -> f64 {
29 self.weight
30 }
31}
32
33pub 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 pub fn sorted_language(&self) -> SortedLanguageIterator<'_> {
59 SortedLanguageIterator::new(self)
60 }
61}
62
63impl<'a> SortedLanguageIterator<'a> {
64 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 pub fn arena(&self) -> &TreeArena<Symbol> {
88 &self.arena
89 }
90
91 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}