#!/usr/bin/env python3 """Evaluate 2x/4x-trained Layer-17 predictors at 1x, 2x, and 4x.""" from __future__ import annotations import argparse import json import os import sys from pathlib import Path def preparse_gpu() -> str: parser = argparse.ArgumentParser(add_help=False) parser.add_argument("--gpu", required=True) args, _ = parser.parse_known_args() os.environ["CUDA_VISIBLE_DEVICES"] = args.gpu return args.gpu GPU = preparse_gpu() import lpips import torch from omegaconf import OmegaConf ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from scripts.evaluate_long_video_fppf import generate_rollout, save_mp4 from scripts.evaluate_single_block_fppf import ( atomic_json, build_pipeline, frame_metrics, load_predictor, load_prompt_metadata, pixels_to_u8, ) from utils.misc import set_seed from utils.wan_wrapper import WanVAEWrapper PREDICTORS = { "trained_2x": ROOT / "outputs/layer17_long_training_four_gpu_v2/2x/predictor_final.safetensors", "trained_4x": ROOT / "outputs/layer17_long_training_four_gpu_v2/4x/predictor_final.safetensors", } def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--gpu", default=GPU) parser.add_argument("--prompt_ids", type=int, nargs="+", required=True) parser.add_argument("--latent_lengths", type=int, nargs="+", default=[21, 42, 84]) parser.add_argument( "--dataset_root", type=Path, default=Path("outputs/predictor_offline_100_all_blocks"), ) parser.add_argument( "--output_dir", type=Path, default=Path("outputs/layer17_long_training_eval"), ) parser.add_argument("--generation_seed", type=int, default=0) parser.add_argument("--metric_batch_size", type=int, default=4) args = parser.parse_args() args.dataset_root = (ROOT / args.dataset_root).resolve() if not args.dataset_root.is_absolute() else args.dataset_root args.output_dir = (ROOT / args.output_dir).resolve() if not args.output_dir.is_absolute() else args.output_dir args.output_dir.mkdir(parents=True, exist_ok=True) device = torch.device("cuda") set_seed(args.generation_seed) config = OmegaConf.merge( OmegaConf.load(ROOT / "configs/default_config.yaml"), OmegaConf.load(ROOT / "configs/self_forcing_sid.yaml"), ) config.model_kwargs.local_attn_size = 21 vae = WanVAEWrapper().to(device=device, dtype=torch.bfloat16).eval() pipeline = build_pipeline( config, ROOT / "checkpoints/self_forcing_dmd.pt", vae, device, ) predictors = { name: load_predictor( pipeline.generator.model, {"source_layer": 17, "weights": path}, device, ) for name, path in PREDICTORS.items() } lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() lpips_model.requires_grad_(False) for prompt_id in args.prompt_ids: prompt = load_prompt_metadata(args.dataset_root, prompt_id)["prompt"] for latent_length in args.latent_lengths: run_dir = args.output_dir / f"latent_{latent_length}" / f"prompt_{prompt_id:04d}" result_path = run_dir / "metrics.json" if result_path.exists(): existing = json.loads(result_path.read_text()) if existing.get("status") == "complete": print(f"[skip] prompt={prompt_id} latent={latent_length}", flush=True) continue print(f"[run] prompt={prompt_id} latent={latent_length} FFFF", flush=True) reference_latent, ffff_counts = generate_rollout( pipeline=pipeline, dataset_root=args.dataset_root, prompt_id=prompt_id, latent_length=latent_length, generation_seed=args.generation_seed, device=device, predictor=None, source_layer=None, schedule="FFFF", ) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): reference_pixels = vae.decode_to_pixel(reference_latent, use_cache=False) reference_u8 = pixels_to_u8(reference_pixels) save_mp4(reference_u8, run_dir / "ffff.mp4") del reference_latent, reference_pixels if hasattr(vae.model, "clear_cache"): vae.model.clear_cache() torch.cuda.empty_cache() results = {} for name, predictor in predictors.items(): print(f"[run] prompt={prompt_id} latent={latent_length} {name}", flush=True) latent, counts = generate_rollout( pipeline=pipeline, dataset_root=args.dataset_root, prompt_id=prompt_id, latent_length=latent_length, generation_seed=args.generation_seed, device=device, predictor=predictor, source_layer=17, schedule="FPPF", ) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): pixels = vae.decode_to_pixel(latent, use_cache=False) prediction_u8 = pixels_to_u8(pixels) save_mp4(prediction_u8, run_dir / f"{name}.mp4") metrics = frame_metrics( reference_u8=reference_u8, prediction_u8=prediction_u8, lpips_model=lpips_model, batch_size=args.metric_batch_size, device=device, ) results[name] = {"fppf": counts, **metrics} print( f"[result] {name} prompt={prompt_id} latent={latent_length} " f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} " f"lpips={metrics['lpips']:.6f}", flush=True, ) del latent, pixels, prediction_u8 if hasattr(vae.model, "clear_cache"): vae.model.clear_cache() torch.cuda.empty_cache() atomic_json(result_path, { "status": "complete", "prompt_id": prompt_id, "prompt": prompt, "latent_length": latent_length, "decoded_frames": next(iter(results.values()))["num_frames"], "reference": "FFFF same prompt/seed/noise", "ffff": ffff_counts, "predictors": results, }) del reference_u8 if __name__ == "__main__": main()