Self-Forcing / scripts /evaluate_trained_long_predictors.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw
History Blame Contribute Delete
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()