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