File size: 4,824 Bytes
bc29ee3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | #!/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()
|