File size: 6,917 Bytes
3acefc3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 | """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))
|