Self-Forcing / scripts /build_vbench8_extended_mapping.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw History Blame Contribute Delete
6.03 kB
#!/usr/bin/env python3
"""Build the auditable VBench-8 extended-prompt subset mapping.
The standard VBench metadata is the source of truth for the 946-item order and
for the prompt-suite membership. Self-Forcing's short prompt file is checked
against that order before the selected indices are transferred to the
extended prompt file.
"""
from __future__ import annotations
import argparse
import copy
import hashlib
import json
from collections import Counter
from pathlib import Path
from typing import Any
REPO_ROOT = Path(__file__).resolve().parents[1]
SELECTED_SUITES = ("subject_consistency", "overall_consistency", "scene")
EXPECTED_COUNTS = {
"subject_consistency": 72,
"overall_consistency": 93,
"scene": 86,
}
def read_prompts(path: Path) -> list[str]:
lines = path.read_text(encoding="utf-8").splitlines()
prompts = [line.strip() for line in lines]
if any(not prompt for prompt in prompts):
raise ValueError(f"Prompt file contains an empty line: {path}")
return prompts
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for block in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def build_mapping(
*,
short_prompts: list[str],
extended_prompts: list[str],
vbench_info: list[dict[str, Any]],
) -> list[dict[str, Any]]:
if len(short_prompts) != 946:
raise ValueError(f"Expected 946 short prompts, got {len(short_prompts)}")
if len(extended_prompts) != 946:
raise ValueError(
f"Expected 946 extended prompts, got {len(extended_prompts)}"
)
if len(vbench_info) != 946:
raise ValueError(f"Expected 946 VBench metadata rows, got {len(vbench_info)}")
canonical = [row.get("prompt_en") for row in vbench_info]
if any(not isinstance(prompt, str) or not prompt.strip() for prompt in canonical):
raise ValueError("VBench metadata contains a missing prompt_en")
mismatches = [
index
for index, (short, official) in enumerate(zip(short_prompts, canonical))
if short != official
]
if mismatches:
preview = mismatches[:10]
raise ValueError(
"Self-Forcing all_dimension.txt does not preserve VBench ordering; "
f"mismatching indices include {preview}"
)
counters: Counter[str] = Counter()
mapping: list[dict[str, Any]] = []
for global_index, row in enumerate(vbench_info):
suites = [suite for suite in SELECTED_SUITES if suite in row["dimension"]]
if not suites:
continue
if len(suites) != 1:
raise ValueError(
f"Metadata row {global_index} belongs to multiple selected suites: {suites}"
)
suite = suites[0]
suite_index = counters[suite]
counters[suite] += 1
item: dict[str, Any] = {
"global_index": global_index,
"prompt_suite": suite,
"suite_index": suite_index,
"original_prompt": short_prompts[global_index],
"extended_prompt": extended_prompts[global_index],
"official_dimensions": list(row["dimension"]),
}
if "auxiliary_info" in row:
item["auxiliary_info"] = copy.deepcopy(row["auxiliary_info"])
mapping.append(item)
if counters != Counter(EXPECTED_COUNTS):
raise ValueError(
f"Unexpected selected-suite counts: {dict(counters)}; "
f"expected {EXPECTED_COUNTS}"
)
if len(mapping) != 251:
raise ValueError(f"Expected 251 selected prompts, got {len(mapping)}")
if len({item["global_index"] for item in mapping}) != len(mapping):
raise ValueError("Duplicate global indices in mapping")
for suite, expected in EXPECTED_COUNTS.items():
indices = [item["suite_index"] for item in mapping if item["prompt_suite"] == suite]
if sorted(indices) != list(range(expected)):
raise ValueError(f"Suite indices for {suite} are not contiguous and unique")
return mapping
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--short-prompts",
type=Path,
default=REPO_ROOT / "prompts/vbench/all_dimension.txt",
)
parser.add_argument(
"--extended-prompts",
type=Path,
default=REPO_ROOT / "prompts/vbench/all_dimension_extended.txt",
)
parser.add_argument(
"--vbench-info",
type=Path,
default=Path(
"/data3/chenzhuo/workspace/HY-WorldPlay-light-interaction-run-DEV/"
".venv-vbench/lib/python3.10/site-packages/vbench/VBench_full_info.json"
),
)
parser.add_argument(
"--output",
type=Path,
default=REPO_ROOT / "assets/vbench8_extended_subset_mapping.json",
)
return parser.parse_args()
def main() -> None:
args = parse_args()
short_path = args.short_prompts.expanduser().resolve()
extended_path = args.extended_prompts.expanduser().resolve()
info_path = args.vbench_info.expanduser().resolve()
output_path = args.output.expanduser().resolve()
mapping = build_mapping(
short_prompts=read_prompts(short_path),
extended_prompts=read_prompts(extended_path),
vbench_info=json.loads(info_path.read_text(encoding="utf-8")),
)
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(
json.dumps(mapping, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
)
print(f"mapping={output_path}")
print(f"counts={dict(Counter(item['prompt_suite'] for item in mapping))}")
print(f"total={len(mapping)}")
print(f"short_sha256={sha256(short_path)}")
print(f"extended_sha256={sha256(extended_path)}")
print(f"vbench_info_sha256={sha256(info_path)}")
print(f"mapping_sha256={sha256(output_path)}")
if __name__ == "__main__":
main()