File size: 5,003 Bytes
fbd9366 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 | #!/usr/bin/env python3
"""Fast startup check for a previously validated T-Rex dataset variant."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Sequence
import numpy as np
import pyarrow.parquet as pq
TRACK_CACHE_NAME = "tracks_trex_track_force_v2"
FORCE_SCHEMA_METADATA_KEY = b"trex_track_force_schema_version"
def _read_json(path: Path) -> dict:
if not path.is_file():
raise FileNotFoundError(path)
return json.loads(path.read_text())
def _jsonl_count(path: Path) -> int:
if not path.is_file():
raise FileNotFoundError(path)
return sum(bool(line.strip()) for line in path.read_text().splitlines())
def _episode_path(root: Path, episode_index: int) -> Path:
return (
root
/ "data"
/ f"chunk-{episode_index // 1000:03d}"
/ f"episode_{episode_index:06d}.parquet"
)
def _track_path(root: Path, episode_index: int) -> Path:
return root / TRACK_CACHE_NAME / f"episode_{episode_index:06d}.npz"
def check_dataset(root: Path, *, require_force: bool) -> dict:
root = root.expanduser().resolve()
ready = _read_json(root / "meta" / "dataset_ready.json")
info = _read_json(root / "meta" / "info.json")
episodes = int(info["total_episodes"])
frames = int(info["total_frames"])
tasks = int(info["total_tasks"])
videos = int(info["total_videos"])
if episodes <= 0 or frames <= 0:
raise ValueError(f"{root}: empty dataset")
if int(ready.get("episodes", -1)) != episodes:
raise ValueError(f"{root}: stale dataset_ready episode count")
if int(ready.get("frames", -1)) != frames:
raise ValueError(f"{root}: stale dataset_ready frame count")
if bool(ready.get("force")) != require_force:
raise ValueError(
f"{root}: force={ready.get('force')} but require_force={require_force}"
)
if _jsonl_count(root / "meta" / "episodes.jsonl") != episodes:
raise ValueError(f"{root}: episodes.jsonl count mismatch")
if _jsonl_count(root / "meta" / "tasks.jsonl") != tasks:
raise ValueError(f"{root}: tasks.jsonl count mismatch")
for required in (
root / "meta" / "stats.json",
root / "meta" / "relative_stats_dreamzero.json",
root / "meta" / "source_episode_index_map.json",
):
if not required.is_file():
raise FileNotFoundError(required)
manifest = None
if require_force:
manifest = _read_json(
root / "meta" / "trex_track_force_manifest.json"
)
if len(manifest.get("episodes", {})) != episodes:
raise ValueError(f"{root}: force manifest count mismatch")
if not (root / TRACK_CACHE_NAME).is_dir():
raise FileNotFoundError(root / TRACK_CACHE_NAME)
sample_indices = sorted({0, episodes // 2, episodes - 1})
for episode_index in sample_indices:
parquet_path = _episode_path(root, episode_index)
parquet_file = pq.ParquetFile(parquet_path)
if int(parquet_file.metadata.num_rows) <= 0:
raise ValueError(f"{parquet_path}: empty parquet")
if require_force:
metadata = parquet_file.schema_arrow.metadata or {}
if FORCE_SCHEMA_METADATA_KEY not in metadata:
raise ValueError(f"{parquet_path}: force schema metadata missing")
entry = manifest["episodes"].get(f"{episode_index:06d}", {})
if entry.get("status") != "complete":
raise ValueError(
f"{root}: force manifest sample {episode_index} incomplete"
)
track_path = _track_path(root, episode_index)
with np.load(track_path, allow_pickle=False) as payload:
if int(np.asarray(payload["episode_index"]).item()) != episode_index:
raise ValueError(f"{track_path}: episode_index mismatch")
if int(np.asarray(payload["num_steps"]).item()) != int(
parquet_file.metadata.num_rows
):
raise ValueError(f"{track_path}: frame count mismatch")
return {
"dataset": str(root),
"episodes": episodes,
"frames": frames,
"tasks": tasks,
"videos": videos,
"force": require_force,
"validated_at": ready.get("validated_at"),
}
def _build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--dataset-root", type=Path, required=True)
parser.add_argument(
"--require-force",
action=argparse.BooleanOptionalAction,
default=False,
)
return parser
def main(argv: Sequence[str] | None = None) -> int:
args = _build_parser().parse_args(argv)
result = check_dataset(
args.dataset_root,
require_force=args.require_force,
)
print(json.dumps(result))
return 0
if __name__ == "__main__":
raise SystemExit(main())
|