Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
77 changes: 64 additions & 13 deletions source/leisaac/leisaac/enhance/datasets/lerobot_dataset_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,14 @@
from isaaclab.utils import configclass
from isaaclab.utils.datasets.dataset_file_handler_base import DatasetFileHandlerBase
from isaaclab.utils.datasets.episode_data import EpisodeData
from lerobot.datasets.lerobot_dataset import LeRobotDataset

try:
from lerobot.datasets.lerobot_dataset import LeRobotDataset

_HAS_LEROBOT = True
except ImportError:
_HAS_LEROBOT = False
LeRobotDataset = None # Placeholder


@configclass
Expand All @@ -30,23 +37,37 @@ def __init__(self, cfg: LeRobotDatasetCfg):
self._env_args = {}

def create(self, file_path: str, env_name: str = None, resume: bool = False):
if resume:
if _HAS_LEROBOT:
# Original LeRobot logic
if resume:
self._lerobot_dataset = LeRobotDataset(repo_id=self._cfg.repo_id)
else:
self._lerobot_dataset = LeRobotDataset.create(
repo_id=self._cfg.repo_id,
fps=self._cfg.fps,
robot_type=self._cfg.robot_type,
features=self._cfg.features,
)
else:
# Decoupled logic: Use a simple dictionary logger
print("LeRobot not found. Using GenericDataRecorder.")
self._lerobot_dataset = GenericDataRecorder(self._cfg.repo_id, self._cfg.fps, self._cfg.features)

def open(self, file_path: str, mode: str = "r"):
"""Opens an existing dataset for reading or appending."""
if _HAS_LEROBOT:
# Standard LeRobot initialization
self._lerobot_dataset = LeRobotDataset(
repo_id=self._cfg.repo_id,
)
else:
self._lerobot_dataset = LeRobotDataset.create(
repo_id=self._cfg.repo_id,
fps=self._cfg.fps,
robot_type=self._cfg.robot_type,
features=self._cfg.features,
# Decoupled logic: Initialize your custom handler
# You can use file_path here to load existing metadata
self._lerobot_dataset = GenericDataRecorder(
repo_id=self._cfg.repo_id, fps=self._cfg.fps, features=self._cfg.features
)
self._env_args["env_name"] = env_name

def open(self, file_path: str, mode: str = "r"):
self._lerobot_dataset = LeRobotDataset(
repo_id=self._cfg.repo_id,
)
if mode == "r":
self._lerobot_dataset.load_from_disk(file_path)

def get_env_name(self) -> str | None:
return self._env_args["env_name"]
Expand Down Expand Up @@ -77,3 +98,33 @@ def load_episode(self, episode_name: str) -> EpisodeData | None:

def get_num_episodes(self) -> int:
raise NotImplementedError("get_num_episodes is not supported for LeRobotDatasetHandler")


# Create a simple fallback recorder if LeRobot is missing
class GenericDataRecorder:
def __init__(self, repo_id, fps, features):
self.repo_id = repo_id
self.fps = fps
self.features = features
self.buffer = []

def add_frame(self, frame):
# Just store in memory or append to a list
self.buffer.append(frame)

def save_episode(self, **kwargs):
# Save to a standard format like pickle or numpy
import pickle

with open(f"{self.repo_id}_episode.pkl", "wb") as f:
pickle.dump(self.buffer, f)
self.buffer = []

def finalize(self):
print("Dataset finalized.")

def load_from_disk(self, path):
# Implementation to load your custom .npz or .json files
# This ensures that even without LeRobot, your 'open' method
# populates the metadata correctly.
pass