| |
| """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 |
|
|
| 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))) |
|
|
| |
| 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) |
|
|
| |
| 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 |
|
|
| |
| 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()) |
|
|