File size: 14,457 Bytes
f70ac4f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
#!/usr/bin/env python
"""Final summary for the Self-Forcing Extended-251 Full Evaluation.

Implements protocol sections 7 (normalize + aggregate), 8.3 (pixel aggregation),
9.4 (matched-FFFF speedup) and 10 (output tables), plus the section 11/15
completeness gate: a strategy that is not a full 251 is never summarised.

    python eval/aggregate.py --out-root eval_out
"""

import argparse
import re
import csv
import glob
import hashlib
import json
import math
import os
import platform
import sys
from collections import defaultdict

ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, ROOT)

MAPPING = os.path.join(ROOT, "assets/vbench8_extended_subset_mapping.json")

DIMENSIONS = ["subject_consistency", "background_consistency", "motion_smoothness",
              "dynamic_degree", "aesthetic_quality", "imaging_quality", "scene",
              "overall_consistency"]

# Protocol section 7.1 -- fixed empirical ranges, never re-estimated from the run.
NORMALIZE_RANGE = {
    "subject_consistency":    (0.1462, 1.0),
    "background_consistency": (0.2615, 1.0),
    "motion_smoothness":      (0.7060, 0.9975),
    "dynamic_degree":         (0.0,    1.0),
    "aesthetic_quality":      (0.0,    1.0),
    "imaging_quality":        (0.0,    1.0),
    "scene":                  (0.0,    0.8222),
    "overall_consistency":    (0.0,    0.3640),
}

REFERENCE_OF = {"sf": "sf_ffff", "cf": "cf_ffff", "cfa": "cfa_ffff"}


def normalize(raw):
    out = {}
    for d, v in raw.items():
        lo, hi = NORMALIZE_RANGE[d]
        out[d] = (v - lo) / (hi - lo)  # protocol does not clip
    return out


def quality_score(n):
    return (n["subject_consistency"] + n["background_consistency"]
            + n["motion_smoothness"] + 0.5 * n["dynamic_degree"]
            + n["aesthetic_quality"] + n["imaging_quality"]) / 5.5


def semantic_score(n):
    return (n["scene"] + n["overall_consistency"]) / 2.0


def selected_score(q, s):
    return (4.0 * q + s) / 5.0


def sha256(path):
    h = hashlib.sha256()
    with open(path, "rb") as f:
        for chunk in iter(lambda: f.read(1 << 20), b""):
            h.update(chunk)
    return h.hexdigest()


def load_per_prompt(out_root, strategy):
    recs = []
    for p in sorted(glob.glob(os.path.join(out_root, "per_prompt", strategy, "*.json"))):
        with open(p) as f:
            r = json.load(f)
        if r.get("status") == "complete":
            recs.append(r)
    return recs


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--out-root", default="eval_out")
    ap.add_argument("--expect", type=int, default=251)
    ap.add_argument("--allow-incomplete", action="store_true",
                    help="Report partial strategies instead of refusing (diagnostics only)")
    ap.add_argument("--latency-subsets", default=os.path.join(ROOT, "results/latency_subsets.json"),
                    help="JSON {strategy_regex: {shard, num_shards, note}}: time these "
                         "strategies on one prompt shard only (the one that ran on an "
                         "uncontended GPU); quality still uses every record")
    ap.add_argument("--only", default=None,
                    help="Regex: summarise only strategies whose name matches (other "
                         "strategies in per_prompt/ are ignored, e.g. while they are "
                         "still being generated)")
    args = ap.parse_args()

    out_root = (args.out_root if os.path.isabs(args.out_root)
                else os.path.join(ROOT, args.out_root))

    with open(MAPPING) as f:
        mapping = json.load(f)

    strategies = sorted(os.path.basename(p) for p in
                        glob.glob(os.path.join(out_root, "per_prompt", "*")))
    if args.only:
        strategies = [s for s in strategies if re.search(args.only, s)]
    if not strategies:
        print("no per-prompt records found")
        return 1

    per_strategy, problems = {}, []
    for s in strategies:
        recs = load_per_prompt(out_root, s)
        if len(recs) != args.expect:
            problems.append(f"{s}: {len(recs)} complete records, expected {args.expect}")
        per_strategy[s] = recs

    if problems and not args.allow_incomplete:
        print("REFUSING TO SUMMARISE -- incomplete strategies (protocol 11.9):")
        for p in problems:
            print("  -", p)
        return 1

    # --- latency, matched to each base model's own FFFF -----------------------
    mean_latency, mean_context, timed_records = {}, {}, {}
    for s, recs in per_strategy.items():
        # Records flagged latency_excluded keep their pixel metrics but carry no
        # usable timing (see their latency_note); they stay out of the means.
        lat = [r for r in recs if not r.get("latency_excluded")]
        if not lat:
            continue
        timed_records[s] = lat
        mean_latency[s] = sum(r["policy_latency_ms"] for r in lat) / len(lat)
        mean_context[s] = sum(r["excluded_context_kv_latency_ms"] for r in lat) / len(lat)

    # Strategies generated while another job shared the GPU are timed only on the
    # prompt shard that ran alone (sharding is rows[shard::num_shards] over the
    # mapping order, see generate_eval.py); their FFFF reference is restricted to
    # the same prompts so the ratio stays paired.
    subsets = {}
    if args.latency_subsets and os.path.exists(args.latency_subsets):
        with open(args.latency_subsets) as f:
            subsets = json.load(f)
    with open(os.path.join(ROOT, "assets/vbench8_extended_subset_mapping.json")) as f:
        mapping_rows = json.load(f)["rows"]
    subset_of = {}
    for s in per_strategy:
        for pat, spec in subsets.items():
            if re.search(pat, s):
                keep = {r["global_index"] for r in mapping_rows[spec["shard"]::spec["num_shards"]]}
                subset_of[s] = (keep, spec)
                break

    rows = []
    for s, recs in per_strategy.items():
        if not recs:
            continue
        ref = recs[0].get("reference_strategy") or REFERENCE_OF[s.split("_")[0]]
        if ref not in mean_latency:
            problems.append(f"{s}: reference {ref} missing")
            continue

        # Ratio of means (protocol 9.4).  Equal sample counts make this the same
        # as the ratio of sums; a mean of per-prompt percentages is not used.
        # A strategy whose records carry the FFFF latency measured in the same
        # process (retime_eval.py) is compared against that, not against the
        # FFFF run's own records: contention differs between runs, and only a
        # same-run pairing is "matched FFFF" in the protocol's sense.
        timed = timed_records[s]
        matched = [r.get("matched_ffff_policy_latency_ms") for r in timed]
        strat_latency = mean_latency[s]
        latency_n = len(timed)
        if s in subset_of:
            keep, spec = subset_of[s]
            sub = [r for r in recs if r["global_index"] in keep]
            ref_sub = [r for r in per_strategy[ref] if r["global_index"] in keep]
            assert sub and ref_sub, f"{s}: empty latency subset"
            strat_latency = sum(r["policy_latency_ms"] for r in sub) / len(sub)
            ref_latency = sum(r["policy_latency_ms"] for r in ref_sub) / len(ref_sub)
            latency_n = len(sub)
            latency_source = (f"generation run, shard {spec['shard']}/{spec['num_shards']} only "
                              f"({len(sub)} prompts; {spec.get('note', '')})")
        elif all(m is not None for m in matched):
            ref_latency = sum(matched) / len(matched)
            latency_source = "paired retime (same process as FFFF)"
        elif all(r.get("latency_retimed") for r in recs):
            ref_latency = mean_latency[ref]
            latency_source = ("strategy retime: all records from one denoise-only pass on an idle GPU "
                              "(FFFF from its own records)")
        elif any(r.get("latency_retimed") for r in recs):
            n_re = sum(1 for r in recs if r.get("latency_retimed"))
            ref_latency = mean_latency[ref]
            latency_source = (f"generation run; {n_re}/{len(recs)} records re-timed "
                              "(strategy-only denoise rerun on an idle GPU; FFFF from its own records)")
        else:
            ref_latency = mean_latency[ref]
            latency_source = "generation run (FFFF from its own records)"
            if len(timed) < len(recs):
                latency_source += (f"; {len(timed)}/{len(recs)} prompts with intact timing "
                                   "(latency_excluded records dropped)")
        speedup_pct = 100.0 * (1.0 - strat_latency / ref_latency)

        # PSNR from the mean of per-prompt mean MSE, then converted once (8.3).
        mses = [r["pixel_metrics_vs_ffff"]["mean_mse"] for r in recs]
        mean_mse = sum(mses) / len(mses)
        psnr = -10.0 * math.log10(max(mean_mse, 1e-12))
        ssim = sum(r["pixel_metrics_vs_ffff"]["ssim"] for r in recs) / len(recs)
        lpips = sum(r["pixel_metrics_vs_ffff"]["lpips"] for r in recs) / len(recs)
        frames = {r["pixel_metrics_vs_ffff"]["num_frames"] for r in recs}
        if frames != {81}:
            problems.append(f"{s}: pixel metrics used frame counts {sorted(frames)}, expected 81")

        row = {
            "strategy": s,
            "base_model": recs[0]["base_model"],
            "method": recs[0]["method"],
            "target_speedup": recs[0]["target_speedup"],
            "schedule": recs[0].get("schedule"),
            "num_inference_steps": recs[0].get("num_inference_steps"),
            "num_videos": len(recs),
            "mp4_reference_records": sum(1 for r in recs if r.get("reference_source") == "ffff_mp4"),
            "policy_latency_ms": strat_latency,
            "latency_num_videos": latency_n,
            "speedup_percent_vs_ffff": 0.0 if s == ref else speedup_pct,
            "latency_ratio_vs_ffff": strat_latency / ref_latency,
            "matched_ffff_policy_latency_ms": ref_latency,
            "latency_source": latency_source,
            "excluded_context_kv_latency_ms": mean_context[s],
            "psnr": psnr,
            "ssim": ssim,
            "lpips": lpips,
            "mean_compute_equivalent_forwards": (
                sum(r["cache_diagnostics"]["compute_equivalent_forwards"] for r in recs)
                / len(recs)),
        }

        score_path = os.path.join(out_root, "vbench", "scores", f"{s}.json")
        if os.path.exists(score_path):
            with open(score_path) as f:
                sc = json.load(f)
            raw = {d: sc["raw"][d] for d in DIMENSIONS if d in sc["raw"]}
            if len(raw) == len(DIMENSIONS):
                bad = {d: v for d, v in raw.items() if not (0.0 <= v <= 1.0)}
                if bad:
                    problems.append(f"{s}: raw scores outside [0,1]: {bad}")
                n = normalize(raw)
                q, sem = quality_score(n), semantic_score(n)
                sel = selected_score(q, sem)
                row.update({f"raw_{d}": raw[d] for d in DIMENSIONS})
                row.update({f"normalized_{d}": n[d] for d in DIMENSIONS})
                row.update({"quality_score": q, "semantic_score": sem,
                            "selected_vbench_score": sel,
                            "selected_vbench_percent": sel * 100.0})
            else:
                problems.append(f"{s}: only {len(raw)}/8 VBench dimensions")
        else:
            problems.append(f"{s}: no VBench scores yet")
        rows.append(row)

    rows.sort(key=lambda r: (r["base_model"], r["method"], r["target_speedup"]))

    summary = {
        "protocol": "Self-Forcing Extended-251 Full Evaluation",
        "vbench_long": False,
        "num_strategies": len(rows),
        "prompts_per_strategy": args.expect,
        "seed": 0,
        "video_spec": {"frames": 81, "height": 480, "width": 832, "fps": 16},
        "precision": "bfloat16",
        "normalize_range": NORMALIZE_RANGE,
        "aggregation": {
            "quality": "(n_sc + n_bc + n_ms + 0.5*n_dd + n_aq + n_iq) / 5.5",
            "semantic": "(n_scene + n_oc) / 2",
            "selected": "(4*quality + semantic) / 5",
            "psnr": "-10*log10(mean(per_prompt_mean_mse)), clamp 1e-12",
            "speedup": "100 * (1 - mean(strategy) / mean(matched_ffff))",
        },
        "latency_includes": "all denoise policy GPU compute incl. cache control modules",
        "latency_excludes": "context/KV-cache DiT, text encoder, VAE decode, metrics, I/O",
        "sources": mapping["sources"],
        "mapping": {"path": MAPPING, "sha256": sha256(MAPPING)},
        "environment": {
            "generation_python": "/local/zoubin/cz/envs/self_forcing/bin/python",
            "vbench_python": "/local/zoubin/cz/envs/vbench_eval/bin/python",
            "vbench_source": "/local/zoubin/cz/projects/VBench",
            "vbench_cache_dir": "/local/zoubin/cz/.cache/vbench",
            "platform": platform.platform(),
        },
        "problems": problems,
        "rows": rows,
    }

    os.makedirs(os.path.join(out_root, "summaries"), exist_ok=True)
    jpath = os.path.join(out_root, "summaries", "final_summary.json")
    with open(jpath, "w") as f:
        json.dump(summary, f, indent=2)

    cpath = os.path.join(out_root, "summaries", "final_summary.csv")
    if rows:
        keys = sorted({k for r in rows for k in r},
                      key=lambda k: (k not in ("strategy", "base_model", "method"), k))
        with open(cpath, "w", newline="") as f:
            w = csv.DictWriter(f, fieldnames=keys)
            w.writeheader()
            w.writerows(rows)

    hdr = (f"{'strategy':24s} {'videos':>6s} {'lat_ms':>9s} {'speedup%':>9s} "
           f"{'psnr':>7s} {'ssim':>7s} {'lpips':>7s} {'sel_vbench%':>12s}")
    print(hdr)
    print("-" * len(hdr))
    for r in rows:
        sel = r.get("selected_vbench_percent")
        print(f"{r['strategy']:24s} {r['num_videos']:6d} {r['policy_latency_ms']:9.1f} "
              f"{r['speedup_percent_vs_ffff']:9.2f} {r['psnr']:7.2f} {r['ssim']:7.4f} "
              f"{r['lpips']:7.4f} " + (f"{sel:12.3f}" if sel is not None else f"{'-':>12s}"))
    if problems:
        print("\nnotes:")
        for p in problems:
            print("  -", p)
    print(f"\nwrote {jpath}\nwrote {cpath}")
    return 0


if __name__ == "__main__":
    sys.exit(main())