Self-Forcing-part-2 / scripts /prepare_moviegen_validation_split.py
Cccccz's picture
Add files using upload-large-folder tool
2847d0b verified
Raw
History Blame Contribute Delete
3.08 kB
#!/usr/bin/env python3
import argparse
import hashlib
import json
import os
from pathlib import Path
def read_nonempty(path: Path) -> list[str]:
return [
line.strip()
for line in path.read_text(encoding="utf-8").splitlines()
if line.strip()
]
def atomic_write(path: Path, text: str) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_name(f".{path.name}.tmp.{os.getpid()}")
temporary.write_text(text, encoding="utf-8")
os.replace(temporary, path)
def main() -> None:
parser = argparse.ArgumentParser(
description="Create the fixed MovieGenVideoBench prompts 101-200 split."
)
parser.add_argument(
"--source",
type=Path,
default=Path("prompts/MovieGenVideoBench_extended.txt"),
)
parser.add_argument(
"--training-cases",
type=Path,
default=Path("data/self_forcing_predictor_v4_1000_seed0/cases.jsonl"),
)
parser.add_argument(
"--output",
type=Path,
default=Path("prompts/MovieGenVideoBench_extended_101_200.txt"),
)
parser.add_argument(
"--manifest",
type=Path,
default=Path("prompts/MovieGenVideoBench_extended_101_200.json"),
)
args = parser.parse_args()
source_prompts = read_nonempty(args.source)
if len(source_prompts) < 200:
raise ValueError(
f"{args.source} has only {len(source_prompts)} non-empty prompts"
)
selected = source_prompts[100:200]
if len(selected) != 100 or len(set(selected)) != 100:
raise ValueError("prompts 101-200 must contain 100 unique prompts")
training_rows = [
json.loads(line)
for line in args.training_cases.read_text(encoding="utf-8").splitlines()
if line.strip()
]
training_prompts = {row["prompt"] for row in training_rows}
overlap = sorted(set(selected).intersection(training_prompts))
if overlap:
raise ValueError(
f"validation split overlaps {len(overlap)} exact training prompts"
)
prompts_text = "".join(f"{prompt}\n" for prompt in selected)
prompt_sha256 = hashlib.sha256(prompts_text.encode("utf-8")).hexdigest()
manifest = {
"source": str(args.source.resolve()),
"source_nonempty_count": len(source_prompts),
"selection": {
"zero_based_slice": [100, 200],
"one_based_prompt_range": [101, 200],
"count": len(selected),
"unique_count": len(set(selected)),
},
"training_cases": str(args.training_cases.resolve()),
"training_prompt_count": len(training_rows),
"exact_training_overlap_count": 0,
"seed": 0,
"output": str(args.output.resolve()),
"output_sha256": prompt_sha256,
}
atomic_write(args.output, prompts_text)
atomic_write(
args.manifest,
json.dumps(manifest, ensure_ascii=False, indent=2) + "\n",
)
print(json.dumps(manifest, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()