VIPRA-reproduce / viper /make_comparison.py
yqi19's picture
VIPER reproduction: code, mini VIPER-19K pipeline, stage-1 checkpoint, qualitative comparisons and eval
e0db531 verified
Raw
History Blame Contribute Delete
8.06 kB
"""Build the qualitative comparison set for a set of validation pairs.
For each pair we save every artifact needed to judge the method by eye:
reference.mp4 the reference video -- the physical process to transfer
target_gt.mp4 the ground-truth target video (what "right" looks like)
target_img.png the target image the generator is actually conditioned on
baseline.mp4 Wan2.2-I2V with the SAME image + prompt and NO reference
(physics tokens zeroed = the untouched base-model context)
viper.mp4 Wan2.2-I2V + VIPER physics tokens from the reference
grid.mp4 all four videos tiled 2x2 with labels, for side-by-side viewing
baseline and viper use identical noise seeds, so any difference between them is
attributable to the reference stream alone.
"""
from __future__ import annotations
import argparse
import json
import shutil
import sys
from pathlib import Path
import numpy as np
import torch
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "third_party" / "Wan2.2"))
from viper.infer import generate, save_video # noqa: E402
from viper.physics_encoder import VisualPhysicsEncoder # noqa: E402
from viper.wan_viper import add_lora, patch_wan_model # noqa: E402
LABELS = {
"reference": "REFERENCE (physics source)",
"target_gt": "TARGET ground truth",
"baseline": "BASELINE Wan2.2 (no reference)",
"viper": "VIPER (reference-conditioned)",
}
def label_frames(frames: np.ndarray, text: str) -> np.ndarray:
"""Burn a caption bar onto the top of every frame."""
import cv2
out = frames.copy()
for i in range(len(out)):
f = np.ascontiguousarray(out[i])
cv2.rectangle(f, (0, 0), (f.shape[1], 28), (0, 0, 0), -1)
cv2.putText(f, text, (8, 20), cv2.FONT_HERSHEY_SIMPLEX, 0.55,
(255, 255, 255), 1, cv2.LINE_AA)
out[i] = f
return out
def make_grid(paths: dict, out_path: str, fps: int = 16):
"""Tile four videos 2x2 into a single labelled mp4."""
import imageio
import imageio.v3 as iio
vids = {}
for k in ("reference", "target_gt", "baseline", "viper"):
v = iio.imread(paths[k], plugin="pyav")
vids[k] = label_frames(v, LABELS[k])
n = min(len(v) for v in vids.values())
H, W = vids["viper"].shape[1:3]
def fit(v):
import cv2
return np.stack([cv2.resize(f, (W, H)) for f in v[:n]])
a, b, c, d = (fit(vids[k]) for k in
("reference", "target_gt", "baseline", "viper"))
top = np.concatenate([a, b], axis=2)
bot = np.concatenate([c, d], axis=2)
grid = np.concatenate([top, bot], axis=1)
imageio.mimsave(out_path, list(grid), fps=fps, quality=7)
print(f" grid -> {out_path}")
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--ckpt", required=True)
ap.add_argument("--pairs", default="data/pairs_val.jsonl")
ap.add_argument("--out_dir", default="results/comparison")
ap.add_argument("--wan_ckpt", default="models/Wan2.2-I2V-A14B")
ap.add_argument("--mllm", default="models/Qwen3-VL-4B-Instruct")
ap.add_argument("--steps", type=int, default=50)
ap.add_argument("--guide_scale", type=float, default=6.0)
ap.add_argument("--shift", type=float, default=5.0)
ap.add_argument("--frames", type=int, default=81)
ap.add_argument("--num_ref_frames", type=int, default=8)
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--limit", type=int, default=0)
ap.add_argument("--shard", type=int, default=0)
ap.add_argument("--num_shards", type=int, default=1)
args = ap.parse_args()
import imageio.v3 as iio
from PIL import Image
from transformers import AutoProcessor, Qwen3VLForConditionalGeneration
from wan.modules.model import WanModel
from wan.modules.t5 import T5EncoderModel
from wan.modules.vae2_1 import Wan2_1_VAE
device = torch.device("cuda")
ckpt_dir = Path(args.wan_ckpt)
rows = [json.loads(l) for l in open(args.pairs)]
if args.limit:
rows = rows[: args.limit]
rows = rows[args.shard::args.num_shards]
print(f"[{args.shard}] {len(rows)} pairs")
vae = Wan2_1_VAE(vae_pth=str(ckpt_dir / "Wan2.1_VAE.pth"), device=device)
t5 = T5EncoderModel(
text_len=512, dtype=torch.bfloat16, device=device,
checkpoint_path=str(ckpt_dir / "models_t5_umt5-xxl-enc-bf16.pth"),
tokenizer_path=str(ckpt_dir / "google/umt5-xxl"))
sd = torch.load(args.ckpt, map_location="cpu")
nq = sd["args"]["num_queries"]
proc = AutoProcessor.from_pretrained(args.mllm)
mllm = Qwen3VLForConditionalGeneration.from_pretrained(
args.mllm, dtype=torch.bfloat16, attn_implementation="sdpa").to(device)
enc = VisualPhysicsEncoder(mllm, nq, out_dim=5120).to(device)
enc.query_tokens.data.copy_(sd["query_tokens"].to(device))
enc.connector.load_state_dict(sd["connector"])
enc.connector.to(device, torch.float32)
enc.eval()
n_gpu = torch.cuda.device_count()
dev_low = torch.device("cuda:0")
dev_high = torch.device(f"cuda:{1 if n_gpu > 1 else 0}")
def load_expert(sub, dev):
m = patch_wan_model(WanModel.from_pretrained(str(ckpt_dir / sub)))
return m.to(device=dev, dtype=torch.bfloat16).eval().requires_grad_(False)
dit_low = load_expert("low_noise_model", dev_low)
dit_high = load_expert("high_noise_model", dev_high)
if "lora" in sd:
for m in (dit_low, dit_high):
add_lora(m, sd["args"]["lora_rank"], sd["args"]["lora_alpha"])
msd = dict(m.named_parameters())
for k, v in sd["lora"].items():
if k in msd:
msd[k].data.copy_(v.to(next(m.parameters()).device))
out_root = Path(args.out_dir)
out_root.mkdir(parents=True, exist_ok=True)
for i, r in enumerate(rows):
tag = f"{r['ref_id']}__{r['tgt_id']}"
d = out_root / tag
d.mkdir(exist_ok=True)
print(f"[{args.shard}] {i+1}/{len(rows)} {tag}", flush=True)
# --- inputs -------------------------------------------------
shutil.copy(r["ref_video"], d / "reference.mp4")
shutil.copy(r["tgt_video"], d / "target_gt.mp4")
tgt = iio.imread(r["tgt_video"], plugin="pyav")
Image.fromarray(tgt[0]).save(d / "target_img.png")
image = (torch.from_numpy(tgt[0].copy()).permute(2, 0, 1)
.float().div_(127.5).sub_(1))
ref = iio.imread(r["ref_video"], plugin="pyav")
idx = np.linspace(0, len(ref) - 1, args.num_ref_frames).astype(int)
ref_frames = ref[idx]
prompt = r.get("tgt_summary") or r.get("tgt_caption", "")[:300]
# --- generate both variants with identical seeds -------------
for name, zero in (("baseline", True), ("viper", False)):
saved = enc.query_tokens.data.clone()
if zero:
enc.query_tokens.data.zero_()
out = generate(dit_low, dit_high, vae, t5, enc, proc, image, prompt,
ref_frames, device, num_queries=nq, steps=args.steps,
guide_scale=args.guide_scale, shift=args.shift,
frames=args.frames, seed=args.seed)
enc.query_tokens.data.copy_(saved)
save_video(out, str(d / f"{name}.mp4"))
paths = {k: str(d / f"{k}.mp4") for k in
("reference", "target_gt", "baseline", "viper")}
make_grid(paths, str(d / "grid.mp4"))
json.dump({
"id": tag, "prompt": prompt,
"physical_impact": r.get("physical_impact"),
"trajectory": r.get("trajectory"),
"transferability": r.get("transferability"),
"judge_reason": r.get("judge_reason"),
"ckpt": args.ckpt, "steps": args.steps,
"guide_scale": args.guide_scale, "seed": args.seed,
}, open(d / "info.json", "w"), indent=2, ensure_ascii=False)
print(f"[{args.shard}] DONE -> {out_root}")
if __name__ == "__main__":
main()