sskj-h3/vbench-base/scripts/h3_vbench_generate.py

413 lines
16 KiB
Python
Executable File

#!/usr/bin/env python3
"""Generate a deterministic VBench-v1 suite through four MiniMax-H3 services.
The program deliberately separates generation from VBench scoring. It creates a
flat ``{prompt}-{sample_index}.mp4`` directory accepted by VBench standard mode,
records a resumable JSONL audit trail, and balances multi-label VBench dimensions
across replicas with a deterministic greedy assignment.
"""
from __future__ import annotations
import argparse
import json
import os
import shutil
import statistics
import threading
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import requests
ALL_DIMENSIONS = [
"subject_consistency",
"background_consistency",
"temporal_flickering",
"motion_smoothness",
"dynamic_degree",
"aesthetic_quality",
"imaging_quality",
"object_class",
"multiple_objects",
"human_action",
"color",
"spatial_relationship",
"scene",
"temporal_style",
"appearance_style",
"overall_consistency",
]
@dataclass(frozen=True)
class Job:
request_id: str
prompt: str
dimensions: tuple[str, ...]
sample_index: int
seed: int
replica_index: int
port: int
@property
def filename(self) -> str:
return f"{self.prompt}-{self.sample_index}.mp4"
def atomic_copy(source: Path, destination: Path) -> None:
destination.parent.mkdir(parents=True, exist_ok=True)
temporary = destination.with_name(destination.name + ".partial")
shutil.copy2(source, temporary)
os.replace(temporary, destination)
def load_prompts(metadata_path: Path, dimensions: set[str]) -> list[dict[str, Any]]:
raw = json.loads(metadata_path.read_text(encoding="utf-8"))
merged: dict[str, set[str]] = {}
order: list[str] = []
for item in raw:
prompt = str(item["prompt_en"])
item_dimensions = set(item.get("dimension", [])) & dimensions
if not item_dimensions:
continue
if prompt not in merged:
merged[prompt] = set()
order.append(prompt)
merged[prompt].update(item_dimensions)
return [{"prompt": prompt, "dimensions": sorted(merged[prompt])} for prompt in order]
def coverage_subset(prompts: list[dict[str, Any]], max_prompts: int) -> list[dict[str, Any]]:
if max_prompts <= 0 or max_prompts >= len(prompts):
return prompts
remaining = list(prompts)
selected: list[dict[str, Any]] = []
counts = {dimension: 0 for dimension in ALL_DIMENSIONS}
while remaining and len(selected) < max_prompts:
# Prefer prompts covering the least represented dimensions; stable order
# is the final tie-breaker, so repeated runs produce the same suite.
best_index = min(
range(len(remaining)),
key=lambda index: (
sum(counts.get(dim, 0) for dim in remaining[index]["dimensions"]),
-len(remaining[index]["dimensions"]),
index,
),
)
item = remaining.pop(best_index)
selected.append(item)
for dimension in item["dimensions"]:
counts[dimension] += 1
return selected
def build_jobs(args: argparse.Namespace) -> list[Job]:
dimensions = set(args.dimensions.split(","))
unknown = dimensions - set(ALL_DIMENSIONS)
if unknown:
raise SystemExit(f"unknown VBench dimensions: {sorted(unknown)}")
prompts = coverage_subset(load_prompts(args.metadata, dimensions), args.max_prompts)
ports = [int(value) for value in args.ports.split(",") if value]
if not ports:
raise SystemExit("at least one port is required")
total_counts = [0] * len(ports)
dimension_counts = [{dimension: 0 for dimension in dimensions} for _ in ports]
jobs: list[Job] = []
for prompt_index, item in enumerate(prompts):
for sample_index in range(args.samples_per_prompt):
item_dimensions = tuple(item["dimensions"])
rotation = (prompt_index + sample_index) % len(ports)
candidates = list(range(len(ports)))
candidates.sort(
key=lambda replica: (
sum(dimension_counts[replica][dim] for dim in item_dimensions),
total_counts[replica],
(replica - rotation) % len(ports),
)
)
replica = candidates[0]
for dimension in item_dimensions:
dimension_counts[replica][dimension] += 1
total_counts[replica] += 1
jobs.append(
Job(
request_id=f"p{prompt_index:04d}-s{sample_index}",
prompt=item["prompt"],
dimensions=item_dimensions,
sample_index=sample_index,
seed=args.seed + prompt_index * args.samples_per_prompt + sample_index,
replica_index=replica,
port=ports[replica],
)
)
return jobs
def make_payload(args: argparse.Namespace, job: Job, steps: int) -> dict[str, Any]:
return {
"model": args.model,
"prompt": job.prompt,
"num_outputs_per_prompt": 1,
"num_inference_steps": steps,
"flow_shift": args.flow_shift,
"audio_flow_shift": args.audio_flow_shift,
"seed": job.seed,
"task": "t2va",
"conditions": [],
"target": {
"short_edge": args.short_edge,
"aspect_ratio": args.aspect_ratio,
"duration_seconds": args.duration_seconds,
},
}
def submit_and_wait(
session: requests.Session, args: argparse.Namespace, job: Job, steps: int
) -> tuple[dict[str, Any], dict[str, Any]]:
response = session.post(
f"http://{args.host}:{job.port}/v1/videos",
json=make_payload(args, job, steps),
timeout=args.submit_timeout,
)
if response.status_code != 200:
raise RuntimeError(f"submit HTTP {response.status_code}: {response.text[:1000]}")
status = response.json()
video_id = status.get("id")
if not video_id:
raise RuntimeError(f"submit response has no id: {status}")
deadline = time.monotonic() + args.request_timeout
while status.get("status") not in {"completed", "failed"}:
if time.monotonic() >= deadline:
raise TimeoutError(f"video job {video_id} exceeded {args.request_timeout}s")
time.sleep(args.poll_interval)
response = session.get(
f"http://{args.host}:{job.port}/v1/videos/{video_id}",
timeout=args.poll_timeout,
)
if response.status_code != 200:
raise RuntimeError(f"poll HTTP {response.status_code}: {response.text[:1000]}")
status = response.json()
if status.get("status") != "completed":
raise RuntimeError(f"job failed: {status.get('error') or status}")
return status, {"video_id": video_id}
def run_job(session: requests.Session, args: argparse.Namespace, job: Job) -> dict[str, Any]:
started_epoch = time.time()
started = time.monotonic()
result: dict[str, Any] = {
**job.__dict__,
"dimensions": list(job.dimensions),
"filename": job.filename,
"started_at_epoch": started_epoch,
"success": False,
"error": None,
}
try:
status, extra = submit_and_wait(session, args, job, args.num_inference_steps)
source_value = status.get("file_path")
if not source_value:
raise RuntimeError(f"completed response has no file_path: {status}")
source = Path(source_value)
if not source.is_file():
raise FileNotFoundError(f"server output does not exist: {source}")
destination = args.videos_dir / job.filename
atomic_copy(source, destination)
result.update(extra)
result.update(
{
"success": True,
"server_file_path": str(source),
"saved_file_path": str(destination),
"inference_time_s": status.get("inference_time_s"),
"peak_memory_mb": status.get("peak_memory_mb"),
"file_size_bytes": destination.stat().st_size,
}
)
except Exception as exc: # Preserve the rest of a long run and make it resumable.
result["error"] = f"{type(exc).__name__}: {exc}"
result["latency_s"] = time.monotonic() - started
result["finished_at_epoch"] = time.time()
return result
def completed_request_ids(results_path: Path, videos_dir: Path) -> set[str]:
completed: set[str] = set()
if not results_path.is_file():
return completed
for line in results_path.read_text(encoding="utf-8").splitlines():
try:
row = json.loads(line)
except json.JSONDecodeError:
continue
video = videos_dir / str(row.get("filename", ""))
if row.get("success") and video.is_file() and video.stat().st_size > 0:
completed.add(str(row["request_id"]))
return completed
def write_manifest(jobs: list[Job], path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as handle:
for job in jobs:
row = {**job.__dict__, "dimensions": list(job.dimensions), "filename": job.filename}
handle.write(json.dumps(row, ensure_ascii=False) + "\n")
def run_worker(
replica: int,
jobs: list[Job],
args: argparse.Namespace,
output_handle: Any,
output_lock: threading.Lock,
) -> int:
failures = 0
if not jobs:
return failures
with requests.Session() as session:
if args.warmup_requests:
warmup = jobs[0]
for index in range(args.warmup_requests):
started = time.monotonic()
try:
submit_and_wait(session, args, warmup, args.warmup_inference_steps)
print(
f"replica={replica} warmup={index + 1}/{args.warmup_requests} "
f"latency={time.monotonic() - started:.2f}s success=true",
flush=True,
)
except Exception as exc:
raise RuntimeError(f"replica {replica} warmup failed: {exc}") from exc
for index, job in enumerate(jobs, start=1):
result = run_job(session, args, job)
with output_lock:
output_handle.write(json.dumps(result, ensure_ascii=False) + "\n")
output_handle.flush()
failures += int(not result["success"])
print(
f"replica={replica} request={index}/{len(jobs)} id={job.request_id} "
f"success={str(result['success']).lower()} latency={result['latency_s']:.2f}s "
f"error={result['error']}",
flush=True,
)
return failures
def summarize(jobs: list[Job], results_path: Path, output_path: Path) -> dict[str, Any]:
latest: dict[str, dict[str, Any]] = {}
if results_path.is_file():
for line in results_path.read_text(encoding="utf-8").splitlines():
try:
row = json.loads(line)
except json.JSONDecodeError:
continue
latest[str(row.get("request_id"))] = row
rows = [latest[job.request_id] for job in jobs if job.request_id in latest]
success = [row for row in rows if row.get("success")]
latencies = [float(row["latency_s"]) for row in success]
by_replica: dict[str, Any] = {}
for replica in sorted({job.replica_index for job in jobs}):
replica_jobs = [job for job in jobs if job.replica_index == replica]
replica_rows = [row for row in rows if int(row["replica_index"]) == replica]
by_dimension = {
dimension: sum(dimension in job.dimensions for job in replica_jobs)
for dimension in ALL_DIMENSIONS
if any(dimension in job.dimensions for job in replica_jobs)
}
by_replica[str(replica)] = {
"expected": len(replica_jobs),
"recorded": len(replica_rows),
"completed": sum(bool(row.get("success")) for row in replica_rows),
"dimension_jobs": by_dimension,
}
summary = {
"expected": len(jobs),
"recorded_latest": len(rows),
"completed": len(success),
"failed_or_missing": len(jobs) - len(success),
"latency_mean_s": statistics.fmean(latencies) if latencies else None,
"latency_median_s": statistics.median(latencies) if latencies else None,
"by_replica": by_replica,
}
output_path.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8")
return summary
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--metadata", type=Path, required=True)
parser.add_argument("--videos-dir", type=Path, required=True)
parser.add_argument("--results", type=Path, required=True)
parser.add_argument("--manifest", type=Path, required=True)
parser.add_argument("--summary", type=Path, required=True)
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--ports", default="30010,30020,30030,30040")
parser.add_argument("--model", default="/data/hf_models/MiniMax-H3")
parser.add_argument("--dimensions", default=",".join(ALL_DIMENSIONS))
parser.add_argument("--samples-per-prompt", type=int, default=1)
parser.add_argument("--max-prompts", type=int, default=0)
parser.add_argument("--num-inference-steps", type=int, default=20)
parser.add_argument("--short-edge", type=int, default=768)
parser.add_argument("--duration-seconds", type=float, default=5.0)
parser.add_argument("--aspect-ratio", default="16:9")
parser.add_argument("--flow-shift", type=float, default=12.0)
parser.add_argument("--audio-flow-shift", type=float, default=3.0)
parser.add_argument("--seed", type=int, default=1101)
parser.add_argument("--warmup-requests", type=int, default=1)
parser.add_argument("--warmup-inference-steps", type=int, default=5)
parser.add_argument("--submit-timeout", type=float, default=120.0)
parser.add_argument("--poll-timeout", type=float, default=30.0)
parser.add_argument("--poll-interval", type=float, default=1.0)
parser.add_argument("--request-timeout", type=float, default=3600.0)
return parser.parse_args()
def main() -> int:
args = parse_args()
if not args.metadata.is_file():
raise SystemExit(f"VBench metadata missing: {args.metadata}")
if args.samples_per_prompt < 1 or args.samples_per_prompt > 5:
raise SystemExit("samples-per-prompt must be between 1 and 5")
args.videos_dir.mkdir(parents=True, exist_ok=True)
args.results.parent.mkdir(parents=True, exist_ok=True)
jobs = build_jobs(args)
write_manifest(jobs, args.manifest)
completed = completed_request_ids(args.results, args.videos_dir)
pending = [job for job in jobs if job.request_id not in completed]
print(
f"suite prompts={len({job.prompt for job in jobs})} jobs={len(jobs)} "
f"completed={len(completed)} pending={len(pending)} replicas={len(set(job.port for job in jobs))}",
flush=True,
)
grouped: dict[int, list[Job]] = {}
for job in pending:
grouped.setdefault(job.replica_index, []).append(job)
failures = 0
lock = threading.Lock()
with args.results.open("a", encoding="utf-8", buffering=1) as output:
with ThreadPoolExecutor(max_workers=max(1, len(grouped))) as executor:
futures = {
executor.submit(run_worker, replica, shard, args, output, lock): replica
for replica, shard in grouped.items()
}
for future in as_completed(futures):
try:
failures += future.result()
except Exception as exc:
print(f"replica={futures[future]} fatal={type(exc).__name__}: {exc}", flush=True)
failures += len(grouped[futures[future]])
summary = summarize(jobs, args.results, args.summary)
print(json.dumps(summary, ensure_ascii=False, indent=2), flush=True)
return int(failures > 0 or summary["failed_or_missing"] > 0)
if __name__ == "__main__":
raise SystemExit(main())