grammar-inference-engine/experiments/context_eval.py

355 lines
13 KiB
Python

"""Context Strategy Experiments — Test context strategies + Reduce merging.
Tests on multiple codebases. Preserves results in experiments/results/.
"""
import json
import sys
import time
from collections import defaultdict
from pathlib import Path
from bex.tag_preprocessor.analyze import (
_preprocess_files, _file_to_package, frequency_filter, scan_directory
)
from bex.twotinf import build_soa
from bex.rwr0 import rwr0
from bex.reduce import reduce_contexts, soa_distance, build_soa_with_support
from bex.grammar import Empty
from bex.gbnf import to_gbnf
RESULTS_DIR = Path(__file__).parent / "results"
CODEBASES = {
"ragsak": {"name": "RAGSAK", "path": "/home/tobi/Desktop/kesai/RAGSAK", "ext": ".kt"},
"flask": {"name": "Flask", "path": "/home/tobi/Desktop/dervish/external_refs/flask", "ext": ".py"},
}
# ---------------------------------------------------------------------------
# Data loading
# ---------------------------------------------------------------------------
def load_data(base, ext):
groups = scan_directory(base)
files = groups.get(ext, [])
sequences, seq_files = _preprocess_files(files)
symbol_seqs = [[text for _, text, _ in seq] for seq in sequences]
packages = [_file_to_package(fp, base) for fp in seq_files]
return symbol_seqs, packages, seq_files
# ---------------------------------------------------------------------------
# Context strategies
# ---------------------------------------------------------------------------
def context_file_path_k(symbol_seqs, packages, k):
"""Option A: Group by last k components of file path."""
contexts = defaultdict(list)
for i, seq in enumerate(symbol_seqs):
pkg = packages[i]
components = pkg.split("/")
ctx = tuple(components[-k:]) if len(components) >= k else tuple(components)
contexts[ctx].append(seq)
return contexts
def context_first_k_symbols(symbol_seqs, k):
"""Option B: Group by first k symbols of call sequence."""
contexts = defaultdict(list)
for seq in symbol_seqs:
if not seq:
contexts[("ε",)].append(seq)
else:
ctx = tuple(seq[:k])
contexts[ctx].append(seq)
return contexts
def context_two_d(symbol_seqs, packages, path_k, sym_k):
"""Option C: Two-dimensional (path_k, first_k_symbols)."""
contexts = defaultdict(list)
for i, seq in enumerate(symbol_seqs):
pkg = packages[i]
components = pkg.split("/")
path_ctx = tuple(components[-path_k:]) if len(components) >= path_k else tuple(components)
if not seq:
sym_ctx = ("ε",)
else:
sym_ctx = tuple(seq[:sym_k])
ctx = path_ctx + sym_ctx
contexts[ctx].append(seq)
return contexts
def context_return_type_heuristic(symbol_seqs):
"""Option H: Group by return type heuristic (based on last symbol)."""
contexts = defaultdict(list)
for seq in symbol_seqs:
if not seq:
contexts[("ε",)].append(seq)
else:
last = seq[-1]
if last.startswith("return"):
ctx = ("RETURN_" + last.split()[1] if len(last.split()) > 1 else "RETURN_OTHER",)
elif last in ("true", "false"):
ctx = ("RETURN_BOOL",)
elif last.startswith("set") or last.startswith("put"):
ctx = ("SIDE_EFFECT",)
else:
ctx = ("RETURN_VALUE",)
contexts[ctx].append(seq)
return contexts
# ---------------------------------------------------------------------------
# Evaluation
# ---------------------------------------------------------------------------
def evaluate_context(contexts, label, min_methods=3, max_methods=50,
max_unique_ratio=0.85, max_alphabet=20, max_soa_edges=100):
"""Evaluate a context grouping: learn SOREs, collect metrics."""
results = {
"strategy": label,
"total_contexts": len(contexts),
"meaningful_contexts": 0,
"total_methods": 0,
"methods_in_good_groups": 0,
"grammar_successes": 0,
"grammar_failures": 0,
"skip_reasons": {},
"groups": [],
}
for ctx, seqs in sorted(contexts.items(), key=lambda x: -len(x[1])):
n = len(seqs)
results["total_methods"] += n
if n < min_methods:
continue
results["meaningful_contexts"] += 1
unique = len(set(tuple(s) for s in seqs))
unique_ratio = unique / n if n else 1.0
if n > max_methods:
results["skip_reasons"]["too_large"] = results["skip_reasons"].get("too_large", 0) + 1
results["groups"].append({
"context": str(ctx), "methods": n, "unique": unique,
"unique_ratio": round(unique_ratio, 3),
"grammar": "SKIP(too_large)", "grammar_success": False,
})
continue
if unique_ratio > max_unique_ratio:
results["skip_reasons"]["too_diverse"] = results["skip_reasons"].get("too_diverse", 0) + 1
results["groups"].append({
"context": str(ctx), "methods": n, "unique": unique,
"unique_ratio": round(unique_ratio, 3),
"grammar": "SKIP(too_diverse)", "grammar_success": False,
})
continue
alphabet = set()
for seq in seqs:
alphabet.update(seq)
if len(alphabet) > max_alphabet:
results["skip_reasons"]["large_alphabet"] = results["skip_reasons"].get("large_alphabet", 0) + 1
results["groups"].append({
"context": str(ctx), "methods": n, "unique": unique,
"unique_ratio": round(unique_ratio, 3),
"grammar": "SKIP(large_alphabet)", "grammar_success": False,
})
continue
filtered = frequency_filter(
[[(j, s, 0) for j, s in enumerate(seq)] for seq in seqs],
min_coverage=0.0,
)
clean = [[text for _, text, _ in r] for r in filtered]
clean = [s for s in clean if s]
if len(clean) < 2:
results["skip_reasons"]["empty_after_filter"] = results["skip_reasons"].get("empty_after_filter", 0) + 1
continue
soa = build_soa(clean)
n_edges = sum(len(v) for v in soa._succ.values())
if n_edges > max_soa_edges:
results["skip_reasons"]["complex_soa"] = results["skip_reasons"].get("complex_soa", 0) + 1
results["groups"].append({
"context": str(ctx), "methods": n, "unique": unique,
"unique_ratio": round(unique_ratio, 3),
"grammar": "SKIP(complex_soa)", "grammar_success": False,
})
continue
grammar = rwr0(soa)
group_info = {
"context": str(ctx),
"methods": n,
"unique": unique,
"unique_ratio": round(unique_ratio, 3),
"grammar": to_gbnf(grammar)[:200] if not isinstance(grammar, Empty) else "",
"grammar_success": not isinstance(grammar, Empty),
}
if not isinstance(grammar, Empty):
results["grammar_successes"] += 1
results["methods_in_good_groups"] += n
else:
results["grammar_failures"] += 1
results["groups"].append(group_info)
results["coverage"] = (
round(results["methods_in_good_groups"] / results["total_methods"] * 100, 1)
if results["total_methods"] > 0
else 0
)
return results
def run_experiment(name, contexts, label, codebase_name, merge_info=None):
"""Run one experiment, print + save results."""
print(f"\n {label}")
t0 = time.time()
results = evaluate_context(contexts, label)
elapsed = time.time() - t0
results["elapsed_seconds"] = round(elapsed, 2)
if merge_info:
results["merge_info"] = merge_info
print(f" Contexts: {results['total_contexts']} Meaningful: {results['meaningful_contexts']} "
f"Grammars: {results['grammar_successes']} Coverage: {results['coverage']}% "
f"Time: {elapsed:.2f}s")
if merge_info:
print(f" Merges: {merge_info['merges']} Threshold: {merge_info['threshold']}")
out_path = RESULTS_DIR / f"{codebase_name}_{name}.json"
with open(out_path, "w") as f:
json.dump(results, f, indent=2)
return results
# ---------------------------------------------------------------------------
# Run all experiments for one codebase
# ---------------------------------------------------------------------------
def run_codebase(codebase_key):
cfg = CODEBASES[codebase_key]
name = cfg["name"]
base = cfg["path"]
ext = cfg["ext"]
print(f"\n{'#' * 70}")
print(f"# {name} ({base})")
print(f"{'#' * 70}")
print(f"Loading {name} data...")
symbol_seqs, packages, seq_files = load_data(base, ext)
print(f"Loaded {len(symbol_seqs)} methods from {len(set(packages))} packages")
all_results = []
# --- Baselines ---
contexts = defaultdict(list)
for i, seq in enumerate(symbol_seqs):
contexts[packages[i]].append(seq)
all_results.append(run_experiment("baseline_package", contexts,
"Baseline: Package grouping", codebase_key))
for k in [1, 2, 3]:
contexts = context_file_path_k(symbol_seqs, packages, k)
all_results.append(run_experiment(f"file_path_k{k}", contexts,
f"Option A: File path k={k}", codebase_key))
for k in [1, 2, 3]:
contexts = context_first_k_symbols(symbol_seqs, k)
all_results.append(run_experiment(f"first_k_sym_{k}", contexts,
f"Option B: First {k} symbols", codebase_key))
for path_k in [1, 2]:
for sym_k in [1, 2]:
contexts = context_two_d(symbol_seqs, packages, path_k, sym_k)
all_results.append(run_experiment(
f"two_d_p{path_k}_s{sym_k}", contexts,
f"Option C: Path k={path_k} + Symbol k={sym_k}", codebase_key))
contexts = context_return_type_heuristic(symbol_seqs)
all_results.append(run_experiment("return_type_heuristic", contexts,
"Option H: Return type heuristic", codebase_key))
# --- Reduce experiments ---
# Best base: first k symbols, k=1,2,3
# Test Reduce on each with thresholds
reduce_thresholds = [0.05, 0.10, 0.15, 0.20, 0.30, 0.40]
print(f"\n--- Reduce experiments ---")
for base_k in [1, 2, 3]:
base_contexts = context_first_k_symbols(symbol_seqs, base_k)
base_eval = evaluate_context(base_contexts, f"First {base_k} symbols (pre-reduce)")
for threshold in reduce_thresholds:
merged, merge_count = reduce_contexts(base_contexts, threshold)
merge_info = {"merges": merge_count, "threshold": threshold,
"contexts_before": len(base_contexts),
"contexts_after": len(merged)}
label = f"Reduce k={base_k} ε={threshold}"
all_results.append(run_experiment(
f"reduce_k{base_k}_e{str(threshold).replace('.', '')}",
merged, label, codebase_key, merge_info))
# --- Summary ---
print(f"\n{'=' * 70}")
print(f" {name} SUMMARY")
print(f"{'=' * 70}")
print(f"{'Strategy':<42} {'Ctx':>5} {'Grms':>5} {'Cov%':>6} {'Merge':>5}")
print("-" * 70)
for r in all_results:
merge_str = ""
if r.get("merge_info"):
merge_str = str(r["merge_info"]["merges"])
print(f"{r['strategy']:<42} {r['meaningful_contexts']:>5} "
f"{r['grammar_successes']:>5} {r['coverage']:>5.1f}% {merge_str:>5}")
summary = [
{k: v for k, v in r.items() if k != "groups"}
for r in all_results
]
summary_path = RESULTS_DIR / f"{codebase_key}_summary.json"
with open(summary_path, "w") as f:
json.dump(summary, f, indent=2)
print(f"\nSummary: {summary_path}")
return all_results
def main():
RESULTS_DIR.mkdir(parents=True, exist_ok=True)
targets = sys.argv[1:] if len(sys.argv) > 1 else list(CODEBASES.keys())
all_results = {}
for key in targets:
if key not in CODEBASES:
print(f"Unknown codebase: {key}. Available: {list(CODEBASES.keys())}")
continue
all_results[key] = run_codebase(key)
# Cross-codebase comparison
print(f"\n{'#' * 70}")
print(f"# CROSS-CODEBASE COMPARISON")
print(f"{'#' * 70}")
for key, results in all_results.items():
cfg = CODEBASES[key]
best = max(results, key=lambda r: r["coverage"])
print(f"\n {cfg['name']}: best = {best['strategy']} ({best['coverage']}%)")
if __name__ == "__main__":
main()