GS-QA-Leaderboard / runner /docker /validate_predictions.py
Zoe's picture
Add public evaluation paths and Docker runner
015e4c5 verified
Raw
History Blame Contribute Delete
4.06 kB
#!/usr/bin/env python3
"""Validate a complete GS-QA2 prediction file before public submission."""
from __future__ import annotations
import argparse
import json
import re
from pathlib import Path
from typing import Any
from run_direct import EXPECTED_ROWS, load_track
SHA_PATTERN = re.compile(r"^[0-9a-fA-F]{40,64}$")
REQUIRED_FIELDS = {
"id",
"template_id",
"question",
"prediction",
"attempted",
"benchmark_revision",
"benchmark_track",
"model_id",
"model_revision",
"official_full_track_run",
}
def require(condition: bool, message: str) -> None:
if not condition:
raise ValueError(message)
def read_records(path: Path) -> list[dict[str, Any]]:
require(path.is_file(), f"Missing predictions file: {path}")
records: list[dict[str, Any]] = []
with path.open(encoding="utf-8") as handle:
for line_number, line in enumerate(handle, 1):
if not line.strip():
continue
try:
record = json.loads(line)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid JSON at line {line_number}") from error
require(isinstance(record, dict), f"Line {line_number} is not an object")
missing = REQUIRED_FIELDS - set(record)
require(not missing, f"Line {line_number} lacks {sorted(missing)}")
records.append(record)
require(bool(records), "Predictions file is empty")
return records
def one_value(records: list[dict[str, Any]], field: str) -> str:
values = {str(record[field]) for record in records}
require(len(values) == 1, f"{field} must be identical in every record")
return values.pop()
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--input",
type=Path,
default=Path("/workspace/results/predictions.jsonl"),
)
parser.add_argument("--track", choices=sorted(EXPECTED_ROWS), default="vector")
args = parser.parse_args()
records = read_records(args.input)
expected_count = EXPECTED_ROWS[args.track]
require(
len(records) == expected_count,
f"{args.track} track requires {expected_count:,} records; found {len(records):,}",
)
track = one_value(records, "benchmark_track")
benchmark_revision = one_value(records, "benchmark_revision")
model_id = one_value(records, "model_id")
model_revision = one_value(records, "model_revision")
require(track == args.track, f"File contains {track!r}, not {args.track!r}")
require(
bool(SHA_PATTERN.fullmatch(benchmark_revision)),
"benchmark_revision must be a full Git commit SHA",
)
require(
bool(SHA_PATTERN.fullmatch(model_revision)),
"model_revision must be a full model commit SHA",
)
require(
all(record["official_full_track_run"] is True for record in records),
"Every record must come from a full run without --limit",
)
ids = [str(record["id"]) for record in records]
require(len(set(ids)) == expected_count, "Prediction IDs must be unique")
benchmark = load_track(args.track, benchmark_revision)
require(
set(ids) == set(map(str, benchmark["id"])),
"Prediction IDs do not match the pinned benchmark",
)
attempted = sum(bool(record["attempted"]) for record in records)
print(
json.dumps(
{
"valid": True,
"track": args.track,
"questions": expected_count,
"attempted": attempted,
"failed_or_empty": expected_count - attempted,
"benchmark_revision": benchmark_revision,
"model_id": model_id,
"model_revision": model_revision,
"predictions": str(args.input),
},
indent=2,
)
)
if __name__ == "__main__":
try:
main()
except ValueError as error:
raise SystemExit(f"Invalid submission: {error}") from error