#!/usr/bin/env python3 """Audit and summarize the matched shared-gate vs KV-head-gate sweep.""" from __future__ import annotations import argparse import csv import glob import hashlib import json import subprocess from collections import defaultdict from pathlib import Path import matplotlib.pyplot as plt SUITES = ("RULER-local-full13", "BabiLong", "HELMET-ICL", "MRCR") THRESHOLDS = (0.45, 0.50, 0.525, 0.55, 0.575, 0.60, 0.625, 0.65) ARMS = ("token", "token_kv_head") SLUGS = { "RULER-local-full13": "ruler", "BabiLong": "babilong", "HELMET-ICL": "helmet", "MRCR": "mrcr", } SCORE_DEFINITIONS = { "RULER-local-full13": "unweighted macro of lm-eval 8192 string-match scores over the local 13-config set", "BabiLong": "unweighted macro exact-match accuracy over qa1-qa5", "HELMET-ICL": "unweighted macro label exact-match over five ICL configurations", "MRCR": "unweighted macro of prefix-check plus SequenceMatcher scores over 2/4/8-needle configurations", } def threshold_slug(value: float) -> str: return { 0.45: "045", 0.50: "050", 0.525: "0525", 0.55: "055", 0.575: "0575", 0.60: "060", 0.625: "0625", 0.65: "065", }[value] def one(pattern: str) -> Path: paths = [Path(path) for path in glob.glob(pattern, recursive=True)] if len(paths) != 1: raise RuntimeError(f"expected one path for {pattern}, found {paths}") return paths[0] def read_jsonl(path: Path) -> list[dict]: return [json.loads(line) for line in path.read_text().splitlines() if line.strip()] def collect_scores(root: Path) -> dict[tuple[str, str], dict]: rows: dict[tuple[str, str], dict] = {} for split in ("a", "b"): result = one(str(root / f"ruler8k_splits/{split}/lm_eval/**/results_*.json")) payload = json.loads(result.read_text()) for config, metrics in payload["results"].items(): if config in payload["n-samples"]: rows[("RULER-local-full13", config)] = { "score": float(metrics["8192,none"]), "n": int(payload["n-samples"][config]["effective"]), } result = one(str(root / "babilong8k_qa1_qa5_n50/**/results_*.json")) payload = json.loads(result.read_text()) for config, metrics in payload["results"].items(): if config in payload["n-samples"]: rows[("BabiLong", config)] = { "score": float(metrics["acc,none"]), "n": int(payload["n-samples"][config]["effective"]), } for suite, filename in ( ("HELMET-ICL", "helmet_icl8k_n50.jsonl"), ("MRCR", "mrcr_8k_2_4_8needle_n10.jsonl"), ): grouped: dict[str, list[dict]] = defaultdict(list) for row in read_jsonl(root / filename): grouped[row["config"]].append(row) for config, values in grouped.items(): rows[(suite, config)] = { "score": sum(float(row["score"]) for row in values) / len(values), "n": len(values), } return rows def sample_signatures(root: Path, score_rows: dict[tuple[str, str], dict]) -> dict[str, str]: signatures: dict[str, str] = {} for suite, rel_dirs in ( ("RULER-local-full13", ("ruler8k_splits/a/lm_eval", "ruler8k_splits/b/lm_eval")), ("BabiLong", ("babilong8k_qa1_qa5_n50",)), ): configs = [config for row_suite, config in score_rows if row_suite == suite] for config in configs: matches = [] for rel_dir in rel_dirs: matches.extend(root.glob(f"{rel_dir}/**/samples_{config}_*.jsonl")) if len(matches) != 1: raise RuntimeError(f"expected one sample log for {suite}/{config}: {matches}") values = [ ( row.get("doc_id"), row.get("doc_hash"), row.get("prompt_hash"), row.get("target_hash"), ) for row in read_jsonl(matches[0]) ] signatures[f"{suite}/{config}"] = hashlib.sha256( json.dumps(values, sort_keys=True).encode() ).hexdigest() for suite, filename in ( ("HELMET-ICL", "helmet_icl8k_n50.jsonl"), ("MRCR", "mrcr_8k_2_4_8needle_n10.jsonl"), ): grouped: dict[str, list[tuple]] = defaultdict(list) for row in read_jsonl(root / filename): grouped[row["config"]].append( (row["row_id"], row.get("prompt_sha256"), row.get("target")) ) for config, values in grouped.items(): signatures[f"{suite}/{config}"] = hashlib.sha256( json.dumps(values, sort_keys=True).encode() ).hexdigest() return signatures def read_sparsity(root: Path, suite: str) -> dict: if suite == "RULER-local-full13": parts = [ json.loads((root / f"sparsity_ruler_{split}.json").read_text()) for split in ("a", "b") ] sparse = sum(int(part["sparse_decisions"]) for part in parts) total = sum(int(part["total_decisions"]) for part in parts) granularities = {part.get("router_granularity") for part in parts} native = sum(int(part.get("native_router_decisions", 0)) for part in parts) phases: dict[str, dict[str, int]] = defaultdict(lambda: {"sparse": 0, "total": 0}) for part in parts: for phase, values in part.get("by_phase", {}).items(): phases[phase]["sparse"] += int(values["sparse_decisions"]) phases[phase]["total"] += int(values["total_decisions"]) else: part = json.loads((root / f"sparsity_{SLUGS[suite]}.json").read_text()) sparse = int(part["sparse_decisions"]) total = int(part["total_decisions"]) granularities = {part.get("router_granularity")} native = int(part.get("native_router_decisions", 0)) phases = { phase: { "sparse": int(values["sparse_decisions"]), "total": int(values["total_decisions"]), } for phase, values in part.get("by_phase", {}).items() } return { "sparse_decisions": sparse, "total_decisions": total, "sparsity": sparse / total, "full_attention_usage": 1.0 - sparse / total, "router_granularities": sorted(value for value in granularities if value), "native_router_decisions": native or None, "by_phase": phases, } def sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as handle: for block in iter(lambda: handle.read(1024 * 1024), b""): digest.update(block) return digest.hexdigest() def command_output(command: list[str], cwd: Path) -> str: try: return subprocess.run( command, cwd=cwd, check=True, text=True, capture_output=True ).stdout.strip() except (OSError, subprocess.CalledProcessError): return "unavailable in downloaded HF snapshot" def main() -> None: parser = argparse.ArgumentParser() parser.add_argument( "--repo", type=Path, default=Path("/data/sjm/AHA/AHA-Qwen3") ) parser.add_argument( "--eval-root", type=Path, default=Path("/data/sjm/AHA/AHA-Qwen3/experiments/qwen3_1p7b_router_granularity_20260714/eval"), ) parser.add_argument( "--baseline-root", type=Path, default=Path("/data/sjm/AHA/AHA-Qwen3/experiments/qwen3_1p7b_fourbench_8k_broad_20260713/vanilla"), ) parser.add_argument("--output-dir", type=Path) parser.add_argument("--input-dir", type=Path) args = parser.parse_args() output = args.output_dir or args.eval_root.parent / "summary" output.mkdir(parents=True, exist_ok=True) baseline = collect_scores(args.baseline_root) baseline_signatures = sample_signatures(args.baseline_root, baseline) curve_rows = [] per_config = [] points = [] alignment = {} done_files = ("ruler_a.DONE", "ruler_b.DONE", "babilong.DONE", "helmet.DONE", "mrcr.DONE") for arm in ARMS: for threshold in THRESHOLDS: method = f"{arm}_t{threshold_slug(threshold)}" root = args.eval_root / method missing = [name for name in done_files if not (root / name).exists()] if missing: raise RuntimeError(f"incomplete {method}: missing {missing}") scores = collect_scores(root) if set(scores) != set(baseline): raise RuntimeError(f"configuration mismatch for {method}") if any(scores[key]["n"] != baseline[key]["n"] for key in baseline): raise RuntimeError(f"sample-count mismatch for {method}") signatures = sample_signatures(root, scores) mismatches = sorted( key for key in baseline_signatures if signatures.get(key) != baseline_signatures[key] ) alignment[method] = {"aligned": not mismatches, "mismatches": mismatches} if mismatches: raise RuntimeError(f"prompt/target hash mismatch for {method}: {mismatches}") suite_rows = [] sparse_sum = total_sum = 0 for suite in SUITES: keys = sorted(key for key in baseline if key[0] == suite) baseline_score = sum(baseline[key]["score"] for key in keys) / len(keys) score = sum(scores[key]["score"] for key in keys) / len(keys) sparsity = read_sparsity(root, suite) expected_granularity = arm if sparsity["router_granularities"] != [expected_granularity]: raise RuntimeError( f"{method}/{suite} recorded {sparsity['router_granularities']}, " f"expected {[expected_granularity]}" ) sparse_sum += sparsity["sparse_decisions"] total_sum += sparsity["total_decisions"] row = { "arm": arm, "router_granularity": arm, "threshold": threshold, "benchmark": suite, "config_count": len(keys), "samples_per_config": ",".join( str(value) for value in sorted({baseline[key]["n"] for key in keys}) ), "total_samples": sum(baseline[key]["n"] for key in keys), "tuned_vanilla_score": baseline_score, "quality_floor_95pct": 0.95 * baseline_score, "score": score, "retention": score / baseline_score if baseline_score else None, "quality_pass": score >= 0.95 * baseline_score, **{key: sparsity[key] for key in ( "sparse_decisions", "total_decisions", "sparsity", "full_attention_usage", "native_router_decisions", "by_phase" )}, } curve_rows.append(row) suite_rows.append(row) for key in keys: per_config.append({ "arm": arm, "threshold": threshold, "benchmark": suite, "config": key[1], "n": baseline[key]["n"], "tuned_vanilla_score": baseline[key]["score"], "score": scores[key]["score"], }) points.append({ "arm": arm, "threshold": threshold, "all_suite_quality_pass": all(row["quality_pass"] for row in suite_rows), "decision_weighted_sparsity": sparse_sum / total_sum, "full_attention_usage": 1.0 - sparse_sum / total_sum, "sparse_decisions": sparse_sum, "total_decisions": total_sum, "suites": suite_rows, }) headline = {} for arm in ARMS: valid = [ point for point in points if point["arm"] == arm and point["all_suite_quality_pass"] ] if valid: best = max(valid, key=lambda point: point["decision_weighted_sparsity"]) headline[arm] = {"found": True, **best} else: headline[arm] = { "found": False, "statement": "No measured threshold preserved at least 95% of tuned vanilla on every suite.", } input_dir = args.input_dir or args.baseline_root.parent / "inputs" frozen_inputs = [ input_dir / "helmet_icl_8k_n50_per_config.jsonl", input_dir / "mrcr_8k_2_4_8needle_n10_per_config.jsonl", ] payload = { "protocol": { "model": "Qwen3-1.7B tuned vanilla", "context_length": 8192, "arms": { "token": "L2A-style shared-gate: one native gate per token/layer", "token_kv_head": "AHA: one native gate per token/KV-head/layer", }, "local_attention": {"sink_tokens": 64, "recent_tokens": 256}, "inference": "strict AHA routing in both prefill and decode; no force-full heads or full-decode fallback", "thresholds": list(THRESHOLDS), "quality_rule": "every suite macro score >= 95% of aligned tuned-vanilla macro score", "headline_rule": "highest decision-weighted measured sparsity among thresholds passing every suite", "sparsity_definition": "hard local routes / token x KV-head x layer effective decisions", "suite_scope": "RULER local full13 means this repository's fixed 13-config set, not complete upstream RULER", }, "score_definitions": SCORE_DEFINITIONS, "headline": headline, "points": points, "curves": curve_rows, "per_config": per_config, "alignment": {"all_aligned": all(value["aligned"] for value in alignment.values()), "methods": alignment}, "appendix_note": ( "Upstream L2A results are not merged into this table because architecture, " "training data, objective, and evaluation protocols differ." ), "reproducibility": { "git_commit": command_output(["git", "rev-parse", "HEAD"], args.repo), "git_status_short": command_output(["git", "status", "--short"], args.repo), "container": "dmtd-repro", "container_image": command_output( ["docker", "inspect", "-f", "{{.Config.Image}}", "dmtd-repro"], args.repo ), "baseline_root": str(args.baseline_root), "eval_root": str(args.eval_root), "frozen_input_sha256": { str(path): sha256(path) for path in frozen_inputs if path.exists() }, }, } (output / "router_granularity_results.json").write_text( json.dumps(payload, indent=2) + "\n" ) flat_curve_rows = [ {key: value for key, value in row.items() if key != "by_phase"} for row in curve_rows ] with (output / "quality_sparsity_curves.csv").open("w", newline="") as handle: writer = csv.DictWriter(handle, fieldnames=list(flat_curve_rows[0])) writer.writeheader() writer.writerows(flat_curve_rows) with (output / "per_config_scores.csv").open("w", newline="") as handle: writer = csv.DictWriter(handle, fieldnames=list(per_config[0])) writer.writeheader() writer.writerows(per_config) figure, axes = plt.subplots(2, 2, figsize=(12, 9), constrained_layout=True) colors = {"token": "#d62728", "token_kv_head": "#1f77b4"} labels = {"token": "L2A-style shared-gate", "token_kv_head": "AHA KV-head gate"} for axis, suite in zip(axes.flat, SUITES): for arm in ARMS: values = sorted( (row for row in curve_rows if row["arm"] == arm and row["benchmark"] == suite), key=lambda row: row["sparsity"], ) axis.plot( [100 * row["sparsity"] for row in values], [100 * row["retention"] for row in values], marker="o", color=colors[arm], label=labels[arm], ) for row in values: axis.annotate(f"{row['threshold']:.3g}", (100 * row["sparsity"], 100 * row["retention"]), fontsize=7) axis.axhline(95, color="black", linestyle="--", linewidth=1) axis.set_title(suite) axis.set_xlabel("Measured effective sparsity (%)") axis.set_ylabel("Retention vs tuned vanilla (%)") axis.grid(alpha=0.25) axes.flat[0].legend() figure.suptitle("Qwen3-1.7B matched router-granularity quality–sparsity curves") figure.savefig(output / "quality_sparsity_curves.png", dpi=180) plt.close(figure) print(json.dumps(headline, indent=2)) if __name__ == "__main__": main()