File size: 3,761 Bytes
38a51ff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
#!/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)