max / openpi /code /scripts /merge_datasets.py
LGG100's picture
Add files using upload-large-folder tool
3acefc3 verified
Raw
History Blame Contribute Delete
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
@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))