diff --git a/scripts/compute_norm_stats.py b/scripts/compute_norm_stats.py index c8aef87222..b39ae9b8e0 100644 --- a/scripts/compute_norm_stats.py +++ b/scripts/compute_norm_stats.py @@ -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, [ @@ -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) @@ -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"] diff --git a/src/openpi/training/data_loader.py b/src/openpi/training/data_loader.py index e2ee7dd06b..4c238e6f6a 100644 --- a/src/openpi/training/data_loader.py +++ b/src/openpi/training/data_loader.py @@ -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 @@ -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 diff --git a/src/openpi/training/data_loader_test.py b/src/openpi/training/data_loader_test.py index d15a73529e..58638d72af 100644 --- a/src/openpi/training/data_loader_test.py +++ b/src/openpi/training/data_loader_test.py @@ -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)