dreamzero / scripts /data /convert_trex_v3_to_v2.py
zhicao's picture
Upload dreamzero source code (no model checkpoints)
fbd9366 verified
Raw
History Blame Contribute Delete
16.6 kB
"""
Convert the T-Rex LeRobot v3.0 dataset to LeRobot v2.1 layout for DreamZero.
LeRobot v3 packs many episodes into shared parquet/video files:
data/chunk-XXX/file-XXX.parquet (rows of many episodes)
videos/{video_key}/chunk-XXX/file-XXX.mp4 (concatenated episodes)
meta/episodes/chunk-XXX/file-XXX.parquet (episode metadata)
meta/tasks.parquet
DreamZero's loader (groot/vla/data/dataset/lerobot.py) expects v2 layout:
data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet
videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4
meta/episodes.jsonl, meta/tasks.jsonl, meta/info.json
RGB cameras are re-encoded to 320x180. Tactile videos keep native resolution
(raw 320x240, deform 240x240) and are re-encoded with libx264 for smaller size.
Usage:
python scripts/data/convert_trex_v3_to_v2.py --phase data
python scripts/data/convert_trex_v3_to_v2.py --phase videos
python scripts/data/convert_trex_v3_to_v2.py --phase videos --include-tactile
python scripts/data/convert_trex_v3_to_v2.py --phase meta --include-tactile
python scripts/data/convert_trex_v3_to_v2.py --phase verify --include-tactile
All phases are resumable: existing valid outputs are skipped.
"""
from __future__ import annotations
import argparse
import json
import logging
import subprocess
from concurrent.futures import ProcessPoolExecutor, as_completed
from pathlib import Path
import numpy as np
import pandas as pd
from tqdm import tqdm
logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s")
log = logging.getLogger(__name__)
DEFAULT_SRC = Path("/scratch1/home/zhicao/dreamzero/data/trex_dataset")
DEFAULT_DST = Path("/scratch1/home/zhicao/dreamzero/data/trex_datasetv2")
SRC = DEFAULT_SRC
DST = DEFAULT_DST
RGB_VIDEO_KEYS = [
"observation.images.head_left",
"observation.images.left_wrist",
"observation.images.right_wrist",
]
RGB_OUT_W, RGB_OUT_H = 320, 180
FPS = 30
CHUNKS_SIZE = 1000
RGB_CRF = 23
TACTILE_CRF = 28
def load_src_info() -> dict:
return json.loads((SRC / "meta" / "info.json").read_text())
def get_tactile_video_keys(src_info: dict | None = None) -> list[str]:
src_info = src_info or load_src_info()
return sorted(
k
for k, v in src_info["features"].items()
if v.get("dtype") == "video" and "tactile" in k
)
def get_output_size(video_key: str, feature: dict) -> tuple[int, int]:
"""Return (width, height) for ffmpeg scale filter."""
if video_key in RGB_VIDEO_KEYS:
return RGB_OUT_W, RGB_OUT_H
shape = feature.get("shape", [])
if len(shape) >= 2:
height, width = int(shape[0]), int(shape[1])
return width, height
info = feature.get("info", {})
return int(info["video.width"]), int(info["video.height"])
def get_crf(video_key: str) -> int:
return RGB_CRF if video_key in RGB_VIDEO_KEYS else TACTILE_CRF
def resolve_video_keys(
include_tactile: bool,
tactile_only: bool,
explicit_keys: list[str] | None,
) -> list[str]:
if explicit_keys:
return explicit_keys
if tactile_only:
return get_tactile_video_keys()
keys = list(RGB_VIDEO_KEYS)
if include_tactile:
keys.extend(get_tactile_video_keys())
return keys
def load_episode_meta() -> pd.DataFrame:
files = sorted(SRC.glob("meta/episodes/chunk-*/file-*.parquet"))
if not files:
raise FileNotFoundError(f"No episode metadata under {SRC / 'meta/episodes'}")
df = pd.concat([pd.read_parquet(f) for f in files], ignore_index=True)
return df.sort_values("episode_index").reset_index(drop=True)
def load_task_map() -> dict[int, str]:
t = pd.read_parquet(SRC / "meta" / "tasks.parquet")
return {int(row.task_index): str(idx) for idx, row in t.iterrows()}
def ep_parquet_path(ep_idx: int) -> Path:
return DST / f"data/chunk-{ep_idx // CHUNKS_SIZE:03d}/episode_{ep_idx:06d}.parquet"
def ep_video_path(ep_idx: int, video_key: str) -> Path:
return DST / f"videos/chunk-{ep_idx // CHUNKS_SIZE:03d}/{video_key}/episode_{ep_idx:06d}.mp4"
def build_video_features(src_info: dict, video_keys: list[str]) -> dict:
features: dict = {}
for k, v in src_info["features"].items():
if v.get("dtype") != "video" or k not in video_keys:
continue
v = dict(v)
out_w, out_h = get_output_size(k, v)
v["shape"] = [out_h, out_w, 3]
info_blk = dict(v.get("info", {}))
info_blk.update(
{
"video.height": out_h,
"video.width": out_w,
"video.codec": "h264",
"video.pix_fmt": "yuv420p",
"video.fps": FPS,
"video.channels": 3,
"has_audio": False,
}
)
v["info"] = info_blk
features[k] = v
return features
def write_info_json(
src_info: dict,
video_keys: list[str],
*,
preserve_existing_features: bool = False,
) -> None:
meta_dir = DST / "meta"
meta_dir.mkdir(parents=True, exist_ok=True)
existing = {}
info_path = meta_dir / "info.json"
if preserve_existing_features and info_path.exists():
existing = json.loads(info_path.read_text())
features = dict(existing.get("features", {}))
for k, v in src_info["features"].items():
if v.get("dtype") != "video":
features[k] = v
features.update(build_video_features(src_info, video_keys))
features["annotation.task"] = {"dtype": "string", "shape": [1], "names": None}
total_episodes = int(src_info["total_episodes"])
all_video_keys = [k for k, v in features.items() if v.get("dtype") == "video"]
info = {
"codebase_version": "v2.1",
"robot_type": src_info.get("robot_type", "dexmate_vega1_and_sharpa_wave"),
"total_episodes": total_episodes,
"total_frames": int(src_info["total_frames"]),
"total_tasks": int(src_info["total_tasks"]),
"total_videos": total_episodes * len(all_video_keys),
"total_chunks": (total_episodes + CHUNKS_SIZE - 1) // CHUNKS_SIZE,
"chunks_size": CHUNKS_SIZE,
"fps": FPS,
"splits": {"train": f"0:{total_episodes}"},
"data_path": "data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet",
"video_path": "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4",
"features": features,
}
with open(info_path, "w") as f:
json.dump(info, f, indent=4)
log.info(
"Wrote meta/info.json with %d video keys (%d total videos)",
len(all_video_keys),
info["total_videos"],
)
# ---------------------------------------------------------------------------
# Phase: data (parquets + meta)
# ---------------------------------------------------------------------------
def convert_data(
ep_meta: pd.DataFrame,
task_map: dict[int, str],
video_keys: list[str],
) -> None:
src_info = load_src_info()
n_done = 0
groups = ep_meta.groupby(["data/chunk_index", "data/file_index"])
for (chunk_idx, file_idx), eps in tqdm(groups, desc="Converting data files"):
src_pq = SRC / f"data/chunk-{int(chunk_idx):03d}/file-{int(file_idx):03d}.parquet"
if not src_pq.exists():
log.warning("Missing source parquet: %s", src_pq)
continue
if all(ep_parquet_path(int(r.episode_index)).exists() for r in eps.itertuples()):
n_done += len(eps)
continue
df = pd.read_parquet(src_pq)
for r in eps.itertuples():
ep_idx = int(r.episode_index)
out = ep_parquet_path(ep_idx)
if out.exists():
n_done += 1
continue
ep_df = df[df["episode_index"] == ep_idx].copy()
assert len(ep_df) == int(r.length), (
f"episode {ep_idx}: rows {len(ep_df)} != meta length {r.length}"
)
task_texts = [str(t) for t in r.tasks]
ep_df["annotation.task"] = task_texts[0] if task_texts else ""
out.parent.mkdir(parents=True, exist_ok=True)
ep_df.to_parquet(out, index=False)
n_done += 1
log.info("Data phase done: %d episode parquets", n_done)
meta_dir = DST / "meta"
meta_dir.mkdir(parents=True, exist_ok=True)
with open(meta_dir / "tasks.jsonl", "w") as f:
for idx in sorted(task_map):
f.write(json.dumps({"task_index": idx, "task": task_map[idx]}) + "\n")
with open(meta_dir / "episodes.jsonl", "w") as f:
for r in ep_meta.itertuples():
f.write(
json.dumps(
{
"episode_index": int(r.episode_index),
"tasks": [str(t) for t in r.tasks],
"length": int(r.length),
}
)
+ "\n"
)
write_info_json(src_info, video_keys, preserve_existing_features=False)
log.info("Wrote meta/episodes.jsonl, meta/tasks.jsonl")
def update_meta(video_keys: list[str]) -> None:
src_info = load_src_info()
write_info_json(src_info, video_keys, preserve_existing_features=True)
# ---------------------------------------------------------------------------
# Phase: videos
# ---------------------------------------------------------------------------
def _cut_one(job: tuple) -> tuple[int, str, bool, str]:
ep_idx, video_key, src_mp4, from_ts, n_frames, out_path, out_w, out_h, crf = job
out = Path(out_path)
out.parent.mkdir(parents=True, exist_ok=True)
tmp = out.with_suffix(".tmp.mp4")
ss = max(0.0, from_ts - 0.5 / FPS)
cmd = [
"ffmpeg",
"-y",
"-loglevel",
"error",
"-ss",
f"{ss:.6f}",
"-i",
src_mp4,
"-frames:v",
str(n_frames),
"-vf",
f"scale={out_w}:{out_h}",
"-c:v",
"libx264",
"-preset",
"veryfast",
"-crf",
str(crf),
"-pix_fmt",
"yuv420p",
"-movflags",
"+faststart",
"-an",
"-threads",
"2",
str(tmp),
]
try:
res = subprocess.run(cmd, capture_output=True, text=True, timeout=600)
if res.returncode != 0:
tmp.unlink(missing_ok=True)
return ep_idx, video_key, False, res.stderr[-500:]
tmp.rename(out)
return ep_idx, video_key, True, ""
except Exception as e: # noqa: BLE001
tmp.unlink(missing_ok=True)
return ep_idx, video_key, False, str(e)
def convert_videos(ep_meta: pd.DataFrame, video_keys: list[str], workers: int) -> None:
src_info = load_src_info()
jobs = []
missing_src = set()
for vk in video_keys:
feature = src_info["features"][vk]
out_w, out_h = get_output_size(vk, feature)
crf = get_crf(vk)
for r in ep_meta.itertuples():
ep_idx = int(r.episode_index)
out = ep_video_path(ep_idx, vk)
if out.exists():
continue
chunk_i = int(ep_meta.loc[r.Index, f"videos/{vk}/chunk_index"])
file_i = int(ep_meta.loc[r.Index, f"videos/{vk}/file_index"])
from_ts = float(ep_meta.loc[r.Index, f"videos/{vk}/from_timestamp"])
src_mp4 = SRC / f"videos/{vk}/chunk-{chunk_i:03d}/file-{file_i:03d}.mp4"
if not src_mp4.exists():
missing_src.add(str(src_mp4))
continue
jobs.append(
(
ep_idx,
vk,
str(src_mp4),
from_ts,
int(r.length),
str(out),
out_w,
out_h,
crf,
)
)
if missing_src:
log.warning(
"%d source videos missing (not yet downloaded?), e.g. %s",
len(missing_src),
sorted(missing_src)[0],
)
log.info("Cutting %d episode videos with %d workers", len(jobs), workers)
failures = []
with ProcessPoolExecutor(max_workers=workers) as pool:
futs = [pool.submit(_cut_one, j) for j in jobs]
for fut in tqdm(as_completed(futs), total=len(futs), desc="Cutting videos"):
ep_idx, vk, ok, err = fut.result()
if not ok:
failures.append((ep_idx, vk, err))
if failures:
log.error("%d failures, first: %s", len(failures), failures[0])
else:
log.info("Video phase done, no failures")
# ---------------------------------------------------------------------------
# Phase: verify
# ---------------------------------------------------------------------------
def _probe_frames(path: Path) -> int:
res = subprocess.run(
[
"ffprobe",
"-v",
"error",
"-count_frames",
"-select_streams",
"v:0",
"-show_entries",
"stream=nb_read_frames",
"-of",
"csv=p=0",
str(path),
],
capture_output=True,
text=True,
timeout=120,
)
return int(res.stdout.strip())
def verify(ep_meta: pd.DataFrame, video_keys: list[str], n_samples: int) -> None:
rng = np.random.default_rng(0)
total = len(ep_meta)
missing_pq = [
int(r.episode_index)
for r in ep_meta.itertuples()
if not ep_parquet_path(int(r.episode_index)).exists()
]
log.info("Parquets: %d/%d present", total - len(missing_pq), total)
for vk in video_keys:
missing = [
int(r.episode_index)
for r in ep_meta.itertuples()
if not ep_video_path(int(r.episode_index), vk).exists()
]
log.info("Videos [%s]: %d/%d present", vk, total - len(missing), total)
sample = rng.choice(total, size=min(n_samples, total), replace=False)
for ep_idx in sample:
ep_idx = int(ep_idx)
row = ep_meta[ep_meta["episode_index"] == ep_idx].iloc[0]
length = int(row["length"])
pq = ep_parquet_path(ep_idx)
if pq.exists():
df = pd.read_parquet(pq)
assert len(df) == length, f"ep {ep_idx}: parquet {len(df)} != {length}"
assert np.asarray(df["action"].iloc[0]).shape == (58,)
assert df["annotation.task"].iloc[0] == str(row["tasks"][0])
for vk in video_keys:
vp = ep_video_path(ep_idx, vk)
if vp.exists():
n = _probe_frames(vp)
assert n == length, f"ep {ep_idx} {vk}: video {n} frames != {length}"
log.info("ep %06d OK (length=%d)", ep_idx, length)
log.info("Verification passed on %d sampled episodes", len(sample))
def main() -> None:
global SRC, DST
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--phase",
choices=["data", "videos", "meta", "verify"],
required=True,
)
parser.add_argument("--src", type=Path, default=DEFAULT_SRC, help="v3 dataset root")
parser.add_argument("--dst", type=Path, default=DEFAULT_DST, help="v2 output root")
parser.add_argument(
"--include-tactile",
action="store_true",
help="Include all 20 tactile video streams",
)
parser.add_argument(
"--tactile-only",
action="store_true",
help="Convert/update only tactile video streams (skip RGB)",
)
parser.add_argument(
"--video-keys",
nargs="+",
default=None,
help="Explicit video keys to convert (overrides --include-tactile default set)",
)
parser.add_argument("--workers", type=int, default=16)
parser.add_argument("--verify-samples", type=int, default=20)
args = parser.parse_args()
SRC = args.src
DST = args.dst
ep_meta = load_episode_meta()
log.info("Loaded %d episodes from v3 metadata", len(ep_meta))
video_keys = resolve_video_keys(args.include_tactile, args.tactile_only, args.video_keys)
log.info("Video keys (%d): %s", len(video_keys), ", ".join(video_keys))
if args.phase == "data":
convert_data(ep_meta, load_task_map(), video_keys)
elif args.phase == "meta":
update_meta(video_keys)
elif args.phase == "videos":
convert_videos(ep_meta, video_keys, args.workers)
elif args.phase == "verify":
verify(ep_meta, video_keys, args.verify_samples)
if __name__ == "__main__":
main()