File size: 7,304 Bytes
f70ac4f | 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 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 | #!/usr/bin/env python
"""Unit tests for the evaluation protocol's section 12 checklist.
python eval/test_protocol.py
"""
import json
import os
import sys
import tempfile
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, ROOT)
from eval import aggregate as agg # noqa: E402
MAPPING = os.path.join(ROOT, "assets/vbench8_extended_subset_mapping.json")
FULL_INFO = "/local/zoubin/cz/projects/VBench/vbench/VBench_full_info.json"
PROMPT_DIR = "/local/zoubin/cz/projects/Self-Forcing/prompts/vbench"
FAILURES = []
def check(name, cond, detail=""):
print(f" {'PASS' if cond else 'FAIL'} {name}" + (f" {detail}" if detail else ""))
if not cond:
FAILURES.append(name)
def test_mapping():
print("946 -> 251 mapping")
with open(MAPPING) as f:
m = json.load(f)
rows = m["rows"]
check("mapping has 251 rows", len(rows) == 251, f"got {len(rows)}")
counts = {s: sum(1 for r in rows if r["prompt_suite"] == s)
for s in ("subject_consistency", "overall_consistency", "scene")}
check("suite counts 72/93/86",
counts == {"subject_consistency": 72, "overall_consistency": 93, "scene": 86},
str(counts))
check("global_index unique",
len({r["global_index"] for r in rows}) == len(rows))
for s in counts:
idx = [r["suite_index"] for r in rows if r["prompt_suite"] == s]
check(f"{s} suite_index contiguous", idx == list(range(len(idx))))
check("every scene row keeps auxiliary_info",
all("auxiliary_info" in r for r in rows if r["prompt_suite"] == "scene"))
check("extended prompts non-empty",
all(r["extended_prompt"].strip() for r in rows))
def test_canonical_order():
print("canonical prompt order is enforced")
with open(FULL_INFO) as f:
info = json.load(f)
with open(os.path.join(PROMPT_DIR, "all_dimension.txt"), encoding="utf-8") as f:
short = [l.rstrip("\n") for l in f]
check("946 short prompts", len(short) == 946, f"got {len(short)}")
check("short prompt order == VBench_full_info",
all(a["prompt_en"].strip() == b.strip() for a, b in zip(info, short)))
# A mismatched prompt file must be refused, not silently fuzzy-matched.
with open(MAPPING) as f:
rows = json.load(f)["rows"]
with open(os.path.join(PROMPT_DIR, "all_dimension_extended.txt"), encoding="utf-8") as f:
ext = [l.rstrip("\n") for l in f]
check("mapping extended prompt matches its global_index",
all(r["extended_prompt"] == ext[r["global_index"]] for r in rows))
check("mapping original prompt matches its global_index",
all(r["original_prompt"] == short[r["global_index"]] for r in rows))
def test_normalize_and_aggregate():
print("normalize + Quality/Semantic/Selected aggregation")
lo_raw = {d: agg.NORMALIZE_RANGE[d][0] for d in agg.DIMENSIONS}
hi_raw = {d: agg.NORMALIZE_RANGE[d][1] for d in agg.DIMENSIONS}
check("normalize maps min -> 0",
all(abs(v) < 1e-12 for v in agg.normalize(lo_raw).values()))
check("normalize maps max -> 1",
all(abs(v - 1.0) < 1e-12 for v in agg.normalize(hi_raw).values()))
n = agg.normalize(hi_raw)
q, s = agg.quality_score(n), agg.semantic_score(n)
check("quality of all-ones == 1", abs(q - 1.0) < 1e-12, f"got {q}")
check("semantic of all-ones == 1", abs(s - 1.0) < 1e-12, f"got {s}")
check("selected of all-ones == 1", abs(agg.selected_score(q, s) - 1.0) < 1e-12)
# dynamic_degree carries weight 0.5, everything else in Quality carries 1.0.
n0 = {d: 0.0 for d in agg.DIMENSIONS}
n_dd = dict(n0, dynamic_degree=1.0)
check("dynamic_degree weight is 0.5/5.5",
abs(agg.quality_score(n_dd) - 0.5 / 5.5) < 1e-12)
n_sc = dict(n0, subject_consistency=1.0)
check("other quality dims weight 1/5.5",
abs(agg.quality_score(n_sc) - 1.0 / 5.5) < 1e-12)
check("selected weights Quality:Semantic 4:1",
abs(agg.selected_score(1.0, 0.0) - 0.8) < 1e-12
and abs(agg.selected_score(0.0, 1.0) - 0.2) < 1e-12)
check("selected is not the plain mean of 8 normalized dims",
abs(agg.selected_score(agg.quality_score(n_dd), agg.semantic_score(n_dd))
- sum(n_dd.values()) / 8.0) > 1e-6)
def _rec(strategy, idx, lat, ctx, mse, frames=81):
return {
"status": "complete", "strategy": strategy, "base_model": "self_forcing",
"method": "none" if strategy.endswith("ffff") else "teacache",
"target_speedup": 1.0, "global_index": idx, "prompt_suite": "scene",
"suite_index": idx, "policy_latency_ms": lat,
"excluded_context_kv_latency_ms": ctx,
"pixel_metrics_vs_ffff": {"mean_mse": mse, "psnr": 0.0, "ssim": 1.0,
"lpips": 0.0, "num_frames": frames},
"cache_diagnostics": {"compute_equivalent_forwards": 28.0},
}
def test_latency_and_completeness():
print("latency accounting and the completeness gate")
import math
# ratio of means, not the mean of per-prompt ratios
ffff = [_rec("sf_ffff", i, lat, 800.0, 1e-12) for i, lat in enumerate([100.0, 300.0])]
cand = [_rec("sf_x", i, lat, 800.0, 1e-4) for i, lat in enumerate([50.0, 250.0])]
ratio_of_means = 100.0 * (1 - (sum(r["policy_latency_ms"] for r in cand) / 2)
/ (sum(r["policy_latency_ms"] for r in ffff) / 2))
mean_of_ratios = 100.0 * (1 - ((50.0 / 100.0) + (250.0 / 300.0)) / 2)
check("speedup uses ratio of means", abs(ratio_of_means - 25.0) < 1e-9,
f"{ratio_of_means:.3f}%")
check("ratio of means differs from mean of ratios",
abs(ratio_of_means - mean_of_ratios) > 1.0)
check("context/KV DiT is a separate field, never inside policy latency",
all("excluded_context_kv_latency_ms" in r
and r["policy_latency_ms"] != r["policy_latency_ms"]
+ r["excluded_context_kv_latency_ms"] for r in cand))
mse = sum(r["pixel_metrics_vs_ffff"]["mean_mse"] for r in cand) / len(cand)
check("PSNR aggregates MSE first, then converts once",
abs(-10.0 * math.log10(max(mse, 1e-12)) - 40.0) < 1e-9)
check("pixel metrics span all 81 frames",
{r["pixel_metrics_vs_ffff"]["num_frames"] for r in cand} == {81})
with tempfile.TemporaryDirectory() as td:
for name, recs in (("sf_ffff", ffff), ("sf_x", cand)):
d = os.path.join(td, "per_prompt", name)
os.makedirs(d)
for r in recs:
with open(os.path.join(d, f"scene_{r['suite_index']:03d}.json"), "w") as f:
json.dump(r, f)
argv = sys.argv
sys.argv = ["aggregate", "--out-root", td, "--expect", "251"]
try:
rc = agg.main()
finally:
sys.argv = argv
check("aggregation refuses a strategy that is not 251", rc == 1, f"rc={rc}")
def main():
for fn in (test_mapping, test_canonical_order, test_normalize_and_aggregate,
test_latency_and_completeness):
fn()
print()
if FAILURES:
print(f"{len(FAILURES)} FAILED: {FAILURES}")
return 1
print("all protocol unit tests passed")
return 0
if __name__ == "__main__":
sys.exit(main())
|