#!/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()