Download scripts/evaluate_layer17_moviebench_step2000.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 10.9 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/evaluate_layer17_moviebench_step2000.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/evaluate_layer17_moviebench_step2000.py
-
curl -L -o evaluate_layer17_moviebench_step2000.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/evaluate_layer17_moviebench_step2000.py
10.9 kB
| #!/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 | |
| 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() | |