grammar-inference-engine/experiments/context_eval.py

301 lines
10 KiB
Python
Raw Normal View History

"""Context Strategy Experiments — Test all context strategies on RAGSAK.
Preserves results in experiments/results/ for future reference.
"""
import json
import os
import time
from collections import defaultdict, Counter
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.crx import CRX
from bex.idregex import idregex
from bex.mdl import lang_size_score
BASE = "/home/tobi/Desktop/kesai/RAGSAK"
RESULTS_DIR = Path(__file__).parent / "results"
def load_data():
"""Load all Kotlin files from RAGSAK."""
groups = scan_directory(BASE)
kt_files = groups.get(".kt", [])
sequences, seq_files = _preprocess_files(kt_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
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:
# Heuristic: last symbol often indicates return type
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
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.
Smart filters before attempting REWRITE:
- max_methods: skip groups too large (will fail, waste time)
- max_unique_ratio: skip if sequences are too diverse
- max_alphabet: skip if too many unique symbols (SORE can't compress)
- max_soa_edges: skip if SOA is too complex (REWRITE will be slow)
"""
results = {
"strategy": label,
"total_contexts": len(contexts),
"meaningful_contexts": 0,
"total_methods": 0,
"methods_in_good_groups": 0,
"sore_successes": 0,
"sore_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
# Filter: too many methods
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),
"sore": "SKIP(too_large)", "sore_success": False,
})
continue
# Filter: too diverse
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),
"sore": "SKIP(too_diverse)", "sore_success": False,
})
continue
# Filter: too many unique symbols
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),
"sore": "SKIP(large_alphabet)", "sore_success": False,
})
continue
# Filter: rare symbols
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
# Build SOA and check complexity before REWRITE
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),
"sore": "SKIP(complex_soa)", "sore_success": False,
})
continue
# Learn SORE
sore = rwr0(soa)
group_info = {
"context": str(ctx),
"methods": n,
"unique": unique,
"unique_ratio": round(unique_ratio, 3),
"sore": sore[:200] if sore not in ("",) else sore,
"sore_success": sore != "",
}
if sore != "":
results["sore_successes"] += 1
results["methods_in_good_groups"] += n
else:
results["sore_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):
"""Run one experiment, save results."""
print(f"\n{'='*60}")
print(f" {label}")
print(f"{'='*60}")
t0 = time.time()
results = evaluate_context(contexts, label)
elapsed = time.time() - t0
results["elapsed_seconds"] = round(elapsed, 2)
print(f" Contexts: {results['total_contexts']}")
print(f" Meaningful (>=3 methods): {results['meaningful_contexts']}")
if results.get("skip_reasons"):
for reason, count in results["skip_reasons"].items():
print(f" Skipped ({reason}): {count}")
print(f" SORE successes: {results['sore_successes']}")
print(f" SORE failures: {results['sore_failures']}")
print(f" Coverage: {results['coverage']}%")
print(f" Time: {elapsed:.1f}s")
# Save
out_path = RESULTS_DIR / f"{name}.json"
with open(out_path, "w") as f:
json.dump(results, f, indent=2)
print(f" Saved: {out_path}")
return results
def main():
print("Loading RAGSAK data...")
symbol_seqs, packages, seq_files = load_data()
print(f"Loaded {len(symbol_seqs)} methods from {len(set(packages))} packages")
all_results = []
# Baseline: package grouping (current approach)
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"))
# Option A: File path k-equivalence
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}"))
# Option B: First k symbols
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"))
# Option C: Two-dimensional
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}"
))
# Option H: Return type heuristic
contexts = context_return_type_heuristic(symbol_seqs)
all_results.append(run_experiment("return_type_heuristic", contexts, "Option H: Return type heuristic"))
# Summary
print("\n" + "=" * 80)
print(" SUMMARY")
print("=" * 80)
print(f"{'Strategy':<40} {'Contexts':>8} {'SORE OK':>8} {'Coverage':>8}")
print("-" * 80)
for r in all_results:
print(f"{r['strategy']:<40} {r['meaningful_contexts']:>8} {r['sore_successes']:>8} {r['coverage']:>7}%")
# Save summary
summary = [
{k: v for k, v in r.items() if k != "groups"}
for r in all_results
]
summary_path = RESULTS_DIR / "summary.json"
with open(summary_path, "w") as f:
json.dump(summary, f, indent=2)
print(f"\nSummary saved: {summary_path}")
if __name__ == "__main__":
main()