Download activity_pass.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 3.76 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/activity_pass.py
- Command line
-
hf download hf://Cccccz/comparison/activity_pass.py
-
curl -L -o activity_pass.py https://huggingface.co/Cccccz/comparison/resolve/main/activity_pass.py
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) | |