grammar-inference-engine/bex/crx.py

171 lines
5.4 KiB
Python

"""CRX — Direct CHARE inference (Algorithm 7, TODS 2010).
Produces AST nodes directly — no SORE string intermediate.
"""
from collections import defaultdict
from .grammar import Symbol, Concat, Alt, Plus, Optional, Star, Epsilon, Empty
class CRX:
"""
|———— Algorithm 7: CRX ————|
Input: sample S (list of token lists)
Output: AST node r such that S ⊆ L(r)
"""
def infer(self, sequences):
S = [list(s) for s in sequences if s]
if not S:
return Epsilon()
sigma = set()
for w in S:
for a in w:
sigma.add(a)
if not sigma:
return Epsilon()
immed = set()
for w in S:
for i in range(len(w) - 1):
immed.add((w[i], w[i + 1]))
closure = self._transitive_closure(sigma, immed)
eq = self._equivalence(sigma, closure)
sym_to_cls = {}
classes = []
for cls_syms in eq:
idx = len(classes)
for sym in cls_syms:
sym_to_cls[sym] = idx
classes.append(set(cls_syms))
changed = True
while changed:
changed = False
singleton_ids = [i for i, c in enumerate(classes) if len(c) == 1]
hs_pred = {}
hs_succ = {}
for i in singleton_ids:
hs_pred[i] = set()
hs_succ[i] = set()
sym_i = next(iter(classes[i]))
for j, c in enumerate(classes):
if i == j:
continue
if any((sym_j, sym_i) in immed for sym_j in c):
hs_pred[i].add(j)
if any((sym_i, sym_j) in immed for sym_j in c):
hs_succ[i].add(j)
groups = defaultdict(list)
for i in singleton_ids:
groups[(frozenset(hs_pred[i]), frozenset(hs_succ[i]))].append(i)
for (pred_set, succ_set), group in groups.items():
if len(group) >= 2:
merged = set()
for i in group:
merged.update(classes[i])
new_id = len(classes)
classes.append(merged)
for i in sorted(group, reverse=True):
classes.pop(i)
changed = True
break
sym_to_cls = {}
for idx, cls in enumerate(classes):
for sym in cls:
sym_to_cls[sym] = idx
adj = {i: set() for i in range(len(classes))}
indeg = {i: 0 for i in range(len(classes))}
for a, b in immed:
ca, cb = sym_to_cls.get(a), sym_to_cls.get(b)
if ca is not None and cb is not None and ca != cb:
if cb not in adj[ca]:
adj[ca].add(cb)
indeg[cb] += 1
order = []
q = [i for i in range(len(classes)) if indeg[i] == 0]
while q:
i = q.pop(0)
order.append(i)
for j in adj[i]:
indeg[j] -= 1
if indeg[j] == 0:
q.append(j)
remaining = set(range(len(classes))) - set(order)
order.extend(remaining)
def count_in_class(w, syms):
return sum(1 for a in w if a in syms)
parts = []
for i in order:
syms = classes[i]
counts = [count_in_class(w, syms) for w in S]
all_exactly_one = all(c == 1 for c in counts)
all_at_most_one = all(c <= 1 for c in counts)
all_at_least_one = all(c >= 1 for c in counts)
some_two_or_more = any(c >= 2 for c in counts)
sym_list = sorted(syms)
if len(sym_list) > 1:
alt_node = Alt([Symbol(s) for s in sym_list])
else:
alt_node = Symbol(sym_list[0])
if all_exactly_one:
parts.append(alt_node)
elif all_at_most_one:
parts.append(Optional(alt_node))
elif all_at_least_one and some_two_or_more:
parts.append(Plus(alt_node))
else:
parts.append(Plus(Optional(alt_node)))
if not parts:
return Epsilon()
if len(parts) == 1:
return parts[0]
return Concat(parts)
def _transitive_closure(self, sigma, immed):
closure = {(a, b) for (a, b) in immed}
for a in sigma:
closure.add((a, a))
changed = True
while changed:
changed = False
for a in sigma:
for b in sigma:
for c in sigma:
if (a, b) in closure and (b, c) in closure and (a, c) not in closure:
closure.add((a, c))
changed = True
return closure
def _equivalence(self, sigma, closure):
remaining = set(sigma)
classes = []
while remaining:
a = remaining.pop()
cls = {a}
added = True
while added:
added = False
for b in list(remaining):
if (a, b) in closure and (b, a) in closure:
if b not in cls:
cls.add(b)
remaining.discard(b)
added = True
classes.append(cls)
return classes