grammar-inference-engine/experiments/coarsen_eval.py

227 lines
7.8 KiB
Python

"""Coarsening Experiment — Compare raw vs coarsened token extraction.
Tests whether collapsing method names → tree-sitter capture categories
reveals cross-package patterns that raw text misses.
"""
import json
import time
from collections import defaultdict
from pathlib import Path
from bex.tag_preprocessor.analyze import (
_preprocess_files, _file_to_package, scan_directory
)
from bex.tag_preprocessor.code import (
_extract_call_tokens, _extract_coarsened_tokens
)
from bex.twotinf import build_soa
from bex.rwr0 import rwr0
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"},
"coroutines": {"name": "Kotlin Coroutines", "path": "/home/tobi/Desktop/dervish/projects/grammar-inference-engine/external_refs/kotlinx.coroutines", "ext": ".kt"},
"fastapi": {"name": "FastAPI", "path": "/home/tobi/Desktop/dervish/projects/grammar-inference-engine/external_refs/fastapi", "ext": ".py"},
}
def load_data(base, ext):
"""Load sequences, extract both raw and coarsened."""
groups = scan_directory(base)
files = groups.get(ext, [])
sequences, seq_files = _preprocess_files(files)
raw_seqs = [_extract_call_tokens(seq) for seq in sequences]
coarse_seqs = [_extract_coarsened_tokens(seq) for seq in sequences]
packages = [_file_to_package(fp, base) for fp in seq_files]
return raw_seqs, coarse_seqs, packages, seq_files
def group_by_context(seqs, packages, k):
"""Group sequences by first k symbols of their sequence."""
contexts = defaultdict(list)
for i, seq in enumerate(seqs):
if not seq:
contexts[("_eps",)].append((i, seq))
else:
ctx = tuple(seq[:k])
contexts[ctx].append((i, seq))
return contexts
def group_by_package(seqs, packages):
"""Group sequences by package."""
contexts = defaultdict(list)
for i, seq in enumerate(seqs):
contexts[packages[i]].append((i, seq))
return contexts
def infer_sore(seqs):
"""Try to infer a SORE from a list of sequences. Returns SORE string or None."""
if len(seqs) < 2:
return None
clean = [s for s in seqs if s]
if len(clean) < 2:
return None
unique = len(set(tuple(s) for s in clean))
if unique / len(clean) > 0.9:
return None
alphabet = set()
for s in clean:
alphabet.update(s)
if len(alphabet) > 20:
return None
try:
soa = build_soa(clean)
grammar = rwr0(soa)
return grammar if not isinstance(grammar, Empty) else None
except Exception:
return None
def measure_group(ctx, items, label):
"""Measure properties of a group."""
seqs = [seq for _, seq in items]
n = len(seqs)
unique = len(set(tuple(s) for s in seqs))
unique_ratio = unique / n if n else 1.0
alphabet = set()
for s in seqs:
alphabet.update(s)
sore = infer_sore(seqs)
return {
"context": str(ctx),
"label": label,
"methods": n,
"unique": unique,
"unique_ratio": round(unique_ratio, 3),
"alphabet_size": len(alphabet),
"grammar": to_gbnf(sore)[:200] if sore else None,
"grammar_success": sore is not None,
}
def find_cross_package(groups_by_ctx, packages):
"""Find contexts that span multiple packages."""
ctx_to_pkgs = defaultdict(set)
for ctx, items in groups_by_ctx.items():
for idx, _ in items:
ctx_to_pkgs[ctx].add(packages[idx])
return {ctx: pkgs for ctx, pkgs in ctx_to_pkgs.items() if len(pkgs) > 1}
def run_comparison(name, raw_seqs, coarse_seqs, packages, k_values=(1, 2, 3)):
"""Run raw vs coarsened comparison for one codebase."""
results = {"name": name, "raw": {}, "coarsened": {}}
for label, seqs in [("raw", raw_seqs), ("coarsened", coarse_seqs)]:
results[label]["seq_count"] = len(seqs)
alphabet = set()
for s in seqs:
alphabet.update(s)
results[label]["alphabet_size"] = len(alphabet)
results[label]["alphabet_sample"] = sorted(alphabet)[:30]
for k in k_values:
t0 = time.time()
groups = group_by_context(seqs, packages, k)
elapsed = time.time() - t0
# Measure all groups
group_results = []
grammar_successes = 0
methods_in_good = 0
total_methods = 0
for ctx, items in sorted(groups.items(), key=lambda x: -len(x[1])):
m = measure_group(ctx, items, f"{label}_k{k}")
group_results.append(m)
total_methods += m["methods"]
if m["grammar_success"]:
grammar_successes += 1
methods_in_good += m["methods"]
# Cross-package analysis
cross_pkg = find_cross_package(groups, packages)
results[label][f"k{k}"] = {
"contexts": len(groups),
"grammar_successes": grammar_successes,
"total_methods": total_methods,
"methods_in_good": methods_in_good,
"coverage": round(methods_in_good / total_methods * 100, 1) if total_methods else 0,
"cross_package_contexts": len(cross_pkg),
"elapsed": round(elapsed, 3),
"top_groups": group_results[:10],
"cross_pkg_examples": [
{"context": str(ctx), "packages": len(pkgs)}
for ctx, pkgs in sorted(cross_pkg.items(), key=lambda x: -len(x[1]))[:10]
],
}
return results
def print_results(results):
"""Pretty-print comparison results."""
name = results["name"]
print(f"\n{'=' * 70}")
print(f" {name}")
print(f"{'=' * 70}")
raw = results["raw"]
coarse = results["coarsened"]
print(f"\n Alphabet size: raw={raw['alphabet_size']} coarsened={coarse['alphabet_size']} "
f"reduction={raw['alphabet_size'] - coarse['alphabet_size']} ({(1 - coarse['alphabet_size']/raw['alphabet_size'])*100:.0f}%)")
print(f" Coarsened categories: {coarse['alphabet_sample']}")
for k in [1, 2, 3]:
rk = raw.get(f"k{k}", {})
ck = coarse.get(f"k{k}", {})
print(f"\n --- k={k} ---")
print(f" {'':20s} {'Raw':>10s} {'Coarsened':>10s}")
print(f" {'Contexts':20s} {rk.get('contexts',0):10d} {ck.get('contexts',0):10d}")
print(f" {'Grammar successes':20s} {rk.get('grammar_successes',0):10d} {ck.get('grammar_successes',0):10d}")
print(f" {'Coverage':20s} {rk.get('coverage',0):9.1f}% {ck.get('coverage',0):9.1f}%")
print(f" {'Cross-pkg contexts':20s} {rk.get('cross_package_contexts',0):10d} {ck.get('cross_package_contexts',0):10d}")
# Show cross-package examples from coarsened
cross_examples = ck.get("cross_pkg_examples", [])
if cross_examples:
print(f" Cross-package shapes:")
for ex in cross_examples[:5]:
print(f" {ex['context']} ({ex['packages']} packages)")
def main():
all_results = {}
for key, cfg in CODEBASES.items():
print(f"\nLoading {cfg['name']}...")
raw_seqs, coarse_seqs, packages, seq_files = load_data(cfg["path"], cfg["ext"])
print(f" {len(raw_seqs)} methods from {len(set(packages))} packages")
results = run_comparison(cfg["name"], raw_seqs, coarse_seqs, packages)
all_results[key] = results
print_results(results)
# Save
out_path = RESULTS_DIR / f"coarsen_{key}.json"
with open(out_path, "w") as f:
json.dump(results, f, indent=2)
print(f"\n Saved to {out_path}")
return all_results
if __name__ == "__main__":
main()