drop kORE from default ensemble, add --kore flag to opt in

- Remove kORE from top-level imports, _ALGORITHMS, and infer_ensemble body
- Add include_kore=False parameter, runs kORE only when opted in
- Add --kore CLI flag threaded through all pipeline layers
- Remove koreinference from --prefer choices
- Keep kore.py module in repo for reference/future use
- Update tests: remove test_prefer_koreinference, update algorithm assertions
This commit is contained in:
tobjend 2026-07-04 02:58:07 +02:00
parent e3ad256321
commit 710da56916
3 changed files with 65 additions and 57 deletions

View file

@ -3,7 +3,6 @@
import re import re
from .crx import CRX from .crx import CRX
from .idregex import idregex from .idregex import idregex
from .kore import kOREInference
from .expr import alphabet from .expr import alphabet
from .mdl import model_cost, mdl_score from .mdl import model_cost, mdl_score
@ -375,8 +374,21 @@ def _run_idregex(sequences, kmax, N):
return None, float('inf') return None, float('inf')
_ALGO_NAMES = {
'crx': 'CRX',
'idregex': 'iDRegEx',
}
_ALGORITHMS = {
'crx': lambda s, k, n: (CRX().infer(s), mdl_score_simple(CRX().infer(s), s)),
'idregex': _run_idregex,
}
def _run_kore(sequences, kmax, N): def _run_kore(sequences, kmax, N):
"""Run kOREInference (Algorithm 4 with MDL), return (grammar, score) or (None, inf).""" """Run kOREInference, return (grammar, score) or (None, inf)."""
from .kore import kOREInference
kore = kOREInference(k_max=kmax, N=N) kore = kOREInference(k_max=kmax, N=N)
result = kore.infer(sequences) result = kore.infer(sequences)
if result: if result:
@ -385,21 +397,7 @@ def _run_kore(sequences, kmax, N):
return None, float('inf') return None, float('inf')
_ALGO_NAMES = { def infer_ensemble(sequences, kmax=2, N=3, prefer=None, min_coverage=1.0, include_kore=False):
'crx': 'CRX',
'idregex': 'iDRegEx',
'koreinference': 'kOREInference',
}
_ALGORITHMS = {
'crx': lambda s, k, n: (CRX().infer(s), mdl_score_simple(CRX().infer(s), s)),
'idregex': _run_idregex,
'koreinference': _run_kore,
}
def infer_ensemble(sequences, kmax=2, N=3, prefer=None, min_coverage=1.0):
"""Run all applicable algorithms and return the best by MDL score. """Run all applicable algorithms and return the best by MDL score.
Args: Args:
@ -450,10 +448,11 @@ def infer_ensemble(sequences, kmax=2, N=3, prefer=None, min_coverage=1.0):
if idr_g: if idr_g:
results.append(('iDRegEx', idr_g, idr_score)) results.append(('iDRegEx', idr_g, idr_score))
# 3. kOREInference (Algorithm 4 with MDL scoring) # 3. kOREInference (opt-in via include_kore=True)
kore_g, kore_score = _run_kore(sequences, kmax, N) if include_kore:
if kore_g: kore_g, kore_score = _run_kore(sequences, kmax, N)
results.append(('kOREInference', kore_g, kore_score)) if kore_g:
results.append(('kOREInference', kore_g, kore_score))
results = [r for r in results if r[1] and r[1] != ''] results = [r for r in results if r[1] and r[1] != '']
if not results: if not results:

View file

@ -14,6 +14,7 @@ import sys
import time import time
from pathlib import Path from pathlib import Path
from collections import Counter from collections import Counter
from concurrent.futures import ProcessPoolExecutor, as_completed
import pathspec import pathspec
@ -203,7 +204,7 @@ def frequency_filter(sequences, min_coverage=0.2):
return filtered return filtered
def analyze_clusters(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3): def analyze_clusters(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, include_kore=False):
"""Run full pipeline: preprocess → frequency filter → ensemble infer. """Run full pipeline: preprocess → frequency filter → ensemble infer.
Returns: Returns:
@ -231,13 +232,25 @@ def analyze_clusters(file_paths, extension, project_root="", min_coverage=DEFAUL
packages = _top_packages(cluster_fps, project_root) packages = _top_packages(cluster_fps, project_root)
symbol_seqs = [[text for _, text, _ in seq] for seq in sequences] symbol_seqs = [[text for _, text, _ in seq] for seq in sequences]
result = infer_ensemble(symbol_seqs, kmax=kmax, N=N, prefer=prefer, min_coverage=min_coverage) result = infer_ensemble(symbol_seqs, kmax=kmax, N=N, prefer=prefer, min_coverage=min_coverage, include_kore=include_kore)
meta = {"files": cluster_fps, "imports": imports, "arg_patterns": arg_patterns, "packages": packages} meta = {"files": cluster_fps, "imports": imports, "arg_patterns": arg_patterns, "packages": packages}
return [("(all methods)", result, len(sequences), meta)] return [("(all methods)", result, len(sequences), meta)]
def analyze_by_package(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, min_pkg_size=3): def _infer_group(label, group_seqs, group_files, project_root, min_coverage, prefer, kmax, N, include_kore=False):
"""Infer grammar for one package group. Module-level for ProcessPoolExecutor."""
filtered = frequency_filter(group_seqs, min_coverage=0.2)
imports = _extract_imports(group_files)
arg_patterns = _build_arg_patterns(group_files)
packages = _top_packages(group_files, project_root)
symbol_seqs = [[text for _, text, _ in seq] for seq in filtered]
result = infer_ensemble(symbol_seqs, kmax=kmax, N=N, prefer=prefer, min_coverage=min_coverage, include_kore=include_kore)
meta = {"files": group_files, "imports": imports, "arg_patterns": arg_patterns, "packages": packages}
return (label, result, len(filtered), meta)
def analyze_by_package(file_paths, extension, project_root="", min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, min_pkg_size=3, include_kore=False):
"""Preprocess and group by package directory, infer per group. """Preprocess and group by package directory, infer per group.
Groups methods by their file's relative directory path, merging Groups methods by their file's relative directory path, merging
@ -273,25 +286,23 @@ def analyze_by_package(file_paths, extension, project_root="", min_coverage=DEFA
_vprint(f" └ (other) ({len(ungrouped)} methods)") _vprint(f" └ (other) ({len(ungrouped)} methods)")
results = [] results = []
for label, indices in groups: n_workers = os.cpu_count()
t1 = time.time() _vprint(f"Inferring {len(groups)} groups across {n_workers} workers ...")
group_seqs = [sequences[i] for i in indices] with ProcessPoolExecutor(max_workers=n_workers) as ex:
group_files = set(seq_files[i] for i in indices) futures = {}
for label, indices in groups:
gs = [sequences[i] for i in indices]
gf = set(seq_files[i] for i in indices)
f = ex.submit(_infer_group, label, gs, gf, project_root,
min_coverage, prefer, kmax, N, include_kore)
futures[f] = label
group_seqs = frequency_filter(group_seqs, min_coverage=0.2) for f in as_completed(futures):
label, result, count, meta = f.result()
results.append((label, result, count, meta))
_vprint(f"Infer {label} ({count} methods) done")
imports = _extract_imports(group_files) results.sort(key=lambda x: x[0])
arg_patterns = _build_arg_patterns(group_files)
packages = _top_packages(group_files, project_root)
group_prefer = prefer
_vprint(f"Infer {label} ({len(group_seqs)} methods, full ensemble) ... ", end="")
symbol_seqs = [[text for _, text, _ in seq] for seq in group_seqs]
result = infer_ensemble(symbol_seqs, kmax=kmax, N=N, prefer=group_prefer, min_coverage=min_coverage)
_vprint(f"done ({time.time()-t1:.1f}s)")
meta = {"files": group_files, "imports": imports, "arg_patterns": arg_patterns, "packages": packages}
results.append((label, result, len(group_seqs), meta))
if ungrouped: if ungrouped:
ungrouped_files = set(seq_files[i] for i in ungrouped) ungrouped_files = set(seq_files[i] for i in ungrouped)
@ -301,7 +312,7 @@ def analyze_by_package(file_paths, extension, project_root="", min_coverage=DEFA
return results return results
def infer(file_paths, extension, min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3): def infer(file_paths, extension, min_coverage=DEFAULT_COVERAGE, prefer=None, kmax=2, N=3, include_kore=False):
"""Run full pipeline: preprocess → frequency filter → ensemble infer. """Run full pipeline: preprocess → frequency filter → ensemble infer.
Args: Args:
@ -397,6 +408,7 @@ def analyze_directory(
slice="flat", slice="flat",
include=None, include=None,
exclude=None, exclude=None,
include_kore=False,
): ):
"""Scan a directory and run analysis for each language found. """Scan a directory and run analysis for each language found.
@ -427,6 +439,7 @@ def analyze_directory(
min_coverage=min_coverage, min_coverage=min_coverage,
prefer=prefer, prefer=prefer,
kmax=kmax, kmax=kmax,
include_kore=include_kore,
) )
else: else:
results[ext] = analyze_clusters( results[ext] = analyze_clusters(
@ -435,6 +448,7 @@ def analyze_directory(
min_coverage=min_coverage, min_coverage=min_coverage,
prefer=prefer, prefer=prefer,
kmax=kmax, kmax=kmax,
include_kore=include_kore,
) )
return results return results
@ -470,9 +484,13 @@ def _parse_args(argv=None):
parser.add_argument("directory", help="Directory to scan") parser.add_argument("directory", help="Directory to scan")
parser.add_argument( parser.add_argument(
"--prefer", "--prefer",
choices=["crx", "idregex", "koreinference"], choices=["crx", "idregex"],
help="Skip ensemble, use only this algorithm", help="Skip ensemble, use only this algorithm",
) )
parser.add_argument(
"--kore", action="store_true",
help="Include kORE in ensemble (off by default for speed)",
)
parser.add_argument( parser.add_argument(
"--kmax", type=int, default=2, "--kmax", type=int, default=2,
help="Maximum k for k-ORE algorithms (default: 2)", help="Maximum k for k-ORE algorithms (default: 2)",
@ -522,6 +540,7 @@ def main():
slice=args.slice, slice=args.slice,
include=args.include, include=args.include,
exclude=args.exclude, exclude=args.exclude,
include_kore=args.kore,
) )
if args.json_flag or args.format == "json": if args.json_flag or args.format == "json":

View file

@ -1,8 +1,7 @@
"""Tests for infer_ensemble — runs CRX, iDRegEx, and kOREInference, picks best by MDL.""" """Tests for infer_ensemble — runs CRX and iDRegEx, picks best by MDL."""
from bex.ensemble import infer_ensemble from bex.ensemble import infer_ensemble
from bex.idregex import is_deterministic from bex.idregex import is_deterministic
from bex.kore import kOREInference
# ── Basic ensemble runs ── # ── Basic ensemble runs ──
@ -21,16 +20,15 @@ def test_ensemble_best_not_none():
result = infer_ensemble(seqs, kmax=2, N=3) result = infer_ensemble(seqs, kmax=2, N=3)
assert result['best'] is not None assert result['best'] is not None
assert result['best']['grammar'] is not None assert result['best']['grammar'] is not None
assert result['best']['algorithm'] in ('CRX', 'iDRegEx', 'kOREInference') assert result['best']['algorithm'] in ('CRX', 'iDRegEx')
assert result['best']['mdl_score'] is not None assert result['best']['mdl_score'] is not None
def test_ensemble_runs_all_three(): def test_ensemble_runs_both():
seqs = [['a', 'b', 'c'], ['a', 'b', 'c', 'd']] seqs = [['a', 'b', 'c'], ['a', 'b', 'c', 'd']]
result = infer_ensemble(seqs, kmax=2, N=3) result = infer_ensemble(seqs, kmax=2, N=3)
algos = {a['algorithm'] for a in result['all']} algos = {a['algorithm'] for a in result['all']}
assert 'CRX' in algos assert 'CRX' in algos
# iDRegEx and kOREInference may fail stochastically, so at least CRX
assert len(result['all']) >= 1 assert len(result['all']) >= 1
@ -67,13 +65,6 @@ def test_prefer_idregex():
assert len(result['all']) == 1 assert len(result['all']) == 1
def test_prefer_koreinference():
seqs = [['a', 'b'], ['a', 'b', 'c']]
result = infer_ensemble(seqs, prefer='koreinference', kmax=2, N=5)
assert result['best']['algorithm'] == 'kOREInference'
assert len(result['all']) == 1
def test_prefer_case_insensitive(): def test_prefer_case_insensitive():
seqs = [['a', 'b']] seqs = [['a', 'b']]
r1 = infer_ensemble(seqs, prefer='CRX') r1 = infer_ensemble(seqs, prefer='CRX')
@ -229,12 +220,11 @@ def run_all():
tests = [ tests = [
test_ensemble_returns_dict, test_ensemble_returns_dict,
test_ensemble_best_not_none, test_ensemble_best_not_none,
test_ensemble_runs_all_three, test_ensemble_runs_both,
test_ensemble_all_results_have_scores, test_ensemble_all_results_have_scores,
test_ensemble_deterministic_results, test_ensemble_deterministic_results,
test_prefer_crx, test_prefer_crx,
test_prefer_idregex, test_prefer_idregex,
test_prefer_koreinference,
test_prefer_case_insensitive, test_prefer_case_insensitive,
test_prefer_unknown_falls_back, test_prefer_unknown_falls_back,
test_ensemble_empty_input, test_ensemble_empty_input,