diff --git a/bex/ensemble.py b/bex/ensemble.py index 93e8cd3..2ba14d5 100644 --- a/bex/ensemble.py +++ b/bex/ensemble.py @@ -3,7 +3,6 @@ import re from .crx import CRX from .idregex import idregex -from .kore import kOREInference from .expr import alphabet from .mdl import model_cost, mdl_score @@ -375,8 +374,21 @@ def _run_idregex(sequences, kmax, N): 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): - """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) result = kore.infer(sequences) if result: @@ -385,21 +397,7 @@ def _run_kore(sequences, kmax, N): return None, float('inf') -_ALGO_NAMES = { - '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): +def infer_ensemble(sequences, kmax=2, N=3, prefer=None, min_coverage=1.0, include_kore=False): """Run all applicable algorithms and return the best by MDL score. Args: @@ -450,10 +448,11 @@ def infer_ensemble(sequences, kmax=2, N=3, prefer=None, min_coverage=1.0): if idr_g: results.append(('iDRegEx', idr_g, idr_score)) - # 3. kOREInference (Algorithm 4 with MDL scoring) - kore_g, kore_score = _run_kore(sequences, kmax, N) - if kore_g: - results.append(('kOREInference', kore_g, kore_score)) + # 3. kOREInference (opt-in via include_kore=True) + if include_kore: + kore_g, kore_score = _run_kore(sequences, kmax, N) + if kore_g: + results.append(('kOREInference', kore_g, kore_score)) results = [r for r in results if r[1] and r[1] != '∅'] if not results: diff --git a/bex/tag_preprocessor/analyze.py b/bex/tag_preprocessor/analyze.py index 2bde248..295a5b7 100644 --- a/bex/tag_preprocessor/analyze.py +++ b/bex/tag_preprocessor/analyze.py @@ -14,6 +14,7 @@ import sys import time from pathlib import Path from collections import Counter +from concurrent.futures import ProcessPoolExecutor, as_completed import pathspec @@ -203,7 +204,7 @@ def frequency_filter(sequences, min_coverage=0.2): 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. Returns: @@ -231,13 +232,25 @@ def analyze_clusters(file_paths, extension, project_root="", min_coverage=DEFAUL packages = _top_packages(cluster_fps, project_root) 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} 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. 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)") results = [] - for label, indices in groups: - t1 = time.time() - group_seqs = [sequences[i] for i in indices] - group_files = set(seq_files[i] for i in indices) + n_workers = os.cpu_count() + _vprint(f"Inferring {len(groups)} groups across {n_workers} workers ...") + with ProcessPoolExecutor(max_workers=n_workers) as ex: + 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) - 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)) + results.sort(key=lambda x: x[0]) if 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 -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. Args: @@ -397,6 +408,7 @@ def analyze_directory( slice="flat", include=None, exclude=None, + include_kore=False, ): """Scan a directory and run analysis for each language found. @@ -427,6 +439,7 @@ def analyze_directory( min_coverage=min_coverage, prefer=prefer, kmax=kmax, + include_kore=include_kore, ) else: results[ext] = analyze_clusters( @@ -435,6 +448,7 @@ def analyze_directory( min_coverage=min_coverage, prefer=prefer, kmax=kmax, + include_kore=include_kore, ) return results @@ -470,9 +484,13 @@ def _parse_args(argv=None): parser.add_argument("directory", help="Directory to scan") parser.add_argument( "--prefer", - choices=["crx", "idregex", "koreinference"], + choices=["crx", "idregex"], 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( "--kmax", type=int, default=2, help="Maximum k for k-ORE algorithms (default: 2)", @@ -522,6 +540,7 @@ def main(): slice=args.slice, include=args.include, exclude=args.exclude, + include_kore=args.kore, ) if args.json_flag or args.format == "json": diff --git a/tests/test_ensemble.py b/tests/test_ensemble.py index db15627..55399a3 100644 --- a/tests/test_ensemble.py +++ b/tests/test_ensemble.py @@ -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.idregex import is_deterministic -from bex.kore import kOREInference # ── Basic ensemble runs ── @@ -21,16 +20,15 @@ def test_ensemble_best_not_none(): result = infer_ensemble(seqs, kmax=2, N=3) assert result['best'] 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 -def test_ensemble_runs_all_three(): +def test_ensemble_runs_both(): seqs = [['a', 'b', 'c'], ['a', 'b', 'c', 'd']] result = infer_ensemble(seqs, kmax=2, N=3) algos = {a['algorithm'] for a in result['all']} assert 'CRX' in algos - # iDRegEx and kOREInference may fail stochastically, so at least CRX assert len(result['all']) >= 1 @@ -67,13 +65,6 @@ def test_prefer_idregex(): 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(): seqs = [['a', 'b']] r1 = infer_ensemble(seqs, prefer='CRX') @@ -229,12 +220,11 @@ def run_all(): tests = [ test_ensemble_returns_dict, test_ensemble_best_not_none, - test_ensemble_runs_all_three, + test_ensemble_runs_both, test_ensemble_all_results_have_scores, test_ensemble_deterministic_results, test_prefer_crx, test_prefer_idregex, - test_prefer_koreinference, test_prefer_case_insensitive, test_prefer_unknown_falls_back, test_ensemble_empty_input,