From 8bda1ec5ee7291f854a3a21e925580b19ce639bc Mon Sep 17 00:00:00 2001 From: joaner <1726541+joaner@users.noreply.github.com> Date: Tue, 22 Sep 2026 15:48:56 +0800 Subject: [PATCH 1/7] Map dataset cameras onto OpenPI slots and reuse v3 conversions. --- README.md | 61 +++++- README.zh-CN.md | 55 ++++- lerobot_v3_compat.py | 409 +++++++++++++++++++++++++++++++++++--- test_lerobot_v3_compat.py | 166 ++++++++++++++++ test_train_lerobot.py | 44 ++++ train_lerobot.py | 305 ++++++++++++++++++++++------ 6 files changed, 948 insertions(+), 92 deletions(-) create mode 100644 test_train_lerobot.py diff --git a/README.md b/README.md index a767cb8..d0edefa 100644 --- a/README.md +++ b/README.md @@ -76,8 +76,17 @@ docker run --rm --gpus all --ipc=host \ | `--ema_decay` | off | e.g. `0.99` | | `--action_horizon` | `50` | | | `--num_workers` | `8` | | +| `--dataset_dir` | `/data/input` | or `$OPENPI_DATASET_DIR` | +| `--output_dir` | `/data/output` | or `$OPENPI_OUTPUT_DIR` | +| `--run_name` | `docker_train` | checkpoint parent directory | +| `--exp_name` | `train` | | +| `--convert_dir` | under `--output_dir` | v3→v2 cache; do not use a small tmpfs | +| `--cameras` | all image keys | comma-separated keys to keep | +| `--drop_cameras` | | key or substring, e.g. `front` | +| `--camera_map` | role-based | `base=key,left_wrist=key,right_wrist=key` | +| `--delta_joint_actions` | off | joint deltas, last dim (gripper) stays absolute | | `--norm_stats_workers` | `min(cpu, 64)` | | -| `--norm_stats_max_frames` | `10000` | | +| `--norm_stats_max_frames` | `0` | `0` reads every state/action row | ## LoRA @@ -86,6 +95,56 @@ VRAM and runs on a 24GB GPU (e.g. RTX 4090). Full fine-tuning needs >70GB. Single-GPU default is on. Disable with `--lora false`. +## Cameras + +Pi0 / Pi0.5 always have three slots: `base_0_rgb`, `left_wrist_0_rgb`, +`right_wrist_0_rgb`. Keys are matched by name (`front` / `base` / `high` / +`exterior` → base, `wrist` → left wrist, `right` + `wrist` → right wrist). +Remaining keys fill empty slots. A missing slot is a zero image with +`image_mask=false`. + +At most three cameras are used. Pass `--cameras` or `--drop_cameras` when a +dataset has more. Example, three cameras then the same cache without the +front camera: + +```bash +docker run --rm --gpus all --shm-size=16g \ + -v /path/to/lerobot_dataset:/data/input:ro \ + -v /path/to/output:/data/output \ + -v /path/to/cache:/data/cache \ + ioaitech/train_openpi:pi05-cuda126 \ + --run_name pi05_my_task_3cam \ + --steps 30000 \ + --save_interval 5000 \ + --convert_dir /data/cache/my_task + +docker run --rm --gpus all --shm-size=16g \ + -v /path/to/lerobot_dataset:/data/input:ro \ + -v /path/to/output:/data/output \ + -v /path/to/cache:/data/cache \ + ioaitech/train_openpi:pi05-cuda126 \ + --run_name pi05_my_task_no_front \ + --drop_cameras front \ + --steps 30000 \ + --save_interval 5000 \ + --convert_dir /data/cache/my_task +``` + +v3 datasets are converted once into `--convert_dir`. A later run with the same +directory and the same videos reuses that tree. A camera subset is a second +directory of symlinks plus a rewritten `meta/info.json`, so LeRobot does not +decode dropped cameras. Put the cache on a real disk. The converter will not +write to `/tmp`. + +Read-only dataset mounts stay read-only. A v2 dataset that cannot be edited is +staged into the cache before metadata fixes. Each run writes +`/.run_manifest.json` with the camera map, prompt, and +whether actions were left absolute. + +The image's entrypoint uses the host `libcuda` when the host driver is newer +than the image's CUDA compat library. That is the usual case for newer +datacenter and workstation GPUs, including compute capability 12.0. + ## License Apache-2.0. See [THIRD_PARTY_NOTICES.md](THIRD_PARTY_NOTICES.md) for upstream diff --git a/README.zh-CN.md b/README.zh-CN.md index 666aff8..03ef15c 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -75,8 +75,17 @@ docker run --rm --gpus all --ipc=host \ | `--ema_decay` | 关闭 | 例如 `0.99` | | `--action_horizon` | `50` | | | `--num_workers` | `8` | | +| `--dataset_dir` | `/data/input` | 或 `$OPENPI_DATASET_DIR` | +| `--output_dir` | `/data/output` | 或 `$OPENPI_OUTPUT_DIR` | +| `--run_name` | `docker_train` | checkpoint 的上一级目录 | +| `--exp_name` | `train` | | +| `--convert_dir` | 在 `--output_dir` 下 | v3→v2 缓存,不要放在小的 tmpfs 上 | +| `--cameras` | 全部图像键 | 逗号分隔,只保留这些键 | +| `--drop_cameras` | | 键名或子串,例如 `front` | +| `--camera_map` | 按名字角色 | `base=键,left_wrist=键,right_wrist=键` | +| `--delta_joint_actions` | 关闭 | 关节用增量,最后一维(夹爪)保持绝对 | | `--norm_stats_workers` | `min(cpu, 64)` | | -| `--norm_stats_max_frames` | `10000` | | +| `--norm_stats_max_frames` | `0` | `0` 表示读完全部状态/动作 | ## LoRA @@ -85,6 +94,50 @@ docker run --rm --gpus all --ipc=host \ 单卡默认开启。关闭:`--lora false`。 +## 摄像头 + +Pi0 / Pi0.5 固定三个槽位:`base_0_rgb`、`left_wrist_0_rgb`、`right_wrist_0_rgb`。 +键名里的 `front` / `base` / `high` / `exterior` 进 base,`wrist` 进左手腕, +同时含 `right` 和 `wrist` 进右手腕。剩下的键按这个顺序填空槽。没有分到的槽位 +是全零图像,并且 `image_mask=false`。 + +最多使用三个摄像头。数据集更多时用 `--cameras` 或 `--drop_cameras`。下面先训 +三个摄像头,再用同一份缓存去掉 front: + +```bash +docker run --rm --gpus all --shm-size=16g \ + -v /path/to/lerobot_dataset:/data/input:ro \ + -v /path/to/output:/data/output \ + -v /path/to/cache:/data/cache \ + ioaitech/train_openpi:pi05-cuda126 \ + --run_name pi05_my_task_3cam \ + --steps 30000 \ + --save_interval 5000 \ + --convert_dir /data/cache/my_task + +docker run --rm --gpus all --shm-size=16g \ + -v /path/to/lerobot_dataset:/data/input:ro \ + -v /path/to/output:/data/output \ + -v /path/to/cache:/data/cache \ + ioaitech/train_openpi:pi05-cuda126 \ + --run_name pi05_my_task_no_front \ + --drop_cameras front \ + --steps 30000 \ + --save_interval 5000 \ + --convert_dir /data/cache/my_task +``` + +v3 数据集只会在 `--convert_dir` 里转换一次。同一目录、同一组视频再次运行会直接 +复用。去掉部分摄像头时,会在旁边做一个符号链接目录,并改写 `meta/info.json`, +LeRobot 就不会去解码被丢掉的摄像头。缓存要放在真实磁盘上,转换器不会再写到 `/tmp`。 + +只读挂载的数据不会被改写。无法写入的 v2 数据会先在缓存里做一个可写视图,再修补 +元数据。每次运行还会在 `/.run_manifest.json` 记下相机映射、 +任务文本,以及动作是否保持绝对量。 + +镜像入口会在宿主机驱动比镜像里的 CUDA compat 库更新时改用宿主机的 `libcuda`。 +较新的数据中心卡和工作站卡(包括 compute capability 12.0)属于这种情况。 + ## 许可证 Apache-2.0。上游 OpenPI 说明见 [THIRD_PARTY_NOTICES.md](THIRD_PARTY_NOTICES.md)。 diff --git a/lerobot_v3_compat.py b/lerobot_v3_compat.py index 005043c..3b756eb 100644 --- a/lerobot_v3_compat.py +++ b/lerobot_v3_compat.py @@ -7,6 +7,7 @@ from __future__ import annotations +import hashlib import json import logging import math @@ -14,7 +15,7 @@ import shutil import subprocess from collections import defaultdict -from collections.abc import Callable, Iterable +from collections.abc import Callable, Iterable, Sequence from pathlib import Path from typing import Any @@ -34,6 +35,17 @@ V2_VIDEO_PATH = "videos/chunk-{episode_chunk:03d}/{video_key}/episode_{episode_index:06d}.mp4" V3_DATA_PATH = "data/chunk-{chunk_index:03d}/file-{file_index:03d}.parquet" V3_VIDEO_PATH = "videos/{video_key}/chunk-{chunk_index:03d}/file-{file_index:03d}.mp4" +CONVERT_STAMP_NAME = ".convert_complete.json" +VIEW_STAMP_NAME = ".view_complete.json" +MODEL_IMAGE_SLOTS = ("base_0_rgb", "left_wrist_0_rgb", "right_wrist_0_rgb") +_SLOT_ALIASES = { + "base": "base_0_rgb", + "base_0_rgb": "base_0_rgb", + "left_wrist": "left_wrist_0_rgb", + "left_wrist_0_rgb": "left_wrist_0_rgb", + "right_wrist": "right_wrist_0_rgb", + "right_wrist_0_rgb": "right_wrist_0_rgb", +} TASK_DESCRIPTION_COLUMNS = ( "task", @@ -179,21 +191,18 @@ def load_tasks(root: Path) -> list[dict[str, Any]]: def tasks_have_text(root: Path) -> bool: - tasks_path = root / "meta" / "tasks.jsonl" - if not tasks_path.is_file(): + """True when ``tasks.parquet`` or ``tasks.jsonl`` contains a non-empty task string. + + A missing task file is not treated as text. ``load_tasks`` invents + ``perform the task`` in that case, which is not a dataset prompt. + """ + root = Path(root) + if not (root / "meta" / "tasks.parquet").is_file() and not (root / "meta" / "tasks.jsonl").is_file(): return False - with tasks_path.open(encoding="utf-8") as fh: - for line in fh: - line = line.strip() - if not line: - continue - try: - row = json.loads(line) - except json.JSONDecodeError: - continue - task = str(row.get("task") or "").strip() - if task: - return True + for row in load_tasks(root): + task = str(row.get("task") or "").strip() + if task: + return True return False @@ -578,9 +587,12 @@ def convert_info( ) -> dict[str, Any]: v2_info = dict(info) features: dict[str, Any] = {} + selected_videos = set(video_keys) for key, feat in (info.get("features") or {}).items(): if isinstance(feat, dict): copied = dict(feat) + if copied.get("dtype") == "video" and key not in selected_videos: + continue # Official v2.1 keeps fps only on video features (see GR00T convert_info). if copied.get("dtype") != "video": copied.pop("fps", None) @@ -654,6 +666,9 @@ def extract_video_segment(src: Path, dst: Path, start: float, end: float) -> Non duration = max(end - start, MIN_VIDEO_DURATION) dst.parent.mkdir(parents=True, exist_ok=True) + # `-ss` before `-i` with `-c copy` leaves the first PTS about one frame + # above 0 when combined with avoid_negative_ts. LeRobot's decoder then + # rejects the clip (tolerance 1e-4). setts moves the copied timestamps to 0. copy_cmd = [ "ffmpeg", "-hide_banner", @@ -667,8 +682,8 @@ def extract_video_segment(src: Path, dst: Path, start: float, end: float) -> Non f"{duration:.6f}", "-c", "copy", - "-avoid_negative_ts", - "1", + "-bsf:v", + "setts=pts=PTS-STARTPTS:dts=DTS-STARTPTS", "-y", str(dst), ] @@ -692,6 +707,8 @@ def extract_video_segment(src: Path, dst: Path, start: float, end: float) -> Non str(src), "-t", f"{duration:.6f}", + "-vf", + "setpts=PTS-STARTPTS", "-c:v", "libx264", "-preset", @@ -868,29 +885,355 @@ def assert_v2_local_files( ) +def directory_is_writable(path: Path) -> bool: + """Return whether ``path`` accepts a new file. + + Mode bits are not enough: a read-only bind mount still looks writable to + root via ``os.access``, and the write then fails with ``EROFS``. + """ + path = Path(path) + if not path.is_dir(): + return False + probe = path / f".write_probe_{os.getpid()}" + try: + with probe.open("x", encoding="utf-8") as fh: + fh.write("") + except OSError: + return False + try: + probe.unlink() + except OSError: + logger.warning("Could not remove write probe %s", probe) + return True + + +def conversion_stamp(src: Path, video_keys: Sequence[str], chunks_size: int) -> dict[str, Any]: + """Identity of a v3 source plus the video keys written into a v2 tree.""" + info_path = Path(src) / "meta" / "info.json" + stat = info_path.stat() + info = load_info(Path(src)) + return { + # 2: episode mp4s are stream-copied with PTS starting at 0. + "stamp_version": 2, + "info_size": stat.st_size, + "info_mtime_ns": stat.st_mtime_ns, + "codebase_version": info.get("codebase_version"), + "total_episodes": info.get("total_episodes"), + "total_frames": info.get("total_frames"), + "video_keys": list(video_keys), + "chunks_size": int(chunks_size), + } + + +def _read_json_stamp(path: Path) -> dict[str, Any] | None: + if not path.is_file(): + return None + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except json.JSONDecodeError: + return None + return payload if isinstance(payload, dict) else None + + +def _write_json_stamp(path: Path, payload: dict[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8") + + +def _resolve_video_keys(info: dict[str, Any], video_keys: Sequence[str] | None) -> list[str]: + available = video_keys_from_info(info) + if video_keys is None: + return available + selected = list(video_keys) + missing = [key for key in selected if key not in available] + if missing: + raise ValueError(f"Unknown video keys {missing}. Available: {available}") + return selected + + +def _reset_dir(path: Path) -> None: + if path.is_symlink() or path.is_file(): + path.unlink() + elif path.exists(): + shutil.rmtree(path) + path.mkdir(parents=True, exist_ok=True) + + +def _symlink_children(src_dir: Path, dest_dir: Path) -> None: + """Symlink each file under ``src_dir`` so replacing one link does not write the source.""" + if not src_dir.is_dir(): + return + for path in sorted(src_dir.rglob("*")): + if not path.is_file() and not path.is_symlink(): + continue + rel = path.relative_to(src_dir) + target = dest_dir / rel + target.parent.mkdir(parents=True, exist_ok=True) + if target.is_symlink() or target.exists(): + target.unlink() + os.symlink(path.resolve(), target) + + +def _filter_stats(stats: dict[str, Any], drop_keys: set[str]) -> dict[str, Any]: + return {key: value for key, value in stats.items() if key not in drop_keys} + + +def stage_writable_dataset(src: Path, dest: Path) -> Path: + """Return ``src`` when it can be edited, otherwise a writable tree of symlinks. + + Parquet files are linked one by one. A later atomic replace updates the + link in ``dest`` and leaves the source dataset untouched. + """ + src = Path(src) + meta = src / "meta" + data = src / "data" + writable = directory_is_writable(meta) and (not data.exists() or directory_is_writable(data)) + if writable: + return src + + dest = Path(dest) + marker = {"source": str(src.resolve())} + existing = _read_json_stamp(dest / "meta" / ".staged.json") + if existing == marker and (dest / "meta" / "info.json").is_file(): + logger.info("Reusing writable dataset view at %s", dest) + return dest + + logger.info("Dataset at %s is read-only; staging a writable view at %s", src, dest) + _reset_dir(dest) + if meta.is_dir(): + shutil.copytree(meta, dest / "meta", symlinks=True, dirs_exist_ok=True) + # copytree preserves a read-only source mode. The copy must accept metadata fixes. + for copied in [dest / "meta", *(dest / "meta").rglob("*")]: + if copied.is_symlink(): + continue + copied.chmod(copied.stat().st_mode | (0o700 if copied.is_dir() else 0o600)) + videos = src / "videos" + if videos.is_dir(): + os.symlink(videos.resolve(), dest / "videos") + _symlink_children(data, dest / "data") + _write_json_stamp(dest / "meta" / ".staged.json", marker) + return dest + + +def make_camera_view(src: Path, dest: Path, keep_video_keys: Sequence[str]) -> Path: + """Point ``dest`` at ``src`` while hiding video features that are not kept. + + LeRobot loads every video feature listed in ``info.json``. Dropping a camera + from the training config is not enough; the view's metadata must omit it. + """ + src = Path(src) + dest = Path(dest) + info = load_info(src) + available = video_keys_from_info(info) + selected = _resolve_video_keys(info, keep_video_keys) + if set(selected) == set(available): + return src + + stamp = { + "stamp_version": 1, + "source": str(src.resolve()), + "video_keys": selected, + } + existing = _read_json_stamp(dest / "meta" / VIEW_STAMP_NAME) + if existing == stamp and (dest / "meta" / "info.json").is_file(): + logger.info("Reusing camera view at %s (%s)", dest, selected) + return dest + + logger.info("Creating camera view at %s with videos %s", dest, selected) + _reset_dir(dest) + _symlink_children(src / "data", dest / "data") + videos = src / "videos" + if videos.is_dir(): + os.symlink(videos.resolve(), dest / "videos") + + meta_src = src / "meta" + meta_dest = dest / "meta" + meta_dest.mkdir(parents=True, exist_ok=True) + for name in ("tasks.jsonl", "tasks.parquet", "episodes.jsonl"): + source_file = meta_src / name + if source_file.is_file() or source_file.is_symlink(): + os.symlink(source_file.resolve(), meta_dest / name) + + drop_keys = set(available) - set(selected) + viewed = convert_info(info, [], selected, int(info.get("chunks_size") or V2_CHUNKS_SIZE)) + viewed["total_episodes"] = info.get("total_episodes", viewed.get("total_episodes")) + viewed["total_frames"] = info.get("total_frames") + viewed["total_videos"] = int(viewed.get("total_episodes") or 0) * len(selected) + with (meta_dest / "info.json").open("w", encoding="utf-8") as fh: + json.dump(viewed, fh, indent=2, ensure_ascii=False) + fh.write("\n") + + stats_path = meta_src / "stats.json" + if stats_path.is_file(): + raw_stats = json.loads(stats_path.read_text(encoding="utf-8")) + if isinstance(raw_stats, dict): + write_v21_stats_json(meta_dest / "stats.json", _filter_stats(raw_stats, drop_keys)) + + episodes_stats = meta_src / "episodes_stats.jsonl" + if episodes_stats.is_file(): + filtered_rows = [] + for row in load_episodes_stats_rows(episodes_stats): + stats = row.get("stats") + if isinstance(stats, dict): + row = dict(row) + row["stats"] = _filter_stats(stats, drop_keys) + filtered_rows.append(row) + _write_jsonl(meta_dest / "episodes_stats.jsonl", filtered_rows) + + _write_json_stamp(meta_dest / VIEW_STAMP_NAME, stamp) + return dest + + +def camera_view_dir(full_dir: Path, keep_video_keys: Sequence[str]) -> Path: + """Stable sibling directory for a camera subset of ``full_dir``.""" + slug = "_".join(key.split(".")[-1] for key in keep_video_keys) or "novideo" + digest = hashlib.sha256("\n".join(keep_video_keys).encode("utf-8")).hexdigest()[:8] + full_dir = Path(full_dir) + return full_dir.parent / f"{full_dir.name}__{slug}_{digest}" + + +def select_image_keys( + image_keys: Sequence[str], + *, + cameras: Sequence[str] | None = None, + drop_cameras: Sequence[str] | None = None, +) -> list[str]: + """Keep an explicit camera list, then drop keys or substrings such as ``front``.""" + available = list(image_keys) + if cameras: + missing = [key for key in cameras if key not in available] + if missing: + raise ValueError(f"Unknown camera keys {missing}. Available: {available}") + selected = [key for key in cameras if key in available] + else: + selected = list(available) + + def _dropped(key: str) -> bool: + if not drop_cameras: + return False + return any(key == token or token in key for token in drop_cameras) + + selected = [key for key in selected if not _dropped(key)] + if not selected: + raise ValueError(f"No image keys left after camera selection. Available: {available}") + if len(selected) > len(MODEL_IMAGE_SLOTS): + raise ValueError( + f"OpenPI has {len(MODEL_IMAGE_SLOTS)} image slots, got {selected}. " + "Pass --cameras or --drop_cameras to choose at most 3." + ) + return selected + + +def parse_camera_map(spec: str | None) -> dict[str, str] | None: + """Parse ``base=key,left_wrist=key,right_wrist=key`` into slot names.""" + if spec is None or not spec.strip(): + return None + mapping: dict[str, str] = {} + for part in spec.split(","): + item = part.strip() + if not item: + continue + if "=" not in item: + raise ValueError(f"camera_map entry {item!r} must look like base=observation.images.front") + slot_name, key = item.split("=", 1) + slot = _SLOT_ALIASES.get(slot_name.strip()) + if slot is None: + raise ValueError( + f"Unknown camera slot {slot_name!r}. Expected base, left_wrist, or right_wrist." + ) + if slot in mapping: + raise ValueError(f"Camera slot {slot_name} was given twice") + mapping[slot] = key.strip() + return mapping + + +def _preferred_slot(key: str) -> str | None: + lower = key.lower() + if "wrist" in lower and "right" in lower: + return "right_wrist_0_rgb" + if "wrist" in lower: + return "left_wrist_0_rgb" + if any(token in lower for token in ("front", "base", "high", "exterior")): + return "base_0_rgb" + return None + + +def assign_camera_slots( + image_keys: Sequence[str], + camera_map: dict[str, str] | None = None, +) -> dict[str, str]: + """Map dataset image keys onto ``base_0_rgb`` / wrist slots. + + Keys that match a role fill that slot. Remaining keys fill empty slots in + slot order. Unfilled slots are omitted so the caller can mask them. + """ + keys = list(image_keys) + if camera_map is not None: + unknown_slots = [slot for slot in camera_map if slot not in MODEL_IMAGE_SLOTS] + if unknown_slots: + raise ValueError(f"Unknown camera slots {unknown_slots}") + missing = [key for key in camera_map.values() if key not in keys] + if missing: + raise ValueError(f"camera_map keys {missing} are not in the selected cameras {keys}") + return dict(camera_map) + + remaining = list(keys) + mapping: dict[str, str] = {} + for key in list(remaining): + slot = _preferred_slot(key) + if slot is None or slot in mapping: + continue + mapping[slot] = key + remaining.remove(key) + for slot in MODEL_IMAGE_SLOTS: + if slot not in mapping and remaining: + mapping[slot] = remaining.pop(0) + return mapping + + def convert_v3_to_v2( src_dir: Path, dest_dir: Path | None = None, *, chunks_size: int = V2_CHUNKS_SIZE, extract_video: ExtractVideoFn = extract_video_segment, + video_keys: Sequence[str] | None = None, ) -> Path: - """Convert a local v3.0 dataset to a v2.1 tree and validate required files.""" + """Convert a local v3.0 dataset to a v2.1 tree and validate required files. + + ``dest_dir`` is required. A previous default of ``/tmp`` is unsafe on hosts + whose ``/tmp`` is a small tmpfs. When ``dest_dir`` already contains a + matching ``meta/.convert_complete.json``, the tree is reused and + ``extract_video`` is not called. + """ src = Path(src_dir) - dest = Path(dest_dir) if dest_dir is not None else Path("/tmp/lerobot_v2_compat") - if dest.exists(): - shutil.rmtree(dest) - dest.mkdir(parents=True, exist_ok=True) + if dest_dir is None: + raise ValueError( + "dest_dir is required. Refusing to write a converted dataset to /tmp; " + "pass an explicit directory on a real filesystem." + ) + dest = Path(dest_dir) info = load_info(src) version = str(info.get("codebase_version", "")).lower() if not (version.startswith("v3") or version.startswith("3")): raise ValueError(f"Expected a v3 dataset, got codebase_version={info.get('codebase_version')!r}") + selected_videos = _resolve_video_keys(info, video_keys) + stamp = conversion_stamp(src, selected_videos, chunks_size) + existing = _read_json_stamp(dest / "meta" / CONVERT_STAMP_NAME) + if existing == stamp: + logger.info("Reusing converted dataset at %s", dest) + return dest + if dest.exists(): + logger.info("Reconverting %s; existing tree at %s does not match this source", src, dest) + _reset_dir(dest) + else: + dest.mkdir(parents=True, exist_ok=True) + episode_records = load_episode_records(src) - video_keys = video_keys_from_info(info) tasks = load_tasks(src) - v2_info = convert_info(info, episode_records, video_keys, chunks_size) + v2_info = convert_info(info, episode_records, selected_videos, chunks_size) meta_dest = dest / "meta" meta_dest.mkdir(parents=True, exist_ok=True) @@ -898,6 +1241,9 @@ def convert_v3_to_v2( json.dump(v2_info, fh, indent=2, ensure_ascii=False) global_stats = load_sanitized_stats_json(src / "meta" / "stats.json") + dropped_videos = set(video_keys_from_info(info)) - set(selected_videos) + if global_stats and dropped_videos: + global_stats = _filter_stats(global_stats, dropped_videos) if global_stats: write_v21_stats_json(meta_dest / "stats.json", global_stats) @@ -907,7 +1253,7 @@ def convert_v3_to_v2( src, dest, episode_records, - video_keys, + selected_videos, chunks_size, extract_video=extract_video, ) @@ -918,16 +1264,27 @@ def convert_v3_to_v2( chunks_size=chunks_size, global_stats=global_stats, ) + if dropped_videos: + stats_path = meta_dest / "episodes_stats.jsonl" + filtered_rows = [] + for row in load_episodes_stats_rows(stats_path): + stats = row.get("stats") + if isinstance(stats, dict): + row = dict(row) + row["stats"] = _filter_stats(stats, dropped_videos) + filtered_rows.append(row) + _write_jsonl(stats_path, filtered_rows) episode_indices = [_as_int(rec["episode_index"]) for rec in episode_records] assert_v2_local_files(dest, v2_info, episode_indices, chunks_size=chunks_size) assert_v21_episodes_stats(dest) + _write_json_stamp(meta_dest / CONVERT_STAMP_NAME, stamp) logger.info("Converted v3 dataset to v2 layout at %s", dest) logger.info( " Episodes: %s, video keys: %s, v2 chunks_size: %s", len(episode_indices), - video_keys, + selected_videos, chunks_size, ) return dest diff --git a/test_lerobot_v3_compat.py b/test_lerobot_v3_compat.py index af83546..90a1c93 100644 --- a/test_lerobot_v3_compat.py +++ b/test_lerobot_v3_compat.py @@ -3,6 +3,8 @@ from __future__ import annotations import json +import shutil +import subprocess from pathlib import Path import pyarrow as pa @@ -15,13 +17,18 @@ from lerobot_v3_compat import _numeric_feature_names from lerobot_v3_compat import assert_v2_local_files from lerobot_v3_compat import assert_v21_episode_stats_rows +from lerobot_v3_compat import assign_camera_slots from lerobot_v3_compat import convert_v3_to_v2 +from lerobot_v3_compat import extract_video_segment from lerobot_v3_compat import episodes_stats_compatible_with_v21 from lerobot_v3_compat import expected_v2_paths from lerobot_v3_compat import generate_episodes_stats_from_parquet from lerobot_v3_compat import load_sanitized_stats_json from lerobot_v3_compat import load_tasks +from lerobot_v3_compat import make_camera_view from lerobot_v3_compat import sanitize_episode_stats +from lerobot_v3_compat import select_image_keys +from lerobot_v3_compat import stage_writable_dataset from lerobot_v3_compat import stats_from_episode_record from lerobot_v3_compat import tasks_have_text from lerobot_v3_compat import unflatten_dict @@ -193,6 +200,8 @@ def test_load_tasks_index_level_column(tmp_path: Path) -> None: ) tasks = load_tasks(tmp_path) assert tasks == [{"task_index": 3, "task": "open the drawer"}] + assert tasks_have_text(tmp_path) + assert not tasks_have_text(tmp_path / "empty") def test_convert_packed_v3_writes_v2_chunks(tmp_path: Path) -> None: @@ -547,3 +556,160 @@ def test_convert_official_v21_from_real_v3_stats_layout(tmp_path: Path) -> None: assert first[CAM_B]["count"] == [2] assert "observation.base_move" in first assert_v21_episode_stats_rows(rows) + + +FRONT = "observation.images.front" +TOP = "observation.images.top" +WRIST = "observation.images.wrist" + + +def test_extract_video_segment_starts_at_zero(tmp_path: Path) -> None: + if shutil.which("ffmpeg") is None or shutil.which("ffprobe") is None: + pytest.skip("ffmpeg is not installed") + src = tmp_path / "src.mp4" + subprocess.run( + [ + "ffmpeg", + "-y", + "-hide_banner", + "-loglevel", + "error", + "-f", + "lavfi", + "-i", + "testsrc=size=64x64:rate=30", + "-t", + "2", + "-c:v", + "libx264", + "-pix_fmt", + "yuv420p", + "-g", + "2", + str(src), + ], + check=True, + ) + dst = tmp_path / "cut.mp4" + extract_video_segment(src, dst, 0.4, 1.3) + start = subprocess.check_output( + [ + "ffprobe", + "-v", + "error", + "-select_streams", + "v:0", + "-show_entries", + "stream=start_time", + "-of", + "default=nw=1:nk=1", + str(dst), + ], + text=True, + ) + first = subprocess.check_output( + [ + "ffprobe", + "-v", + "error", + "-select_streams", + "v:0", + "-show_frames", + "-show_entries", + "frame=best_effort_timestamp_time", + "-of", + "csv=p=0", + str(dst), + ], + text=True, + ).splitlines()[0] + assert float(start.strip()) < 1e-3 + assert abs(float(first.split(",")[0])) < 1e-3 + + +def test_camera_slots_follow_role_names() -> None: + three = [FRONT, TOP, WRIST] + slots = assign_camera_slots(three) + assert slots == { + "base_0_rgb": FRONT, + "left_wrist_0_rgb": WRIST, + "right_wrist_0_rgb": TOP, + } + kept = select_image_keys(three, drop_cameras=["front"]) + assert kept == [TOP, WRIST] + two = assign_camera_slots(kept) + assert two == {"base_0_rgb": TOP, "left_wrist_0_rgb": WRIST} + assert "right_wrist_0_rgb" not in two + + +def test_convert_reuses_completed_tree_and_can_limit_cameras(tmp_path: Path) -> None: + src = build_packed_v3(tmp_path / "v3") + calls = {"n": 0} + + def _counting(src_path: Path, dst: Path, start: float, end: float) -> None: + calls["n"] += 1 + _fake_extract(src_path, dst, start, end) + + dest = convert_v3_to_v2(src, tmp_path / "v2", chunks_size=2, extract_video=_counting) + assert calls["n"] > 0 + seen = calls["n"] + + def _forbidden(*_args, **_kwargs): + raise AssertionError("cached conversion must not extract video again") + + again = convert_v3_to_v2(src, dest, chunks_size=2, extract_video=_forbidden) + assert again == dest + assert calls["n"] == seen + + with pytest.raises(ValueError, match="/tmp"): + convert_v3_to_v2(src) + + limited = convert_v3_to_v2( + src, + tmp_path / "wrist-only", + chunks_size=2, + extract_video=_fake_extract, + video_keys=[CAM_B], + ) + info = json.loads((limited / "meta" / "info.json").read_text(encoding="utf-8")) + assert CAM_B in info["features"] + assert CAM_A not in info["features"] + assert not any(path.name.startswith("episode_") and CAM_A in str(path) for path in (limited / "videos").rglob("*")) + + +def test_camera_view_hides_dropped_camera_and_reuses(tmp_path: Path) -> None: + src = build_packed_v3(tmp_path / "v3") + full = convert_v3_to_v2(src, tmp_path / "v2", chunks_size=2, extract_video=_fake_extract) + view = make_camera_view(full, tmp_path / "view", [CAM_B]) + info = json.loads((view / "meta" / "info.json").read_text(encoding="utf-8")) + assert CAM_A not in info["features"] + assert CAM_B in info["features"] + episode = view / "data" / "chunk-000" / "episode_000000.parquet" + assert episode.is_file() + assert episode.is_symlink() + reused = make_camera_view(full, view, [CAM_B]) + assert reused == view + assert make_camera_view(full, tmp_path / "unused", [CAM_A, CAM_B]) == full + + +def test_stage_writable_dataset_does_not_touch_readonly_source(tmp_path: Path) -> None: + src = tmp_path / "src" + meta = src / "meta" + data = src / "data" / "chunk-000" + data.mkdir(parents=True) + (meta).mkdir(parents=True) + (meta / "info.json").write_text("{}\n", encoding="utf-8") + _write_table(data / "episode_000000.parquet", pa.table({"action": [[0.0, 1.0]]})) + before = (data / "episode_000000.parquet").read_bytes() + for path in (src, meta, src / "data", data): + path.chmod(0o555) + try: + staged = stage_writable_dataset(src, tmp_path / "work") + assert staged != src + (staged / "meta" / "episodes_stats.jsonl").write_text("{}\n", encoding="utf-8") + assert not (meta / "episodes_stats.jsonl").exists() + assert (data / "episode_000000.parquet").read_bytes() == before + assert stage_writable_dataset(src, tmp_path / "work") == staged + finally: + for path in (data, src / "data", meta, src): + path.chmod(0o755) diff --git a/test_train_lerobot.py b/test_train_lerobot.py new file mode 100644 index 0000000..37c817b --- /dev/null +++ b/test_train_lerobot.py @@ -0,0 +1,44 @@ +"""Tests for camera wiring and norm-stat sampling that do not need JAX.""" + +from pathlib import Path + +import numpy as np + +from train_lerobot import GenericLeRobotInputs +from train_lerobot import select_parquet_files + + +def test_missing_camera_slot_is_masked() -> None: + transform = GenericLeRobotInputs( + camera_slots=( + ("base_0_rgb", "observation.images.top"), + ("left_wrist_0_rgb", "observation.images.wrist"), + ), + state_key="observation.state", + action_key="action", + ) + top = np.zeros((4, 5, 3), dtype=np.uint8) + wrist = np.ones((4, 5, 3), dtype=np.uint8) + out = transform( + { + "observation.images.top": top, + "observation.images.wrist": wrist, + "observation.state": np.zeros(6, dtype=np.float32), + "action": np.ones(6, dtype=np.float32), + } + ) + assert bool(out["image_mask"]["base_0_rgb"]) + assert bool(out["image_mask"]["left_wrist_0_rgb"]) + assert not bool(out["image_mask"]["right_wrist_0_rgb"]) + assert out["image"]["right_wrist_0_rgb"].shape == (4, 5, 3) + assert out["image"]["right_wrist_0_rgb"].sum() == 0 + assert np.array_equal(out["actions"], np.ones(6, dtype=np.float32)) + + +def test_norm_stats_frame_cap_defaults_to_every_file() -> None: + files = [Path(f"file-{idx:03d}.parquet") for idx in range(10)] + assert select_parquet_files(files, None, 100) == files + assert select_parquet_files(files, 0, 100) == files + first = select_parquet_files(files, 100, 100, seed=0) + assert len(first) == 2 + assert select_parquet_files(files, 100, 100, seed=0) == first diff --git a/train_lerobot.py b/train_lerobot.py index 66012a5..4c60076 100644 --- a/train_lerobot.py +++ b/train_lerobot.py @@ -10,10 +10,12 @@ import argparse import dataclasses +import hashlib import json import logging import os import pathlib +import random import shutil import sys @@ -33,13 +35,22 @@ import numpy as np +from lerobot_v3_compat import assign_camera_slots +from lerobot_v3_compat import camera_view_dir +from lerobot_v3_compat import conversion_stamp from lerobot_v3_compat import convert_v3_to_v2 from lerobot_v3_compat import episodes_stats_compatible_with_v21 from lerobot_v3_compat import generate_episodes_stats_from_parquet +from lerobot_v3_compat import load_tasks +from lerobot_v3_compat import make_camera_view +from lerobot_v3_compat import parse_camera_map +from lerobot_v3_compat import select_image_keys +from lerobot_v3_compat import stage_writable_dataset from lerobot_v3_compat import tasks_have_text +from lerobot_v3_compat import video_keys_from_info -DATASET_DIR = pathlib.Path("/data/input") -OUTPUT_DIR = pathlib.Path("/data/output") +DEFAULT_DATASET_DIR = pathlib.Path(os.environ.get("OPENPI_DATASET_DIR", "/data/input")) +DEFAULT_OUTPUT_DIR = pathlib.Path(os.environ.get("OPENPI_OUTPUT_DIR", "/data/output")) logger = logging.getLogger("train_lerobot") @@ -48,8 +59,34 @@ # CLI # --------------------------------------------------------------------------- +def _csv_list(value: str | None) -> list[str] | None: + if value is None or not value.strip(): + return None + items = [part.strip() for part in value.split(",") if part.strip()] + return items or None + + def parse_args(): parser = argparse.ArgumentParser(description="Train OpenPI on a mounted LeRobot dataset") + parser.add_argument("--dataset_dir", type=pathlib.Path, default=DEFAULT_DATASET_DIR, + help="LeRobot dataset root (default: $OPENPI_DATASET_DIR or /data/input)") + parser.add_argument("--output_dir", type=pathlib.Path, default=DEFAULT_OUTPUT_DIR, + help="Checkpoint root (default: $OPENPI_OUTPUT_DIR or /data/output)") + parser.add_argument("--run_name", type=str, default="docker_train", + help="TrainConfig name. Checkpoints land in //") + parser.add_argument("--exp_name", type=str, default="train") + parser.add_argument("--convert_dir", type=pathlib.Path, default=None, + help="Writable v3->v2 cache. Default: /.v21_cache/. " + "Do not point this at a small tmpfs.") + parser.add_argument("--cameras", type=str, default=None, + help="Comma-separated image keys to keep, e.g. observation.images.top,observation.images.wrist") + parser.add_argument("--drop_cameras", type=str, default=None, + help="Comma-separated image keys or substrings to drop, e.g. front") + parser.add_argument("--camera_map", type=str, default=None, + help="Explicit slots: base=key,left_wrist=key,right_wrist=key") + parser.add_argument("--delta_joint_actions", action="store_true", + help="Train joint dims as deltas and keep the last dim (gripper) absolute. " + "Default keeps the dataset action values unchanged.") parser.add_argument("--batch_size", type=int, default=1) parser.add_argument("--steps", type=int, default=1000) parser.add_argument("--gpus", type=str, default="all", @@ -72,10 +109,10 @@ def parse_args(): parser.add_argument("--norm_stats_workers", type=int, default=_default_workers, help=f"Parallel workers for fast norm-stats parquet reading " f"(default: auto = min(cpu_count, 64), currently {_default_workers})") - parser.add_argument("--norm_stats_max_frames", type=int, default=10000, - help="Limit frames sampled for norm-stats slow-path fallback " - "(default: auto-cap at 200,000 when dataset > 500,000 frames). " - "The fast parquet path always reads all frames regardless.") + parser.add_argument("--norm_stats_max_frames", type=int, default=0, + help="Fast-path frame cap. 0 (default) reads every state/action row. " + "A positive value samples about that many frames from a seeded file subset. " + "Slow-path fallback still auto-caps very large datasets at 200,000 frames.") return parser.parse_args() @@ -306,9 +343,13 @@ def analyze_features(info: dict) -> dict: @dataclasses.dataclass(frozen=True) class GenericLeRobotInputs: - """Map any LeRobot schema to the three-image format expected by OpenPI models.""" + """Map dataset image keys onto OpenPI's three camera slots. + + ``camera_slots`` is ``(slot_name, dataset_key)`` pairs. Slots that are + absent are zeros with ``image_mask=False``. + """ - image_keys: tuple + camera_slots: tuple state_key: str action_key: str = "action" @@ -323,13 +364,15 @@ def _parse_image(img): def __call__(self, data: dict) -> dict: model_keys = ["base_0_rgb", "left_wrist_0_rgb", "right_wrist_0_rgb"] + slot_to_key = dict(self.camera_slots) images: dict[str, np.ndarray] = {} image_masks: dict[str, np.bool_] = {} ref_shape = None - for i, mkey in enumerate(model_keys): - if i < len(self.image_keys) and self.image_keys[i] in data: - img = self._parse_image(data[self.image_keys[i]]) + for mkey in model_keys: + dataset_key = slot_to_key.get(mkey) + if dataset_key is not None and dataset_key in data: + img = self._parse_image(data[dataset_key]) images[mkey] = img image_masks[mkey] = np.True_ if ref_shape is None: @@ -371,6 +414,33 @@ def __call__(self, data: dict) -> dict: # Normalization statistics # --------------------------------------------------------------------------- +def _rows_per_parquet(path: pathlib.Path, column: str) -> int: + import pyarrow.parquet as pq + + try: + return int(pq.read_table(path, columns=[column]).num_rows) + except Exception: + return 0 + + +def select_parquet_files( + files: list[pathlib.Path], + max_frames: int | None, + rows_per_file: int, + seed: int = 0, +) -> list[pathlib.Path]: + """Return every file when ``max_frames`` is unset, else a seeded subset.""" + ordered = list(files) + if max_frames is None or max_frames <= 0 or rows_per_file <= 0 or not ordered: + return ordered + n_files = max(1, max_frames // max(1, rows_per_file) + 1) + if n_files >= len(ordered): + return ordered + rng = random.Random(seed) + rng.shuffle(ordered) + return ordered[:n_files] + + def _compute_norm_stats_fast( config, dataset_dir: pathlib.Path, @@ -385,12 +455,12 @@ def _compute_norm_stats_fast( RunningStats.update() reshapes input to (-1, last_dim), so per-frame parquet data [N, feat_dim] produces statistically equivalent results to the full pipeline's - [N, action_horizon, feat_dim] batches. + [N, action_horizon, feat_dim] batches. ``max_frames is None`` reads every row. + A positive cap samples a seeded subset of files. Returns True on success, False if the fast path cannot be used. """ import concurrent.futures - import random import pyarrow.parquet as pq import openpi.shared.normalize as normalize @@ -420,16 +490,13 @@ def _compute_norm_stats_fast( logger.warning(f"Failed to read parquet schema: {e}; skipping fast norm-stats path.") return False - files_to_process: list[pathlib.Path] = list(parquet_files) - if max_frames is not None: - # Estimate average frames per file from first file, then take a random subset - try: - n_sample = pq.read_table(parquet_files[0], columns=[state_col]).num_rows - n_files_needed = max(1, max_frames // max(1, n_sample) + 1) - except Exception: - n_files_needed = len(files_to_process) - random.shuffle(files_to_process) - files_to_process = files_to_process[:n_files_needed] + files_to_process = list(parquet_files) + if max_frames is not None and max_frames > 0: + files_to_process = select_parquet_files( + files_to_process, + max_frames, + _rows_per_parquet(parquet_files[0], state_col), + ) logger.info( f"Fast norm-stats: {len(files_to_process)}/{len(parquet_files)} parquet files, " @@ -613,51 +680,114 @@ def __call__(self, x): # Main # --------------------------------------------------------------------------- -def main(): - args = parse_args() - - logging.basicConfig( - level=logging.INFO, - format="%(asctime)s.%(msecs)03d [%(levelname)s] %(message)s", - datefmt="%H:%M:%S", - ) - - # ---- validate dataset mount ---- - if not DATASET_DIR.exists(): - logger.error( - "Dataset not found at /data/input. " - "Mount your dataset: docker run -v /path/to/dataset:/data/input …" +def _cache_dir_for(args, dataset_dir: pathlib.Path, info: dict) -> pathlib.Path: + if args.convert_dir is not None: + return pathlib.Path(args.convert_dir) + stamp = conversion_stamp(dataset_dir, video_keys_from_info(info), 1000) + digest = hashlib.sha256( + json.dumps(stamp, sort_keys=True, default=str).encode("utf-8") + ).hexdigest()[:16] + return pathlib.Path(args.output_dir) / ".v21_cache" / digest + + +def prepare_dataset(args) -> dict: + """Resolve the on-disk tree training should read, including camera filters.""" + dataset_dir = pathlib.Path(args.dataset_dir) + output_dir = pathlib.Path(args.output_dir) + if not (dataset_dir / "meta" / "info.json").is_file(): + raise FileNotFoundError( + f"No LeRobot dataset at {dataset_dir}. " + "Expected meta/info.json. Mount with -v /path/to/dataset:/data/input" ) - sys.exit(1) - - OUTPUT_DIR.mkdir(parents=True, exist_ok=True) - - # ---- discover dataset ---- - info = discover_dataset(DATASET_DIR) - dataset_version = info.get("codebase_version", "unknown") - expected_version = os.environ.get("LEROBOT_DATASET_VERSION", "") - logger.info(f"Dataset codebase_version : {dataset_version}") - logger.info(f"Image built for LeRobot : {expected_version or 'auto'}") + output_dir.mkdir(parents=True, exist_ok=True) + + info = discover_dataset(dataset_dir) + raw_images = analyze_features(info)["image_keys"] + selected = select_image_keys( + raw_images, + cameras=_csv_list(args.cameras), + drop_cameras=_csv_list(args.drop_cameras), + ) + dropped = [key for key in raw_images if key not in selected] + cache_dir = _cache_dir_for(args, dataset_dir, info) + version = str(info.get("codebase_version", "")) + logger.info(f"Dataset codebase_version : {version or 'unknown'}") + logger.info(f"Image built for LeRobot : {os.environ.get('LEROBOT_DATASET_VERSION') or 'auto'}") - # ---- v3 -> v2 conversion if needed ---- - effective_dir = DATASET_DIR - if dataset_version.startswith("v3") or dataset_version.startswith("3"): + if version.startswith("v3") or version.startswith("3"): logger.info("Detected v3.0 dataset – converting to v2.1 layout for compatibility …") - effective_dir = convert_v3_to_v2(DATASET_DIR) - info = discover_dataset(effective_dir) - logger.info(f"Converted dataset version: {info.get('codebase_version', 'unknown')}") + full_dir = convert_v3_to_v2(dataset_dir, cache_dir) + logger.info(f"Converted dataset version: {discover_dataset(full_dir).get('codebase_version')}") + else: + full_dir = stage_writable_dataset(dataset_dir, cache_dir / "writable") - effective_version = str(info.get("codebase_version", "")).lower() - if effective_version.startswith("v2"): + full_videos = set(video_keys_from_info(discover_dataset(full_dir))) + if set(selected) != full_videos: + effective_dir = make_camera_view(full_dir, camera_view_dir(cache_dir, selected), selected) + else: + effective_dir = full_dir + + info = discover_dataset(effective_dir) + if str(info.get("codebase_version", "")).lower().startswith("v2"): ensure_v21_episodes_stats(effective_dir, info) normalize_parquet_hf_metadata(effective_dir) schema = analyze_features(info) + slots = assign_camera_slots(schema["image_keys"], parse_camera_map(args.camera_map)) has_task_text = bool(schema["has_tasks"] and tasks_have_text(effective_dir)) + task_texts = [str(row["task"]) for row in load_tasks(effective_dir)] if has_task_text else [] logger.info(f" image keys : {schema['image_keys']}") + logger.info(f" dropped : {dropped or 'none'}") + logger.info(f" camera slots: {slots}") logger.info(f" state : {schema['state_key']} dim={schema['state_dim']}") logger.info(f" action : {schema['action_key']} dim={schema['action_dim']}") - logger.info(f" has tasks : {schema['has_tasks']} task text: {has_task_text}") + logger.info(f" has tasks : {schema['has_tasks']} task text: {task_texts or has_task_text}") + return { + "dataset_dir": dataset_dir, + "output_dir": output_dir, + "effective_dir": effective_dir, + "schema": schema, + "slots": slots, + "dropped": dropped, + "has_task_text": has_task_text, + "task_texts": task_texts, + } + + +def write_run_manifest(path: pathlib.Path, payload: dict) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8") + + +def copy_manifest_into_steps(checkpoint_dir: pathlib.Path) -> None: + manifest = checkpoint_dir / "run_manifest.json" + if not manifest.is_file(): + return + for child in checkpoint_dir.iterdir(): + if child.is_dir() and child.name.isdigit(): + shutil.copyfile(manifest, child / "run_manifest.json") + + +def main(): + args = parse_args() + + logging.basicConfig( + level=logging.INFO, + format="%(asctime)s.%(msecs)03d [%(levelname)s] %(message)s", + datefmt="%H:%M:%S", + ) + + try: + prepared = prepare_dataset(args) + except (FileNotFoundError, ValueError) as exc: + logger.error("%s", exc) + sys.exit(1) + + output_dir = prepared["output_dir"] + effective_dir = prepared["effective_dir"] + schema = prepared["schema"] + slots = prepared["slots"] + has_task_text = prepared["has_task_text"] # ---- link dataset into LeRobot cache ---- repo_id = setup_dataset_link(effective_dir) @@ -765,11 +895,24 @@ def main(): # ---- transforms ---- generic_inputs = GenericLeRobotInputs( - image_keys=tuple(schema["image_keys"]), + camera_slots=tuple(slots.items()), state_key=schema["state_key"], action_key=schema["action_key"], ) generic_outputs = GenericLeRobotOutputs(action_dim=schema["action_dim"]) + input_transforms = [generic_inputs] + output_transforms = [generic_outputs] + action_mode = "absolute" + delta_mask: tuple | None = None + if args.delta_joint_actions: + if schema["action_dim"] < 1: + logger.error("delta_joint_actions requires a positive action dimension") + sys.exit(1) + delta_mask = _transforms.make_bool_mask(max(schema["action_dim"] - 1, 0), -1) + input_transforms.append(_transforms.DeltaActions(delta_mask)) + output_transforms = [_transforms.AbsoluteActions(delta_mask), *output_transforms] + action_mode = "delta_joints" + logger.info(f"delta joint actions enabled, mask={delta_mask}") default_prompt = args.prompt or os.environ.get("DEFAULT_PROMPT", "perform the task") @@ -777,8 +920,8 @@ def main(): repo_id=repo_id, assets=_config.AssetsConfig(asset_id="training_dataset"), data_transforms=lambda _mc: _transforms.Group( - inputs=[generic_inputs], - outputs=[generic_outputs], + inputs=input_transforms, + outputs=output_transforms, ), model_transforms=_config.ModelTransformFactory( default_prompt=None if has_task_text else default_prompt, @@ -798,19 +941,21 @@ def main(): ) # ---- assemble TrainConfig ---- + peak_lr = args.learning_rate or 2.5e-5 config = _config.TrainConfig( - name="docker_train", + name=args.run_name, model=model_config, data=data_factory, weight_loader=weight_loaders.CheckpointWeightLoader(weight_path), batch_size=args.batch_size, num_train_steps=args.steps, - checkpoint_base_dir=str(OUTPUT_DIR), + checkpoint_base_dir=str(output_dir), assets_base_dir="/workspace/assets", - exp_name="train", + exp_name=args.exp_name, overwrite=True, wandb_enabled=False, save_interval=args.save_interval, + keep_period=5000, lr_schedule=lr_schedule, num_workers=args.num_workers, fsdp_devices=fsdp_devices, @@ -822,19 +967,51 @@ def main(): logger.info(f"checkpoint_dir = {config.checkpoint_dir}") logger.info(f"weight source = {weight_path}") + norm_max_frames = args.norm_stats_max_frames if args.norm_stats_max_frames > 0 else None + manifest = { + "model_type": model_type, + "lora": use_lora, + "weight_path": weight_path, + "openpi_git_ref": os.environ.get("OPENPI_GIT_REF", ""), + "run_name": args.run_name, + "exp_name": args.exp_name, + "cameras": list(schema["image_keys"]), + "camera_slots": slots, + "dropped_cameras": prepared["dropped"], + "prompt_from_task": has_task_text, + "task_texts": prepared["task_texts"], + "default_prompt": None if has_task_text else default_prompt, + "action_horizon": args.action_horizon, + "action_mode": action_mode, + "delta_mask": list(delta_mask) if delta_mask is not None else None, + "batch_size": args.batch_size, + "steps": args.steps, + "learning_rate": peak_lr, + "save_interval": args.save_interval, + "norm_stats_max_frames": args.norm_stats_max_frames, + "dataset_dir": str(prepared["dataset_dir"]), + "effective_dataset_dir": str(effective_dir), + } + # Keep the manifest beside the checkpoint directory. Training with + # overwrite=True deletes checkpoint_dir itself before the first step. + manifest_sidecar = config.checkpoint_dir.parent / f"{config.checkpoint_dir.name}.run_manifest.json" + write_run_manifest(manifest_sidecar, manifest) + # ---- step 1: normalization statistics ---- compute_norm_stats( config, dataset_dir=effective_dir, schema=schema, - max_frames=args.norm_stats_max_frames, + max_frames=norm_max_frames, num_workers=args.norm_stats_workers, ) # ---- step 2: train ---- logger.info("Starting training …") train_main(config) - logger.info(f"Training complete. Checkpoints → {OUTPUT_DIR}") + write_run_manifest(config.checkpoint_dir / "run_manifest.json", manifest) + copy_manifest_into_steps(config.checkpoint_dir) + logger.info(f"Training complete. Checkpoints → {config.checkpoint_dir}") if __name__ == "__main__": From fd5de5df95dfca6deb73ee5879f47544d648e408 Mon Sep 17 00:00:00 2001 From: joaner <1726541+joaner@users.noreply.github.com> Date: Tue, 22 Sep 2026 16:59:48 +0800 Subject: [PATCH 2/7] Keep every saved checkpoint and allow resuming a run. --- README.md | 2 ++ README.zh-CN.md | 2 ++ test_train_lerobot.py | 10 ++++++++++ train_lerobot.py | 30 ++++++++++++++++++++++++++++-- 4 files changed, 42 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index d0edefa..10f10eb 100644 --- a/README.md +++ b/README.md @@ -70,6 +70,8 @@ docker run --rm --gpus all --ipc=host \ | `--gpus` | `all` | or `0,1` | | `--prompt` | | used only when the dataset has no task text | | `--save_interval` | `500` | | +| `--keep_period` | `--save_interval` | steps divisible by this are never pruned; `0` keeps only the newest | +| `--resume` | off | continue from the newest checkpoint in the run directory | | `--learning_rate` | `2.5e-5` | | | `--fsdp_devices` | `auto` | GPU count when >= 2 | | `--lora` | `auto` | `true` / `false` | diff --git a/README.zh-CN.md b/README.zh-CN.md index 03ef15c..132f1ad 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -69,6 +69,8 @@ docker run --rm --gpus all --ipc=host \ | `--gpus` | `all` | 或 `0,1` | | `--prompt` | | 仅当数据集没有 task 文本时使用 | | `--save_interval` | `500` | | +| `--keep_period` | 等于 `--save_interval` | 步数能被它整除的 checkpoint 不会被清理;`0` 表示只保留最新一个 | +| `--resume` | 关闭 | 从 run 目录里最新的 checkpoint 继续 | | `--learning_rate` | `2.5e-5` | | | `--fsdp_devices` | `auto` | GPU >= 2 时等于卡数 | | `--lora` | `auto` | `true` / `false` | diff --git a/test_train_lerobot.py b/test_train_lerobot.py index 37c817b..f397801 100644 --- a/test_train_lerobot.py +++ b/test_train_lerobot.py @@ -5,6 +5,7 @@ import numpy as np from train_lerobot import GenericLeRobotInputs +from train_lerobot import resolve_keep_period from train_lerobot import select_parquet_files @@ -42,3 +43,12 @@ def test_norm_stats_frame_cap_defaults_to_every_file() -> None: first = select_parquet_files(files, 100, 100, seed=0) assert len(first) == 2 assert select_parquet_files(files, 100, 100, seed=0) == first + + +def test_keep_period_follows_save_interval_by_default() -> None: + # max_to_keep=1 prunes anything outside keep_period, so saves must be covered. + assert resolve_keep_period(4000, None) == 4000 + assert resolve_keep_period(500, None) == 500 + assert resolve_keep_period(4000, 20000) == 20000 + assert resolve_keep_period(4000, 0) is None + assert resolve_keep_period(0, None) is None diff --git a/train_lerobot.py b/train_lerobot.py index 4c60076..4fcb326 100644 --- a/train_lerobot.py +++ b/train_lerobot.py @@ -66,6 +66,20 @@ def _csv_list(value: str | None) -> list[str] | None: return items or None +def resolve_keep_period(save_interval: int, keep_period: int | None) -> int | None: + """Pick ``keep_period`` so a run keeps the checkpoints it writes. + + OpenPI runs the checkpoint manager with ``max_to_keep=1``, so anything not + covered by ``keep_period`` is deleted when the next checkpoint lands. A + ``keep_period`` of 0 means "keep only the newest". + """ + if keep_period is None: + return save_interval or None + if keep_period <= 0: + return None + return keep_period + + def parse_args(): parser = argparse.ArgumentParser(description="Train OpenPI on a mounted LeRobot dataset") parser.add_argument("--dataset_dir", type=pathlib.Path, default=DEFAULT_DATASET_DIR, @@ -94,6 +108,13 @@ def parse_args(): parser.add_argument("--prompt", type=str, default=None, help="Default language prompt when dataset has no tasks") parser.add_argument("--save_interval", type=int, default=500) + parser.add_argument("--keep_period", type=int, default=None, + help="Checkpoints at steps divisible by this are never pruned. " + "Default keeps every saved checkpoint (same as --save_interval); " + "0 keeps only the most recent one.") + parser.add_argument("--resume", action="store_true", + help="Continue from the newest checkpoint in " + "// instead of starting over") parser.add_argument("--learning_rate", type=float, default=None) parser.add_argument("--fsdp_devices", type=str, default="auto", help="FSDP device count: 'auto' (=GPU count when >=2), or integer") @@ -942,6 +963,7 @@ def main(): # ---- assemble TrainConfig ---- peak_lr = args.learning_rate or 2.5e-5 + keep_period = resolve_keep_period(args.save_interval, args.keep_period) config = _config.TrainConfig( name=args.run_name, model=model_config, @@ -952,10 +974,11 @@ def main(): checkpoint_base_dir=str(output_dir), assets_base_dir="/workspace/assets", exp_name=args.exp_name, - overwrite=True, + overwrite=not args.resume, + resume=args.resume, wandb_enabled=False, save_interval=args.save_interval, - keep_period=5000, + keep_period=keep_period, lr_schedule=lr_schedule, num_workers=args.num_workers, fsdp_devices=fsdp_devices, @@ -966,6 +989,7 @@ def main(): logger.info(f"batch_size={args.batch_size} steps={args.steps}") logger.info(f"checkpoint_dir = {config.checkpoint_dir}") logger.info(f"weight source = {weight_path}") + logger.info(f"save_interval = {args.save_interval} keep_period = {keep_period} resume = {args.resume}") norm_max_frames = args.norm_stats_max_frames if args.norm_stats_max_frames > 0 else None manifest = { @@ -988,6 +1012,8 @@ def main(): "steps": args.steps, "learning_rate": peak_lr, "save_interval": args.save_interval, + "keep_period": keep_period, + "resume": args.resume, "norm_stats_max_frames": args.norm_stats_max_frames, "dataset_dir": str(prepared["dataset_dir"]), "effective_dataset_dir": str(effective_dir), From f40940e75f0e471235ac2b5710e2714af0bbefb1 Mon Sep 17 00:00:00 2001 From: joaner <1726541+joaner@users.noreply.github.com> Date: Tue, 22 Sep 2026 19:06:37 +0800 Subject: [PATCH 3/7] Stream per-step metrics instead of block-buffering them. --- test_train_lerobot.py | 6 ++++++ train_lerobot.py | 15 +++++++++++++++ 2 files changed, 21 insertions(+) diff --git a/test_train_lerobot.py b/test_train_lerobot.py index f397801..2185383 100644 --- a/test_train_lerobot.py +++ b/test_train_lerobot.py @@ -5,6 +5,7 @@ import numpy as np from train_lerobot import GenericLeRobotInputs +from train_lerobot import enable_line_buffered_stdout from train_lerobot import resolve_keep_period from train_lerobot import select_parquet_files @@ -52,3 +53,8 @@ def test_keep_period_follows_save_interval_by_default() -> None: assert resolve_keep_period(4000, 20000) == 20000 assert resolve_keep_period(4000, 0) is None assert resolve_keep_period(0, None) is None + + +def test_enable_line_buffered_stdout_tolerates_unbuffered_streams() -> None: + # pytest replaces sys.stdout with a stream that cannot be reconfigured. + enable_line_buffered_stdout() diff --git a/train_lerobot.py b/train_lerobot.py index 4fcb326..bb9e501 100644 --- a/train_lerobot.py +++ b/train_lerobot.py @@ -701,6 +701,20 @@ def __call__(self, x): # Main # --------------------------------------------------------------------------- +def enable_line_buffered_stdout() -> None: + """Flush progress lines as they are written. + + Per-step metrics are emitted with ``tqdm.write``, which does not flush. + When stdout is a pipe (`docker logs`, a redirected log file) Python + block-buffers it, so ``Step N: loss=...`` can stay invisible for thousands + of steps. Progress bars are unaffected because they go through logging. + """ + try: + sys.stdout.reconfigure(line_buffering=True) + except (AttributeError, OSError, ValueError): + pass + + def _cache_dir_for(args, dataset_dir: pathlib.Path, info: dict) -> pathlib.Path: if args.convert_dir is not None: return pathlib.Path(args.convert_dir) @@ -797,6 +811,7 @@ def main(): format="%(asctime)s.%(msecs)03d [%(levelname)s] %(message)s", datefmt="%H:%M:%S", ) + enable_line_buffered_stdout() try: prepared = prepare_dataset(args) From 1d9f892e834bdf01578b5e9a64cdd533ef87fd71 Mon Sep 17 00:00:00 2001 From: joaner <1726541+joaner@users.noreply.github.com> Date: Tue, 22 Sep 2026 21:32:06 +0800 Subject: [PATCH 4/7] Keep converted datasets readable outside the container that built them. --- README.md | 5 +++ README.zh-CN.md | 4 ++ lerobot_v3_compat.py | 83 +++++++++++++++++++++++++++++++++------ test_lerobot_v3_compat.py | 43 ++++++++++++++++++-- 4 files changed, 120 insertions(+), 15 deletions(-) diff --git a/README.md b/README.md index 10f10eb..4576d4b 100644 --- a/README.md +++ b/README.md @@ -138,6 +138,11 @@ directory of symlinks plus a rewritten `meta/info.json`, so LeRobot does not decode dropped cameras. Put the cache on a real disk. The converter will not write to `/tmp`. +The cache and the camera view are self-contained: links between them are +relative and resolve outside the container that built them. Episodes that +cover a whole source video are hardlinked, or copied when the source and the +cache are separate mounts, so no file points back at the dataset mount. + Read-only dataset mounts stay read-only. A v2 dataset that cannot be edited is staged into the cache before metadata fixes. Each run writes `/.run_manifest.json` with the camera map, prompt, and diff --git a/README.zh-CN.md b/README.zh-CN.md index 132f1ad..c9178d7 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -133,6 +133,10 @@ v3 数据集只会在 `--convert_dir` 里转换一次。同一目录、同一组 复用。去掉部分摄像头时,会在旁边做一个符号链接目录,并改写 `meta/info.json`, LeRobot 就不会去解码被丢掉的摄像头。缓存要放在真实磁盘上,转换器不会再写到 `/tmp`。 +缓存和相机视图都是自包含的:它们之间的链接是相对路径,在容器外也能解析。整段 +等于某个源视频的 episode 用硬链接,源与缓存是不同挂载点时就退化为复制,因此没有 +任何文件指回数据集挂载点。 + 只读挂载的数据不会被改写。无法写入的 v2 数据会先在缓存里做一个可写视图,再修补 元数据。每次运行还会在 `/.run_manifest.json` 记下相机映射、 任务文本,以及动作是否保持绝对量。 diff --git a/lerobot_v3_compat.py b/lerobot_v3_compat.py index 3b756eb..5c4ae44 100644 --- a/lerobot_v3_compat.py +++ b/lerobot_v3_compat.py @@ -778,11 +778,41 @@ def convert_videos( if unique_owner: if dest_path.exists() or dest_path.is_symlink(): dest_path.unlink() - os.symlink(str(src_path.resolve()), str(dest_path)) + _link_or_copy_whole_file(src_path, dest_path) continue extract_video(src_path, dest_path, start, end) +def _relative_link_target(target: Path, link: Path) -> str: + """Relative path from the directory holding ``link`` to ``target``. + + Relative links resolve identically on the host and inside a container as + long as the two trees keep their relative layout. An absolute link records + the container-side mount path, so it dangles everywhere else and ties the + cache to one mount layout. + """ + return os.path.relpath(Path(target).resolve(), Path(link).parent.resolve()) + + +def _link_or_copy_whole_file(src_file: Path, dest_file: Path) -> str: + """Publish a source video as an episode clip without re-encoding it. + + A hardlink shares the bytes, so the cache stays self-contained without + duplicating storage. It fails with ``EXDEV`` across mount points (a source + dataset and the cache are usually separate bind mounts), and a read-only + source rejects it too, so fall back to a copy. A symlink is not used here: + the two trees are mounted at unrelated paths, so no relative target exists + and an absolute one would only work under the original mounts. + """ + try: + os.link(src_file, dest_file) + return "hardlink" + except OSError as exc: + logger.debug("Hardlinking %s failed (%s); copying instead", dest_file, exc) + shutil.copy2(src_file, dest_file) + return "copy" + + def _normalize_tasks_list(record: dict[str, Any], task_by_index: dict[int, str]) -> list[str]: raw = record.get("tasks") if isinstance(raw, str) and raw.strip(): @@ -960,7 +990,12 @@ def _reset_dir(path: Path) -> None: def _symlink_children(src_dir: Path, dest_dir: Path) -> None: - """Symlink each file under ``src_dir`` so replacing one link does not write the source.""" + """Link each file under ``src_dir`` into ``dest_dir`` with relative targets. + + Replacing one entry in ``dest_dir`` then leaves the source untouched, and + because the targets are relative the tree also resolves outside the + container that created it. + """ if not src_dir.is_dir(): return for path in sorted(src_dir.rglob("*")): @@ -971,7 +1006,19 @@ def _symlink_children(src_dir: Path, dest_dir: Path) -> None: target.parent.mkdir(parents=True, exist_ok=True) if target.is_symlink() or target.exists(): target.unlink() - os.symlink(path.resolve(), target) + os.symlink(_relative_link_target(path, target), target) + + +def _copy_children(src_dir: Path, dest_dir: Path) -> None: + """Copy each file under ``src_dir`` into ``dest_dir``.""" + if not src_dir.is_dir(): + return + for path in sorted(src_dir.rglob("*")): + if not path.is_file(): + continue + target = dest_dir / path.relative_to(src_dir) + target.parent.mkdir(parents=True, exist_ok=True) + shutil.copy2(path, target) def _filter_stats(stats: dict[str, Any], drop_keys: set[str]) -> dict[str, Any]: @@ -979,10 +1026,13 @@ def _filter_stats(stats: dict[str, Any], drop_keys: set[str]) -> dict[str, Any]: def stage_writable_dataset(src: Path, dest: Path) -> Path: - """Return ``src`` when it can be edited, otherwise a writable tree of symlinks. + """Return ``src`` when it can be edited, otherwise a writable copy to edit. - Parquet files are linked one by one. A later atomic replace updates the - link in ``dest`` and leaves the source dataset untouched. + Metadata and frame data are copied, so the staged tree is writable and + self-contained. Only ``videos/`` is linked, because video data is large; + that link is relative when the source sits inside the same tree as ``dest`` + and absolute otherwise, in which case the staged tree stays valid only + while the source is mounted at the same path. """ src = Path(src) meta = src / "meta" @@ -1009,8 +1059,18 @@ def stage_writable_dataset(src: Path, dest: Path) -> Path: copied.chmod(copied.stat().st_mode | (0o700 if copied.is_dir() else 0o600)) videos = src / "videos" if videos.is_dir(): - os.symlink(videos.resolve(), dest / "videos") - _symlink_children(data, dest / "data") + link_target = _relative_link_target(videos, dest / "videos") + # Climbing several levels means the source sits in a different tree, + # which is normally a different mount. + if link_target.count("..") > 1: + logger.warning( + "Staged dataset links videos with %r; that tree is only valid while the source " + "dataset is mounted at %s", + link_target, + videos, + ) + os.symlink(link_target, dest / "videos") + _copy_children(data, dest / "data") _write_json_stamp(dest / "meta" / ".staged.json", marker) return dest @@ -1030,7 +1090,8 @@ def make_camera_view(src: Path, dest: Path, keep_video_keys: Sequence[str]) -> P return src stamp = { - "stamp_version": 1, + # 2: links are relative, so the view also resolves outside the container. + "stamp_version": 2, "source": str(src.resolve()), "video_keys": selected, } @@ -1044,7 +1105,7 @@ def make_camera_view(src: Path, dest: Path, keep_video_keys: Sequence[str]) -> P _symlink_children(src / "data", dest / "data") videos = src / "videos" if videos.is_dir(): - os.symlink(videos.resolve(), dest / "videos") + os.symlink(_relative_link_target(videos, dest / "videos"), dest / "videos") meta_src = src / "meta" meta_dest = dest / "meta" @@ -1052,7 +1113,7 @@ def make_camera_view(src: Path, dest: Path, keep_video_keys: Sequence[str]) -> P for name in ("tasks.jsonl", "tasks.parquet", "episodes.jsonl"): source_file = meta_src / name if source_file.is_file() or source_file.is_symlink(): - os.symlink(source_file.resolve(), meta_dest / name) + os.symlink(_relative_link_target(source_file, meta_dest / name), meta_dest / name) drop_keys = set(available) - set(selected) viewed = convert_info(info, [], selected, int(info.get("chunks_size") or V2_CHUNKS_SIZE)) diff --git a/test_lerobot_v3_compat.py b/test_lerobot_v3_compat.py index 90a1c93..2a29985 100644 --- a/test_lerobot_v3_compat.py +++ b/test_lerobot_v3_compat.py @@ -3,6 +3,7 @@ from __future__ import annotations import json +import os import shutil import subprocess from pathlib import Path @@ -14,6 +15,7 @@ import numpy as np from lerobot_v3_compat import DatasetLayoutError +from lerobot_v3_compat import _link_or_copy_whole_file from lerobot_v3_compat import _numeric_feature_names from lerobot_v3_compat import assert_v2_local_files from lerobot_v3_compat import assert_v21_episode_stats_rows @@ -239,13 +241,19 @@ def test_convert_packed_v3_writes_v2_chunks(tmp_path: Path) -> None: cam_a_ep1 = (dest / "videos" / "chunk-000" / CAM_A / "episode_000001.mp4").read_bytes() assert cam_a_ep0 == b"file-000.mp4:0.000:1.000" assert cam_a_ep1 == b"file-000.mp4:1.000:2.000" + # A whole-file episode is published as a real file, never a symlink: the source + # dataset and the cache are separate trees, so no relative link exists and an + # absolute one would only work under the mount layout that created it. cam_a_ep2 = dest / "videos" / "chunk-001" / CAM_A / "episode_000002.mp4" - assert cam_a_ep2.is_symlink() - assert cam_a_ep2.resolve().name == "file-001.mp4" + assert cam_a_ep2.is_file() + assert not cam_a_ep2.is_symlink() + assert cam_a_ep2.read_bytes() == (src / "videos" / CAM_A / "chunk-000" / "file-001.mp4").read_bytes() # cam_b file_index is independent of the single data parquet (always file-000). - assert (dest / "videos" / "chunk-000" / CAM_B / "episode_000001.mp4").is_symlink() - assert (dest / "videos" / "chunk-000" / CAM_B / "episode_000001.mp4").resolve().name == "file-001.mp4" + cam_b_ep1 = dest / "videos" / "chunk-000" / CAM_B / "episode_000001.mp4" + assert cam_b_ep1.is_file() + assert not cam_b_ep1.is_symlink() + assert cam_b_ep1.read_bytes() == (src / "videos" / CAM_B / "chunk-000" / "file-001.mp4").read_bytes() tasks = [ json.loads(line) @@ -692,6 +700,33 @@ def test_camera_view_hides_dropped_camera_and_reuses(tmp_path: Path) -> None: assert make_camera_view(full, tmp_path / "unused", [CAM_A, CAM_B]) == full +def test_view_and_cache_links_are_relative(tmp_path: Path) -> None: + """Absolute links would record the container's mount path and dangle outside it.""" + src = build_packed_v3(tmp_path / "v3") + full = convert_v3_to_v2(src, tmp_path / "v2", chunks_size=2, extract_video=_fake_extract) + view = make_camera_view(full, tmp_path / "view", [CAM_B]) + + for tree in (full, view): + links = [p for p in tree.rglob("*") if p.is_symlink()] + absolute = [str(p) for p in links if os.path.isabs(os.readlink(p))] + assert absolute == [], absolute + dangling = [str(p) for p in links if not p.resolve().exists()] + assert dangling == [], dangling + assert (view / "videos").is_dir() + assert list((view / "videos").rglob("*.mp4")) + + +def test_whole_file_video_is_published_without_a_symlink(tmp_path: Path) -> None: + """A cross-mount symlink cannot be relative, so whole files are linked or copied.""" + src = tmp_path / "source.mp4" + src.write_bytes(b"video-bytes" * 8) + dest = tmp_path / "dest.mp4" + outcome = _link_or_copy_whole_file(src, dest) + assert outcome in {"hardlink", "copy"} + assert not dest.is_symlink() + assert dest.read_bytes() == src.read_bytes() + + def test_stage_writable_dataset_does_not_touch_readonly_source(tmp_path: Path) -> None: src = tmp_path / "src" meta = src / "meta" From 1e04cea554006eee2d7af03257a6f103484121a3 Mon Sep 17 00:00:00 2001 From: joaner <1726541+joaner@users.noreply.github.com> Date: Wed, 23 Sep 2026 00:30:57 +0800 Subject: [PATCH 5/7] Fold norm-stats results in a fixed order so runs are reproducible. --- test_train_lerobot.py | 16 ++++++++++++++++ train_lerobot.py | 18 +++++++++++++++--- 2 files changed, 31 insertions(+), 3 deletions(-) diff --git a/test_train_lerobot.py b/test_train_lerobot.py index 2185383..848bf92 100644 --- a/test_train_lerobot.py +++ b/test_train_lerobot.py @@ -5,6 +5,7 @@ import numpy as np from train_lerobot import GenericLeRobotInputs +from train_lerobot import _consume_in_order from train_lerobot import enable_line_buffered_stdout from train_lerobot import resolve_keep_period from train_lerobot import select_parquet_files @@ -58,3 +59,18 @@ def test_keep_period_follows_save_interval_by_default() -> None: def test_enable_line_buffered_stdout_tolerates_unbuffered_streams() -> None: # pytest replaces sys.stdout with a stream that cannot be reconfigured. enable_line_buffered_stdout() + + +def test_fold_results_in_submission_order() -> None: + """Norm stats are order-sensitive, so thread timing must not reorder them.""" + import concurrent.futures + import time + + def finish_later(value: int, delay: float) -> int: + time.sleep(delay) + return value + + with concurrent.futures.ThreadPoolExecutor(max_workers=4) as pool: + # Earliest submissions finish last, so completion order is reversed. + futures = [pool.submit(finish_later, value, 0.02 * (4 - value)) for value in range(5)] + assert list(_consume_in_order(futures)) == [0, 1, 2, 3, 4] diff --git a/train_lerobot.py b/train_lerobot.py index bb9e501..a651603 100644 --- a/train_lerobot.py +++ b/train_lerobot.py @@ -462,6 +462,19 @@ def select_parquet_files( return ordered[:n_files] +def _consume_in_order(futures): + """Yield future results in submission order. + + ``RunningStats`` is order-sensitive twice over: its mean is a running + average, and its quantiles come from a histogram whose grid is anchored on + the first batch it sees. Folding results in completion order therefore made + the statistics depend on thread timing, so two runs over identical data + could bake slightly different norm stats into their checkpoints. + """ + for future in futures: + yield future.result() + + def _compute_norm_stats_fast( config, dataset_dir: pathlib.Path, @@ -543,12 +556,11 @@ def _read_file(pq_path: pathlib.Path): with concurrent.futures.ThreadPoolExecutor(max_workers=num_workers) as pool: future_list = [pool.submit(_read_file, p) for p in files_to_process] pbar = tqdm( - concurrent.futures.as_completed(future_list), + _consume_in_order(future_list), total=len(future_list), desc="norm-stats (fast)", ) - for fut in pbar: - state_arr, action_arr = fut.result() + for state_arr, action_arr in pbar: if state_arr is None: continue state_stats.update(state_arr) From 37a7c2a565358d382cc5828bc5802ce3d2f119cd Mon Sep 17 00:00:00 2001 From: joaner <1726541+joaner@users.noreply.github.com> Date: Wed, 23 Sep 2026 01:30:12 +0800 Subject: [PATCH 6/7] Trigger image CI for any test module, not one filename. --- .github/workflows/docker.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/docker.yml b/.github/workflows/docker.yml index 211efa6..f2873ef 100644 --- a/.github/workflows/docker.yml +++ b/.github/workflows/docker.yml @@ -9,7 +9,7 @@ on: - entrypoint.sh - train_lerobot.py - lerobot_v3_compat.py - - test_lerobot_v3_compat.py + - test_*.py - requirements-test.txt - .dockerignore - .github/workflows/docker.yml @@ -19,7 +19,7 @@ on: - entrypoint.sh - train_lerobot.py - lerobot_v3_compat.py - - test_lerobot_v3_compat.py + - test_*.py - requirements-test.txt - .dockerignore - .github/workflows/docker.yml From 015345e5f1e1aefe66ab1bf0f24cd43b56e5f547 Mon Sep 17 00:00:00 2001 From: joaner <1726541+joaner@users.noreply.github.com> Date: Wed, 23 Sep 2026 18:15:10 +0800 Subject: [PATCH 7/7] Keep long dual-gripper episodes trainable without changing old flags. --delta_joint_actions still leaves only the last dimension absolute. --absolute_action_dims replaces that set for robots with more than one gripper. Delta runs store norm stats of the chunk deltas the model sees, resume reuses checkpoint stats, and constant quantile bins are widened. Multi-hour episodes get exact per-frame timestamps, and the loader tolerance is half a frame so float32 rounding does not abort training while a dropped frame still does. --- README.md | 3 +- README.zh-CN.md | 3 +- lerobot_v3_compat.py | 73 +++++++- test_lerobot_v3_compat.py | 49 +++++ test_train_lerobot.py | 126 +++++++++++++ train_lerobot.py | 379 ++++++++++++++++++++++++++++++++++---- 6 files changed, 592 insertions(+), 41 deletions(-) diff --git a/README.md b/README.md index 4576d4b..46cbd84 100644 --- a/README.md +++ b/README.md @@ -86,7 +86,8 @@ docker run --rm --gpus all --ipc=host \ | `--cameras` | all image keys | comma-separated keys to keep | | `--drop_cameras` | | key or substring, e.g. `front` | | `--camera_map` | role-based | `base=key,left_wrist=key,right_wrist=key` | -| `--delta_joint_actions` | off | joint deltas, last dim (gripper) stays absolute | +| `--delta_joint_actions` | off | joint deltas; only the last dim stays absolute | +| `--absolute_action_dims` | | with `--delta_joint_actions`, indices or action names that stay absolute. Replaces the last-dim default, e.g. `right_gripper,left_gripper` | | `--norm_stats_workers` | `min(cpu, 64)` | | | `--norm_stats_max_frames` | `0` | `0` reads every state/action row | diff --git a/README.zh-CN.md b/README.zh-CN.md index c9178d7..8efe761 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -85,7 +85,8 @@ docker run --rm --gpus all --ipc=host \ | `--cameras` | 全部图像键 | 逗号分隔,只保留这些键 | | `--drop_cameras` | | 键名或子串,例如 `front` | | `--camera_map` | 按名字角色 | `base=键,left_wrist=键,right_wrist=键` | -| `--delta_joint_actions` | 关闭 | 关节用增量,最后一维(夹爪)保持绝对 | +| `--delta_joint_actions` | 关闭 | 关节用增量;不传下一参数时仍只有最后一维保持绝对 | +| `--absolute_action_dims` | | 与 `--delta_joint_actions` 一起用。逗号分隔的下标或动作名保持绝对,并替换「只保留最后一维」的默认,例如 `right_gripper,left_gripper` | | `--norm_stats_workers` | `min(cpu, 64)` | | | `--norm_stats_max_frames` | `0` | `0` 表示读完全部状态/动作 | diff --git a/lerobot_v3_compat.py b/lerobot_v3_compat.py index 5c4ae44..2f941a5 100644 --- a/lerobot_v3_compat.py +++ b/lerobot_v3_compat.py @@ -613,11 +613,75 @@ def convert_info( return v2_info +# LeRobotDataset rejects a step that misses 1/fps by more than this. +TIMESTAMP_TOLERANCE_S = 1e-4 + + +def timestamps_within_tolerance( + values: np.ndarray, + fps: float, + tolerance_s: float = TIMESTAMP_TOLERANCE_S, +) -> bool: + """True when consecutive timestamps match ``1/fps`` inside ``tolerance_s``.""" + series = np.asarray(values, dtype=np.float64).reshape(-1) + if series.size < 2: + return True + if fps <= 0: + return False + return bool(np.all(np.abs(np.diff(series) - (1.0 / float(fps))) <= tolerance_s)) + + +def rewrite_episode_timestamps(table: pa.Table, fps: float) -> pa.Table: + """Replace timestamps that float32 can no longer space at ``1/fps``. + + A multi-hour episode stored as float32 drifts by more than LeRobot's + ``1e-4`` s once the clock is large. ``i / fps`` in float64 stays exact. + Episodes that already pass are returned unchanged. + """ + if "timestamp" not in table.column_names or fps <= 0: + return table + current = np.asarray(table.column("timestamp").to_pylist(), dtype=np.float64).reshape(-1) + if timestamps_within_tolerance(current, fps): + return table + fresh = np.arange(table.num_rows, dtype=np.float64) / float(fps) + index = table.schema.get_field_index("timestamp") + return table.set_column(index, "timestamp", pa.array(fresh, type=pa.float64())) + + +def repair_episode_timestamps(dataset_dir: Path, fps: float) -> int: + """Rewrite episode parquet files whose timestamps miss LeRobot's tolerance. + + Returns the number of files changed. Safe to repeat. + """ + data_root = Path(dataset_dir) / "data" + if fps <= 0 or not data_root.is_dir(): + return 0 + rewritten = 0 + for path in sorted(data_root.glob("**/*.parquet")): + table = pq.read_table(path) + updated = rewrite_episode_timestamps(table, fps) + if updated is table: + continue + tmp_path = path.with_suffix(path.suffix + ".tmp") + pq.write_table(updated, tmp_path) + os.replace(tmp_path, path) + rewritten += 1 + if rewritten: + logger.info( + "Rewrote timestamps in %s episode file(s) at %s so steps stay within %.1e s of 1/fps", + rewritten, + dataset_dir, + TIMESTAMP_TOLERANCE_S, + ) + return rewritten + + def convert_data( src: Path, dest: Path, episode_records: list[dict[str, Any]], chunks_size: int, + fps: float | None = None, ) -> None: grouped = _group_episodes_by_data_file(episode_records) for (chunk_idx, file_idx), records in grouped.items(): @@ -652,6 +716,9 @@ def convert_data( if episode_table.num_rows <= 0: raise ValueError(f"No rows for episode_index={episode_index} in {source_path}") + if fps is not None and fps > 0: + episode_table = rewrite_episode_timestamps(episode_table, fps) + dest_path = dest / v2_data_relpath(episode_index, chunks_size) dest_path.parent.mkdir(parents=True, exist_ok=True) pq.write_table(episode_table, dest_path) @@ -1178,8 +1245,10 @@ def _dropped(key: str) -> bool: if not selected: raise ValueError(f"No image keys left after camera selection. Available: {available}") if len(selected) > len(MODEL_IMAGE_SLOTS): + unmapped = [key for key in selected if _preferred_slot(key) is None] + detail = f" Keys without a base/wrist role: {unmapped}." if unmapped else "" raise ValueError( - f"OpenPI has {len(MODEL_IMAGE_SLOTS)} image slots, got {selected}. " + f"OpenPI has {len(MODEL_IMAGE_SLOTS)} image slots, got {selected}.{detail} " "Pass --cameras or --drop_cameras to choose at most 3." ) return selected @@ -1309,7 +1378,7 @@ def convert_v3_to_v2( write_v21_stats_json(meta_dest / "stats.json", global_stats) _write_jsonl(meta_dest / "tasks.jsonl", tasks) - convert_data(src, dest, episode_records, chunks_size) + convert_data(src, dest, episode_records, chunks_size, fps=float(info.get("fps") or 0)) convert_videos( src, dest, diff --git a/test_lerobot_v3_compat.py b/test_lerobot_v3_compat.py index 2a29985..58c998e 100644 --- a/test_lerobot_v3_compat.py +++ b/test_lerobot_v3_compat.py @@ -28,8 +28,11 @@ from lerobot_v3_compat import load_sanitized_stats_json from lerobot_v3_compat import load_tasks from lerobot_v3_compat import make_camera_view +from lerobot_v3_compat import repair_episode_timestamps +from lerobot_v3_compat import rewrite_episode_timestamps from lerobot_v3_compat import sanitize_episode_stats from lerobot_v3_compat import select_image_keys +from lerobot_v3_compat import timestamps_within_tolerance from lerobot_v3_compat import stage_writable_dataset from lerobot_v3_compat import stats_from_episode_record from lerobot_v3_compat import tasks_have_text @@ -635,6 +638,52 @@ def test_extract_video_segment_starts_at_zero(tmp_path: Path) -> None: assert abs(float(first.split(",")[0])) < 1e-3 +def test_long_float32_episode_timestamps_are_rewritten() -> None: + fps = 30.0 + length = 40_000 + coarse = (np.arange(length, dtype=np.float64) / fps).astype(np.float32) + assert not timestamps_within_tolerance(coarse, fps) + table = pa.table({"timestamp": pa.array(coarse, type=pa.float32())}) + updated = rewrite_episode_timestamps(table, fps) + values = np.asarray(updated.column("timestamp").to_pylist(), dtype=np.float64) + assert timestamps_within_tolerance(values, fps) + assert values[0] == 0.0 + assert abs(values[1] - (1.0 / fps)) < 1e-12 + + short = pa.table({"timestamp": pa.array(np.array([0.0, 1.0 / fps], dtype=np.float64))}) + assert rewrite_episode_timestamps(short, fps) is short + + +def test_repair_episode_timestamps_rewrites_only_bad_files(tmp_path: Path) -> None: + fps = 30.0 + data = tmp_path / "data" / "chunk-000" + data.mkdir(parents=True) + good = np.arange(4, dtype=np.float64) / fps + bad = (np.arange(40_000, dtype=np.float64) / fps).astype(np.float32) + pq.write_table(pa.table({"timestamp": good}), data / "episode_000000.parquet") + pq.write_table(pa.table({"timestamp": pa.array(bad, type=pa.float32())}), data / "episode_000001.parquet") + assert repair_episode_timestamps(tmp_path, fps) == 1 + kept = pq.read_table(data / "episode_000000.parquet").column("timestamp").to_pylist() + assert np.allclose(kept, good) + fixed = np.asarray( + pq.read_table(data / "episode_000001.parquet").column("timestamp").to_pylist(), + dtype=np.float64, + ) + assert timestamps_within_tolerance(fixed, fps) + assert repair_episode_timestamps(tmp_path, fps) == 0 + + +def test_too_many_cameras_names_the_unmapped_key() -> None: + keys = [ + "observation.images.camera_high", + "observation.images.camera_low", + "observation.images.camera_left_wrist", + "observation.images.camera_right_wrist", + ] + with pytest.raises(ValueError, match="camera_low"): + select_image_keys(keys) + + def test_camera_slots_follow_role_names() -> None: three = [FRONT, TOP, WRIST] slots = assign_camera_slots(three) diff --git a/test_train_lerobot.py b/test_train_lerobot.py index 848bf92..5a3328d 100644 --- a/test_train_lerobot.py +++ b/test_train_lerobot.py @@ -1,14 +1,23 @@ """Tests for camera wiring and norm-stat sampling that do not need JAX.""" +import dataclasses from pathlib import Path import numpy as np +import pytest from train_lerobot import GenericLeRobotInputs from train_lerobot import _consume_in_order +from train_lerobot import delta_action_chunks from train_lerobot import enable_line_buffered_stdout +from train_lerobot import install_resume_norm_stats +from train_lerobot import newest_checkpoint_norm_stats +from train_lerobot import relaxed_timestamp_tolerance +from train_lerobot import require_delta_for_absolute_dims +from train_lerobot import resolve_delta_mask from train_lerobot import resolve_keep_period from train_lerobot import select_parquet_files +from train_lerobot import stabilize_norm_stats def test_missing_camera_slot_is_masked() -> None: @@ -61,6 +70,123 @@ def test_enable_line_buffered_stdout_tolerates_unbuffered_streams() -> None: enable_line_buffered_stdout() +WIPE_NAMES = [ + "left_joint1", + "left_joint2", + "left_joint3", + "left_joint4", + "left_joint5", + "left_joint6", + "right_joint1", + "right_joint2", + "right_joint3", + "right_joint4", + "right_joint5", + "right_joint6", + "right_gripper", + "left_gripper", +] + + +def test_timestamp_tolerance_covers_float32_but_not_a_dropped_frame() -> None: + fps = 30.0 + tolerance = relaxed_timestamp_tolerance(fps) + float32_error = 0.033447265625 - (1.0 / fps) + assert float32_error < tolerance + assert (1.0 / fps) > tolerance + assert relaxed_timestamp_tolerance(fps, 0.05) == 0.05 + assert relaxed_timestamp_tolerance(None) == 1e-4 + + +def test_delta_mask_keeps_only_the_last_dim_by_default() -> None: + mask = resolve_delta_mask(14, None, WIPE_NAMES) + assert mask == (True,) * 13 + (False,) + assert resolve_delta_mask(6, None, None) == (True, True, True, True, True, False) + + +def test_absolute_action_dims_replace_the_last_dim_default() -> None: + expected = (True,) * 12 + (False, False) + assert resolve_delta_mask(14, "right_gripper,left_gripper", WIPE_NAMES) == expected + assert resolve_delta_mask(14, "12,13", WIPE_NAMES) == expected + + +def test_absolute_action_dims_require_delta_and_known_names() -> None: + with pytest.raises(ValueError, match="requires --delta_joint_actions"): + require_delta_for_absolute_dims(False, "right_gripper") + require_delta_for_absolute_dims(True, "right_gripper") + require_delta_for_absolute_dims(False, None) + with pytest.raises(ValueError, match="Unknown action dim"): + resolve_delta_mask(14, "not_a_joint", WIPE_NAMES) + + +def test_delta_chunks_subtract_state_and_clamp_the_episode_end() -> None: + state = np.array( + [[0, 0, 10], [1, 0, 10], [2, 0, 10], [3, 0, 10], [4, 0, 10]], + dtype=np.float32, + ) + actions = np.array( + [[0, 1, 10], [1, 1, 11], [2, 1, 12], [3, 1, 13], [4, 1, 14]], + dtype=np.float32, + ) + chunks = delta_action_chunks(state, actions, (True, True, False), 3) + assert chunks.shape == (5, 3, 3) + assert chunks[0, 0].tolist() == [0, 1, 10] + assert chunks[0, 1].tolist() == [1, 1, 11] + assert chunks[0, 2].tolist() == [2, 1, 12] + assert chunks[4, 0].tolist() == [0, 1, 14] + assert chunks[4, 2].tolist() == [0, 1, 14] + raw_mean = actions.mean(axis=0) + chunk_mean = chunks.reshape(-1, 3).mean(axis=0) + assert not np.allclose(raw_mean, chunk_mean) + + +@dataclasses.dataclass +class _Stats: + mean: np.ndarray + std: np.ndarray + q01: np.ndarray | None = None + q99: np.ndarray | None = None + + +def test_quantile_guard_widens_only_constant_dimensions() -> None: + constant = _Stats( + mean=np.array([0.0, 5.0]), + std=np.array([1.0, 2.0]), + q01=np.array([0.0, 0.0]), + q99=np.array([0.0, 4.0]), + ) + updated = stabilize_norm_stats({"actions": constant})["actions"] + assert updated.std is constant.std + assert np.allclose(updated.mean, constant.mean) + assert np.allclose(updated.q01, [-0.5, 0.0]) + assert np.allclose(updated.q99, [0.5, 4.0]) + + varying = _Stats( + mean=np.array([1.0, 2.0]), + std=np.array([0.5, 0.5]), + q01=np.array([0.0, 0.0]), + q99=np.array([2.0, 4.0]), + ) + assert stabilize_norm_stats({"state": varying})["state"] is varying + + +def test_resume_reuses_the_newest_checkpoint_norm_stats(tmp_path: Path) -> None: + older = tmp_path / "1000" / "assets" / "training_dataset" + newer = tmp_path / "2000" / "assets" / "training_dataset" + older.mkdir(parents=True) + newer.mkdir(parents=True) + (older / "norm_stats.json").write_text("old", encoding="utf-8") + (newer / "norm_stats.json").write_text("new", encoding="utf-8") + found = newest_checkpoint_norm_stats(tmp_path, "training_dataset") + assert found is not None + assert found.read_text(encoding="utf-8") == "new" + assets = tmp_path / "assets" / "run" + installed = install_resume_norm_stats(tmp_path, assets, "training_dataset") + assert installed is not None + assert installed.read_text(encoding="utf-8") == "new" + assert newest_checkpoint_norm_stats(tmp_path / "missing", "training_dataset") is None + + def test_fold_results_in_submission_order() -> None: """Norm stats are order-sensitive, so thread timing must not reorder them.""" import concurrent.futures diff --git a/train_lerobot.py b/train_lerobot.py index a651603..4de710a 100644 --- a/train_lerobot.py +++ b/train_lerobot.py @@ -44,6 +44,7 @@ from lerobot_v3_compat import load_tasks from lerobot_v3_compat import make_camera_view from lerobot_v3_compat import parse_camera_map +from lerobot_v3_compat import repair_episode_timestamps from lerobot_v3_compat import select_image_keys from lerobot_v3_compat import stage_writable_dataset from lerobot_v3_compat import tasks_have_text @@ -101,6 +102,10 @@ def parse_args(): parser.add_argument("--delta_joint_actions", action="store_true", help="Train joint dims as deltas and keep the last dim (gripper) absolute. " "Default keeps the dataset action values unchanged.") + parser.add_argument("--absolute_action_dims", type=str, default=None, + help="With --delta_joint_actions, comma-separated indices or action names " + "that stay absolute. Replaces the default of keeping only the last dim. " + "Example: right_gripper,left_gripper or 12,13.") parser.add_argument("--batch_size", type=int, default=1) parser.add_argument("--steps", type=int, default=1000) parser.add_argument("--gpus", type=str, default="all", @@ -321,6 +326,7 @@ def analyze_features(info: dict) -> dict: action_key = None state_dim = 0 action_dim = 0 + action_names = None has_tasks = "task_index" in features for key, feat in features.items(): @@ -339,6 +345,13 @@ def analyze_features(info: dict) -> dict: elif key in ("action", "actions") and action_key is None: action_key = key action_dim = shape[-1] if shape else 0 + raw_names = feat.get("names") + if ( + isinstance(raw_names, list) + and len(raw_names) == action_dim + and all(isinstance(name, str) for name in raw_names) + ): + action_names = list(raw_names) if not image_keys: raise ValueError("No image features found in the dataset") @@ -353,6 +366,7 @@ def analyze_features(info: dict) -> dict: action_key=action_key, state_dim=state_dim, action_dim=action_dim, + action_names=action_names, has_tasks=has_tasks, fps=info.get("fps", 50), ) @@ -432,9 +446,204 @@ def __call__(self, data: dict) -> dict: # --------------------------------------------------------------------------- -# Normalization statistics +# Action deltas and normalization statistics # --------------------------------------------------------------------------- +_DEGENERATE_QUANTILE_SPAN = 1e-6 + + +def relaxed_timestamp_tolerance(fps: float | None, tolerance_s: float = 1e-4) -> float: + """Widen LeRobot's 1e-4 s check just enough for float32 clocks. + + ``torch.tensor`` casts timestamps to float32 before the check. Past a few + thousand seconds the float32 step no longer matches ``1/fps`` within 1e-4, + so a multi-hour episode is rejected even when every frame is present. + Half a frame is still stricter than a dropped frame, whose error is ``1/fps``. + An explicit tolerance above 1e-4 is left alone. + """ + if fps is None or fps <= 0 or tolerance_s > 1e-4: + return tolerance_s + return max(tolerance_s, 0.5 / float(fps)) + + +def relax_lerobot_timestamp_tolerance() -> None: + """Teach the pinned LeRobot to accept float32 timestamp rounding.""" + import lerobot.common.datasets.lerobot_dataset as lerobot_dataset + + if getattr(lerobot_dataset.LeRobotDataset.__init__, "_openpi_relaxed_tolerance", False): + return + original = lerobot_dataset.LeRobotDataset.__init__ + + def _init(self, repo_id, root=None, *args, tolerance_s: float = 1e-4, **kwargs): + info_root = pathlib.Path(root) if root is not None else ( + pathlib.Path(os.environ.get("HF_LEROBOT_HOME", pathlib.Path.home() / ".cache" / "lerobot")) + / str(repo_id) + ) + fps = None + info_path = info_root / "meta" / "info.json" + if info_path.is_file(): + try: + fps = float(json.loads(info_path.read_text(encoding="utf-8")).get("fps") or 0) + except (OSError, json.JSONDecodeError, TypeError, ValueError): + fps = None + tolerance_s = relaxed_timestamp_tolerance(fps, tolerance_s) + return original(self, repo_id, root=root, *args, tolerance_s=tolerance_s, **kwargs) + + _init._openpi_relaxed_tolerance = True + lerobot_dataset.LeRobotDataset.__init__ = _init + + +def require_delta_for_absolute_dims(delta_joint_actions: bool, absolute_action_dims: str | None) -> None: + if absolute_action_dims and absolute_action_dims.strip() and not delta_joint_actions: + raise ValueError("--absolute_action_dims requires --delta_joint_actions") + + +def resolve_delta_mask( + action_dim: int, + absolute_action_dims: str | None, + action_names: list[str] | None, +) -> tuple[bool, ...]: + """Mask for ``DeltaActions``. True means delta, False means absolute. + + Omitting ``absolute_action_dims`` keeps only the last dimension absolute. + A provided list replaces that default instead of adding to it. + """ + if action_dim < 1: + raise ValueError("delta_joint_actions requires a positive action dimension") + if absolute_action_dims is None or not absolute_action_dims.strip(): + mask = [True] * action_dim + mask[-1] = False + return tuple(mask) + + names = list(action_names or []) + absolute: set[int] = set() + for token in absolute_action_dims.split(","): + item = token.strip() + if not item: + continue + if item.lstrip("-").isdigit(): + index = int(item) + if index < 0: + index += action_dim + if not 0 <= index < action_dim: + raise ValueError(f"absolute action index {item} is outside 0..{action_dim - 1}") + absolute.add(index) + continue + matches = [index for index, name in enumerate(names) if name == item] + if not matches: + raise ValueError(f"Unknown action dim {item!r}. Names: {names or 'none'}") + absolute.update(matches) + if not absolute: + raise ValueError("--absolute_action_dims did not select any dimension") + return tuple(index not in absolute for index in range(action_dim)) + + +def delta_action_chunks( + state: np.ndarray, + actions: np.ndarray, + delta_mask: tuple[bool, ...] | list[bool], + horizon: int, +) -> np.ndarray: + """Chunk-relative deltas for one episode, clamped at the last frame. + + ``state`` and ``actions`` are ``(T, D)`` in time order. Masked dimensions of + each ``actions[t:t+horizon]`` subtract ``state[t]``. Past the episode, the + last frame is repeated, which is how LeRobot builds an action chunk. + """ + state = np.asarray(state, dtype=np.float32) + actions = np.asarray(actions, dtype=np.float32) + if state.ndim != 2 or actions.shape != state.shape: + raise ValueError( + f"state and actions must share shape (T, D), got {state.shape} and {actions.shape}" + ) + length, dim = actions.shape + if length < 1: + raise ValueError("episode has no frames") + if horizon < 1: + raise ValueError("action horizon must be positive") + mask = np.asarray(list(delta_mask), dtype=bool) + if mask.shape != (dim,): + raise ValueError(f"delta mask length {mask.shape[0]} != action dim {dim}") + offsets = np.arange(length)[:, None] + np.arange(horizon)[None, :] + np.minimum(offsets, length - 1, out=offsets) + chunks = np.array(actions[offsets], dtype=np.float32, copy=True) + if bool(mask.any()): + chunks[:, :, mask] -= state[:, None, mask] + return chunks + + +def widen_degenerate_quantile_bounds(mean, q01, q99, min_span: float = _DEGENERATE_QUANTILE_SPAN): + """Widen only dimensions whose quantile span is below ``min_span``. + + Unchanged dimensions keep their original values. When nothing is degenerate + the original ``q01`` and ``q99`` objects are returned. + """ + mean_arr = np.asarray(mean, dtype=np.float64) + q01_arr = np.asarray(q01, dtype=np.float64) + q99_arr = np.asarray(q99, dtype=np.float64) + if mean_arr.shape != q01_arr.shape or q01_arr.shape != q99_arr.shape: + raise ValueError("mean, q01, and q99 must share a shape") + degenerate = (q99_arr - q01_arr) < min_span + if not np.any(degenerate): + return q01, q99 + q01_out = np.array(q01_arr, copy=True) + q99_out = np.array(q99_arr, copy=True) + q01_out[degenerate] = mean_arr[degenerate] - 0.5 + q99_out[degenerate] = mean_arr[degenerate] + 0.5 + return q01_out, q99_out + + +def stabilize_norm_stats(norm_stats: dict) -> dict: + """Leave mean and std alone. Widen constant quantile bounds before saving.""" + updated = {} + for key, stats in norm_stats.items(): + q01 = getattr(stats, "q01", None) + q99 = getattr(stats, "q99", None) + if q01 is None or q99 is None: + updated[key] = stats + continue + new_q01, new_q99 = widen_degenerate_quantile_bounds(stats.mean, q01, q99) + if new_q01 is q01 and new_q99 is q99: + updated[key] = stats + continue + updated[key] = dataclasses.replace(stats, q01=new_q01, q99=new_q99) + return updated + + +def newest_checkpoint_norm_stats(checkpoint_dir: pathlib.Path, asset_id: str) -> pathlib.Path | None: + """``norm_stats.json`` from the highest numbered checkpoint step, if present.""" + checkpoint_dir = pathlib.Path(checkpoint_dir) + if not checkpoint_dir.is_dir(): + return None + found: list[tuple[int, pathlib.Path]] = [] + for child in checkpoint_dir.iterdir(): + if not child.is_dir() or not child.name.isdigit(): + continue + stats = child / "assets" / asset_id / "norm_stats.json" + if stats.is_file(): + found.append((int(child.name), stats)) + if not found: + return None + found.sort() + return found[-1][1] + + +def install_resume_norm_stats( + checkpoint_dir: pathlib.Path, + assets_dir: pathlib.Path, + asset_id: str, +) -> pathlib.Path | None: + """Copy saved norm stats so a resumed run keeps the scale it already trained with.""" + source = newest_checkpoint_norm_stats(checkpoint_dir, asset_id) + if source is None: + return None + dest_dir = pathlib.Path(assets_dir) / asset_id + dest_dir.mkdir(parents=True, exist_ok=True) + dest = dest_dir / "norm_stats.json" + shutil.copy2(source, dest) + return dest + + def _rows_per_parquet(path: pathlib.Path, column: str) -> int: import pyarrow.parquet as pq @@ -475,12 +684,65 @@ def _consume_in_order(futures): yield future.result() +def _fold_absolute_rows(state_stats, action_stats, rows) -> int: + """Update stats once per file, in file order. That anchors quantile bins.""" + total = 0 + for state_arr, action_arr, _episode, _frame in rows: + if state_arr is None or len(state_arr) == 0: + continue + state_stats.update(state_arr) + action_stats.update(action_arr) + total += len(state_arr) + return total + + +def _fold_delta_rows(state_stats, action_stats, rows, delta_mask, action_horizon: int) -> int: + """Update action stats with chunk deltas, episodes in index order.""" + states = [] + actions = [] + episodes = [] + frames = [] + for state_arr, action_arr, episode, frame in rows: + if state_arr is None or len(state_arr) == 0: + continue + if episode is None or frame is None: + raise ValueError("delta norm stats require episode_index and frame_index") + states.append(np.asarray(state_arr, dtype=np.float32)) + actions.append(np.asarray(action_arr, dtype=np.float32)) + episodes.append(np.asarray(episode, dtype=np.int64).reshape(-1)) + frames.append(np.asarray(frame, dtype=np.int64).reshape(-1)) + if not states: + return 0 + state = np.concatenate(states, axis=0) + action = np.concatenate(actions, axis=0) + episode = np.concatenate(episodes, axis=0) + frame = np.concatenate(frames, axis=0) + if not (len(state) == len(action) == len(episode) == len(frame)): + raise ValueError("state, action, episode_index, and frame_index row counts disagree") + order = np.lexsort((frame, episode)) + state = state[order] + action = action[order] + episode = episode[order] + boundaries = np.flatnonzero(np.diff(episode)) + 1 + total = 0 + for index in np.split(np.arange(len(episode)), boundaries): + if len(index) == 0: + continue + chunks = delta_action_chunks(state[index], action[index], delta_mask, action_horizon) + action_stats.update(chunks) + state_stats.update(state[index]) + total += len(index) + return total + + def _compute_norm_stats_fast( config, dataset_dir: pathlib.Path, schema: dict, max_frames: int | None, num_workers: int, + delta_mask: tuple[bool, ...] | None = None, + action_horizon: int = 50, ) -> bool: """Compute norm-stats by reading state/action columns directly from parquet files. @@ -489,8 +751,12 @@ def _compute_norm_stats_fast( RunningStats.update() reshapes input to (-1, last_dim), so per-frame parquet data [N, feat_dim] produces statistically equivalent results to the full pipeline's - [N, action_horizon, feat_dim] batches. ``max_frames is None`` reads every row. - A positive cap samples a seeded subset of files. + [N, action_horizon, feat_dim] batches when actions are left absolute. + ``max_frames is None`` reads every row. A positive cap samples a seeded subset of files. + + When ``delta_mask`` is set, action stats are computed on the same chunk deltas the + training transforms produce. Missing episode columns fail this path instead of + saving absolute-action statistics. Returns True on success, False if the fast path cannot be used. """ @@ -514,9 +780,13 @@ def _compute_norm_stats_fast( try: sample_schema = pq.read_schema(parquet_files[0]) available = sample_schema.names - if state_col not in available or action_col not in available: + required = [state_col, action_col] + if delta_mask is not None: + required.extend(["episode_index", "frame_index"]) + missing = [name for name in required if name not in available] + if missing: logger.warning( - f"Parquet columns '{state_col}' or '{action_col}' not found " + f"Parquet columns {missing} not found " f"(available: {available}); skipping fast norm-stats path." ) return False @@ -538,52 +808,54 @@ def _compute_norm_stats_fast( + (f", max_frames={max_frames}" if max_frames else "") ) + read_columns = [state_col, action_col] + if delta_mask is not None: + read_columns.extend(["episode_index", "frame_index"]) + def _read_file(pq_path: pathlib.Path): try: - table = pq.read_table(pq_path, columns=[state_col, action_col]) - state_arr = np.array(table.column(state_col).to_pylist(), dtype=np.float32) - action_arr = np.array(table.column(action_col).to_pylist(), dtype=np.float32) - return state_arr, action_arr + table = pq.read_table(pq_path, columns=read_columns) + state_arr = np.asarray(table.column(state_col).to_pylist(), dtype=np.float32) + action_arr = np.asarray(table.column(action_col).to_pylist(), dtype=np.float32) + if delta_mask is None: + return state_arr, action_arr, None, None + episode = np.asarray(table.column("episode_index").to_pylist(), dtype=np.int64) + frame = np.asarray(table.column("frame_index").to_pylist(), dtype=np.int64) + return state_arr, action_arr, episode, frame except Exception as e: logger.warning(f"Skipping {pq_path.name}: {e}") - return None, None + return None, None, None, None state_stats = normalize.RunningStats() action_stats = normalize.RunningStats() - total_frames = 0 + loaded = [] if num_workers > 1: with concurrent.futures.ThreadPoolExecutor(max_workers=num_workers) as pool: future_list = [pool.submit(_read_file, p) for p in files_to_process] - pbar = tqdm( + loaded = list(tqdm( _consume_in_order(future_list), total=len(future_list), desc="norm-stats (fast)", - ) - for state_arr, action_arr in pbar: - if state_arr is None: - continue - state_stats.update(state_arr) - action_stats.update(action_arr) - total_frames += len(state_arr) - pbar.set_postfix(frames=total_frames) + )) + else: + loaded = [_read_file(pq_path) for pq_path in tqdm(files_to_process, desc="norm-stats (fast)")] + + if delta_mask is None: + total_frames = _fold_absolute_rows(state_stats, action_stats, loaded) else: - for pq_path in tqdm(files_to_process, desc="norm-stats (fast)"): - state_arr, action_arr = _read_file(pq_path) - if state_arr is None: - continue - state_stats.update(state_arr) - action_stats.update(action_arr) - total_frames += len(state_arr) + total_frames = _fold_delta_rows( + state_stats, action_stats, loaded, delta_mask, action_horizon + ) if state_stats._count < 2 or action_stats._count < 2: logger.warning("Not enough frames for fast norm-stats; will fall back to slow path.") return False - norm_stats = { + norm_stats = stabilize_norm_stats({ "state": state_stats.get_statistics(), "actions": action_stats.get_statistics(), - } + }) data_config = config.data.create(config.assets_dirs, config.model) out = config.assets_dirs / data_config.asset_id @@ -598,6 +870,8 @@ def compute_norm_stats( schema: dict, max_frames: int | None = None, num_workers: int = 0, + delta_mask: tuple[bool, ...] | None = None, + action_horizon: int = 50, ) -> None: """Compute and save normalization stats if they don't already exist. @@ -621,7 +895,13 @@ def compute_norm_stats( # --- Fast path: direct parquet reads (skips video decoding entirely) --- try: success = _compute_norm_stats_fast( - config, dataset_dir, schema, max_frames, num_workers + config, + dataset_dir, + schema, + max_frames, + num_workers, + delta_mask=delta_mask, + action_horizon=action_horizon, ) if success: return @@ -706,7 +986,7 @@ def __call__(self, x): out = config.assets_dirs / data_config.asset_id logger.info(f"Saving norm stats → {out}") - normalize.save(out, norm_stats) + normalize.save(out, stabilize_norm_stats(norm_stats)) # --------------------------------------------------------------------------- @@ -768,6 +1048,10 @@ def prepare_dataset(args) -> dict: else: full_dir = stage_writable_dataset(dataset_dir, cache_dir / "writable") + dataset_fps = float(discover_dataset(full_dir).get("fps") or 0) + if dataset_fps > 0: + repair_episode_timestamps(full_dir, dataset_fps) + full_videos = set(video_keys_from_info(discover_dataset(full_dir))) if set(selected) != full_videos: effective_dir = make_camera_view(full_dir, camera_view_dir(cache_dir, selected), selected) @@ -787,7 +1071,10 @@ def prepare_dataset(args) -> dict: logger.info(f" dropped : {dropped or 'none'}") logger.info(f" camera slots: {slots}") logger.info(f" state : {schema['state_key']} dim={schema['state_dim']}") - logger.info(f" action : {schema['action_key']} dim={schema['action_dim']}") + logger.info( + f" action : {schema['action_key']} dim={schema['action_dim']}" + f" names={schema.get('action_names')}" + ) logger.info(f" has tasks : {schema['has_tasks']} task text: {task_texts or has_task_text}") return { "dataset_dir": dataset_dir, @@ -846,6 +1133,7 @@ def main(): logger.info(f"CUDA_VISIBLE_DEVICES = {args.gpus}") # ---- delayed heavy imports ---- + relax_lerobot_timestamp_tolerance() import jax import flax.nnx as nnx import openpi.models.pi0_config as pi0_config @@ -952,11 +1240,18 @@ def main(): output_transforms = [generic_outputs] action_mode = "absolute" delta_mask: tuple | None = None - if args.delta_joint_actions: - if schema["action_dim"] < 1: - logger.error("delta_joint_actions requires a positive action dimension") - sys.exit(1) - delta_mask = _transforms.make_bool_mask(max(schema["action_dim"] - 1, 0), -1) + try: + require_delta_for_absolute_dims(args.delta_joint_actions, args.absolute_action_dims) + if args.delta_joint_actions: + delta_mask = resolve_delta_mask( + schema["action_dim"], + args.absolute_action_dims, + schema.get("action_names"), + ) + except ValueError as exc: + logger.error("%s", exc) + sys.exit(1) + if delta_mask is not None: input_transforms.append(_transforms.DeltaActions(delta_mask)) output_transforms = [_transforms.AbsoluteActions(delta_mask), *output_transforms] action_mode = "delta_joints" @@ -1034,6 +1329,7 @@ def main(): "default_prompt": None if has_task_text else default_prompt, "action_horizon": args.action_horizon, "action_mode": action_mode, + "absolute_action_dims": args.absolute_action_dims, "delta_mask": list(delta_mask) if delta_mask is not None else None, "batch_size": args.batch_size, "steps": args.steps, @@ -1051,12 +1347,21 @@ def main(): write_run_manifest(manifest_sidecar, manifest) # ---- step 1: normalization statistics ---- + asset_id = "training_dataset" + if args.resume: + restored = install_resume_norm_stats(config.checkpoint_dir, config.assets_dirs, asset_id) + if restored is not None: + logger.info("Reusing checkpoint norm stats at %s", restored) + else: + logger.info("Resume requested but no checkpoint norm stats were found; recomputing") compute_norm_stats( config, dataset_dir=effective_dir, schema=schema, max_frames=norm_max_frames, num_workers=args.norm_stats_workers, + delta_mask=delta_mask, + action_horizon=args.action_horizon, ) # ---- step 2: train ----