dreamzero / scripts /data /build_trex_mini_force.py
zhicao's picture
Upload dreamzero source code (no model checkpoints)
fbd9366 verified
Raw
History Blame Contribute Delete
6.77 kB
#!/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())