comparison / eval /build_mapping.py
Cccccz's picture
Add files using upload-large-folder tool
f70ac4f verified
Raw History Blame Contribute Delete
5.11 kB
#!/usr/bin/env python
"""Build and validate assets/vbench8_extended_subset_mapping.json (Extended-251).
Protocol section 4: take the subject_consistency (72), overall_consistency (93)
and scene (86) suites out of the canonical 946-prompt VBench order, and carry the
Self-Forcing rewritten extended prompt for each.
Selection is strictly by index into VBench_full_info.json -- the short prompt list
contains two duplicate texts, so text matching is not safe, and the protocol
forbids it anyway.
python eval/build_mapping.py --out assets/vbench8_extended_subset_mapping.json
"""
import argparse
import hashlib
import json
import os
import sys
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
SUITES = ["subject_consistency", "overall_consistency", "scene"]
EXPECTED = {"subject_consistency": 72, "overall_consistency": 93, "scene": 86}
DEFAULT_FULL_INFO = "/local/zoubin/cz/projects/VBench/vbench/VBench_full_info.json"
DEFAULT_PROMPT_DIR = "/local/zoubin/cz/projects/Self-Forcing/prompts/vbench"
def sha256(path):
h = hashlib.sha256()
with open(path, "rb") as f:
for chunk in iter(lambda: f.read(1 << 20), b""):
h.update(chunk)
return h.hexdigest()
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--full-info", default=DEFAULT_FULL_INFO)
ap.add_argument("--prompt-dir", default=DEFAULT_PROMPT_DIR)
ap.add_argument("--out", default=os.path.join(ROOT, "assets/vbench8_extended_subset_mapping.json"))
args = ap.parse_args()
short_path = os.path.join(args.prompt_dir, "all_dimension.txt")
ext_path = os.path.join(args.prompt_dir, "all_dimension_extended.txt")
with open(args.full_info) as f:
info = json.load(f)
with open(short_path, encoding="utf-8") as f:
short = [l.rstrip("\n") for l in f]
with open(ext_path, encoding="utf-8") as f:
ext = [l.rstrip("\n") for l in f]
errors = []
if len(short) != 946:
errors.append(f"short prompts == {len(short)}, expected 946")
if len(ext) != 946:
errors.append(f"extended prompts == {len(ext)}, expected 946")
if len(info) != 946:
errors.append(f"VBench_full_info rows == {len(info)}, expected 946")
for i, (row, s) in enumerate(zip(info, short)):
if row["prompt_en"].strip() != s.strip():
errors.append(f"short prompt order differs from VBench_full_info at line {i}")
break
for i, e in enumerate(ext):
if not e.strip():
errors.append(f"empty extended prompt at line {i}")
break
rows = []
for suite in SUITES:
suite_index = 0
for gi, row in enumerate(info):
if suite not in row["dimension"]:
continue
entry = {
"global_index": gi,
"prompt_suite": suite,
"suite_index": suite_index,
"original_prompt": short[gi],
"extended_prompt": ext[gi],
"official_dimensions": list(row["dimension"]),
}
# Standard VBench 0.1.5 needs the official scene keywords to score `scene`.
for key in ("auxiliary_info",):
if key in row:
entry[key] = row[key]
rows.append(entry)
suite_index += 1
counts = {s: sum(1 for r in rows if r["prompt_suite"] == s) for s in SUITES}
for s, n in EXPECTED.items():
if counts.get(s) != n:
errors.append(f"{s} == {counts.get(s)}, expected {n}")
if len(rows) != 251:
errors.append(f"mapping rows == {len(rows)}, expected 251")
if len({r["global_index"] for r in rows}) != len(rows):
errors.append("duplicate global_index")
for s in SUITES:
idx = [r["suite_index"] for r in rows if r["prompt_suite"] == s]
if idx != list(range(len(idx))):
errors.append(f"{s} suite_index not contiguous")
if errors:
print("MAPPING VALIDATION FAILED:")
for e in errors:
print(" -", e)
return 1
payload = {
"protocol": "Self-Forcing Extended-251 Full Evaluation",
"num_rows": len(rows),
"counts": counts,
"sources": {
"vbench_full_info": {"path": args.full_info, "sha256": sha256(args.full_info)},
"short_prompts": {"path": short_path, "sha256": sha256(short_path)},
"extended_prompts": {"path": ext_path, "sha256": sha256(ext_path)},
},
"rows": rows,
}
os.makedirs(os.path.dirname(args.out), exist_ok=True)
with open(args.out, "w") as f:
json.dump(payload, f, indent=2)
print("mapping OK")
print(f" rows {len(rows)}")
for s in SUITES:
print(f" {s:22s} {counts[s]}")
for k, v in payload["sources"].items():
print(f" sha256 {k:18s} {v['sha256'][:16]}...")
print(f" scene rows with auxiliary_info: "
f"{sum(1 for r in rows if r['prompt_suite'] == 'scene' and 'auxiliary_info' in r)}")
print(f"wrote {args.out}")
return 0
if __name__ == "__main__":
sys.exit(main())