#!/usr/bin/env python3 """Evaluate a trained Self-Forcing Layer-17 predictor on MovieBench prompts.""" from __future__ import annotations import argparse import json import os import sys import time 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 predictor_training.offline_data import TOKENS_PER_CHUNK from scripts.evaluate_single_block_fppf import ( FinalHiddenCapture, atomic_json, build_pipeline, frame_metrics, load_predictor, pixels_to_u8, predictor_step, save_mp4, ) from utils.misc import set_seed from utils.wan_wrapper import WanTextEncoder, WanVAEWrapper NUM_CHUNKS = 7 FRAMES_PER_CHUNK = 3 NUM_STEPS = 4 LATENT_CHANNELS = 16 LATENT_HEIGHT = 60 LATENT_WIDTH = 104 def read_lines(path: Path) -> list[str]: return [line.strip() for line in path.read_text(encoding="utf-8").splitlines() if line.strip()] def reset_caches(pipeline, device: torch.device) -> None: if pipeline.kv_cache1 is None: pipeline._initialize_kv_cache(1, torch.bfloat16, device) pipeline._initialize_crossattn_cache(1, torch.bfloat16, device) for cache in pipeline.kv_cache1: cache["global_end_index"].zero_() cache["local_end_index"].zero_() for cache in pipeline.crossattn_cache: cache["is_init"] = False @torch.inference_mode() def rollout(pipeline, conditional_dict, seed: int, device: torch.device, predictor=None): reset_caches(pipeline, device) set_seed(seed) noise = torch.randn( 1, NUM_CHUNKS * FRAMES_PER_CHUNK, LATENT_CHANNELS, LATENT_HEIGHT, LATENT_WIDTH, dtype=torch.bfloat16, device=device, ) teacher = pipeline.generator.model timesteps = pipeline.denoising_step_list.to(device=device) outputs = [] previous_chunk_hidden = None capture = FinalHiddenCapture(teacher) full_calls = predictor_calls = 0 started = time.perf_counter() try: for chunk in range(NUM_CHUNKS): noisy_input = noise[:, chunk * 3 : (chunk + 1) * 3] current_hidden = [None] * NUM_STEPS denoised_pred = timestep = None for step, current_timestep in enumerate(timesteps): timestep = torch.ones([1, 3], dtype=torch.int64, device=device) * current_timestep use_predictor = predictor is not None and chunk > 0 and step in {1, 2} if use_predictor: hidden, flow, _ = predictor_step( predictor=predictor, teacher=teacher, noisy_input=noisy_input, timestep=timestep, anchor_hidden=current_hidden[step - 1], previous_hidden=previous_chunk_hidden[step], history_cache=pipeline.kv_cache1[17], cross_cache=pipeline.crossattn_cache[17], current_start=chunk * TOKENS_PER_CHUNK, ) denoised_pred = pipeline.generator._convert_flow_pred_to_x0( flow_pred=flow.flatten(0, 1), xt=noisy_input.flatten(0, 1), timestep=timestep.flatten(0, 1), ).unflatten(0, flow.shape[:2]) current_hidden[step] = hidden predictor_calls += 1 else: capture.start() _, denoised_pred = pipeline.generator( noisy_image_or_video=noisy_input, conditional_dict=conditional_dict, timestep=timestep, kv_cache=pipeline.kv_cache1, crossattn_cache=pipeline.crossattn_cache, current_start=chunk * TOKENS_PER_CHUNK, ) current_hidden[step] = capture.finish() full_calls += 1 if step < NUM_STEPS - 1: flat = denoised_pred.flatten(0, 1) noisy_input = pipeline.scheduler.add_noise( flat, torch.randn_like(flat), timesteps[step + 1] * torch.ones([3], dtype=torch.long, device=device), ).unflatten(0, denoised_pred.shape[:2]) outputs.append(denoised_pred) pipeline.generator( noisy_image_or_video=denoised_pred, conditional_dict=conditional_dict, timestep=torch.ones_like(timestep) * pipeline.args.context_noise, kv_cache=pipeline.kv_cache1, crossattn_cache=pipeline.crossattn_cache, current_start=chunk * TOKENS_PER_CHUNK, ) previous_chunk_hidden = current_hidden finally: capture.close() torch.cuda.synchronize() return torch.cat(outputs, dim=1), { "generation_time_s": time.perf_counter() - started, "full_calls": full_calls, "predictor_calls": predictor_calls, } 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("--original_prompts", type=Path, required=True) parser.add_argument("--extended_prompts", type=Path, required=True) parser.add_argument("--output_dir", type=Path, required=True) parser.add_argument("--weights", type=Path, required=True) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--metric_batch_size", type=int, default=4) args = parser.parse_args() args.output_dir.mkdir(parents=True, exist_ok=True) original = read_lines(args.original_prompts) extended = read_lines(args.extended_prompts) if len(original) != len(extended) or min(args.prompt_ids) < 0 or max(args.prompt_ids) >= len(original): raise ValueError("MovieBench original/extended prompt pairing is invalid") atomic_json( args.output_dir / "manifest.json", { "status": "running", "gpu": args.gpu, "prompt_ids": args.prompt_ids, "weights": str(args.weights.resolve()), "teacher": str((ROOT / "checkpoints/self_forcing_dmd.pt").resolve()), "generation_seed_reset_per_prompt": args.seed, "generation_prompts": str(args.extended_prompts.resolve()), "evaluation_prompts": str(args.original_prompts.resolve()), "schedule": "chunk0=FFFF; chunks1-6=FPPF", "metrics": ["PSNR", "SSIM", "LPIPS", "rollout-only PSNR/SSIM/LPIPS"], }, ) device = torch.device("cuda") torch.set_grad_enabled(False) config = OmegaConf.merge( OmegaConf.load(ROOT / "configs/default_config.yaml"), OmegaConf.load(ROOT / "configs/self_forcing_sid.yaml"), ) vae = WanVAEWrapper().to(device=device, dtype=torch.bfloat16).eval() pipeline = build_pipeline(config, ROOT / "checkpoints/self_forcing_dmd.pt", vae, device) text_encoder = WanTextEncoder().to(device=device, dtype=torch.bfloat16).eval() text_encoder.requires_grad_(False) predictor = load_predictor( pipeline.generator.model, {"source_layer": 17, "weights": args.weights, "gate_mode": "baseline"}, device, ) lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() lpips_model.requires_grad_(False) reference_dir = args.output_dir / "videos" / "ffff" prediction_dir = args.output_dir / "videos" / "fppf_step2000" for offset, prompt_id in enumerate(args.prompt_ids, start=1): result_path = args.output_dir / "per_prompt" / f"prompt_{prompt_id:04d}.json" if ( result_path.exists() and (reference_dir / f"{prompt_id:05d}.mp4").exists() and (prediction_dir / f"{prompt_id:05d}.mp4").exists() ): print(f"[skip] {offset}/{len(args.prompt_ids)} id={prompt_id}", flush=True) continue print(f"[encode] {offset}/{len(args.prompt_ids)} id={prompt_id}", flush=True) conditional = text_encoder(text_prompts=[extended[prompt_id]]) reference_latent, ffff_counts = rollout(pipeline, conditional, args.seed, device) prediction_latent, fppf_counts = rollout( pipeline, conditional, args.seed, device, predictor=predictor ) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): reference_pixels = vae.decode_to_pixel(reference_latent, use_cache=False) prediction_pixels = vae.decode_to_pixel(prediction_latent, use_cache=False) reference_u8 = pixels_to_u8(reference_pixels) prediction_u8 = pixels_to_u8(prediction_pixels) save_mp4(reference_u8, reference_dir / f"{prompt_id:05d}.mp4") save_mp4(prediction_u8, prediction_dir / f"{prompt_id:05d}.mp4") metrics = frame_metrics( reference_u8=reference_u8, prediction_u8=prediction_u8, lpips_model=lpips_model, batch_size=args.metric_batch_size, device=device, ) atomic_json( result_path, { "status": "complete", "prompt_id": prompt_id, "original_prompt": original[prompt_id], "generation_prompt": extended[prompt_id], "seed": args.seed, "latent_frames": NUM_CHUNKS * FRAMES_PER_CHUNK, "decoded_frames": metrics["num_frames"], "schedule": "chunk0=FFFF; chunks1-6=FPPF", "ffff": ffff_counts, "fppf": fppf_counts, **metrics, }, ) print( f"[result] id={prompt_id} psnr={metrics['psnr']:.4f} " f"ssim={metrics['ssim']:.6f} lpips={metrics['lpips']:.6f}", flush=True, ) if hasattr(vae.model, "clear_cache"): vae.model.clear_cache() del conditional, reference_latent, prediction_latent, reference_pixels del prediction_pixels, reference_u8, prediction_u8 torch.cuda.empty_cache() manifest = json.loads((args.output_dir / "manifest.json").read_text(encoding="utf-8")) manifest["status"] = "complete" atomic_json(args.output_dir / "manifest.json", manifest) if __name__ == "__main__": main()