File size: 6,769 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
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
#!/usr/bin/env python3
"""Build one train-ready T-Rex Track-Force demo with fresh statistics."""

from __future__ import annotations

import argparse
import copy
import hashlib
import json
import os
import shutil
import sys
from pathlib import Path
from typing import Sequence

_ROOT = Path(__file__).resolve().parents[2]
_DATA_SCRIPTS = _ROOT / "scripts" / "data"
if str(_DATA_SCRIPTS) not in sys.path:
    sys.path.insert(0, str(_DATA_SCRIPTS))

import build_trex_track_force_v2 as force_builder  # noqa: E402
import check_trex_dataset_ready as ready_check  # noqa: E402
import rebuild_trex_dataset_variants as variants  # noqa: E402


def _read_json(path: Path) -> dict:
    return json.loads(path.read_text())


def _read_jsonl(path: Path) -> list[dict]:
    return [
        json.loads(line)
        for line in path.read_text().splitlines()
        if line.strip()
    ]


def _sha256_or_empty(path: Path) -> str:
    if path.is_file():
        return variants._sha256(path)
    return hashlib.sha256(b"[]\n").hexdigest()


def build_mini_dataset(
    *,
    source_root: Path,
    output_root: Path,
    source_episode: int,
    workers: int,
) -> dict:
    source_root = source_root.expanduser().resolve()
    output_root = output_root.expanduser().resolve()
    if source_root == output_root:
        raise ValueError("source and output dataset roots must differ")

    ready_check.check_dataset(source_root, require_force=True)
    if output_root.exists():
        result = ready_check.check_dataset(output_root, require_force=True)
        variant = _read_json(output_root / "meta" / "dataset_variant.json")
        if int(variant.get("source_episode", -1)) != source_episode:
            raise ValueError(
                f"{output_root} already uses source episode "
                f"{variant.get('source_episode')}, not {source_episode}"
            )
        return result

    source_info = _read_json(source_root / "meta" / "info.json")
    source_modality = _read_json(source_root / "meta" / "modality.json")
    source_embodiment = _read_json(source_root / "meta" / "embodiment.json")
    source_episodes = _read_jsonl(source_root / "meta" / "episodes.jsonl")
    source_tasks = _read_jsonl(source_root / "meta" / "tasks.jsonl")
    selected = [
        episode
        for episode in source_episodes
        if int(episode["episode_index"]) == source_episode
    ]
    if len(selected) != 1:
        raise ValueError(
            f"source episode {source_episode} was not found exactly once"
        )
    registry, tasks = variants._build_registry(
        selected,
        source_tasks,
        excluded=set(),
        source_limit=None,
    )
    if len(registry) != 1 or registry[0].episode_index != 0:
        raise AssertionError("mini dataset must contain dense episode 0")

    stage_root = output_root.parent / f".{output_root.name}.staging"
    if stage_root.exists():
        shutil.rmtree(stage_root)
    video_keys = variants._video_keys(source_info)
    blacklist = source_root / "meta" / "excluded_source_episode_indices.json"
    variants._write_variant_metadata_skeleton(
        variant_root=stage_root,
        final_root=output_root,
        source_info=source_info,
        source_modality=source_modality,
        source_embodiment=source_embodiment,
        registry=registry,
        tasks=tasks,
        video_keys=video_keys,
        excluded_source_indices=[],
        force=True,
        source_dataset=source_root,
        blacklist_sha256=_sha256_or_empty(blacklist),
    )
    variant_path = stage_root / "meta" / "dataset_variant.json"
    variant = _read_json(variant_path)
    variant["source_episode"] = source_episode
    variants._atomic_write_json(variant_path, variant)

    variants._rewrite_parquets(
        source_root=source_root,
        destination_root=stage_root,
        registry=registry,
        workers=workers,
        force=True,
    )
    variants._link_videos(
        source_root=source_root,
        destination_root=stage_root,
        registry=registry,
        video_keys=video_keys,
        source_uses_new_indices=False,
        workers=workers,
    )
    track_hashes = variants._rewrite_tracks(
        source_root=source_root,
        destination_root=stage_root,
        registry=registry,
        workers=workers,
    )

    source_manifest = _read_json(
        source_root / "meta" / "trex_track_force_manifest.json"
    )
    source_entries = {
        int(index): copy.deepcopy(entry)
        for index, entry in source_manifest["episodes"].items()
    }
    variants._build_force_manifest(
        variant_root=stage_root,
        final_root=output_root,
        registry=registry,
        source_entries=source_entries,
        track_hashes=track_hashes,
    )

    variants._atomic_write_json(
        stage_root / "meta" / "stats.json",
        variants._compute_base_stats(stage_root, registry),
    )
    metadata_result = force_builder.update_metadata(
        stage_root,
        assume_all_converted=True,
    )
    manifest_path = stage_root / "meta" / "trex_track_force_manifest.json"
    manifest = _read_json(manifest_path)
    manifest["metadata"] = metadata_result
    manifest["updated_at"] = variants._utc_now()
    variants._atomic_write_json(manifest_path, manifest)
    force_builder.validate_metadata(stage_root)
    for backup in (stage_root / "meta").glob("*.trex_track_force.bak"):
        backup.unlink()

    variants._validate_variant(
        root=stage_root,
        expected_registry=registry,
        video_keys=video_keys,
        force=True,
    )
    os.replace(stage_root, output_root)
    variants._validate_variant(
        root=output_root,
        expected_registry=registry,
        video_keys=video_keys,
        force=True,
    )
    return ready_check.check_dataset(output_root, require_force=True)


def _parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--source-root",
        type=Path,
        default=_ROOT / "data" / "trex_full_force",
    )
    parser.add_argument(
        "--output-root",
        type=Path,
        default=_ROOT / "data" / "trex_mini_force",
    )
    parser.add_argument("--source-episode", type=int, default=49)
    parser.add_argument("--workers", type=int, default=4)
    return parser.parse_args(argv)


def main(argv: Sequence[str] | None = None) -> int:
    args = _parse_args(argv)
    result = build_mini_dataset(
        source_root=args.source_root,
        output_root=args.output_root,
        source_episode=args.source_episode,
        workers=max(1, int(args.workers)),
    )
    print(json.dumps(result))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())