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()