Skip to content
Open
Show file tree
Hide file tree
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
19 changes: 16 additions & 3 deletions scripts/compute_norm_stats.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,11 +27,18 @@ def create_torch_dataloader(
batch_size: int,
model_config: _model.BaseModelConfig,
num_workers: int,
*,
max_frames: int | None = None,
skip_video_decoding: bool = False,
) -> tuple[_data_loader.Dataset, int]:
if data_config.repo_id is None:
raise ValueError("Data config must have a repo_id")
dataset = _data_loader.create_torch_dataset(data_config, action_horizon, model_config)
dataset = _data_loader.create_torch_dataset(
data_config,
action_horizon,
model_config,
skip_video_decoding=skip_video_decoding,
)
dataset = _data_loader.TransformedDataset(
dataset,
[
Expand Down Expand Up @@ -86,7 +93,7 @@ def create_rlds_dataloader(
return data_loader, num_batches


def main(config_name: str, max_frames: int | None = None):
def main(config_name: str, *, max_frames: int | None = None, skip_video_decoding: bool = False):
config = _config.get_config(config_name)
data_config = config.data.create(config.assets_dirs, config.model)

Expand All @@ -96,7 +103,13 @@ def main(config_name: str, max_frames: int | None = None):
)
else:
data_loader, num_batches = create_torch_dataloader(
data_config, config.model.action_horizon, config.batch_size, config.model, config.num_workers, max_frames
data_config,
config.model.action_horizon,
config.batch_size,
config.model,
config.num_workers,
max_frames=max_frames,
skip_video_decoding=skip_video_decoding,
)

keys = ["state", "actions"]
Expand Down
24 changes: 22 additions & 2 deletions src/openpi/training/data_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,8 +127,27 @@ def __len__(self) -> int:
return self._num_samples


class LeRobotDatasetWithDummyVideos(lerobot_dataset.LeRobotDataset):
"""LeRobot dataset that replaces decoded video frames with tiny placeholders."""

def __init__(self, *args, **kwargs):
kwargs["download_videos"] = False
super().__init__(*args, **kwargs)

def _query_videos(self, query_timestamps: dict[str, list[float]], ep_idx: int) -> dict[str, torch.Tensor]:
del ep_idx
return {
key: torch.zeros((len(timestamps), 3, 1, 1), dtype=torch.uint8).squeeze(0)
for key, timestamps in query_timestamps.items()
}


def create_torch_dataset(
data_config: _config.DataConfig, action_horizon: int, model_config: _model.BaseModelConfig
data_config: _config.DataConfig,
action_horizon: int,
model_config: _model.BaseModelConfig,
*,
skip_video_decoding: bool = False,
) -> Dataset:
"""Create a dataset for training."""
repo_id = data_config.repo_id
Expand All @@ -138,7 +157,8 @@ def create_torch_dataset(
return FakeDataset(model_config, num_samples=1024)

dataset_meta = lerobot_dataset.LeRobotDatasetMetadata(repo_id)
dataset = lerobot_dataset.LeRobotDataset(
dataset_cls = LeRobotDatasetWithDummyVideos if skip_video_decoding else lerobot_dataset.LeRobotDataset
dataset = dataset_cls(
data_config.repo_id,
delta_timestamps={
key: [t / dataset_meta.fps for t in range(action_horizon)] for key in data_config.action_sequence_keys
Expand Down
18 changes: 18 additions & 0 deletions src/openpi/training/data_loader_test.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,30 @@
import dataclasses

import jax
import torch

from openpi.models import pi0_config
from openpi.training import config as _config
from openpi.training import data_loader as _data_loader


def test_dummy_video_dataset_skips_video_download_and_decoding(monkeypatch):
captured_kwargs = {}

def mock_init(self, *args, **kwargs):
del self, args
captured_kwargs.update(kwargs)

monkeypatch.setattr(_data_loader.lerobot_dataset.LeRobotDataset, "__init__", mock_init)
dataset = _data_loader.LeRobotDatasetWithDummyVideos("test")

frames = dataset._query_videos({"camera": [0.0, 0.1]}, ep_idx=0) # noqa: SLF001

assert captured_kwargs["download_videos"] is False
assert frames["camera"].shape == (2, 3, 1, 1)
assert frames["camera"].dtype == torch.uint8


def test_torch_data_loader():
config = pi0_config.Pi0Config(action_dim=24, action_horizon=50, max_token_len=48)
dataset = _data_loader.FakeDataset(config, 16)
Expand Down