"""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: # Source LeRobot dataset directories, concatenated in this order. srcs: list[str] # Output directory for the merged dataset. out: str # Copy videos instead of hard-linking them (needed across filesystems). 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: # different filesystem, or links unsupported 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"] # Shared task table: first-seen order, deduplicated across sources. 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))