#!/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 --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_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 ""))