Download eval/generate_eval.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 11.6 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/eval/generate_eval.py
- Command line
-
hf download hf://Cccccz/comparison/eval/generate_eval.py
-
curl -L -o generate_eval.py https://huggingface.co/Cccccz/comparison/resolve/main/eval/generate_eval.py
11.6 kB
| #!/usr/bin/env python | |
| """Generate the Extended-251 videos for every strategy of one base model. | |
| Work is sharded by prompt, not by strategy, because the protocol requires each | |
| strategy's clip to be paired with the FFFF clip for the *same* prompt and seed | |
| (section 8.1). A shard therefore generates FFFF first, keeps its decoded frames, | |
| and immediately scores every cache strategy against them -- no 25 GB of reference | |
| frames on disk, and the pairing cannot get mismatched. | |
| CUDA_VISIBLE_DEVICES=0 python eval/generate_eval.py \ | |
| --base self_forcing --shard 0 --num-shards 4 --out-root eval_out | |
| Re-running skips prompts whose per-strategy record is already complete. | |
| """ | |
| import argparse | |
| import json | |
| import os | |
| import sys | |
| import time | |
| ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| sys.path.insert(0, ROOT) | |
| from eval.strategies import load_strategies # noqa: E402 | |
| from harness import denoising_steps_for, BASES, enter_base, load_pipeline, use_denoising_steps # noqa: E402 | |
| MAPPING = os.path.join(ROOT, "assets/vbench8_extended_subset_mapping.json") | |
| def parse_args(): | |
| p = argparse.ArgumentParser() | |
| p.add_argument("--base", choices=sorted(BASES), required=True) | |
| p.add_argument("--shard", type=int, default=0) | |
| p.add_argument("--num-shards", type=int, default=1) | |
| p.add_argument("--out-root", default="eval_out") | |
| p.add_argument("--seed", type=int, default=0) | |
| p.add_argument("--num-latent-frames", type=int, default=21) | |
| p.add_argument("--fps", type=int, default=16) | |
| p.add_argument("--limit", type=int, default=None, help="Smoke-test: first N prompts") | |
| p.add_argument("--only-strategy", default=None, help="Smoke-test: one strategy name") | |
| p.add_argument("--strategies", default=None, | |
| help="Comma-separated subset to generate (the FFFF reference is always kept)") | |
| p.add_argument("--overwrite", action="store_true") | |
| p.add_argument("--regenerate-ffff", action="store_true", | |
| help="Regenerate the FFFF reference instead of reading its MP4") | |
| return p.parse_args() | |
| def main(): | |
| args = 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) | |
| rows = mapping["rows"] | |
| if args.limit: | |
| rows = rows[:args.limit] | |
| shard_rows = rows[args.shard::args.num_shards] | |
| strategies = load_strategies(args.base) | |
| if args.only_strategy: | |
| keep = {strategies[0]["name"], args.only_strategy} | |
| strategies = [s for s in strategies if s["name"] in keep] | |
| if args.strategies: | |
| keep = set(args.strategies.split(",")) | |
| strategies = [s for s in strategies if s["is_reference"] or s["name"] in keep] | |
| ref_name = strategies[0]["name"] | |
| assert strategies[0]["is_reference"], "first strategy must be FFFF" | |
| print(f"base={args.base} shard {args.shard}/{args.num_shards}: " | |
| f"{len(shard_rows)} prompts x {len(strategies)} strategies", flush=True) | |
| base = enter_base(args.base) | |
| import torch | |
| from einops import rearrange | |
| from torchvision.io import read_video, write_video | |
| from utils.misc import set_seed | |
| from cachelib import (CacheController, build_method, cached_inference, install, | |
| parse_schedule) | |
| from eval.pixel_metrics import PixelMetrics | |
| torch.set_grad_enabled(False) | |
| device = torch.device("cuda") | |
| pipeline = load_pipeline(base) | |
| num_steps = len(pipeline.denoising_step_list) | |
| metrics = PixelMetrics(device) | |
| def record_path(strategy, row): | |
| d = os.path.join(out_root, "per_prompt", strategy) | |
| os.makedirs(d, exist_ok=True) | |
| return os.path.join(d, f"{row['prompt_suite']}_{row['suite_index']:03d}.json") | |
| def video_path(strategy, row): | |
| d = os.path.join(out_root, "generated_videos", strategy, row["prompt_suite"]) | |
| os.makedirs(d, exist_ok=True) | |
| return os.path.join(d, f"{row['suite_index']:03d}.mp4") | |
| def ref_cache_path(row): | |
| d = os.path.join(out_root, "ref_cache", args.base) | |
| os.makedirs(d, exist_ok=True) | |
| return os.path.join(d, f"{row['prompt_suite']}_{row['suite_index']:03d}.pt") | |
| def is_done(strategy, row): | |
| p = record_path(strategy, row) | |
| if args.overwrite or not os.path.exists(p): | |
| return False | |
| try: | |
| with open(p) as f: | |
| rec = json.load(f) | |
| return rec.get("status") == "complete" and os.path.exists(rec.get("video", "")) | |
| except Exception: | |
| return False | |
| def write_json(path, payload): | |
| tmp = path + ".tmp" | |
| with open(tmp, "w") as f: | |
| json.dump(payload, f, indent=2) | |
| os.replace(tmp, path) # atomic: a partial file is never seen as complete | |
| def run_one(strategy, row): | |
| kw = dict(indicator=strategy["indicator"], coefficients=strategy["coefficients"]) | |
| if strategy["method"] != "none" and strategy.get("param"): | |
| kw[strategy["param"]] = strategy["value"] | |
| n_steps = strategy.get("num_inference_steps") or num_steps | |
| original_steps = use_denoising_steps(pipeline, n_steps) | |
| sched = strategy.get("schedule") or ("F" * n_steps) | |
| first = strategy.get("first_chunk_schedule") | |
| # Chunk 0 may have a longer schedule than the rest (e.g. FFFF over naive 2-step). | |
| first_steps = (denoising_steps_for(original_steps, len(first)) | |
| if first and len(first) != n_steps else None) | |
| ctrl = CacheController( | |
| build_method(strategy["method"], **kw), num_steps=n_steps, | |
| forced_steps=parse_schedule(sched, n_steps), | |
| first_chunk_forced_steps=(parse_schedule(first, len(first)) if first else None)) | |
| install(pipeline.generator.model, ctrl) | |
| # Identical noise for every strategy: same seed, same draw order, same shape. | |
| set_seed(args.seed) | |
| noise = torch.randn([1, args.num_latent_frames, 16, 60, 104], | |
| device=device, dtype=torch.bfloat16) | |
| wall0 = time.time() | |
| video, latents, timing = cached_inference( | |
| pipeline, ctrl, noise, [row["extended_prompt"]], decode=True, | |
| first_chunk_steps=first_steps, sampler=strategy.get("sampler", "renoise")) | |
| wall = time.time() - wall0 | |
| pipeline.vae.model.clear_cache() | |
| pipeline.denoising_step_list = original_steps | |
| return video, timing, ctrl.summary(), wall | |
| # Warm up every strategy once on a throwaway prompt and discard the result. | |
| # Without this the first recorded prompt carries one-off costs -- notably | |
| # TaylorSeer's torch.compile of the forecast, which made its first video look | |
| # slower than FFFF. | |
| if shard_rows: | |
| print("warm-up pass (not recorded)", flush=True) | |
| for strategy in strategies: | |
| run_one(strategy, shard_rows[0]) | |
| torch.cuda.empty_cache() | |
| t_start = time.time() | |
| for n, row in enumerate(shard_rows): | |
| todo = [s for s in strategies if not is_done(s["name"], row)] | |
| if not todo: | |
| continue | |
| # The reference is needed for pixel metrics even when only some strategies | |
| # are outstanding. | |
| need_ref = any(not s["is_reference"] for s in todo) | |
| ref_frames = None | |
| # The FFFF reference frames (pre-MP4, fp16) are cached on disk after their | |
| # first generation, so later strategy runs do not regenerate FFFF at all. | |
| rc = ref_cache_path(row) | |
| ref_from_mp4 = False | |
| if need_ref and os.path.exists(rc): | |
| ref_frames = torch.load(rc, map_location=device) | |
| need_ref = False | |
| elif need_ref and os.path.exists(video_path(ref_name, row)) and not args.regenerate_ffff: | |
| # Reference from the FFFF run's MP4 (H.264, ~40+ dB above the 12-20 dB | |
| # strategy PSNRs, so the compression adds <1 % to the MSE); never | |
| # regenerate FFFF just to compare against it. | |
| vid, _, _ = read_video(video_path(ref_name, row), pts_unit="sec", output_format="TCHW") | |
| ref_frames = (vid.to(device, torch.float16) / 255.0) | |
| ref_from_mp4 = True | |
| need_ref = False | |
| for strategy in strategies: | |
| if strategy["is_reference"]: | |
| if strategy not in todo and not need_ref: | |
| continue | |
| elif strategy not in todo: | |
| continue | |
| video, timing, summary, wall = run_one(strategy, row) | |
| frames = video[0] # [T, C, H, W] in [0, 1] | |
| if strategy["is_reference"]: | |
| ref_frames = frames.detach().to(torch.float16) | |
| if not os.path.exists(rc): | |
| torch.save(ref_frames.cpu(), rc + ".tmp") | |
| os.replace(rc + ".tmp", rc) | |
| pix = PixelMetrics.identity(frames.shape[0]) | |
| else: | |
| pix = metrics.compute(frames, ref_frames) | |
| if strategy in todo: | |
| vp = video_path(strategy["name"], row) | |
| arr = (255.0 * rearrange(frames, "t c h w -> t h w c")).clamp(0, 255) | |
| write_video(vp, arr.to(torch.uint8).cpu(), fps=args.fps) | |
| write_json(record_path(strategy["name"], row), { | |
| "status": "complete", | |
| "protocol": "Self-Forcing Extended-251 Full Evaluation", | |
| "strategy": strategy["name"], | |
| "base_model": args.base, | |
| "method": strategy["method"], | |
| "param": strategy["param"], | |
| "param_value": strategy["value"], | |
| "target_speedup": strategy["target"], | |
| "schedule": strategy.get("schedule") or ("F" * (strategy.get("num_inference_steps") or 4)), | |
| "num_inference_steps": strategy.get("num_inference_steps") or 4, | |
| "first_chunk_schedule": strategy.get("first_chunk_schedule"), | |
| "sampler": strategy.get("sampler", "renoise"), | |
| "global_index": row["global_index"], | |
| "prompt_suite": row["prompt_suite"], | |
| "suite_index": row["suite_index"], | |
| "prompt": row["extended_prompt"], | |
| "original_prompt": row["original_prompt"], | |
| "seed": args.seed, | |
| "video": vp, | |
| "num_frames": int(frames.shape[0]), | |
| "height": int(frames.shape[2]), | |
| "width": int(frames.shape[3]), | |
| "fps": args.fps, | |
| "policy_latency_ms": timing["denoise_dit_ms"], | |
| "excluded_context_kv_latency_ms": timing["context_kv_dit_ms"], | |
| "wall_generation_s": wall, | |
| "reference_strategy": ref_name, | |
| "pixel_metrics_vs_ffff": pix, | |
| "reference_source": "ffff_mp4" if ref_from_mp4 else "ffff_frames", | |
| "cache_diagnostics": dict(summary), | |
| }) | |
| del frames, video | |
| ref_frames = None | |
| done = n + 1 | |
| rate = (time.time() - t_start) / done | |
| print(f"[{args.base} shard {args.shard}] {done}/{len(shard_rows)} prompts " | |
| f"({row['prompt_suite']}/{row['suite_index']:03d}) " | |
| f"{rate:.1f}s/prompt eta {(len(shard_rows) - done) * rate / 60:.0f} min", | |
| flush=True) | |
| print(f"shard {args.shard} done in {(time.time() - t_start) / 60:.1f} min", flush=True) | |
| if __name__ == "__main__": | |
| main() | |