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())