171 lines
5.4 KiB
Python
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
|