File size: 7,938 Bytes
f87692b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python
"""Import a HY-WorldPlay-DEV-Predictor validation25 evaluation into eval_out_hy as a
strategy: symlink its 100 videos and write per-video records in the layout
hy_aggregate.py / hy_vbench.sh expect.

Latency: the DEV generator profiles host wall-clock per stage (HY_PROFILE_TIMING);
``ar_step_transformer + ar_step_predictor`` is the denoising path (no context /
history KV passes), the same boundary as hycache's ``denoise_dit_ms``.  It is host
timing on a shared 2-worker run, so the row is marked provisional.  Compute: full
DiT forwards + predictor calls at 1/54 of a forward (one of 54 blocks).

    python hy_import_dev_eval.py --eval-root <DEV eval dir> --strategy hy_atc_s1_fppf \
        --schedule FPPF --model atc_stage1
"""
import argparse, glob, json, os, re
ROOT = os.path.dirname(os.path.abspath(__file__))
ap = argparse.ArgumentParser()
ap.add_argument("--eval-root", required=True)
ap.add_argument("--strategy", required=True)
ap.add_argument("--schedule", required=True, help="FPPF or FPPP (chunk pattern)")
ap.add_argument("--model", required=True)
ap.add_argument("--out-root", default=os.path.join(ROOT, "eval_out_hy"))
ap.add_argument("--num-chunks", type=int, default=8)
ap.add_argument("--num-blocks", type=int, default=54)
ap.add_argument("--predictor-blocks", type=int, default=1, help="Teacher blocks' worth of compute per Predictor call (DisCa: 2)")
ap.add_argument("--first-chunk", default=None, help="FFFF when the DEV run used --force_first_chunk_full")
ap.add_argument("--num-frames", type=int, default=125)
ap.add_argument("--reference-strategy", default="hy_ffff",
                help="Reference FFFF strategy in eval_out_hy (hy_<n>c_ffff for long videos)")
ap.add_argument("--metrics-from-videos", action="store_true",
                help="Compute PSNR/SSIM/LPIPS here against the reference strategy's MP4s "
                     "(eval/pixel_metrics, on the GPU) instead of reading the DEV run's "
                     "metrics.json -- the DEV metrics tool only accepts 125-frame videos")
a = ap.parse_args()

cases = {}
dev_root = os.path.dirname(os.path.dirname(a.eval_root.rstrip("/")))
for line in open(os.path.join(dev_root, "validation", "cases.jsonl")):
    c = json.loads(line); cases[int(c.get("split_case_id", c.get("case_id")))] = c
if a.metrics_from_videos:
    import sys, torch
    sys.path.insert(0, ROOT)
    from torchvision.io import read_video
    from eval.pixel_metrics import PixelMetrics
    pm = PixelMetrics(torch.device("cuda"))

    def video_metrics(vp, cid, ai):
        ref_vp = os.path.join(a.out_root, "generated_videos", a.reference_strategy, f"case_{cid:04d}_action_{ai:02d}.mp4")
        x = read_video(vp, pts_unit="sec", output_format="TCHW")[0].to(torch.float16) / 255.0
        y = read_video(ref_vp, pts_unit="sec", output_format="TCHW")[0].to(torch.float16) / 255.0
        assert x.shape[0] == a.num_frames and y.shape[0] == a.num_frames, (vp, x.shape, y.shape)
        r = pm.compute(x, y)
        return {"pixel_mse": r["mean_mse"], "psnr_db": r["psnr"], "ssim": r["ssim"], "lpips_alex": r["lpips"]}
    metrics = None
else:
    metrics = {(r["case_id"], r["action_id"]): r for r in json.load(open(os.path.join(a.eval_root, "metrics.json")))["records"]}
timing = {}
for lf in glob.glob(os.path.join(a.eval_root, "logs", "generate_*worker_*.log")):
    for line in open(lf):
        if line.startswith("{"):
            r = json.loads(line); timing[(r["case_id"], r["action_id"])] = r
# Forward counts come from the generator's own stage timing (one entry per call):
# the DEV rollout always denoises chunk 0 in full (a Predictor needs a previous
# chunk), so FPPF is 18 full + 14 Predictor over 8 chunks, FPPP 11 + 21.
def counts(tm):
    return tm["ar_step_transformer"]["count"], tm.get("ar_step_predictor", {"count": 0})["count"]
full_fw, pred_fw = counts(next(iter(timing.values()))["timing"])
assert all(counts(t["timing"]) == (full_fw, pred_fw) for t in timing.values()), "forward counts differ across videos"
# A video whose generation log line was lost (a resumed run re-opened the log) gets
# the mean stage timing of the videos that do have one; the record says so.
def stage_total(tm, k):
    return tm.get(k, {"total_s": 0.0})["total_s"]
mean_tm = {k: sum(stage_total(t["timing"], k) for t in timing.values()) / len(timing)
           for k in ("ar_step_transformer", "ar_step_predictor", "ar_history_kv_cache")}
assert len(timing) >= 20, f"only {len(timing)} timing entries"
n_imputed = 0
compute = full_fw + pred_fw * a.predictor_blocks / a.num_blocks
first_chunk = "FFFF" if full_fw == 4 + (a.num_chunks - 1) * a.schedule.upper().count("F") else a.first_chunk
vdir = os.path.join(a.out_root, "generated_videos", a.strategy)
rdir = os.path.join(a.out_root, "per_prompt", a.strategy)
os.makedirs(vdir, exist_ok=True); os.makedirs(rdir, exist_ok=True)
n = 0
for vp in sorted(glob.glob(os.path.join(a.eval_root, "predictor", "case_*_action_*.mp4"))):
    m = re.match(r"case_(\d+)_action_(\d+)\.mp4", os.path.basename(vp))
    cid, ai = int(m.group(1)), int(m.group(2))
    mt = video_metrics(vp, cid, ai) if metrics is None else metrics[(cid, ai)]
    imputed = (cid, ai) not in timing
    if imputed:
        n_imputed += 1
        tm = {k: {"total_s": v} for k, v in mean_tm.items()}
        tinfo = next(iter(timing.values()))
    else:
        tm = timing[(cid, ai)]["timing"]; tinfo = timing[(cid, ai)]
    denoise_s = tm["ar_step_transformer"]["total_s"] + tm.get("ar_step_predictor", {"total_s": 0.0})["total_s"]
    link = os.path.join(vdir, os.path.basename(vp))
    if os.path.islink(link) or os.path.exists(link):
        os.remove(link)
    os.symlink(vp, link)
    case = cases.get(cid, {})
    rec = {"status": "complete", "protocol": "HY-WorldPlay VBench-I2V-100 (validation25 x 4 actions)",
           "strategy": a.strategy, "base_model": "hy_worldplay", "method": f"predictor_{a.model}",
           "param": "pattern", "param_value": a.schedule.upper(), "target_speedup": 32.0 / compute,
           "schedule": a.schedule.upper(), "num_inference_steps": 4, "first_chunk_schedule": first_chunk,
           "case_id": cid, "split_case_id": cid, "action_id": ai, "action_name": tinfo.get("action_name"),
           "prompt": case.get("caption"), "image_path": case.get("image_path"), "seed": tinfo.get("seed", 0),
           "video": link, "num_frames": a.num_frames, "height": 480, "width": 832, "fps": 24,
           "policy_latency_ms": 1000.0 * denoise_s,
           "excluded_context_kv_latency_ms": 1000.0 * tm.get("ar_history_kv_cache", {"total_s": 0.0})["total_s"],
           "matched_ffff_policy_latency_ms": None,
           "latency_source": "DEV generator host-side stage timing (ar_step_transformer + ar_step_predictor); FFFF from its own records",
           "peak_mem_gib": tinfo.get("peak_memory_gib"),
           "latency_imputed": imputed,
           "reference_strategy": a.reference_strategy,
           "reference_source": f"{a.reference_strategy}_mp4" if a.metrics_from_videos else "dev_validation_full_mp4",
           "pixel_metrics_vs_ffff": {"mean_mse": mt["pixel_mse"], "psnr": mt["psnr_db"], "ssim": mt["ssim"], "lpips": mt["lpips_alex"]},
           "cache_diagnostics": {"denoise_forwards": full_fw + pred_fw, "full_forwards": full_fw,
                                 "predictor_forwards": pred_fw, "predictor_compute_equivalent": a.predictor_blocks / a.num_blocks,
                                 "compute_equivalent_forwards": compute, "middle_steps": a.num_chunks * 2,
                                 "middle_compute_equivalent": compute - a.num_chunks * 2},
           "dev_eval_root": a.eval_root}
    json.dump(rec, open(os.path.join(rdir, f"case_{cid:04d}_action_{ai:02d}.json"), "w"), indent=2)
    n += 1
print(f"{a.strategy}: imported {n} videos/records from {a.eval_root}"
      + (f" ({n_imputed} with mean stage timing: log lines lost)" if n_imputed else ""))