diff --git a/docs/source/datasets/index.rst b/docs/source/datasets/index.rst index af2db8a..4272537 100644 --- a/docs/source/datasets/index.rst +++ b/docs/source/datasets/index.rst @@ -77,6 +77,7 @@ Available Datasets :caption: Video Datasets something_something_v2 + minerl_treechop .. note:: Documentation is being added progressively, as datasets are ready for usage. Please only use datasets found in the documentation. diff --git a/docs/source/datasets/minerl_treechop.rst b/docs/source/datasets/minerl_treechop.rst new file mode 100644 index 0000000..02154b2 --- /dev/null +++ b/docs/source/datasets/minerl_treechop.rst @@ -0,0 +1,23 @@ +MineRL Treechop +============== + +``MineRLTreechop`` builds a transition-level dataset from locally acquired +``MineRLTreechop-v0`` demonstrations. It does not download or redistribute +MineRL data. Each row stores an RGB POV image, a deterministically flattened +action vector, reward, terminal flag, episode ID, trajectory name, and timestep. + +.. code-block:: python + + from stable_datasets.video import MineRLTreechop + + dataset = MineRLTreechop( + split="train", + data_dir="/path/to/minerl-data", + ).with_format("torch") + + sample = dataset[0] + # image: C x H x W; action: flattened MineRL controls + +The episode and timestep fields preserve trajectory boundaries, allowing +downstream world-model code to construct temporal windows explicitly. Install +MineRL and acquire the source demonstrations under its own terms before use. diff --git a/stable_datasets/tests/video/test_minerl.py b/stable_datasets/tests/video/test_minerl.py new file mode 100644 index 0000000..76c0bd7 --- /dev/null +++ b/stable_datasets/tests/video/test_minerl.py @@ -0,0 +1,54 @@ +"""Tests for MineRL transition caching without a MineRL installation.""" + +from __future__ import annotations + +import numpy as np +import pytest + +from stable_datasets.video.minerl import MineRLTreechop, _flatten_numeric + + +def _transition(value: int, *, done: bool = False): + state = {"pov": np.full((4, 5, 3), value, dtype=np.uint8)} + action = { + "camera": np.array([value, -value], dtype=np.float32), + "forward": value % 2, + "jump": bool(value % 2), + } + return state, action, float(value), state, done + + +class _Pipeline: + def get_trajectory_names(self): + return ["second", "first"] + + def load_data(self, name): + return [_transition(1), _transition(2, done=name == "second")] + + +def test_minerl_treechop_caches_transition_rows(tmp_path, monkeypatch): + monkeypatch.setattr("stable_datasets.video.minerl._make_pipeline", lambda *args: _Pipeline()) + + ds = MineRLTreechop(split="train", processed_cache_dir=tmp_path / "processed") + + assert len(ds) == 4 + sample = ds.with_format("numpy")[0] + assert set(sample) == { + "image", + "action", + "reward", + "done", + "episode_id", + "trajectory_name", + "timestep", + } + assert sample["trajectory_name"] == "first" + assert sample["image"].shape == (4, 5, 3) + np.testing.assert_allclose(sample["action"], [1, -1, 1, 1]) + assert sample["episode_id"] == 0 + assert sample["timestep"] == 0 + + +def test_minerl_actions_reject_non_numeric_values(): + with pytest.raises(TypeError, match="numeric"): + _flatten_numeric({"forward": "yes"}) diff --git a/stable_datasets/video/__init__.py b/stable_datasets/video/__init__.py index c9a531c..d35f109 100644 --- a/stable_datasets/video/__init__.py +++ b/stable_datasets/video/__init__.py @@ -1,7 +1,8 @@ """Video dataset builders.""" from .moving_mnist import MovingMNIST +from .minerl import MineRLTreechop from .something_something_v2 import SomethingSomethingV2, SSv2 -__all__ = ["MovingMNIST", "SSv2", "SomethingSomethingV2"] +__all__ = ["MineRLTreechop", "MovingMNIST", "SSv2", "SomethingSomethingV2"] diff --git a/stable_datasets/video/minerl.py b/stable_datasets/video/minerl.py new file mode 100644 index 0000000..8a100e0 --- /dev/null +++ b/stable_datasets/video/minerl.py @@ -0,0 +1,174 @@ +"""MineRL trajectory builders. + +The MineRL package owns acquisition and decoding of its demonstration archive. +This module deliberately imports it lazily, then turns each transition into a +cacheable StableDataset row. The resulting rows retain episode and timestep +metadata, so downstream users can form temporal windows without guessing at +trajectory boundaries. +""" + +from __future__ import annotations + +from collections.abc import Iterable, Mapping +from pathlib import Path +from typing import Any, Protocol + +import numpy as np + +from stable_datasets.schema import DatasetInfo, DatasetSource, Features, Image, Sequence, Value, Version +from stable_datasets.splits import Split, SplitGenerator +from stable_datasets.utils import BaseDatasetBuilder + + +class MineRLDataPipeline(Protocol): + """The small MineRL data API required by the dataset builder.""" + + def get_trajectory_names(self) -> list[str]: ... + + def load_data(self, stream_name: str) -> Iterable[tuple[Any, ...]]: ... + + +def _flatten_numeric(value: Any) -> np.ndarray: + """Flatten nested numeric MineRL actions in a stable key order.""" + if isinstance(value, Mapping): + parts = [_flatten_numeric(value[key]) for key in sorted(value)] + return np.concatenate(parts) if parts else np.empty(0, dtype=np.float32) + if isinstance(value, tuple | list): + parts = [_flatten_numeric(item) for item in value] + return np.concatenate(parts) if parts else np.empty(0, dtype=np.float32) + + array = np.asarray(value) + if array.dtype.kind not in "biuf": + raise TypeError(f"MineRL action values must be numeric, got {array.dtype!s}.") + return array.astype(np.float32, copy=False).reshape(-1) + + +def _rgb_pov(state: Mapping[str, Any]) -> np.ndarray: + """Return MineRL's ``pov`` observation as HWC uint8 RGB.""" + if "pov" not in state: + raise KeyError("MineRL state has no 'pov' observation.") + frame = np.asarray(state["pov"]) + if frame.ndim != 3: + raise ValueError(f"MineRL state['pov'] must be rank-3, got {frame.shape}.") + if frame.shape[0] in (1, 3) and frame.shape[-1] not in (1, 3): + frame = np.moveaxis(frame, 0, -1) + if frame.shape[-1] != 3: + raise ValueError(f"MineRL state['pov'] must have three channels, got {frame.shape}.") + return frame.astype(np.uint8, copy=False) + + +def _make_pipeline(environment: str, data_dir: Path | None) -> MineRLDataPipeline: + try: + import minerl + except ImportError as exc: + raise ImportError( + "MineRLTreechop requires MineRL. Install MineRL and download " + "MineRLTreechop-v0 before constructing this dataset." + ) from exc + return minerl.data.make(environment, data_dir=str(data_dir) if data_dir else None, num_workers=1) + + +class MineRLTreechop(BaseDatasetBuilder): + """MineRLTreechop-v0 human demonstrations as transition-level rows. + + MineRL is not redistributed. ``data_dir`` must point to a locally acquired + MineRL data root (or MineRL's configured default data directory is used). + Every row is one ``(observation_t, action_t, reward_t, done_t)`` record; + ``episode_id`` and ``timestep`` make temporal reconstruction explicit. + """ + + VERSION = Version("1.0.0") + SOURCE = DatasetSource( + homepage="https://minerl.readthedocs.io/", + assets={}, + license="MineRL data terms apply. This builder does not redistribute demonstrations.", + citation="""@inproceedings{guss2019minerl, + title={MineRL: A Large-Scale Dataset of Minecraft Demonstrations}, + author={Guss, William H. and others}, + booktitle={IJCAI}, + year={2019} +}""", + ) + + def __init__(self, config_name: str | None = None, data_dir: str | Path | None = None, **kwargs): + self.data_dir = Path(data_dir).expanduser() if data_dir is not None else None + super().__init__(config_name=config_name, **kwargs) + + def _info(self) -> DatasetInfo: + return DatasetInfo( + description=( + "MineRLTreechop-v0 transitions with first-person RGB observations and " + "deterministically flattened MineRL actions." + ), + features=Features( + { + "image": Image(), + "action": Sequence(Value("float32")), + "reward": Value("float32"), + "done": Value("bool"), + "episode_id": Value("int32"), + "trajectory_name": Value("string"), + "timestep": Value("int32"), + } + ), + supervised_keys=None, + homepage=self.SOURCE["homepage"], + license=self.SOURCE["license"], + citation=self.SOURCE["citation"], + ) + + def _candidate_splits(self) -> list: + return [Split.TRAIN] + + def _split_generators(self) -> list[SplitGenerator]: + pipeline = _make_pipeline("MineRLTreechop-v0", self.data_dir) + names = sorted(pipeline.get_trajectory_names()) + if not names: + raise ValueError("No MineRLTreechop-v0 trajectories were found.") + return [SplitGenerator(name=Split.TRAIN, gen_kwargs={"pipeline": pipeline, "trajectory_names": names})] + + def _generate_examples( + self, + pipeline: MineRLDataPipeline, + trajectory_names: list[str], + ): + action_dim: int | None = None + try: + for episode_id, trajectory_name in enumerate(trajectory_names): + for timestep, transition in enumerate(pipeline.load_data(trajectory_name)): + if len(transition) < 5: + raise ValueError( + "MineRL transitions must be (state, action, reward, next_state, done)." + ) + state, action, reward, _next_state, done = transition[:5] + if not isinstance(state, Mapping): + raise TypeError("MineRL state must be a mapping containing 'pov'.") + action_vector = _flatten_numeric(action) + if action_dim is None: + action_dim = int(action_vector.size) + elif action_vector.size != action_dim: + raise ValueError( + "MineRL action schema changed within the dataset: " + f"expected {action_dim} values, got {action_vector.size}." + ) + key = f"{episode_id}:{timestep}" + yield key, { + "image": _rgb_pov(state), + "action": action_vector, + "reward": np.float32(reward), + "done": bool(done), + "episode_id": np.int32(episode_id), + "trajectory_name": trajectory_name, + "timestep": np.int32(timestep), + } + finally: + # MineRL's legacy DataPipeline owns a multiprocessing Pool but has + # no public close method. Clean it up after cache construction so + # consumers do not see an interpreter-shutdown warning. + pool = getattr(pipeline, "processing_pool", None) + if pool is not None: + pool.close() + pool.join() + + +__all__ = ["MineRLTreechop"]