Download openpi/code/scripts/trim_static_frames.py from LGG100/max: direct link, hf CLI and curl.
- Browser
- Download file 6.32 kB
-
https://huggingface.co/LGG100/max/resolve/main/openpi/code/scripts/trim_static_frames.py
- Command line
-
hf download hf://LGG100/max/openpi/code/scripts/trim_static_frames.py
-
curl -L -o trim_static_frames.py https://huggingface.co/LGG100/max/resolve/main/openpi/code/scripts/trim_static_frames.py
6.32 kB
| """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 | |
| 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)) | |