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))