Self-Forcing / scripts /aggregate_layer17_dynamic_runs.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw History Blame Contribute Delete
4.82 kB
#!/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()