File size: 6,493 Bytes
bc29ee3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 | #!/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()
|