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