| |
| """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() |
|
|