File size: 6,319 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
"""Trim the static head and tail of every episode in a LeRobot (v2.1) dataset.

A frame is "moving" if either observation.state or action changes between it and
the next frame. Everything before the first moving frame is dropped, and the
static tail after the last moving frame is cut down to `--keep-tail` frames, so
the policy still sees the arm come to rest. Short pauses in the middle of an
episode are kept: cutting them would break the timing of the action chunks.

Videos are hard-linked, not re-encoded. Each kept row keeps its original
`timestamp` and `frame_index`, so LeRobot still decodes the matching frame from
the untouched per-episode mp4 (it only checks that timestamps within an episode
are 1/fps apart, not that they start at 0). `index` is renumbered globally and
episodes.jsonl / episodes_stats.jsonl are rewritten for the new lengths.

Example:
    python scripts/trim_static_frames.py \
        --src ~/workspace/yuhao/pi/data/rm65/4tasks_v1 \
        --out ~/workspace/yuhao/pi/data/rm65/4tasks_v1_trim
"""

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 v2.1 dataset directory.
    src: str
    # Output directory for the trimmed dataset.
    out: str
    # Static frames kept after the last moving frame.
    keep_tail: int = 10
    # Absolute change below which a state/action value counts as unchanged.
    eps: float = 1e-6


def _link_or_copy(src: pathlib.Path, dst: pathlib.Path) -> None:
    dst.parent.mkdir(parents=True, exist_ok=True)
    try:
        dst.hardlink_to(src)
    except OSError:  # different filesystem, or links unsupported
        shutil.copy2(src, dst)


def _vector_stats(x: np.ndarray) -> dict:
    return {
        "min": x.min(0).tolist(),
        "max": x.max(0).tolist(),
        "mean": x.mean(0).tolist(),
        "std": x.std(0).tolist(),
        "count": [len(x)],
    }


def main(args: Args) -> None:
    src = pathlib.Path(args.src).expanduser().resolve()
    out = pathlib.Path(args.out).expanduser().resolve()
    if out == src:
        raise SystemExit("--out must differ from --src.")
    if out.exists():
        raise SystemExit(f"--out already exists: {out}")

    info = json.loads((src / "meta/info.json").read_text())
    if info["codebase_version"] != "v2.1":
        raise SystemExit(f"Expected a v2.1 dataset, got {info['codebase_version']}")
    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"]

    eps = [json.loads(l) for l in (src / "meta/episodes.jsonl").open()]
    stats = {json.loads(l)["episode_index"]: json.loads(l) for l in (src / "meta/episodes_stats.jsonl").open()}

    (out / "meta").mkdir(parents=True)
    shutil.copy2(src / "meta/tasks.jsonl", out / "meta/tasks.jsonl")
    new_eps, new_stats = [], []
    running, dropped_head, dropped_tail = 0, 0, 0

    for ep in eps:
        ep_idx = ep["episode_index"]
        chunk = ep_idx // chunks_size
        table = pq.read_table(src / data_tmpl.format(episode_chunk=chunk, episode_index=ep_idx))
        state = np.stack(table.column("observation.state").to_numpy(zero_copy_only=False))
        action = np.stack(table.column("action").to_numpy(zero_copy_only=False))
        n = len(state)

        moving = (np.abs(np.diff(state, axis=0)).max(1) > args.eps) | (
            np.abs(np.diff(action, axis=0)).max(1) > args.eps
        )
        if not moving.any():
            raise SystemExit(f"Episode {ep_idx} never moves; refusing to trim it to nothing.")
        start = int(np.argmax(moving))  # first frame that moves to the next one
        last_move = n - 1 - int(np.argmax(moving[::-1]))  # frame reached by the last movement
        end = min(n, last_move + 1 + args.keep_tail)
        dropped_head += start
        dropped_tail += n - end

        table = table.slice(start, end - start)
        m = table.num_rows
        table = table.set_column(
            table.schema.get_field_index("index"), "index", pa.array(np.arange(running, running + m), pa.int64())
        )
        dst_pq = out / data_tmpl.format(episode_chunk=chunk, episode_index=ep_idx)
        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=chunk, video_key=vk, episode_index=ep_idx),
                out / video_tmpl.format(episode_chunk=chunk, video_key=vk, episode_index=ep_idx),
            )

        # action_config frame ranges are relative to the episode start.
        new_ep = {**ep, "length": m}
        if "action_config" in ep:
            new_ep["action_config"] = [
                {**c, "start_frame": max(0, c["start_frame"] - start), "end_frame": min(m, c["end_frame"] - start)}
                for c in ep["action_config"]
                if c["end_frame"] > start and c["start_frame"] - start < m
            ]
        new_eps.append(new_ep)

        # Recompute the per-episode stats that depend on the kept rows; video stats are kept as-is.
        st = dict(stats[ep_idx]["stats"])
        st["observation.state"] = _vector_stats(state[start:end])
        st["action"] = _vector_stats(action[start:end])
        for key in ("timestamp", "frame_index", "episode_index", "index", "task_index"):
            col = table.column(key).to_numpy().astype(np.float64)[:, None]
            st[key] = _vector_stats(col)
        new_stats.append({"episode_index": ep_idx, "stats": st})
        running += m

    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)
    (out / "meta/info.json").write_text(json.dumps({**info, "total_frames": running}, indent=4))

    total = running + dropped_head + dropped_tail
    print(f"Trimmed -> {out}")
    print(f"  {len(new_eps)} episodes, {total} -> {running} frames "
          f"(dropped {dropped_head} head + {dropped_tail} tail, {1 - running / total:.1%})")


if __name__ == "__main__":
    main(tyro.cli(Args))