Download openpi/code/scripts/merge_datasets.py from LGG100/max: direct link, hf CLI and curl.
- Browser
- Download file 6.92 kB
-
https://huggingface.co/LGG100/max/resolve/main/openpi/code/scripts/merge_datasets.py
- Command line
-
hf download hf://LGG100/max/openpi/code/scripts/merge_datasets.py
-
curl -L -o merge_datasets.py https://huggingface.co/LGG100/max/resolve/main/openpi/code/scripts/merge_datasets.py
6.92 kB
| """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 | |
| 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)) | |