Download scripts/aggregate_layer17_dynamic_runs.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 4.82 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/aggregate_layer17_dynamic_runs.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/aggregate_layer17_dynamic_runs.py
-
curl -L -o aggregate_layer17_dynamic_runs.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/aggregate_layer17_dynamic_runs.py
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() | |