Download hy_import_dev_eval.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 7.94 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/hy_import_dev_eval.py
- Command line
-
hf download hf://Cccccz/comparison/hy_import_dev_eval.py
-
curl -L -o hy_import_dev_eval.py https://huggingface.co/Cccccz/comparison/resolve/main/hy_import_dev_eval.py
7.94 kB
| #!/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 "")) | |