"""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 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) sore = rwr0(soa) return sore if sore != "∅" 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), "sore": sore[:200] if sore else None, "sore_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 = [] sore_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["sore_success"]: sore_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), "sore_successes": sore_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" {'SORE successes':20s} {rk.get('sore_successes',0):10d} {ck.get('sore_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()