CoolFace
Modelpublic

Cccccz/HY

sourceHugging Faceupdated 11d agoView on Hugging Face
0likes
validate_predictor_dataset.py53 linesDownload Raw Back to tools
1#!/usr/bin/env python32"""Validate Predictor safetensors files without loading the Teacher."""3 4from __future__ import annotations5 6import argparse7import json8from pathlib import Path9 10from safetensors.torch import load_file11 12from predictor_data.schema import SCHEMA_VERSION, validate_case_tensors, validate_chunk_tensors13 14 15def main() -> None:16    parser = argparse.ArgumentParser(description=__doc__)17    parser.add_argument("--root", default="datasets/predictor_v1")18    parser.add_argument("--expected_cases", type=int, default=None)19    parser.add_argument("--expected_chunks_per_case", type=int, default=None)20    args = parser.parse_args()21 22    root = Path(args.root).resolve()23    manifest = root / "manifest.jsonl"24    if not manifest.is_file():25        raise FileNotFoundError(manifest)26    records = [json.loads(line) for line in manifest.read_text(encoding="utf-8").splitlines() if line]27    if len({(r["case_id"], r["seed"], r["chunk_id"]) for r in records}) != len(records):28        raise ValueError("Duplicate manifest keys")29 30    seen_cases = set()31    per_case: dict[int, int] = {}32    for record in records:33        if record["schema_version"] != SCHEMA_VERSION:34            raise ValueError(f"Unexpected schema: {record['schema_version']}")35        case_id = int(record["case_id"])36        if case_id not in seen_cases:37            validate_case_tensors(load_file(str(root / record["case_tensor_file"]), device="cpu"))38            seen_cases.add(case_id)39        validate_chunk_tensors(load_file(str(root / record["tensor_file"]), device="cpu"))40        per_case[case_id] = per_case.get(case_id, 0) + 141 42    if args.expected_cases is not None and len(seen_cases) != args.expected_cases:43        raise ValueError(f"Expected {args.expected_cases} cases, found {len(seen_cases)}")44    if args.expected_chunks_per_case is not None:45        bad = {case: count for case, count in per_case.items() if count != args.expected_chunks_per_case}46        if bad:47            raise ValueError(f"Unexpected chunks per case: {bad}")48    print(json.dumps({"records": len(records), "cases": len(seen_cases), "chunks_per_case": per_case}, indent=2))49 50 51if __name__ == "__main__":52    main()53