| """Merge several LeRobot (v2.1) datasets into one, re-indexed and with a shared |
| task table. The multi-task counterpart of split_dataset.py. |
| |
| Episodes are concatenated in the order the sources are given. Each episode is |
| re-numbered, its parquet's `episode_index` / `index` / `task_index` columns are |
| rewritten, and its videos are hard-linked (same filesystem) or copied. Task |
| strings are deduplicated across sources, so two datasets sharing a prompt share |
| one task_index. |
| |
| Example: |
| python scripts/merge_datasets.py \ |
| --srcs ~/workspace/yuhao/pi/data/item-classification-eef \ |
| ~/workspace/yuhao/pi/data/sort-tools-eef \ |
| ~/workspace/yuhao/pi/data/stack-cube-eef \ |
| --out ~/workspace/yuhao/pi/data/g1-multitask-eef |
| |
| Sources are left untouched. Only the metadata every LeRobot loader needs is |
| written (info/tasks/episodes/episodes_stats); consumer-specific extras that some |
| sources happen to carry (modality.json, stats.json, relative_stats.json) are NOT |
| merged, because they describe one source's statistics and would be wrong for the |
| union — regenerate them for the merged set if a downstream tool needs them. |
| """ |
|
|
| import dataclasses |
| import json |
| import pathlib |
| import shutil |
|
|
| import numpy as np |
| import pyarrow as pa |
| import pyarrow.parquet as pq |
| import tyro |
|
|
|
|
| @dataclasses.dataclass |
| class Args: |
| |
| srcs: list[str] |
| |
| out: str |
| |
| copy: bool = False |
|
|
|
|
| def _link_or_copy(src: pathlib.Path, dst: pathlib.Path, force_copy: bool) -> None: |
| dst.parent.mkdir(parents=True, exist_ok=True) |
| if force_copy: |
| shutil.copy2(src, dst) |
| return |
| try: |
| dst.hardlink_to(src) |
| except OSError: |
| shutil.copy2(src, dst) |
|
|
|
|
| def _load_jsonl(path: pathlib.Path) -> dict: |
| return {json.loads(l)["episode_index"]: json.loads(l) for l in path.open()} |
|
|
|
|
| def _check_compatible(infos: list[dict], srcs: list[pathlib.Path]) -> None: |
| """Every field that must agree for the episodes to be interchangeable. |
| |
| `features` is compared whole: it pins dtypes, shapes, joint/EEF column names |
| and the video codec settings in one go. A mismatch here means the merged |
| dataset would feed a model two different action spaces or two different |
| decoders, which no downstream loader would catch. |
| """ |
| keys = ["codebase_version", "robot_type", "fps", "chunks_size", |
| "data_path", "video_path", "features"] |
| for key in keys: |
| values = [info.get(key) for info in infos] |
| if any(v != values[0] for v in values): |
| detail = "\n".join(f" {s.name}: {json.dumps(v)[:200]}" |
| for s, v in zip(srcs, values)) |
| raise SystemExit(f"Sources disagree on info.json['{key}']:\n{detail}") |
|
|
|
|
| def main(args: Args) -> None: |
| srcs = [pathlib.Path(s).expanduser().resolve() for s in args.srcs] |
| out = pathlib.Path(args.out).expanduser().resolve() |
| if len(srcs) < 2: |
| raise SystemExit("Need at least two --srcs to merge.") |
| if out in srcs: |
| raise SystemExit("--out must differ from every source.") |
| if out.exists(): |
| raise SystemExit(f"--out already exists: {out}") |
|
|
| infos = [json.loads((s / "meta/info.json").read_text()) for s in srcs] |
| _check_compatible(infos, srcs) |
| info = infos[0] |
| chunks_size = info["chunks_size"] |
| data_tmpl, video_tmpl = info["data_path"], info["video_path"] |
| video_keys = [k for k, v in info["features"].items() if v["dtype"] == "video"] |
|
|
| |
| task_ids: dict[str, int] = {} |
| remaps = [] |
| for src in srcs: |
| remap = {} |
| for line in (src / "meta/tasks.jsonl").open(): |
| row = json.loads(line) |
| remap[row["task_index"]] = task_ids.setdefault(row["task"], len(task_ids)) |
| remaps.append(remap) |
|
|
| (out / "meta").mkdir(parents=True, exist_ok=True) |
| new_eps, new_stats = [], [] |
| new_ep, running = 0, 0 |
|
|
| for src, src_info, remap in zip(srcs, infos, remaps): |
| eps = _load_jsonl(src / "meta/episodes.jsonl") |
| stats = _load_jsonl(src / "meta/episodes_stats.jsonl") |
| n_eps = src_info["total_episodes"] |
| print(f"{src.name}: {n_eps} episodes, {src_info['total_frames']} frames " |
| f"-> episodes {new_ep}..{new_ep + n_eps - 1}") |
|
|
| for old in range(n_eps): |
| src_chunk, dst_chunk = old // chunks_size, new_ep // chunks_size |
| table = pq.read_table( |
| src / data_tmpl.format(episode_chunk=src_chunk, episode_index=old)) |
| n = table.num_rows |
|
|
| for name, values in ( |
| ("episode_index", pa.array([new_ep] * n, pa.int64())), |
| ("index", pa.array(np.arange(running, running + n), pa.int64())), |
| ("task_index", pa.array( |
| [remap[t] for t in table.column("task_index").to_pylist()], pa.int64())), |
| ): |
| table = table.set_column(table.schema.get_field_index(name), name, values) |
|
|
| dst_pq = out / data_tmpl.format(episode_chunk=dst_chunk, episode_index=new_ep) |
| dst_pq.parent.mkdir(parents=True, exist_ok=True) |
| pq.write_table(table, dst_pq) |
|
|
| for vk in video_keys: |
| _link_or_copy( |
| src / video_tmpl.format(episode_chunk=src_chunk, video_key=vk, episode_index=old), |
| out / video_tmpl.format(episode_chunk=dst_chunk, video_key=vk, episode_index=new_ep), |
| args.copy) |
|
|
| new_eps.append({**eps[old], "episode_index": new_ep}) |
| new_stats.append({**stats[old], "episode_index": new_ep}) |
| new_ep += 1 |
| running += n |
|
|
| with (out / "meta/episodes.jsonl").open("w") as f: |
| f.writelines(json.dumps(e) + "\n" for e in new_eps) |
| with (out / "meta/episodes_stats.jsonl").open("w") as f: |
| f.writelines(json.dumps(s) + "\n" for s in new_stats) |
| with (out / "meta/tasks.jsonl").open("w") as f: |
| f.writelines(json.dumps({"task_index": i, "task": t}) + "\n" |
| for t, i in sorted(task_ids.items(), key=lambda kv: kv[1])) |
|
|
| (out / "meta/info.json").write_text(json.dumps({ |
| **info, |
| "total_episodes": new_ep, |
| "total_frames": running, |
| "total_tasks": len(task_ids), |
| "total_videos": new_ep * len(video_keys), |
| "total_chunks": (new_ep - 1) // chunks_size + 1 if new_ep else 0, |
| "splits": {"train": f"0:{new_ep}"}, |
| }, indent=4)) |
|
|
| print(f"\nMerged -> {out}") |
| print(f" {new_ep} episodes, {running} frames, {len(task_ids)} tasks, " |
| f"{new_ep * len(video_keys)} videos") |
|
|
|
|
| if __name__ == "__main__": |
| main(tyro.cli(Args)) |
|
|