#!/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)