comparison / activity_pass.py
Cccccz's picture
Add files using upload-large-folder tool
38a51ff verified
Raw History Blame Contribute Delete
3.76 kB
#!/usr/bin/env python
"""Backfill token-activity statistics into per-prompt records generated before the
controller recorded them.
Re-runs only the denoising loop (no VAE decode, no video write) for the given
strategies, which is deterministic, and merges ``active_step_ratio`` /
``empty_step_ratio`` / ``full_step_ratio`` / ``mean_selected_fraction_active``
into each record's ``cache_diagnostics``.
CUDA_VISIBLE_DEVICES=4 python activity_pass.py --base self_forcing \
--strategies sf_motioncache_x1.75,sf_motioncache_Fxxx_x2.85 --shard 0 --num-shards 4
"""
import argparse, json, os, sys, time
ROOT = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, ROOT)
from eval.strategies import load_strategies # noqa: E402
from harness import BASES, enter_base, load_pipeline # noqa: E402
MAPPING = os.path.join(ROOT, "assets/vbench8_extended_subset_mapping.json")
ap = argparse.ArgumentParser()
ap.add_argument("--base", choices=sorted(BASES), required=True)
ap.add_argument("--strategies", required=True)
ap.add_argument("--shard", type=int, default=0)
ap.add_argument("--num-shards", type=int, default=1)
ap.add_argument("--out-root", default="eval_out")
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--num-latent-frames", type=int, default=21)
args = ap.parse_args()
out_root = args.out_root if os.path.isabs(args.out_root) else os.path.join(ROOT, args.out_root)
rows = json.load(open(MAPPING))["rows"][args.shard::args.num_shards]
wanted = set(args.strategies.split(","))
strategies = [s for s in load_strategies(args.base) if s["name"] in wanted]
if not strategies:
sys.exit(f"no strategies matched {sorted(wanted)}")
print(f"shard {args.shard}/{args.num_shards}: {len(rows)} prompts x {len(strategies)} strategies", flush=True)
base = enter_base(args.base)
import torch
from utils.misc import set_seed
from cachelib import CacheController, build_method, cached_inference, install, parse_schedule
torch.set_grad_enabled(False)
pipe = load_pipeline(base)
num_steps = len(pipe.denoising_step_list)
t0 = time.time()
for n, row in enumerate(rows):
for s in strategies:
rp = os.path.join(out_root, "per_prompt", s["name"],
f"{row['prompt_suite']}_{row['suite_index']:03d}.json")
if not os.path.exists(rp):
continue
rec = json.load(open(rp))
if "active_step_ratio" in rec.get("cache_diagnostics", {}):
continue
kw = dict(indicator=s["indicator"], coefficients=s["coefficients"])
if s["method"] != "none":
kw[s["param"]] = s["value"]
ctrl = CacheController(build_method(s["method"], **kw), num_steps=num_steps,
forced_steps=parse_schedule(s.get("schedule", "FxxF"), num_steps))
install(pipe.generator.model, ctrl)
set_seed(args.seed)
noise = torch.randn([1, args.num_latent_frames, 16, 60, 104],
device=torch.device("cuda"), dtype=torch.bfloat16)
cached_inference(pipe, ctrl, noise, [rec["prompt"]], decode=False)
summary = ctrl.summary()
assert abs(summary["compute_equivalent_forwards"]
- rec["cache_diagnostics"]["compute_equivalent_forwards"]) < 1e-6, \
f"compute mismatch on {rp}"
rec["cache_diagnostics"] = dict(summary)
tmp = rp + ".tmp"
json.dump(rec, open(tmp, "w"), indent=2)
os.replace(tmp, rp)
if (n + 1) % 20 == 0:
rate = (time.time() - t0) / (n + 1)
print(f"[{args.base} activity shard {args.shard}] {n+1}/{len(rows)} "
f"{rate:.1f}s/prompt eta {(len(rows)-n-1)*rate/60:.0f} min", flush=True)
print(f"shard {args.shard} done in {(time.time()-t0)/60:.1f} min", flush=True)