Download scripts/build_vbench8_extended_mapping.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 6.03 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/build_vbench8_extended_mapping.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/build_vbench8_extended_mapping.py
-
curl -L -o build_vbench8_extended_mapping.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/build_vbench8_extended_mapping.py
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() | |