| """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: |
| |
| src: str |
| |
| out: str |
| |
| keep_tail: int = 10 |
| |
| 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: |
| 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)) |
| last_move = n - 1 - int(np.argmax(moving[::-1])) |
| 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), |
| ) |
|
|
| |
| 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) |
|
|
| |
| 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)) |
|
|