Source code for bigym.loco.demos.dataset

"""Read a LeRobot v3 lossless task export back into replay-format episodes.

This is the reader half of :mod:`bigym.loco.demos.lerobot_export`. It reads
the on-disk format directly (pyarrow + PNG decode); the ``lerobot`` package
is never imported, so it works on every Python the core package supports.

Only lossless PNG-mode datasets are accepted: video-mode cameras are lossy
and would silently change training inputs.

Alignment: the export (``meta/alignment.json``, version 2) stores
transition-scoped features shifted so LeRobot frame k pairs obs[k] with the
action executed FROM it; this reader undoes the shift (replay index 0 comes
back from the ``first_transition`` sidecar, the repeated final frame is
dropped), so the episodes come back exactly as the collector wrote them:
row t holds the observation at step t together with the action, reward and
discount of the transition that PRODUCED it (row 0 is the reset row with a
zero action).
"""

from __future__ import annotations

import io
import json
from pathlib import Path
from typing import Any

import numpy as np

from bigym.loco.action_representation import (
    ABSOLUTE,
    convert_episode_action_representation,
    validate_action_representation,
)
from bigym.loco.demos.lerobot_export import ACTION_ALIGNMENT_VERSION

INDEX_FEATURES = ("timestamp", "frame_index", "episode_index", "index", "task_index")
IMAGE_PREFIX = "observation.images."


def load_task_metadata(task_dir: Path) -> dict[str, Any]:
    """Return the collector metadata stored with a task export (root ``metadata.json``)."""
    return json.loads((Path(task_dir) / "metadata.json").read_text())


def load_info(task_dir: Path) -> dict[str, Any]:
    """Return the LeRobot ``meta/info.json`` of a task export."""
    return json.loads((Path(task_dir) / "meta" / "info.json").read_text())


def load_transition_keys(task_dir: Path) -> list[str]:
    """Return the export's transition-scoped column names (``meta/alignment.json``).

    Raises:
        RuntimeError: The export uses another action alignment version.
    """
    alignment = json.loads((Path(task_dir) / "meta" / "alignment.json").read_text())
    version = int(alignment["action_alignment_version"])
    if version != ACTION_ALIGNMENT_VERSION:
        raise RuntimeError(
            f"{task_dir}: action alignment v{version}; this reader reads "
            f"v{ACTION_ALIGNMENT_VERSION}"
        )
    return list(alignment["transition_keys"])


def undo_transition_shift(
    episode: dict[str, np.ndarray],
    first_transition: dict[str, Any],
    transition_keys: list[str],
    where: str = "",
) -> None:
    """Shift an exported episode's transition columns back, in place.

    LeRobot frame k stores the transition executed FROM obs[k]; the
    collector's row t holds the one that PRODUCED obs[t]. Row 0 comes back
    from the ``first_transition`` sidecar and the repeated final frame is
    dropped. Keys the episode does not carry are skipped.

    Args:
        episode: Column name to ``[T, ...]`` array, as read from the export.
        first_transition: The episode's ``first_transition`` entry of
            ``meta/episode_init_states.json``.
        transition_keys: The export's transition-scoped column names.
        where: What to name in the error message.
    """
    for k in transition_keys:
        if k not in episode:
            continue
        rec = first_transition.get(k)
        if rec is None:
            raise RuntimeError(
                f"{where}: alignment v2 needs first_transition[{k!r}] in "
                "episode_init_states.json"
            )
        first_row = np.asarray(rec["data"], dtype=rec["dtype"]).reshape(rec["shape"])
        episode[k] = np.concatenate([first_row[None], episode[k][:-1]], axis=0)


[docs] def load_episodes( task_dir: Path, max_episodes: int = -1, *, action_representation: str = ABSOLUTE, upper_delta_scale_rad: float | None = None, ) -> list[tuple[str, dict[str, np.ndarray]]]: """Reconstruct replay-format episode dicts from one task's LeRobot export. Returns ``(source_name, episode)`` pairs in dataset order. Each episode maps feature names to ``[T, ...]`` arrays: ``rgb_obs`` ``[T, cams, 3, H, W]`` uint8 (camera order from the collector metadata), ``low_dim_obs`` ``[T, D]``, ``action`` ``[T, A]`` (normalized outer action), ``reward``, ``discount``, ``demo``, ``is_expert``, ``event_progress`` ``[T, 1]`` plus any per-step extras the collector stored. With the default ``action_representation="absolute"`` the arrays are the collector's bit for bit. ``"upper_delta"`` derives a training view in memory; the dataset is never modified. """ import pyarrow as pa import pyarrow.parquet as pq from PIL import Image task_dir = Path(task_dir) info = load_info(task_dir) version = str(info.get("codebase_version", "")) if not version.startswith("v3"): raise RuntimeError( f"LeRobot dataset {task_dir} is format {version or 'unknown'}; this " "reader is pinned to v3 (bigym-export-lerobot writes v3)" ) features = info["features"] if any(v.get("dtype") == "video" for v in features.values()): raise RuntimeError( f"LeRobot dataset {task_dir} stores cameras as lossy video; demos " "must be the lossless PNG-mode export (bigym-export-lerobot " "without --videos)" ) transition_keys = load_transition_keys(task_dir) metadata = load_task_metadata(task_dir) cams = list(metadata["task"]["camera_keys"]) action_representation = validate_action_representation(action_representation) init_states_path = task_dir / "meta" / "episode_init_states.json" init_states = ( json.loads(init_states_path.read_text())["episodes"] if init_states_path.exists() else {} ) data_paths = sorted(task_dir.glob("data/chunk-*/file-*.parquet")) if not data_paths: raise RuntimeError(f"No data parquet files under {task_dir}") table = pa.concat_tables([pq.read_table(path) for path in data_paths]) ep_index = table.column("episode_index").to_numpy() episodes: list[tuple[str, dict[str, np.ndarray]]] = [] for ep in np.unique(ep_index): if max_episodes > 0 and len(episodes) >= max_episodes: break rows = table.slice( int(np.searchsorted(ep_index, ep)), int((ep_index == ep).sum()) ) num_rows = rows.num_rows episode: dict[str, np.ndarray] = {} rgb = None for name, spec in features.items(): if name in INDEX_FEATURES: continue if name.startswith(IMAGE_PREFIX): cam_index = cams.index(name[len(IMAGE_PREFIX) :]) col = rows.column(name).to_pylist() if rgb is None: channels, height, width = spec["shape"] rgb = np.empty( (num_rows, len(cams), channels, height, width), dtype=np.uint8 ) for t in range(num_rows): img = np.asarray(Image.open(io.BytesIO(col[t]["bytes"]))) rgb[t, cam_index] = img.transpose(2, 0, 1) continue key = {"observation.state": "low_dim_obs"}.get(name, name) arr = np.stack([np.asarray(v) for v in rows.column(name).to_pylist()]) episode[key] = arr.astype(spec["dtype"]).reshape((num_rows, *spec["shape"])) episode["rgb_obs"] = rgb # ty: ignore[invalid-assignment] ep_extra = init_states.get(str(int(ep)), {}) undo_transition_shift( episode, ep_extra.get("first_transition", {}), transition_keys, where=f"{task_dir} episode {int(ep)}", ) for k, rec in ep_extra.get("arrays", {}).items(): episode[k] = np.asarray(rec["data"], dtype=rec["dtype"]).reshape( rec["shape"] ) episode = convert_episode_action_representation( episode, metadata, action_representation=action_representation, upper_delta_scale_rad=upper_delta_scale_rad, ) episodes.append((ep_extra.get("source_file", f"episode_{int(ep)}"), episode)) return episodes