During the AST migration, _count_concat lost its @lru_cache (it was replaced by an uncached recursive version). Distributing a length L across concat parts then revisited the same (remaining_parts, length) states exponentially -> 16M function calls for a single 4-method group, and several RAGSAK groups took 18-52s (looked hung at the tail). Restore memoization: _count_concat is now @lru_cache-keyed on (tuple(parts), length). Also fix a stray blank line between the @lru_cache decorator and count_words. Impact (RAGSAK, --slice package --min-structure 0.5): straggler groups 18-52s -> <0.5s; full run 54.9s -> 6.7s. Adds a regression test asserting a length-20 (a|b)* string counts in <2s.
269 lines
8.1 KiB
Python
269 lines
8.1 KiB
Python
"""AST — canonical grammar representation.
|
|
|
|
Node types: Symbol, Concat, Alt, Plus, Optional, Star, Epsilon, Empty.
|
|
AST is the ONLY representation. No SORE strings exist anywhere.
|
|
"""
|
|
|
|
import math
|
|
from functools import lru_cache
|
|
|
|
|
|
class Symbol:
|
|
__slots__ = ('value',)
|
|
def __init__(self, value): self.value = value
|
|
def __eq__(self, other): return isinstance(other, Symbol) and self.value == other.value
|
|
def __hash__(self): return hash(('Sym', self.value))
|
|
def __repr__(self): return f"Symbol({self.value!r})"
|
|
|
|
|
|
class Concat:
|
|
__slots__ = ('parts',)
|
|
def __init__(self, parts): self.parts = list(parts)
|
|
def __eq__(self, other): return isinstance(other, Concat) and self.parts == other.parts
|
|
def __hash__(self): return hash(('Concat', tuple(self.parts)))
|
|
def __repr__(self): return f"Concat({self.parts!r})"
|
|
|
|
|
|
class Alt:
|
|
__slots__ = ('parts',)
|
|
def __init__(self, parts): self.parts = list(parts)
|
|
def __eq__(self, other): return isinstance(other, Alt) and self.parts == other.parts
|
|
def __hash__(self): return hash(('Alt', tuple(self.parts)))
|
|
def __repr__(self): return f"Alt({self.parts!r})"
|
|
|
|
|
|
class Plus:
|
|
__slots__ = ('child',)
|
|
def __init__(self, child): self.child = child
|
|
def __eq__(self, other): return isinstance(other, Plus) and self.child == other.child
|
|
def __hash__(self): return hash(('Plus', self.child))
|
|
def __repr__(self): return f"Plus({self.child!r})"
|
|
|
|
|
|
class Optional:
|
|
__slots__ = ('child',)
|
|
def __init__(self, child): self.child = child
|
|
def __eq__(self, other): return isinstance(other, Optional) and self.child == other.child
|
|
def __hash__(self): return hash(('Optional', self.child))
|
|
def __repr__(self): return f"Optional({self.child!r})"
|
|
|
|
|
|
class Star:
|
|
__slots__ = ('child',)
|
|
def __init__(self, child): self.child = child
|
|
def __eq__(self, other): return isinstance(other, Star) and self.child == other.child
|
|
def __hash__(self): return hash(('Star', self.child))
|
|
def __repr__(self): return f"Star({self.child!r})"
|
|
|
|
|
|
class Epsilon:
|
|
__slots__ = ()
|
|
def __eq__(self, other): return isinstance(other, Epsilon)
|
|
def __hash__(self): return hash('Epsilon')
|
|
def __repr__(self): return 'Epsilon()'
|
|
|
|
|
|
class Empty:
|
|
__slots__ = ()
|
|
def __eq__(self, other): return isinstance(other, Empty)
|
|
def __hash__(self): return hash('Empty')
|
|
def __repr__(self): return 'Empty()'
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# AST operations
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def alphabet(node):
|
|
"""Collect all Symbol values from an AST."""
|
|
if isinstance(node, Symbol):
|
|
return {node.value}
|
|
if isinstance(node, (Epsilon, Empty)):
|
|
return set()
|
|
if isinstance(node, (Plus, Optional, Star)):
|
|
return alphabet(node.child)
|
|
if isinstance(node, (Concat, Alt)):
|
|
result = set()
|
|
for p in node.parts:
|
|
result |= alphabet(p)
|
|
return result
|
|
return set()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Matching
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def match(node, seq):
|
|
"""Check if seq matches the grammar defined by node."""
|
|
ends = _match_set(node, seq, 0)
|
|
return len(seq) in ends
|
|
|
|
|
|
def _match_set(node, seq, pos):
|
|
"""Return set of positions reachable from pos after matching node."""
|
|
if isinstance(node, Symbol):
|
|
if pos < len(seq) and seq[pos] == node.value:
|
|
return {pos + 1}
|
|
return set()
|
|
if isinstance(node, Epsilon):
|
|
return {pos}
|
|
if isinstance(node, Empty):
|
|
return set()
|
|
if isinstance(node, Concat):
|
|
current = {pos}
|
|
for part in node.parts:
|
|
next_set = set()
|
|
for p in current:
|
|
next_set |= _match_set(part, seq, p)
|
|
current = next_set
|
|
if not current:
|
|
break
|
|
return current
|
|
if isinstance(node, Alt):
|
|
result = set()
|
|
for part in node.parts:
|
|
result |= _match_set(part, seq, pos)
|
|
return result
|
|
if isinstance(node, Plus):
|
|
return _match_rep(node.child, seq, pos, min_rep=1)
|
|
if isinstance(node, Optional):
|
|
return _match_set(node.child, seq, pos) | {pos}
|
|
if isinstance(node, Star):
|
|
return _match_rep(node.child, seq, pos, min_rep=0)
|
|
return set()
|
|
|
|
|
|
def _match_rep(child, seq, pos, min_rep):
|
|
"""Match child repeated min_rep or more times."""
|
|
if min_rep == 0:
|
|
accept = {pos}
|
|
else:
|
|
accept = set()
|
|
current = {pos}
|
|
for _ in range(min_rep):
|
|
next_set = set()
|
|
for p in current:
|
|
next_set |= _match_set(child, seq, p)
|
|
current = next_set
|
|
if not current:
|
|
break
|
|
if min_rep == 0:
|
|
accept |= current
|
|
seen = set()
|
|
frontier = current
|
|
while frontier:
|
|
frontier_next = set()
|
|
for p in frontier:
|
|
if p in seen:
|
|
continue
|
|
seen.add(p)
|
|
accept.add(p)
|
|
frontier_next |= _match_set(child, seq, p)
|
|
frontier = frontier_next - seen
|
|
if min_rep > 0:
|
|
accept |= current
|
|
return accept
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Counting (for MDL scoring)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
_COUNT_CAP = 10**12
|
|
|
|
@lru_cache(maxsize=None)
|
|
def count_words(node, length):
|
|
"""Count how many words of exactly `length` are in L(node).
|
|
|
|
Capped at _COUNT_CAP to prevent combinatorial explosion on
|
|
deeply nested CRX grammars with large alphabets.
|
|
"""
|
|
if length < 0:
|
|
return 0
|
|
if isinstance(node, Symbol):
|
|
return 1 if length == 1 else 0
|
|
if isinstance(node, Epsilon):
|
|
return 1 if length == 0 else 0
|
|
if isinstance(node, Empty):
|
|
return 0
|
|
if isinstance(node, Concat):
|
|
return _count_concat(tuple(node.parts), length)
|
|
if isinstance(node, Alt):
|
|
total = 0
|
|
for p in node.parts:
|
|
total += count_words(p, length)
|
|
if total >= _COUNT_CAP:
|
|
return _COUNT_CAP
|
|
return total
|
|
if isinstance(node, Plus):
|
|
return _count_rep(node.child, length, 1)
|
|
if isinstance(node, Optional):
|
|
return count_words(node.child, length) + (1 if length == 0 else 0)
|
|
if isinstance(node, Star):
|
|
return _count_rep(node.child, length, 0)
|
|
return 0
|
|
|
|
|
|
@lru_cache(maxsize=None)
|
|
def _count_concat(parts, length):
|
|
if not parts:
|
|
return 1 if length == 0 else 0
|
|
first = parts[0]
|
|
rest = parts[1:]
|
|
total = 0
|
|
for take in range(length + 1):
|
|
cnt = count_words(first, take)
|
|
if cnt:
|
|
total += cnt * _count_concat(rest, length - take)
|
|
if total >= _COUNT_CAP:
|
|
return _COUNT_CAP
|
|
return total
|
|
|
|
|
|
@lru_cache(maxsize=None)
|
|
def _count_rep(child, length, min_rep):
|
|
total = 0
|
|
for rep in range(min_rep, length + 1):
|
|
total += _count_repeat(child, rep, length)
|
|
if total >= _COUNT_CAP:
|
|
return _COUNT_CAP
|
|
return total
|
|
|
|
|
|
@lru_cache(maxsize=None)
|
|
def _count_repeat(child, rep, length):
|
|
if rep == 0:
|
|
return 1 if length == 0 else 0
|
|
total = 0
|
|
for take in range(length + 1):
|
|
cnt = count_words(child, take)
|
|
if cnt:
|
|
total += cnt * _count_repeat(child, rep - 1, length - take)
|
|
if total >= _COUNT_CAP:
|
|
return _COUNT_CAP
|
|
return total
|
|
|
|
|
|
def lang_size(node, n=None):
|
|
"""|L(r)≤n| — number of words of length ≤ n."""
|
|
if isinstance(node, Empty):
|
|
return 0
|
|
if isinstance(node, Epsilon):
|
|
return 1
|
|
if n is None:
|
|
n = 2 * model_cost(node) + 1
|
|
return sum(count_words(node, l) for l in range(n + 1))
|
|
|
|
|
|
def model_cost(node):
|
|
"""|r| — number of alphabet symbol occurrences in expression."""
|
|
if isinstance(node, Symbol):
|
|
return 1
|
|
if isinstance(node, (Epsilon, Empty)):
|
|
return 0
|
|
if isinstance(node, (Plus, Optional, Star)):
|
|
return model_cost(node.child)
|
|
if isinstance(node, (Concat, Alt)):
|
|
return sum(model_cost(p) for p in node.parts)
|
|
return 0
|