comparison / hy_import_dev_eval.py
Cccccz's picture
Add files using upload-large-folder tool
f87692b verified
Raw History Blame Contribute Delete
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 ""))