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:
parent
e3ad256321
commit
710da56916
3 changed files with 65 additions and 57 deletions
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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":
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue