"""Shared model loading / evaluation so a sweep can reuse one loaded pipeline. Import this only *after* chdir-ing into the chosen base repo (``enter_base`` does that), because both base repos load ``wan_models/...`` by relative path. """ import os import statistics import sys ROOT = os.path.dirname(os.path.abspath(__file__)) BASES = { "self_forcing": { "repo": os.path.join(ROOT, "repos/Self-Forcing"), "config": "configs/self_forcing_dmd.yaml", "checkpoint": "checkpoints/self_forcing_dmd.pt", }, "causal_forcing": { "repo": os.path.join(ROOT, "repos/Causal-Forcing"), "config": "configs/causal_forcing_dmd_chunkwise.yaml", "checkpoint": "checkpoints/chunkwise/causal_forcing.pt", }, } def enter_base(name): base = BASES[name] os.chdir(base["repo"]) sys.path.insert(0, base["repo"]) sys.path.insert(0, ROOT) return base def load_pipeline(base, device=None): import torch from omegaconf import OmegaConf from pipeline import CausalInferencePipeline device = device or torch.device("cuda") config = OmegaConf.merge(OmegaConf.load("configs/default_config.yaml"), OmegaConf.load(base["config"])) pipeline = CausalInferencePipeline(config, device=device) sd = torch.load(base["checkpoint"], map_location="cpu") # self_forcing_dmd.pt ships only EMA weights, causal_forcing.pt only raw ones. gen_sd = sd["generator"] if "generator" in sd else sd["generator_ema"] try: pipeline.generator.load_state_dict(gen_sd) except RuntimeError: pipeline.generator.load_state_dict( {k.replace("model._fsdp_wrapped_module.", "model.", 1): v for k, v in gen_sd.items()}, strict=False) pipeline = pipeline.to(dtype=torch.bfloat16) pipeline.text_encoder.to(device) pipeline.generator.to(device) pipeline.vae.to(device) return pipeline def load_prompts(n, offset=0): from utils.dataset import TextDataset path = ("prompts/MovieGenVideoBench.txt" if os.path.exists("prompts/MovieGenVideoBench.txt") else "prompts/demos.txt") ds = TextDataset(prompt_path=path) idx = range(offset, min(offset + n, len(ds))) return [ds[i]["prompts"] for i in idx] def free_gpus(min_free_mb=70000, max_util=5): """Indices of GPUs with no meaningful other tenant, most-free first. This node is shared, and a neighbour saturating the SMs inflates timings by 30%+ and drifts them within a single sweep, so runs must be placed carefully and speedups measured pairwise. """ import subprocess try: out = subprocess.check_output( ["nvidia-smi", "--query-gpu=index,utilization.gpu,memory.used,memory.total", "--format=csv,noheader,nounits"], text=True) except Exception: return [] rows = [] for line in out.strip().splitlines(): idx, util, used, total = [int(v.strip()) for v in line.split(",")] if util <= max_util and (total - used) >= min_free_mb: rows.append((total - used, idx)) return [i for _, i in sorted(rows, reverse=True)] def paired_evaluate(pipeline, make_baseline_ctrl, make_method_ctrl, prompts, num_output_frames=21, seed=0, warmup=1, save_video_dir=None): """Interleave baseline and method on the same prompt, one pair at a time. Returns per-prompt ratios. Absolute timings drift with GPU contention, but the two halves of a pair are seconds apart, so their ratio is stable in a way a start-of-sweep baseline is not. """ import statistics import torch from utils.misc import set_seed from cachelib import cached_inference ratios, base_ms, meth_ms, meth_ce, meth_fw = [], [], [], [], [] for i, prompt in enumerate(prompts): row = {} for tag, make in (("base", make_baseline_ctrl), ("method", make_method_ctrl)): ctrl = make() set_seed(seed) noise = torch.randn([1, num_output_frames, 16, 60, 104], device=torch.device("cuda"), dtype=torch.bfloat16) video, latents, timing = cached_inference( pipeline, ctrl, noise, [prompt], decode=bool(save_video_dir) and tag == "method") row[tag] = timing["denoise_dit_ms"] row[tag + "_summary"] = ctrl.summary() if tag == "method" and save_video_dir and i >= warmup: from einops import rearrange from torchvision.io import write_video os.makedirs(save_video_dir, exist_ok=True) frames = 255.0 * rearrange(video, "b t c h w -> b t h w c").cpu() write_video(os.path.join(save_video_dir, f"{i:04d}.mp4"), frames[0], fps=16) pipeline.vae.model.clear_cache() if i < warmup: continue ratios.append(row["base"] / row["method"]) base_ms.append(row["base"]) meth_ms.append(row["method"]) # Compute fraction of the very runs that were timed. Tuning a threshold on # one prompt subset and timing it on another silently disagrees whenever the # threshold sits near a decision boundary, so the two must come from the # same runs to be comparable. summ = row["method_summary"] meth_ce.append(summ["compute_equivalent_forwards"]) meth_fw.append(summ["denoise_forwards"]) # Contention is one-sided: a neighbouring job can only ever make a run slower. # The fastest observed run is therefore the best estimate of the uncontended # time, and min(baseline)/min(method) is the headline speedup. The paired # median is kept alongside, but it is biased *upward* here because the # baseline run is longer and so more exposed to being hit by a spike -- which # is why it sometimes exceeds the method's arithmetic ceiling. ce = statistics.fmean(meth_ce) fw = statistics.fmean(meth_fw) return { "speedup_from_minima": min(base_ms) / min(meth_ms), "paired_flops_speedup": fw / ce if ce else None, "paired_compute_equivalent": ce, "paired_denoise_forwards": fw, "paired_speedup_median": statistics.median(ratios), "paired_speedup_mean": statistics.fmean(ratios), "paired_speedup_stdev": statistics.stdev(ratios) if len(ratios) > 1 else 0.0, "paired_speedup_min": min(ratios), "paired_speedup_max": max(ratios), "min_baseline_ms": min(base_ms), "min_method_ms": min(meth_ms), "median_baseline_ms": statistics.median(base_ms), "median_method_ms": statistics.median(meth_ms), "num_pairs": len(ratios), "ratios": ratios, "baseline_ms": base_ms, "method_ms": meth_ms, } def evaluate(pipeline, controller, prompts, num_output_frames=21, seed=0, warmup=1, save_video_dir=None, save_latents_dir=None): """Run ``prompts`` and return timing / compute-fraction statistics.""" import torch from utils.misc import set_seed from cachelib import cached_inference records = [] for i, prompt in enumerate(prompts): set_seed(seed) noise = torch.randn([1, num_output_frames, 16, 60, 104], device=torch.device("cuda"), dtype=torch.bfloat16) video, latents, timing = cached_inference( pipeline, controller, noise, [prompt], decode=bool(save_video_dir)) rec = dict(timing) rec.update(controller.summary()) rec["prompt_index"] = i rec["warmup"] = i < warmup records.append(rec) if i >= warmup and save_latents_dir: os.makedirs(save_latents_dir, exist_ok=True) torch.save(latents.float().cpu(), os.path.join(save_latents_dir, f"{i:04d}.pt")) if i >= warmup and save_video_dir: from einops import rearrange from torchvision.io import write_video os.makedirs(save_video_dir, exist_ok=True) frames = 255.0 * rearrange(video, "b t c h w -> b t h w c").cpu() write_video(os.path.join(save_video_dir, f"{i:04d}.mp4"), frames[0], fps=16) pipeline.vae.model.clear_cache() timed = [r for r in records if not r["warmup"]] or records ms = sorted(r["denoise_dit_ms"] for r in timed) compute = statistics.fmean(r["compute_equivalent_forwards"] for r in timed) forwards = statistics.fmean(r["denoise_forwards"] for r in timed) return { "median_denoise_dit_ms": statistics.median(ms), "mean_denoise_dit_ms": statistics.fmean(ms), "stdev_denoise_dit_ms": statistics.stdev(ms) if len(ms) > 1 else 0.0, "min_denoise_dit_ms": ms[0], "max_denoise_dit_ms": ms[-1], "mean_compute_equivalent_forwards": compute, "mean_denoise_forwards": forwards, "flops_speedup_estimate": forwards / compute if compute else None, "mean_first_chunk_denoise_ms": statistics.fmean( r["first_chunk_denoise_ms"] for r in timed), "num_timed": len(timed), "records": records, } def denoising_steps_for(original, num_steps): """The ``num_steps`` subset of ``original`` that use_denoising_steps would install.""" import torch if num_steps is None or num_steps == len(original): return original n_full = len(original) idx = sorted({int(round(i * (n_full - 1) / max(num_steps - 1, 1))) for i in range(num_steps)}) if num_steps > 1 else [0] return original[torch.tensor(idx)] def use_denoising_steps(pipeline, num_steps): """Temporarily reduce the model's denoising schedule to ``num_steps``. These are step-distilled models: the generator was only ever trained at the timesteps in ``denoising_step_list`` (1000/750/500/250 warped), so a shorter schedule has to be a *subset* of them rather than a freshly spaced grid. The subset is uniform and keeps both ends, so the last prediction is still made at the finest trained timestep (n=3 -> 1000/500/250, n=2 -> 1000/250, n=1 -> 1000). Returns the original list; the caller restores it. """ import torch original = pipeline.denoising_step_list if num_steps is None or num_steps == len(original): return original n_full = len(original) idx = sorted({int(round(i * (n_full - 1) / max(num_steps - 1, 1))) for i in range(num_steps)}) if num_steps > 1 else [0] pipeline.denoising_step_list = original[torch.tensor(idx)] return original