comparison / harness.py
Cccccz's picture
Add files using upload-large-folder tool
38a51ff verified
Raw History Blame Contribute Delete
10.7 kB
"""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