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