#!/usr/bin/env python3 """Aggregate per-run dynamic-gate JSON files without loading GPU models.""" from __future__ import annotations import argparse import csv import json from pathlib import Path NUMERIC_FIELDS = [ "accepted_predictor_calls", "full_calls", "predictor_calls", "full_dit_time_ms", "predictor_time_ms", "confidence_head_time_ms", "context_dit_time_ms", "actual_dit_time_ms", "model_path_time_ms", "generation_time_s", "total_time_s", "latent_nrmse", "latent_tail_nrmse", "psnr", "ssim", "lpips", "tail_psnr", "tail_ssim", "tail_lpips", ] def write_csv(path: Path, rows: list[dict], fields: list[str]) -> None: path.parent.mkdir(parents=True, exist_ok=True) temporary = path.with_suffix(path.suffix + ".tmp") with temporary.open("w", encoding="utf-8", newline="") as handle: writer = csv.DictWriter(handle, fieldnames=fields) writer.writeheader() writer.writerows(rows) temporary.replace(path) def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--test_dir", type=Path, required=True) parser.add_argument("--expected_prompts", type=int, default=10) parser.add_argument( "--select_targets", type=int, nargs="*", default=None, help="Also select the best dynamic validation row for each target budget.", ) args = parser.parse_args() per_run = args.test_dir / "per_run" records = [ json.loads(path.read_text(encoding="utf-8")) for path in sorted(per_run.glob("*/prompt_*.json")) ] if not records: raise ValueError(f"No per-run JSON files under {per_run}") flat_fields = sorted({key for row in records for key in row if key != "decisions"}) write_csv( args.test_dir / "runs.csv", [{key: row.get(key) for key in flat_fields} for row in records], flat_fields, ) summary = [] for name in sorted({str(row["config_name"]) for row in records}): selected = [row for row in records if row["config_name"] == name] if len(selected) != args.expected_prompts: raise ValueError( f"{name}: expected {args.expected_prompts} prompts, got {len(selected)}" ) first = selected[0] item = { "config_name": name, "policy": first["policy"], "beta": first["beta"], "target_accepts": first["target_accepts"], "threshold": first["threshold"], "num_prompts": len(selected), } for field in NUMERIC_FIELDS: item[field] = sum(float(row[field]) for row in selected) / len(selected) summary.append(item) fields = [ "config_name", "policy", "beta", "target_accepts", "threshold", "num_prompts", *NUMERIC_FIELDS, ] write_csv(args.test_dir / "summary.csv", summary, fields) if args.select_targets: selected_dynamic = [] for target in args.select_targets: candidates = [ row for row in summary if row["policy"] == "dynamic" and int(row["target_accepts"]) == target ] if not candidates: raise ValueError(f"No dynamic candidates for target {target}") same_budget = [ row for row in candidates if abs(float(row["accepted_predictor_calls"]) - target) <= 0.5 + 1e-8 ] if not same_budget: closest = min( abs(float(row["accepted_predictor_calls"]) - target) for row in candidates ) same_budget = [ row for row in candidates if abs( abs(float(row["accepted_predictor_calls"]) - target) - closest ) <= 1e-8 ] same_budget.sort( key=lambda row: ( float(row["tail_lpips"]), abs(float(row["accepted_predictor_calls"]) - target), float(row["beta"]), ) ) selected_dynamic.append(same_budget[0]) (args.test_dir / "selected.json").write_text( json.dumps( { "selection_rule": ( "within target accepted calls +/-0.5, lowest validation " "tail LPIPS; then budget distance and lower beta" ), "selected_dynamic": selected_dynamic, }, indent=2, ) + "\n", encoding="utf-8", ) print(json.dumps(summary, indent=2)) if __name__ == "__main__": main()