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